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:
Brace Sproul 2025-07-24 16:01:14 -07:00 • committed by GitHub
parent aec3a8fc1b
commit 4490393494
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 382 additions and 45 deletions

View file

@ -4,12 +4,16 @@ import {
GraphState,
GraphUpdate,
PlanItem,
TaskPlan,
} from "@open-swe/shared/open-swe/types";
import {
checkoutBranchAndCommit,
getChangedFilesStatus,
} 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 { z } from "zod";
import {
@ -25,10 +29,18 @@ import {
getSandboxWithErrorHandling,
} from "../../../utils/sandbox.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 { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools";
import { trackCachePerformance } from "../../../utils/caching.js";
import {
GitHubPullRequest,
GitHubPullRequestList,
GitHubPullRequestUpdate,
} from "../../../utils/github/types.js";
const logger = createLogger(LogLevel.INFO, "Open PR");
@ -81,19 +93,24 @@ export async function openPullRequest(
sandbox,
);
let branchName = state.branchName;
let updatedTaskPlan: TaskPlan | undefined;
if (changedFiles.length > 0) {
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
changedFiles,
});
branchName = await checkoutBranchAndCommit(
const result = await checkoutBranchAndCommit(
config,
state.targetRepository,
sandbox,
{
branchName,
githubInstallationToken,
taskPlan: state.taskPlan,
githubIssueId: state.githubIssueId,
},
);
branchName = result.branchName;
updatedTaskPlan = result.updatedTaskPlan;
}
const openPrTool = createOpenPrToolFields();
@ -132,18 +149,39 @@ export async function openPullRequest(
const { title, body } = toolCall.args as z.infer<typeof openPrTool.schema>;
const pr = await createPullRequest({
owner,
repo,
headBranch: branchName,
title,
body: `Fixes #${state.githubIssueId}\n\n${body}`,
githubInstallationToken,
baseBranch: state.targetRepository.branch,
});
const prForTask = getPullRequestNumberFromActiveTask(
updatedTaskPlan ?? state.taskPlan,
);
let pullRequest:
| GitHubPullRequest
| GitHubPullRequestList[number]
| GitHubPullRequestUpdate
| 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;
if (pr) {
if (pullRequest) {
// Delete the sandbox.
sandboxDeleted = await deleteSandbox(sandboxSessionId);
}
@ -161,12 +199,12 @@ export async function openPullRequest(
new ToolMessage({
id: uuidv4(),
tool_call_id: toolCall.id ?? "",
content: pr
? `Created pull request: ${pr.html_url}`
: "Failed to create pull request.",
content: pullRequest
? `Marked pull request as ready for review: ${pullRequest.html_url}`
: "Failed to mark pull request as ready for review.",
name: toolCall.name,
additional_kwargs: {
pull_request: pr,
pull_request: pullRequest,
},
}),
];
@ -182,5 +220,6 @@ export async function openPullRequest(
...(codebaseTree && { codebaseTree }),
...(dependenciesInstalled !== null && { dependenciesInstalled }),
tokenData: trackCachePerformance(response),
...(updatedTaskPlan && { taskPlan: updatedTaskPlan }),
};
}

View file

@ -11,6 +11,7 @@ import {
GraphState,
GraphConfig,
GraphUpdate,
TaskPlan,
} from "@open-swe/shared/open-swe/types";
import {
checkoutBranchAndCommit,
@ -34,6 +35,8 @@ import { getMcpTools } from "../../../utils/mcp-client.js";
import { shouldDiagnoseError } from "../../../utils/tool-message-error.js";
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.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");
@ -200,20 +203,29 @@ export async function takeAction(
);
let branchName: string | undefined = state.branchName;
let pullRequestNumber: number | undefined;
let updatedTaskPlan: TaskPlan | undefined;
if (changedFiles.length > 0) {
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
changedFiles,
});
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
branchName = await checkoutBranchAndCommit(
const result = await checkoutBranchAndCommit(
config,
state.targetRepository,
sandbox,
{
branchName,
githubInstallationToken,
taskPlan: state.taskPlan,
githubIssueId: state.githubIssueId,
},
);
branchName = result.branchName;
pullRequestNumber = result.updatedTaskPlan
? getActiveTask(result.updatedTaskPlan)?.pullRequestNumber
: undefined;
updatedTaskPlan = result.updatedTaskPlan;
}
const shouldRouteDiagnoseNode = shouldDiagnoseError([
@ -236,10 +248,24 @@ export async function takeAction(
? dependenciesInstalled
: 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 = {
messages: toolCallResults,
messages: userFacingMessagesUpdate,
internalMessages: toolCallResults,
...(branchName && { branchName }),
...(updatedTaskPlan && {
taskPlan: updatedTaskPlan,
}),
codebaseTree: codebaseTreeToReturn,
sandboxSessionId: sandbox.id,
...(dependenciesInstalledUpdate !== null && {

View file

@ -8,7 +8,7 @@ import {
createInstallDependenciesTool,
createShellTool,
} from "../../../tools/index.js";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { GraphConfig, TaskPlan } from "@open-swe/shared/open-swe/types";
import {
ReviewerGraphState,
ReviewerGraphUpdate,
@ -28,6 +28,8 @@ import { Command } from "@langchain/langgraph";
import { shouldDiagnoseError } from "../../../utils/tool-message-error.js";
import { filterHiddenMessages } from "../../../utils/message/filter-hidden.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");
@ -136,21 +138,32 @@ export async function takeReviewerActions(
const toolCallResults = await Promise.all(toolCallResultsPromise);
const repoPath = getRepoAbsolutePath(state.targetRepository);
const changedFiles = await getChangedFilesStatus(repoPath, sandbox);
let branchName: string | undefined = state.branchName;
let pullRequestNumber: number | undefined;
let updatedTaskPlan: TaskPlan | undefined;
if (changedFiles.length > 0) {
logger.info(`Has ${changedFiles.length} changed files. Committing.`, {
changedFiles,
});
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
branchName = await checkoutBranchAndCommit(
const result = await checkoutBranchAndCommit(
config,
state.targetRepository,
sandbox,
{
branchName,
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;
@ -175,10 +188,23 @@ export async function takeReviewerActions(
})),
});
const userFacingMessagesUpdate = [
...toolCallResults,
...(updatedTaskPlan && pullRequestNumber
? createPullRequestToolCallMessage(
state.targetRepository,
pullRequestNumber,
true,
)
: []),
];
const commandUpdate: ReviewerGraphUpdate = {
messages: toolCallResults,
messages: userFacingMessagesUpdate,
reviewerMessages: toolCallResults,
...(branchName && { branchName }),
...(updatedTaskPlan && {
taskPlan: updatedTaskPlan,
}),
...(codebaseTree ? { codebaseTree } : {}),
...(dependenciesInstalledUpdate !== null && {
dependenciesInstalled: dependenciesInstalledUpdate,

View file

@ -1,6 +1,12 @@
import { Octokit } from "@octokit/rest";
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 { getInstallationToken } from "@open-swe/shared/github/auth";
import { getConfig } from "@langchain/langgraph";
@ -93,7 +99,7 @@ async function getExistingPullRequest(
branchName: string,
githubToken: string,
numRetries = 1,
) {
): Promise<GitHubPullRequestList[number] | null> {
return withGitHubRetry(
async (token: string) => {
const octokit = new Octokit({
@ -123,6 +129,8 @@ export async function createPullRequest({
body = "",
githubInstallationToken,
baseBranch,
draft = false,
nullOnError = false,
}: {
owner: string;
repo: string;
@ -131,7 +139,9 @@ export async function createPullRequest({
body?: string;
githubInstallationToken: string;
baseBranch?: string;
}) {
draft?: boolean;
nullOnError?: boolean;
}): Promise<GitHubPullRequest | GitHubPullRequestList[number] | null> {
const octokit = new Octokit({
auth: githubInstallationToken,
});
@ -175,10 +185,12 @@ export async function createPullRequest({
try {
logger.info(
`Creating pull request against default branch: ${repoBaseBranch}`,
{ nullOnError },
);
// Step 2: Create the pull request
const { data: pullRequestData } = await octokit.pulls.create({
draft,
owner,
repo,
title,
@ -190,9 +202,19 @@ export async function createPullRequest({
pullRequest = pullRequestData;
logger.info(`🐙 Pull request created: ${pullRequest.html_url}`);
} catch (error) {
if (nullOnError) {
logger.info("Pull request creation failed, returning null", {
nullOnError,
});
return null;
}
if (error instanceof Error && error.message.includes("already exists")) {
logger.info(
"Pull request already exists. Getting existing pull request...",
{
nullOnError,
},
);
return getExistingPullRequest(
owner,
@ -231,6 +253,37 @@ export async function createPullRequest({
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({
owner,
repo,

View file

@ -1,11 +1,22 @@
import { Sandbox } from "@daytonaio/sdk";
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 { getSandboxErrorFields } from "../sandbox-error-fields.js";
import { getRepoAbsolutePath } from "@open-swe/shared/git";
import { ExecuteResponse } from "@daytonaio/sdk/src/types/ExecuteResponse.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");
@ -84,8 +95,10 @@ export async function checkoutBranchAndCommit(
options: {
branchName?: string;
githubInstallationToken: string;
taskPlan: TaskPlan;
githubIssueId: number;
},
): Promise<string> {
): Promise<{ branchName: string; updatedTaskPlan?: TaskPlan }> {
const absoluteRepoDir = getRepoAbsolutePath(targetRepository);
const branchName = options.branchName || getBranchName(config);
@ -127,13 +140,49 @@ export async function checkoutBranchAndCommit(
};
logger.error("Failed to push changes", errorFields);
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.", {
commitAuthor: userName,
});
return branchName;
return { branchName, updatedTaskPlan };
}
export async function pullLatestChanges(
@ -252,16 +301,11 @@ async function performClone(
branch: branchName,
});
const setUpstreamBranchRes = await sandbox.process.executeCommand(
`git branch --set-upstream-to=origin/${branchName}`,
absoluteRepoDir,
);
if (setUpstreamBranchRes.exitCode !== 0) {
logger.error("Failed to set upstream branch", {
setUpstreamBranchRes,
});
}
logger.info("Set upstream branch");
// push an empty commit so that the branch exists in the remote
await sandbox.git.push(absoluteRepoDir, "git", githubInstallationToken);
logger.info("Pushed empty commit to remote", {
branch: branchName,
});
return branchName;
} catch {
@ -283,8 +327,9 @@ async function performClone(
logger.error("Failed to set upstream branch", {
setUpstreamBranchRes,
});
} else {
logger.info("Set upstream branch");
}
logger.info("Set upstream branch");
return branchName;
}

View file

@ -8,3 +8,9 @@ export type GitHubIssueComment =
export type GitHubPullRequest =
RestEndpointMethodTypes["pulls"]["create"]["response"]["data"];
export type GitHubPullRequestUpdate =
RestEndpointMethodTypes["pulls"]["update"]["response"]["data"];
export type GitHubPullRequestList =
RestEndpointMethodTypes["pulls"]["list"]["response"]["data"];

View 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",
}),
];
}

View file

@ -8,7 +8,15 @@ import {
ChevronDown,
ChevronUp,
ExternalLink,
GitPullRequestDraft,
} from "lucide-react";
import { cn } from "@/lib/utils";
import {
Tooltip,
TooltipContent,
TooltipProvider,
TooltipTrigger,
} from "../ui/tooltip";
type PullRequestOpenedProps = {
status: "loading" | "generating" | "done";
@ -18,6 +26,7 @@ type PullRequestOpenedProps = {
prNumber?: number;
branch?: string;
targetBranch?: string;
isDraft?: boolean;
};
export function PullRequestOpened({
@ -28,6 +37,7 @@ export function PullRequestOpened({
prNumber,
branch,
targetBranch = "main",
isDraft = false,
}: PullRequestOpenedProps) {
const [expanded, setExpanded] = useState(false);
@ -63,10 +73,42 @@ export function PullRequestOpened({
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 (
<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">
<GitPullRequest className="mr-2 h-3.5 w-3.5 text-gray-500 dark:text-gray-400" />
{getIconWithTooltip()}
<div className="flex-1">
{title && status === "done" && (
<div className="mb-0.5 text-xs font-normal text-gray-800 dark:text-gray-200">

View file

@ -516,14 +516,15 @@ export function AssistantMessage({
const status = correspondingToolResult ? "done" : "generating";
const content = correspondingToolResult
? getContentString(correspondingToolResult.content)
: "";
// Extract PR URL from the tool message content
// Format: "Created pull request: https://github.com/owner/repo/pull/123"
let prUrl: string | undefined = undefined;
if (correspondingToolResult) {
const content = getContentString(correspondingToolResult.content);
if (content.includes("Created pull request: ")) {
prUrl = content.split("Created pull request: ")[1].trim();
}
if (content && content.includes("pull request: ")) {
prUrl = content.split("pull request: ")[1].trim();
}
// Extract PR number from URL if available
@ -545,6 +546,7 @@ export function AssistantMessage({
prNumber={prNumber}
branch={branch}
targetBranch={targetBranch}
isDraft={content.includes("Opened draft")}
/>
</div>
);

View file

@ -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
*

View file

@ -120,6 +120,10 @@ export type Task = {
* Optional parent task id if this task was derived from another task
*/
parentTaskId?: string;
/**
* The pull request number associated with this task
*/
pullRequestNumber?: number;
};
export type TaskPlan = {