From 4490393494a780f8821f3cd2ad3ef88dad1646bc Mon Sep 17 00:00:00 2001 From: Brace Sproul Date: Thu, 24 Jul 2025 16:01:14 -0700 Subject: [PATCH] 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] --- .../src/graphs/programmer/nodes/open-pr.ts | 73 ++++++++++++++----- .../graphs/programmer/nodes/take-action.ts | 30 +++++++- .../reviewer/nodes/take-review-action.ts | 32 +++++++- apps/open-swe/src/utils/github/api.ts | 59 ++++++++++++++- apps/open-swe/src/utils/github/git.ts | 73 +++++++++++++++---- apps/open-swe/src/utils/github/types.ts | 6 ++ .../src/utils/message/create-pr-message.ts | 45 ++++++++++++ .../components/gen-ui/pull-request-opened.tsx | 44 ++++++++++- .../web/src/components/thread/messages/ai.tsx | 12 +-- packages/shared/src/open-swe/tasks.ts | 49 +++++++++++++ packages/shared/src/open-swe/types.ts | 4 + 11 files changed, 382 insertions(+), 45 deletions(-) create mode 100644 apps/open-swe/src/utils/message/create-pr-message.ts diff --git a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts index 1d604314..90bf253e 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts @@ -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; - 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 }), }; } diff --git a/apps/open-swe/src/graphs/programmer/nodes/take-action.ts b/apps/open-swe/src/graphs/programmer/nodes/take-action.ts index 55a3dab8..fc6472b0 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/take-action.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/take-action.ts @@ -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 && { diff --git a/apps/open-swe/src/graphs/reviewer/nodes/take-review-action.ts b/apps/open-swe/src/graphs/reviewer/nodes/take-review-action.ts index 213ee35c..f0ff111f 100644 --- a/apps/open-swe/src/graphs/reviewer/nodes/take-review-action.ts +++ b/apps/open-swe/src/graphs/reviewer/nodes/take-review-action.ts @@ -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, diff --git a/apps/open-swe/src/utils/github/api.ts b/apps/open-swe/src/utils/github/api.ts index 2420a3dd..f2b9bbee 100644 --- a/apps/open-swe/src/utils/github/api.ts +++ b/apps/open-swe/src/utils/github/api.ts @@ -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 { 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 { 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 { + 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, diff --git a/apps/open-swe/src/utils/github/git.ts b/apps/open-swe/src/utils/github/git.ts index 3aa2a974..818b6f2d 100644 --- a/apps/open-swe/src/utils/github/git.ts +++ b/apps/open-swe/src/utils/github/git.ts @@ -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 { +): 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; } diff --git a/apps/open-swe/src/utils/github/types.ts b/apps/open-swe/src/utils/github/types.ts index f159de3a..2f889530 100644 --- a/apps/open-swe/src/utils/github/types.ts +++ b/apps/open-swe/src/utils/github/types.ts @@ -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"]; diff --git a/apps/open-swe/src/utils/message/create-pr-message.ts b/apps/open-swe/src/utils/message/create-pr-message.ts new file mode 100644 index 00000000..d709ddb2 --- /dev/null +++ b/apps/open-swe/src/utils/message/create-pr-message.ts @@ -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 = { + 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", + }), + ]; +} diff --git a/apps/web/src/components/gen-ui/pull-request-opened.tsx b/apps/web/src/components/gen-ui/pull-request-opened.tsx index 46981967..95bce078 100644 --- a/apps/web/src/components/gen-ui/pull-request-opened.tsx +++ b/apps/web/src/components/gen-ui/pull-request-opened.tsx @@ -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 ( + + + + {isDraft && ( + + )} + {!isDraft && ( + + )} + + + {isDraft ? "Opened draft pull request" : "Opened pull request"} + + + + ); + }; + return (
- + {getIconWithTooltip()}
{title && status === "done" && (
diff --git a/apps/web/src/components/thread/messages/ai.tsx b/apps/web/src/components/thread/messages/ai.tsx index 99132744..63722031 100644 --- a/apps/web/src/components/thread/messages/ai.tsx +++ b/apps/web/src/components/thread/messages/ai.tsx @@ -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")} />
); diff --git a/packages/shared/src/open-swe/tasks.ts b/packages/shared/src/open-swe/tasks.ts index 7e53e920..b257270e 100644 --- a/packages/shared/src/open-swe/tasks.ts +++ b/packages/shared/src/open-swe/tasks.ts @@ -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 * diff --git a/packages/shared/src/open-swe/types.ts b/packages/shared/src/open-swe/types.ts index c45a0006..56de7210 100644 --- a/packages/shared/src/open-swe/types.ts +++ b/packages/shared/src/open-swe/types.ts @@ -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 = {