mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-04 16:02:13 +00:00
fix: Open draft PRs after first commit (#521)
* Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * format * cr --------- Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
parent
aec3a8fc1b
commit
4490393494
11 changed files with 382 additions and 45 deletions
|
|
@ -4,12 +4,16 @@ import {
|
||||||
GraphState,
|
GraphState,
|
||||||
GraphUpdate,
|
GraphUpdate,
|
||||||
PlanItem,
|
PlanItem,
|
||||||
|
TaskPlan,
|
||||||
} from "@open-swe/shared/open-swe/types";
|
} from "@open-swe/shared/open-swe/types";
|
||||||
import {
|
import {
|
||||||
checkoutBranchAndCommit,
|
checkoutBranchAndCommit,
|
||||||
getChangedFilesStatus,
|
getChangedFilesStatus,
|
||||||
} from "../../../utils/github/git.js";
|
} from "../../../utils/github/git.js";
|
||||||
import { createPullRequest } from "../../../utils/github/api.js";
|
import {
|
||||||
|
createPullRequest,
|
||||||
|
markPullRequestReadyForReview,
|
||||||
|
} from "../../../utils/github/api.js";
|
||||||
import { createLogger, LogLevel } from "../../../utils/logger.js";
|
import { createLogger, LogLevel } from "../../../utils/logger.js";
|
||||||
import { z } from "zod";
|
import { z } from "zod";
|
||||||
import {
|
import {
|
||||||
|
|
@ -25,10 +29,18 @@ import {
|
||||||
getSandboxWithErrorHandling,
|
getSandboxWithErrorHandling,
|
||||||
} from "../../../utils/sandbox.js";
|
} from "../../../utils/sandbox.js";
|
||||||
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
||||||
import { getActivePlanItems } from "@open-swe/shared/open-swe/tasks";
|
import {
|
||||||
|
getActivePlanItems,
|
||||||
|
getPullRequestNumberFromActiveTask,
|
||||||
|
} from "@open-swe/shared/open-swe/tasks";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.js";
|
import { trackCachePerformance } from "../../../utils/caching.js";
|
||||||
|
import {
|
||||||
|
GitHubPullRequest,
|
||||||
|
GitHubPullRequestList,
|
||||||
|
GitHubPullRequestUpdate,
|
||||||
|
} from "../../../utils/github/types.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "Open PR");
|
const logger = createLogger(LogLevel.INFO, "Open PR");
|
||||||
|
|
||||||
|
|
@ -81,19 +93,24 @@ export async function openPullRequest(
|
||||||
sandbox,
|
sandbox,
|
||||||
);
|
);
|
||||||
let branchName = state.branchName;
|
let branchName = state.branchName;
|
||||||
|
let updatedTaskPlan: TaskPlan | undefined;
|
||||||
if (changedFiles.length > 0) {
|
if (changedFiles.length > 0) {
|
||||||
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
|
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
|
||||||
changedFiles,
|
changedFiles,
|
||||||
});
|
});
|
||||||
branchName = await checkoutBranchAndCommit(
|
const result = await checkoutBranchAndCommit(
|
||||||
config,
|
config,
|
||||||
state.targetRepository,
|
state.targetRepository,
|
||||||
sandbox,
|
sandbox,
|
||||||
{
|
{
|
||||||
branchName,
|
branchName,
|
||||||
githubInstallationToken,
|
githubInstallationToken,
|
||||||
|
taskPlan: state.taskPlan,
|
||||||
|
githubIssueId: state.githubIssueId,
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
|
branchName = result.branchName;
|
||||||
|
updatedTaskPlan = result.updatedTaskPlan;
|
||||||
}
|
}
|
||||||
|
|
||||||
const openPrTool = createOpenPrToolFields();
|
const openPrTool = createOpenPrToolFields();
|
||||||
|
|
@ -132,18 +149,39 @@ export async function openPullRequest(
|
||||||
|
|
||||||
const { title, body } = toolCall.args as z.infer<typeof openPrTool.schema>;
|
const { title, body } = toolCall.args as z.infer<typeof openPrTool.schema>;
|
||||||
|
|
||||||
const pr = await createPullRequest({
|
const prForTask = getPullRequestNumberFromActiveTask(
|
||||||
owner,
|
updatedTaskPlan ?? state.taskPlan,
|
||||||
repo,
|
);
|
||||||
headBranch: branchName,
|
let pullRequest:
|
||||||
title,
|
| GitHubPullRequest
|
||||||
body: `Fixes #${state.githubIssueId}\n\n${body}`,
|
| GitHubPullRequestList[number]
|
||||||
githubInstallationToken,
|
| GitHubPullRequestUpdate
|
||||||
baseBranch: state.targetRepository.branch,
|
| null = null;
|
||||||
});
|
if (!prForTask) {
|
||||||
|
// No PR created yet. Shouldn't be possible, but we have a condition here anyway
|
||||||
|
pullRequest = await createPullRequest({
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
headBranch: branchName,
|
||||||
|
title,
|
||||||
|
body: `Fixes #${state.githubIssueId}\n\n${body}`,
|
||||||
|
githubInstallationToken,
|
||||||
|
baseBranch: state.targetRepository.branch,
|
||||||
|
});
|
||||||
|
} else {
|
||||||
|
// Ensure the PR is ready for review
|
||||||
|
pullRequest = await markPullRequestReadyForReview({
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
title,
|
||||||
|
body: `Fixes #${state.githubIssueId}\n\n${body}`,
|
||||||
|
pullNumber: prForTask,
|
||||||
|
githubInstallationToken,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
let sandboxDeleted = false;
|
let sandboxDeleted = false;
|
||||||
if (pr) {
|
if (pullRequest) {
|
||||||
// Delete the sandbox.
|
// Delete the sandbox.
|
||||||
sandboxDeleted = await deleteSandbox(sandboxSessionId);
|
sandboxDeleted = await deleteSandbox(sandboxSessionId);
|
||||||
}
|
}
|
||||||
|
|
@ -161,12 +199,12 @@ export async function openPullRequest(
|
||||||
new ToolMessage({
|
new ToolMessage({
|
||||||
id: uuidv4(),
|
id: uuidv4(),
|
||||||
tool_call_id: toolCall.id ?? "",
|
tool_call_id: toolCall.id ?? "",
|
||||||
content: pr
|
content: pullRequest
|
||||||
? `Created pull request: ${pr.html_url}`
|
? `Marked pull request as ready for review: ${pullRequest.html_url}`
|
||||||
: "Failed to create pull request.",
|
: "Failed to mark pull request as ready for review.",
|
||||||
name: toolCall.name,
|
name: toolCall.name,
|
||||||
additional_kwargs: {
|
additional_kwargs: {
|
||||||
pull_request: pr,
|
pull_request: pullRequest,
|
||||||
},
|
},
|
||||||
}),
|
}),
|
||||||
];
|
];
|
||||||
|
|
@ -182,5 +220,6 @@ export async function openPullRequest(
|
||||||
...(codebaseTree && { codebaseTree }),
|
...(codebaseTree && { codebaseTree }),
|
||||||
...(dependenciesInstalled !== null && { dependenciesInstalled }),
|
...(dependenciesInstalled !== null && { dependenciesInstalled }),
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response),
|
||||||
|
...(updatedTaskPlan && { taskPlan: updatedTaskPlan }),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ import {
|
||||||
GraphState,
|
GraphState,
|
||||||
GraphConfig,
|
GraphConfig,
|
||||||
GraphUpdate,
|
GraphUpdate,
|
||||||
|
TaskPlan,
|
||||||
} from "@open-swe/shared/open-swe/types";
|
} from "@open-swe/shared/open-swe/types";
|
||||||
import {
|
import {
|
||||||
checkoutBranchAndCommit,
|
checkoutBranchAndCommit,
|
||||||
|
|
@ -34,6 +35,8 @@ import { getMcpTools } from "../../../utils/mcp-client.js";
|
||||||
import { shouldDiagnoseError } from "../../../utils/tool-message-error.js";
|
import { shouldDiagnoseError } from "../../../utils/tool-message-error.js";
|
||||||
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
||||||
import { processToolCallContent } from "../../../utils/tool-output-processing.js";
|
import { processToolCallContent } from "../../../utils/tool-output-processing.js";
|
||||||
|
import { getActiveTask } from "@open-swe/shared/open-swe/tasks";
|
||||||
|
import { createPullRequestToolCallMessage } from "../../../utils/message/create-pr-message.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
||||||
|
|
||||||
|
|
@ -200,20 +203,29 @@ export async function takeAction(
|
||||||
);
|
);
|
||||||
|
|
||||||
let branchName: string | undefined = state.branchName;
|
let branchName: string | undefined = state.branchName;
|
||||||
|
let pullRequestNumber: number | undefined;
|
||||||
|
let updatedTaskPlan: TaskPlan | undefined;
|
||||||
if (changedFiles.length > 0) {
|
if (changedFiles.length > 0) {
|
||||||
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
|
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
|
||||||
changedFiles,
|
changedFiles,
|
||||||
});
|
});
|
||||||
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
|
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
|
||||||
branchName = await checkoutBranchAndCommit(
|
const result = await checkoutBranchAndCommit(
|
||||||
config,
|
config,
|
||||||
state.targetRepository,
|
state.targetRepository,
|
||||||
sandbox,
|
sandbox,
|
||||||
{
|
{
|
||||||
branchName,
|
branchName,
|
||||||
githubInstallationToken,
|
githubInstallationToken,
|
||||||
|
taskPlan: state.taskPlan,
|
||||||
|
githubIssueId: state.githubIssueId,
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
|
branchName = result.branchName;
|
||||||
|
pullRequestNumber = result.updatedTaskPlan
|
||||||
|
? getActiveTask(result.updatedTaskPlan)?.pullRequestNumber
|
||||||
|
: undefined;
|
||||||
|
updatedTaskPlan = result.updatedTaskPlan;
|
||||||
}
|
}
|
||||||
|
|
||||||
const shouldRouteDiagnoseNode = shouldDiagnoseError([
|
const shouldRouteDiagnoseNode = shouldDiagnoseError([
|
||||||
|
|
@ -236,10 +248,24 @@ export async function takeAction(
|
||||||
? dependenciesInstalled
|
? dependenciesInstalled
|
||||||
: null;
|
: null;
|
||||||
|
|
||||||
|
// Add the tool call messages for the draft PR to the user facing messages if a draft PR was opened
|
||||||
|
const userFacingMessagesUpdate = [
|
||||||
|
...toolCallResults,
|
||||||
|
...(updatedTaskPlan && pullRequestNumber
|
||||||
|
? createPullRequestToolCallMessage(
|
||||||
|
state.targetRepository,
|
||||||
|
pullRequestNumber,
|
||||||
|
true,
|
||||||
|
)
|
||||||
|
: []),
|
||||||
|
];
|
||||||
const commandUpdate: GraphUpdate = {
|
const commandUpdate: GraphUpdate = {
|
||||||
messages: toolCallResults,
|
messages: userFacingMessagesUpdate,
|
||||||
internalMessages: toolCallResults,
|
internalMessages: toolCallResults,
|
||||||
...(branchName && { branchName }),
|
...(branchName && { branchName }),
|
||||||
|
...(updatedTaskPlan && {
|
||||||
|
taskPlan: updatedTaskPlan,
|
||||||
|
}),
|
||||||
codebaseTree: codebaseTreeToReturn,
|
codebaseTree: codebaseTreeToReturn,
|
||||||
sandboxSessionId: sandbox.id,
|
sandboxSessionId: sandbox.id,
|
||||||
...(dependenciesInstalledUpdate !== null && {
|
...(dependenciesInstalledUpdate !== null && {
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,7 @@ import {
|
||||||
createInstallDependenciesTool,
|
createInstallDependenciesTool,
|
||||||
createShellTool,
|
createShellTool,
|
||||||
} from "../../../tools/index.js";
|
} from "../../../tools/index.js";
|
||||||
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
import { GraphConfig, TaskPlan } from "@open-swe/shared/open-swe/types";
|
||||||
import {
|
import {
|
||||||
ReviewerGraphState,
|
ReviewerGraphState,
|
||||||
ReviewerGraphUpdate,
|
ReviewerGraphUpdate,
|
||||||
|
|
@ -28,6 +28,8 @@ import { Command } from "@langchain/langgraph";
|
||||||
import { shouldDiagnoseError } from "../../../utils/tool-message-error.js";
|
import { shouldDiagnoseError } from "../../../utils/tool-message-error.js";
|
||||||
import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js";
|
import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js";
|
||||||
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
||||||
|
import { getActiveTask } from "@open-swe/shared/open-swe/tasks";
|
||||||
|
import { createPullRequestToolCallMessage } from "../../../utils/message/create-pr-message.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "TakeReviewAction");
|
const logger = createLogger(LogLevel.INFO, "TakeReviewAction");
|
||||||
|
|
||||||
|
|
@ -136,21 +138,32 @@ export async function takeReviewerActions(
|
||||||
const toolCallResults = await Promise.all(toolCallResultsPromise);
|
const toolCallResults = await Promise.all(toolCallResultsPromise);
|
||||||
const repoPath = getRepoAbsolutePath(state.targetRepository);
|
const repoPath = getRepoAbsolutePath(state.targetRepository);
|
||||||
const changedFiles = await getChangedFilesStatus(repoPath, sandbox);
|
const changedFiles = await getChangedFilesStatus(repoPath, sandbox);
|
||||||
|
|
||||||
let branchName: string | undefined = state.branchName;
|
let branchName: string | undefined = state.branchName;
|
||||||
|
let pullRequestNumber: number | undefined;
|
||||||
|
let updatedTaskPlan: TaskPlan | undefined;
|
||||||
|
|
||||||
if (changedFiles.length > 0) {
|
if (changedFiles.length > 0) {
|
||||||
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
|
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
|
||||||
changedFiles,
|
changedFiles,
|
||||||
});
|
});
|
||||||
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
|
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
|
||||||
branchName = await checkoutBranchAndCommit(
|
const result = await checkoutBranchAndCommit(
|
||||||
config,
|
config,
|
||||||
state.targetRepository,
|
state.targetRepository,
|
||||||
sandbox,
|
sandbox,
|
||||||
{
|
{
|
||||||
branchName,
|
branchName,
|
||||||
githubInstallationToken,
|
githubInstallationToken,
|
||||||
|
taskPlan: state.taskPlan,
|
||||||
|
githubIssueId: state.githubIssueId,
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
|
branchName = result.branchName;
|
||||||
|
pullRequestNumber = result.updatedTaskPlan
|
||||||
|
? getActiveTask(result.updatedTaskPlan)?.pullRequestNumber
|
||||||
|
: undefined;
|
||||||
|
updatedTaskPlan = result.updatedTaskPlan;
|
||||||
}
|
}
|
||||||
|
|
||||||
let wereDependenciesInstalled: boolean | null = null;
|
let wereDependenciesInstalled: boolean | null = null;
|
||||||
|
|
@ -175,10 +188,23 @@ export async function takeReviewerActions(
|
||||||
})),
|
})),
|
||||||
});
|
});
|
||||||
|
|
||||||
|
const userFacingMessagesUpdate = [
|
||||||
|
...toolCallResults,
|
||||||
|
...(updatedTaskPlan && pullRequestNumber
|
||||||
|
? createPullRequestToolCallMessage(
|
||||||
|
state.targetRepository,
|
||||||
|
pullRequestNumber,
|
||||||
|
true,
|
||||||
|
)
|
||||||
|
: []),
|
||||||
|
];
|
||||||
const commandUpdate: ReviewerGraphUpdate = {
|
const commandUpdate: ReviewerGraphUpdate = {
|
||||||
messages: toolCallResults,
|
messages: userFacingMessagesUpdate,
|
||||||
reviewerMessages: toolCallResults,
|
reviewerMessages: toolCallResults,
|
||||||
...(branchName && { branchName }),
|
...(branchName && { branchName }),
|
||||||
|
...(updatedTaskPlan && {
|
||||||
|
taskPlan: updatedTaskPlan,
|
||||||
|
}),
|
||||||
...(codebaseTree ? { codebaseTree } : {}),
|
...(codebaseTree ? { codebaseTree } : {}),
|
||||||
...(dependenciesInstalledUpdate !== null && {
|
...(dependenciesInstalledUpdate !== null && {
|
||||||
dependenciesInstalled: dependenciesInstalledUpdate,
|
dependenciesInstalled: dependenciesInstalledUpdate,
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,12 @@
|
||||||
import { Octokit } from "@octokit/rest";
|
import { Octokit } from "@octokit/rest";
|
||||||
import { createLogger, LogLevel } from "../logger.js";
|
import { createLogger, LogLevel } from "../logger.js";
|
||||||
import { GitHubIssue, GitHubIssueComment, GitHubPullRequest } from "./types.js";
|
import {
|
||||||
|
GitHubIssue,
|
||||||
|
GitHubIssueComment,
|
||||||
|
GitHubPullRequest,
|
||||||
|
GitHubPullRequestList,
|
||||||
|
GitHubPullRequestUpdate,
|
||||||
|
} from "./types.js";
|
||||||
import { getOpenSWELabel } from "./label.js";
|
import { getOpenSWELabel } from "./label.js";
|
||||||
import { getInstallationToken } from "@open-swe/shared/github/auth";
|
import { getInstallationToken } from "@open-swe/shared/github/auth";
|
||||||
import { getConfig } from "@langchain/langgraph";
|
import { getConfig } from "@langchain/langgraph";
|
||||||
|
|
@ -93,7 +99,7 @@ async function getExistingPullRequest(
|
||||||
branchName: string,
|
branchName: string,
|
||||||
githubToken: string,
|
githubToken: string,
|
||||||
numRetries = 1,
|
numRetries = 1,
|
||||||
) {
|
): Promise<GitHubPullRequestList[number] | null> {
|
||||||
return withGitHubRetry(
|
return withGitHubRetry(
|
||||||
async (token: string) => {
|
async (token: string) => {
|
||||||
const octokit = new Octokit({
|
const octokit = new Octokit({
|
||||||
|
|
@ -123,6 +129,8 @@ export async function createPullRequest({
|
||||||
body = "",
|
body = "",
|
||||||
githubInstallationToken,
|
githubInstallationToken,
|
||||||
baseBranch,
|
baseBranch,
|
||||||
|
draft = false,
|
||||||
|
nullOnError = false,
|
||||||
}: {
|
}: {
|
||||||
owner: string;
|
owner: string;
|
||||||
repo: string;
|
repo: string;
|
||||||
|
|
@ -131,7 +139,9 @@ export async function createPullRequest({
|
||||||
body?: string;
|
body?: string;
|
||||||
githubInstallationToken: string;
|
githubInstallationToken: string;
|
||||||
baseBranch?: string;
|
baseBranch?: string;
|
||||||
}) {
|
draft?: boolean;
|
||||||
|
nullOnError?: boolean;
|
||||||
|
}): Promise<GitHubPullRequest | GitHubPullRequestList[number] | null> {
|
||||||
const octokit = new Octokit({
|
const octokit = new Octokit({
|
||||||
auth: githubInstallationToken,
|
auth: githubInstallationToken,
|
||||||
});
|
});
|
||||||
|
|
@ -175,10 +185,12 @@ export async function createPullRequest({
|
||||||
try {
|
try {
|
||||||
logger.info(
|
logger.info(
|
||||||
`Creating pull request against default branch: ${repoBaseBranch}`,
|
`Creating pull request against default branch: ${repoBaseBranch}`,
|
||||||
|
{ nullOnError },
|
||||||
);
|
);
|
||||||
|
|
||||||
// Step 2: Create the pull request
|
// Step 2: Create the pull request
|
||||||
const { data: pullRequestData } = await octokit.pulls.create({
|
const { data: pullRequestData } = await octokit.pulls.create({
|
||||||
|
draft,
|
||||||
owner,
|
owner,
|
||||||
repo,
|
repo,
|
||||||
title,
|
title,
|
||||||
|
|
@ -190,9 +202,19 @@ export async function createPullRequest({
|
||||||
pullRequest = pullRequestData;
|
pullRequest = pullRequestData;
|
||||||
logger.info(`🐙 Pull request created: ${pullRequest.html_url}`);
|
logger.info(`🐙 Pull request created: ${pullRequest.html_url}`);
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
|
if (nullOnError) {
|
||||||
|
logger.info("Pull request creation failed, returning null", {
|
||||||
|
nullOnError,
|
||||||
|
});
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
|
||||||
if (error instanceof Error && error.message.includes("already exists")) {
|
if (error instanceof Error && error.message.includes("already exists")) {
|
||||||
logger.info(
|
logger.info(
|
||||||
"Pull request already exists. Getting existing pull request...",
|
"Pull request already exists. Getting existing pull request...",
|
||||||
|
{
|
||||||
|
nullOnError,
|
||||||
|
},
|
||||||
);
|
);
|
||||||
return getExistingPullRequest(
|
return getExistingPullRequest(
|
||||||
owner,
|
owner,
|
||||||
|
|
@ -231,6 +253,37 @@ export async function createPullRequest({
|
||||||
return pullRequest;
|
return pullRequest;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export async function markPullRequestReadyForReview({
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
pullNumber,
|
||||||
|
title,
|
||||||
|
body,
|
||||||
|
githubInstallationToken,
|
||||||
|
}: {
|
||||||
|
owner: string;
|
||||||
|
repo: string;
|
||||||
|
pullNumber: number;
|
||||||
|
title: string;
|
||||||
|
body: string;
|
||||||
|
githubInstallationToken: string;
|
||||||
|
}): Promise<GitHubPullRequestUpdate> {
|
||||||
|
const octokit = new Octokit({
|
||||||
|
auth: githubInstallationToken,
|
||||||
|
});
|
||||||
|
|
||||||
|
const { data: updatedPR } = await octokit.pulls.update({
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
pull_number: pullNumber,
|
||||||
|
title,
|
||||||
|
body,
|
||||||
|
draft: false,
|
||||||
|
});
|
||||||
|
logger.info(`Pull request #${pullNumber} marked as ready for review.`);
|
||||||
|
return updatedPR;
|
||||||
|
}
|
||||||
|
|
||||||
export async function getIssue({
|
export async function getIssue({
|
||||||
owner,
|
owner,
|
||||||
repo,
|
repo,
|
||||||
|
|
|
||||||
|
|
@ -1,11 +1,22 @@
|
||||||
import { Sandbox } from "@daytonaio/sdk";
|
import { Sandbox } from "@daytonaio/sdk";
|
||||||
import { createLogger, LogLevel } from "../logger.js";
|
import { createLogger, LogLevel } from "../logger.js";
|
||||||
import { GraphConfig, TargetRepository } from "@open-swe/shared/open-swe/types";
|
import {
|
||||||
|
GraphConfig,
|
||||||
|
TargetRepository,
|
||||||
|
TaskPlan,
|
||||||
|
} from "@open-swe/shared/open-swe/types";
|
||||||
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
||||||
import { getSandboxErrorFields } from "../sandbox-error-fields.js";
|
import { getSandboxErrorFields } from "../sandbox-error-fields.js";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { ExecuteResponse } from "@daytonaio/sdk/src/types/ExecuteResponse.js";
|
import { ExecuteResponse } from "@daytonaio/sdk/src/types/ExecuteResponse.js";
|
||||||
import { withRetry } from "../retry.js";
|
import { withRetry } from "../retry.js";
|
||||||
|
import {
|
||||||
|
addPullRequestNumberToActiveTask,
|
||||||
|
getActiveTask,
|
||||||
|
getPullRequestNumberFromActiveTask,
|
||||||
|
} from "@open-swe/shared/open-swe/tasks";
|
||||||
|
import { createPullRequest } from "./api.js";
|
||||||
|
import { addTaskPlanToIssue } from "./issue-task.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "GitHub-Git");
|
const logger = createLogger(LogLevel.INFO, "GitHub-Git");
|
||||||
|
|
||||||
|
|
@ -84,8 +95,10 @@ export async function checkoutBranchAndCommit(
|
||||||
options: {
|
options: {
|
||||||
branchName?: string;
|
branchName?: string;
|
||||||
githubInstallationToken: string;
|
githubInstallationToken: string;
|
||||||
|
taskPlan: TaskPlan;
|
||||||
|
githubIssueId: number;
|
||||||
},
|
},
|
||||||
): Promise<string> {
|
): Promise<{ branchName: string; updatedTaskPlan?: TaskPlan }> {
|
||||||
const absoluteRepoDir = getRepoAbsolutePath(targetRepository);
|
const absoluteRepoDir = getRepoAbsolutePath(targetRepository);
|
||||||
const branchName = options.branchName || getBranchName(config);
|
const branchName = options.branchName || getBranchName(config);
|
||||||
|
|
||||||
|
|
@ -127,13 +140,49 @@ export async function checkoutBranchAndCommit(
|
||||||
};
|
};
|
||||||
logger.error("Failed to push changes", errorFields);
|
logger.error("Failed to push changes", errorFields);
|
||||||
throw new Error("Failed to push changes");
|
throw new Error("Failed to push changes");
|
||||||
|
} else {
|
||||||
|
logger.info("Successfully pushed changes");
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if the active task has a PR associated with it. If not, create a draft PR.
|
||||||
|
let updatedTaskPlan: TaskPlan | undefined;
|
||||||
|
const activeTask = getActiveTask(options.taskPlan);
|
||||||
|
const prForTask = getPullRequestNumberFromActiveTask(options.taskPlan);
|
||||||
|
if (!prForTask) {
|
||||||
|
logger.info("First commit detected, creating a draft pull request.");
|
||||||
|
const pullRequest = await createPullRequest({
|
||||||
|
owner: targetRepository.owner,
|
||||||
|
repo: targetRepository.repo,
|
||||||
|
headBranch: branchName,
|
||||||
|
title: `[WIP]: ${activeTask?.title ?? "Open SWE task"}`,
|
||||||
|
body: `**WORK IN PROGRESS OPEN SWE PR**\n\nFixes: #${options.githubIssueId}`,
|
||||||
|
githubInstallationToken: options.githubInstallationToken,
|
||||||
|
draft: true,
|
||||||
|
baseBranch: targetRepository.branch,
|
||||||
|
nullOnError: true,
|
||||||
|
});
|
||||||
|
if (pullRequest) {
|
||||||
|
updatedTaskPlan = addPullRequestNumberToActiveTask(
|
||||||
|
options.taskPlan,
|
||||||
|
pullRequest.number,
|
||||||
|
);
|
||||||
|
await addTaskPlanToIssue(
|
||||||
|
{
|
||||||
|
githubIssueId: options.githubIssueId,
|
||||||
|
targetRepository,
|
||||||
|
},
|
||||||
|
config,
|
||||||
|
updatedTaskPlan,
|
||||||
|
);
|
||||||
|
logger.info(`Draft pull request created: #${pullRequest.number}`);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.info("Successfully checked out & committed changes.", {
|
logger.info("Successfully checked out & committed changes.", {
|
||||||
commitAuthor: userName,
|
commitAuthor: userName,
|
||||||
});
|
});
|
||||||
|
|
||||||
return branchName;
|
return { branchName, updatedTaskPlan };
|
||||||
}
|
}
|
||||||
|
|
||||||
export async function pullLatestChanges(
|
export async function pullLatestChanges(
|
||||||
|
|
@ -252,16 +301,11 @@ async function performClone(
|
||||||
branch: branchName,
|
branch: branchName,
|
||||||
});
|
});
|
||||||
|
|
||||||
const setUpstreamBranchRes = await sandbox.process.executeCommand(
|
// push an empty commit so that the branch exists in the remote
|
||||||
`git branch --set-upstream-to=origin/${branchName}`,
|
await sandbox.git.push(absoluteRepoDir, "git", githubInstallationToken);
|
||||||
absoluteRepoDir,
|
logger.info("Pushed empty commit to remote", {
|
||||||
);
|
branch: branchName,
|
||||||
if (setUpstreamBranchRes.exitCode !== 0) {
|
});
|
||||||
logger.error("Failed to set upstream branch", {
|
|
||||||
setUpstreamBranchRes,
|
|
||||||
});
|
|
||||||
}
|
|
||||||
logger.info("Set upstream branch");
|
|
||||||
|
|
||||||
return branchName;
|
return branchName;
|
||||||
} catch {
|
} catch {
|
||||||
|
|
@ -283,8 +327,9 @@ async function performClone(
|
||||||
logger.error("Failed to set upstream branch", {
|
logger.error("Failed to set upstream branch", {
|
||||||
setUpstreamBranchRes,
|
setUpstreamBranchRes,
|
||||||
});
|
});
|
||||||
|
} else {
|
||||||
|
logger.info("Set upstream branch");
|
||||||
}
|
}
|
||||||
logger.info("Set upstream branch");
|
|
||||||
|
|
||||||
return branchName;
|
return branchName;
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -8,3 +8,9 @@ export type GitHubIssueComment =
|
||||||
|
|
||||||
export type GitHubPullRequest =
|
export type GitHubPullRequest =
|
||||||
RestEndpointMethodTypes["pulls"]["create"]["response"]["data"];
|
RestEndpointMethodTypes["pulls"]["create"]["response"]["data"];
|
||||||
|
|
||||||
|
export type GitHubPullRequestUpdate =
|
||||||
|
RestEndpointMethodTypes["pulls"]["update"]["response"]["data"];
|
||||||
|
|
||||||
|
export type GitHubPullRequestList =
|
||||||
|
RestEndpointMethodTypes["pulls"]["list"]["response"]["data"];
|
||||||
|
|
|
||||||
45
apps/open-swe/src/utils/message/create-pr-message.ts
Normal file
45
apps/open-swe/src/utils/message/create-pr-message.ts
Normal file
|
|
@ -0,0 +1,45 @@
|
||||||
|
import { v4 as uuidv4 } from "uuid";
|
||||||
|
import { AIMessage, BaseMessage, ToolMessage } from "@langchain/core/messages";
|
||||||
|
import { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
|
import { z } from "zod";
|
||||||
|
import { TargetRepository } from "@open-swe/shared/open-swe/types";
|
||||||
|
|
||||||
|
function constructPullRequestUrl(
|
||||||
|
targetRepository: TargetRepository,
|
||||||
|
number: number,
|
||||||
|
) {
|
||||||
|
return `https://github.com/${targetRepository.owner}/${targetRepository.repo}/pull/${number}`;
|
||||||
|
}
|
||||||
|
|
||||||
|
export function createPullRequestToolCallMessage(
|
||||||
|
targetRepository: TargetRepository,
|
||||||
|
number: number,
|
||||||
|
isDraft?: boolean,
|
||||||
|
): BaseMessage[] {
|
||||||
|
const openPrTool = createOpenPrToolFields();
|
||||||
|
const openPrToolArgs: z.infer<typeof openPrTool.schema> = {
|
||||||
|
title: "",
|
||||||
|
body: "",
|
||||||
|
};
|
||||||
|
const toolCallId = uuidv4();
|
||||||
|
return [
|
||||||
|
new AIMessage({
|
||||||
|
id: uuidv4(),
|
||||||
|
content: "",
|
||||||
|
tool_calls: [
|
||||||
|
{
|
||||||
|
name: openPrTool.name,
|
||||||
|
args: openPrToolArgs,
|
||||||
|
id: toolCallId,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}),
|
||||||
|
new ToolMessage({
|
||||||
|
id: uuidv4(),
|
||||||
|
tool_call_id: toolCallId,
|
||||||
|
content: `${isDraft ? "Opened draft" : "Opened"} pull request: ${constructPullRequestUrl(targetRepository, number)}`,
|
||||||
|
name: openPrTool.name,
|
||||||
|
status: "success",
|
||||||
|
}),
|
||||||
|
];
|
||||||
|
}
|
||||||
|
|
@ -8,7 +8,15 @@ import {
|
||||||
ChevronDown,
|
ChevronDown,
|
||||||
ChevronUp,
|
ChevronUp,
|
||||||
ExternalLink,
|
ExternalLink,
|
||||||
|
GitPullRequestDraft,
|
||||||
} from "lucide-react";
|
} from "lucide-react";
|
||||||
|
import { cn } from "@/lib/utils";
|
||||||
|
import {
|
||||||
|
Tooltip,
|
||||||
|
TooltipContent,
|
||||||
|
TooltipProvider,
|
||||||
|
TooltipTrigger,
|
||||||
|
} from "../ui/tooltip";
|
||||||
|
|
||||||
type PullRequestOpenedProps = {
|
type PullRequestOpenedProps = {
|
||||||
status: "loading" | "generating" | "done";
|
status: "loading" | "generating" | "done";
|
||||||
|
|
@ -18,6 +26,7 @@ type PullRequestOpenedProps = {
|
||||||
prNumber?: number;
|
prNumber?: number;
|
||||||
branch?: string;
|
branch?: string;
|
||||||
targetBranch?: string;
|
targetBranch?: string;
|
||||||
|
isDraft?: boolean;
|
||||||
};
|
};
|
||||||
|
|
||||||
export function PullRequestOpened({
|
export function PullRequestOpened({
|
||||||
|
|
@ -28,6 +37,7 @@ export function PullRequestOpened({
|
||||||
prNumber,
|
prNumber,
|
||||||
branch,
|
branch,
|
||||||
targetBranch = "main",
|
targetBranch = "main",
|
||||||
|
isDraft = false,
|
||||||
}: PullRequestOpenedProps) {
|
}: PullRequestOpenedProps) {
|
||||||
const [expanded, setExpanded] = useState(false);
|
const [expanded, setExpanded] = useState(false);
|
||||||
|
|
||||||
|
|
@ -63,10 +73,42 @@ export function PullRequestOpened({
|
||||||
return status === "done" && description;
|
return status === "done" && description;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
const iconClassName = "mr-2 h-3.5 w-3.5";
|
||||||
|
|
||||||
|
const getIconWithTooltip = () => {
|
||||||
|
return (
|
||||||
|
<TooltipProvider>
|
||||||
|
<Tooltip>
|
||||||
|
<TooltipTrigger>
|
||||||
|
{isDraft && (
|
||||||
|
<GitPullRequestDraft
|
||||||
|
className={cn(
|
||||||
|
iconClassName,
|
||||||
|
"text-gray-500 dark:text-gray-400",
|
||||||
|
)}
|
||||||
|
/>
|
||||||
|
)}
|
||||||
|
{!isDraft && (
|
||||||
|
<GitPullRequest
|
||||||
|
className={cn(
|
||||||
|
iconClassName,
|
||||||
|
"text-green-500 dark:text-green-400",
|
||||||
|
)}
|
||||||
|
/>
|
||||||
|
)}
|
||||||
|
</TooltipTrigger>
|
||||||
|
<TooltipContent>
|
||||||
|
{isDraft ? "Opened draft pull request" : "Opened pull request"}
|
||||||
|
</TooltipContent>
|
||||||
|
</Tooltip>
|
||||||
|
</TooltipProvider>
|
||||||
|
);
|
||||||
|
};
|
||||||
|
|
||||||
return (
|
return (
|
||||||
<div className="overflow-hidden rounded-md border border-gray-200 dark:border-gray-700">
|
<div className="overflow-hidden rounded-md border border-gray-200 dark:border-gray-700">
|
||||||
<div className="flex items-center border-b border-gray-200 bg-gray-50 p-2 dark:border-gray-700 dark:bg-gray-800">
|
<div className="flex items-center border-b border-gray-200 bg-gray-50 p-2 dark:border-gray-700 dark:bg-gray-800">
|
||||||
<GitPullRequest className="mr-2 h-3.5 w-3.5 text-gray-500 dark:text-gray-400" />
|
{getIconWithTooltip()}
|
||||||
<div className="flex-1">
|
<div className="flex-1">
|
||||||
{title && status === "done" && (
|
{title && status === "done" && (
|
||||||
<div className="mb-0.5 text-xs font-normal text-gray-800 dark:text-gray-200">
|
<div className="mb-0.5 text-xs font-normal text-gray-800 dark:text-gray-200">
|
||||||
|
|
|
||||||
|
|
@ -516,14 +516,15 @@ export function AssistantMessage({
|
||||||
|
|
||||||
const status = correspondingToolResult ? "done" : "generating";
|
const status = correspondingToolResult ? "done" : "generating";
|
||||||
|
|
||||||
|
const content = correspondingToolResult
|
||||||
|
? getContentString(correspondingToolResult.content)
|
||||||
|
: "";
|
||||||
|
|
||||||
// Extract PR URL from the tool message content
|
// Extract PR URL from the tool message content
|
||||||
// Format: "Created pull request: https://github.com/owner/repo/pull/123"
|
// Format: "Created pull request: https://github.com/owner/repo/pull/123"
|
||||||
let prUrl: string | undefined = undefined;
|
let prUrl: string | undefined = undefined;
|
||||||
if (correspondingToolResult) {
|
if (content && content.includes("pull request: ")) {
|
||||||
const content = getContentString(correspondingToolResult.content);
|
prUrl = content.split("pull request: ")[1].trim();
|
||||||
if (content.includes("Created pull request: ")) {
|
|
||||||
prUrl = content.split("Created pull request: ")[1].trim();
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Extract PR number from URL if available
|
// Extract PR number from URL if available
|
||||||
|
|
@ -545,6 +546,7 @@ export function AssistantMessage({
|
||||||
prNumber={prNumber}
|
prNumber={prNumber}
|
||||||
branch={branch}
|
branch={branch}
|
||||||
targetBranch={targetBranch}
|
targetBranch={targetBranch}
|
||||||
|
isDraft={content.includes("Opened draft")}
|
||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
);
|
);
|
||||||
|
|
|
||||||
|
|
@ -110,6 +110,55 @@ export function updateTaskPlanItems(
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Adds a pull request number to the active task in the task plan.
|
||||||
|
*
|
||||||
|
* @param taskPlan The task plan to update
|
||||||
|
* @param pullRequestNumber The pull request number to add
|
||||||
|
* @returns The updated task plan
|
||||||
|
* @throws Error if the task ID doesn't exist
|
||||||
|
*/
|
||||||
|
export function addPullRequestNumberToActiveTask(
|
||||||
|
taskPlan: TaskPlan,
|
||||||
|
pullRequestNumber: number,
|
||||||
|
): TaskPlan {
|
||||||
|
const activeTaskIndex = taskPlan.activeTaskIndex;
|
||||||
|
const activeTask = taskPlan.tasks[activeTaskIndex];
|
||||||
|
|
||||||
|
if (!activeTask) {
|
||||||
|
throw new Error(`Task with index ${activeTaskIndex} not found`);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create an updated task marked as completed
|
||||||
|
const updatedTask: Task = {
|
||||||
|
...activeTask,
|
||||||
|
pullRequestNumber,
|
||||||
|
};
|
||||||
|
|
||||||
|
// Create a new array of tasks with the updated task
|
||||||
|
const updatedTasks = [...taskPlan.tasks];
|
||||||
|
updatedTasks[activeTaskIndex] = updatedTask;
|
||||||
|
|
||||||
|
// Return the updated task plan
|
||||||
|
return {
|
||||||
|
...taskPlan,
|
||||||
|
tasks: updatedTasks,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets the pull request number from the active task in the task plan.
|
||||||
|
*
|
||||||
|
* @param taskPlan The task plan
|
||||||
|
* @returns The pull request number of the active task, or undefined if the active task has no pull request number
|
||||||
|
*/
|
||||||
|
export function getPullRequestNumberFromActiveTask(
|
||||||
|
taskPlan: TaskPlan,
|
||||||
|
): number | undefined {
|
||||||
|
const activeTask = getActiveTask(taskPlan);
|
||||||
|
return activeTask.pullRequestNumber;
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Helper function to get the active task from a TaskPlan
|
* Helper function to get the active task from a TaskPlan
|
||||||
*
|
*
|
||||||
|
|
|
||||||
|
|
@ -120,6 +120,10 @@ export type Task = {
|
||||||
* Optional parent task id if this task was derived from another task
|
* Optional parent task id if this task was derived from another task
|
||||||
*/
|
*/
|
||||||
parentTaskId?: string;
|
parentTaskId?: string;
|
||||||
|
/**
|
||||||
|
* The pull request number associated with this task
|
||||||
|
*/
|
||||||
|
pullRequestNumber?: number;
|
||||||
};
|
};
|
||||||
|
|
||||||
export type TaskPlan = {
|
export type TaskPlan = {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue