From 2ee04861be04b904e3d1477276a5a3d9d8003817 Mon Sep 17 00:00:00 2001 From: Brace Sproul Date: Sun, 22 Jun 2025 16:35:25 -0700 Subject: [PATCH] feat: Add opened PR tool (#292) * feat: Add opened PR tool * cr --- .../src/graphs/planner/nodes/generate-plan.ts | 8 +- .../src/graphs/planner/nodes/proposed-plan.ts | 14 +- .../src/graphs/programmer/nodes/open-pr.ts | 36 ++--- .../components/gen-ui/pull-request-opened.tsx | 128 ++++++++++++++++++ .../web/src/components/thread/messages/ai.tsx | 61 +++++++++ packages/shared/src/open-swe/planner/types.ts | 7 + packages/shared/src/open-swe/tasks.ts | 15 +- packages/shared/src/open-swe/tools.ts | 27 ++++ packages/shared/src/open-swe/types.ts | 4 + 9 files changed, 270 insertions(+), 30 deletions(-) create mode 100644 apps/web/src/components/gen-ui/pull-request-opened.tsx diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-plan.ts b/apps/open-swe/src/graphs/planner/nodes/generate-plan.ts index b8843bea..8ebd31be 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-plan.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-plan.ts @@ -13,6 +13,7 @@ import { } from "../utils/followup.js"; import { stopSandbox } from "../../../utils/sandbox.js"; import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js"; +import { z } from "zod"; const systemPrompt = `You are a terminal-based agentic coding assistant built by LangChain, designed to enable natural language interaction with local codebases through wrapped LLM models. @@ -123,9 +124,14 @@ export async function generatePlan( newSessionId = await stopSandbox(state.sandboxSessionId); } + const proposedPlanArgs = response.tool_calls[0].args as z.infer< + typeof sessionPlanTool.schema + >; + return { messages: [response], - proposedPlan: response.tool_calls[0].args.plan, + proposedPlanTitle: proposedPlanArgs.title, + proposedPlan: proposedPlanArgs.plan, ...(newSessionId && { sandboxSessionId: newSessionId }), }; } diff --git a/apps/open-swe/src/graphs/planner/nodes/proposed-plan.ts b/apps/open-swe/src/graphs/planner/nodes/proposed-plan.ts index 043c891d..951e302a 100644 --- a/apps/open-swe/src/graphs/planner/nodes/proposed-plan.ts +++ b/apps/open-swe/src/graphs/planner/nodes/proposed-plan.ts @@ -92,7 +92,12 @@ export async function interruptProposedPlan( completed: false, })); - runInput.taskPlan = createNewTask(userRequest, planItems, state.taskPlan); + runInput.taskPlan = createNewTask( + userRequest, + state.proposedPlanTitle, + planItems, + { existingTaskPlan: state.taskPlan }, + ); } else if (interruptRes.type === "edit") { const editedPlan = (interruptRes.args as ActionRequest).args.plan .split(PLAN_INTERRUPT_DELIMITER) @@ -104,7 +109,12 @@ export async function interruptProposedPlan( completed: false, })); - runInput.taskPlan = createNewTask(userRequest, planItems, state.taskPlan); + runInput.taskPlan = createNewTask( + userRequest, + state.proposedPlanTitle, + planItems, + { existingTaskPlan: state.taskPlan }, + ); } else { throw new Error("Unknown interrupt type." + interruptRes.type); } 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 98a2a000..b9aabc77 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts @@ -14,11 +14,12 @@ import { z } from "zod"; import { loadModel, Task } from "../../../utils/load-model.js"; import { formatPlanPromptWithSummaries } from "../../../utils/plan-prompt.js"; import { getUserRequest } from "../../../utils/user-request.js"; -import { ToolMessage } from "@langchain/core/messages"; +import { AIMessage, ToolMessage } from "@langchain/core/messages"; import { daytonaClient, deleteSandbox } from "../../../utils/sandbox.js"; import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js"; import { getActivePlanItems } from "@open-swe/shared/open-swe/tasks"; import { getRepoAbsolutePath } from "@open-swe/shared/git"; +import { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools"; const logger = createLogger(LogLevel.INFO, "Open PR"); @@ -34,26 +35,6 @@ And here is the user's original request: With all of this in mind, please use the \`open_pr\` tool to open a pull request.`; -const openPrToolSchema = z.object({ - title: z - .string() - .describe( - "The title of the pull request. Ensure this is a concise and thoughtful title. You should follow conventional commit title format (e.g. prefixing with '[open-swe] fix:', '[open-swe] feat:', '[open-swe] chore:', etc.). Remember to include the '[open-swe]' prefix in the title.", - ), - body: z - .string() - .optional() - .describe( - "The body of the pull request. This should provide a concise description what the PR changes. Do not over-explain, or add technical details unless they're the absolute minimum needed. The user should be able to quickly read your description, and understand what the PR does. Remember: if they want the technical details they can read the changed files, so you don't need to go into great detail here.", - ), -}); - -const openPrTool = { - name: "open_pr", - schema: openPrToolSchema, - description: "Use this tool to open a pull request.", -}; - const formatPrompt = (taskPlan: PlanItem[], userRequest: string): string => { const completedTasks = taskPlan.filter((task) => task.completed); return openPrSysPrompt @@ -102,6 +83,7 @@ export async function openPullRequest( ); } + const openPrTool = createOpenPrToolFields(); const model = await loadModel(config, Task.SUMMARIZER); const modelWithTool = model.bindTools([openPrTool], { tool_choice: openPrTool.name, @@ -124,7 +106,7 @@ export async function openPullRequest( ); } - const { title, body } = toolCall.args as z.infer; + const { title, body } = toolCall.args as z.infer; const pr = await createPullRequest({ owner, @@ -143,7 +125,15 @@ export async function openPullRequest( } const newMessages = [ - response, + new AIMessage({ + ...response, + additional_kwargs: { + ...response.additional_kwargs, + // Required for the UI to render these fields. + branch: branchName, + targetBranch: state.targetRepository.branch, + }, + }), new ToolMessage({ tool_call_id: toolCall.id ?? "", content: pr diff --git a/apps/web/src/components/gen-ui/pull-request-opened.tsx b/apps/web/src/components/gen-ui/pull-request-opened.tsx new file mode 100644 index 00000000..5dc3167c --- /dev/null +++ b/apps/web/src/components/gen-ui/pull-request-opened.tsx @@ -0,0 +1,128 @@ +"use client"; + +import { useState } from "react"; +import { + GitPullRequest, + Loader2, + CheckCircle, + ChevronDown, + ChevronUp, + ExternalLink, +} from "lucide-react"; + +type PullRequestOpenedProps = { + status: "loading" | "generating" | "done"; + title?: string; + description?: string; + url?: string; + prNumber?: number; + branch?: string; + targetBranch?: string; +}; + +export function PullRequestOpened({ + status, + title, + description, + url, + prNumber, + branch, + targetBranch = "main", +}: PullRequestOpenedProps) { + const [expanded, setExpanded] = useState(false); + + const getStatusIcon = () => { + switch (status) { + case "loading": + return ( +
+ ); + case "generating": + return ; + case "done": + return ; + } + }; + + const getStatusText = () => { + switch (status) { + case "loading": + return "Preparing pull request..."; + case "generating": + return "Opening pull request..."; + case "done": + return prNumber + ? `Pull request #${prNumber} opened` + : "Pull request opened"; + } + }; + + const shouldShowToggle = () => { + return status === "done" && description; + }; + + return ( +
+
+ +
+ {title && status === "done" && ( +
+ {title} +
+ )} + {branch && status === "done" && ( +
+ {branch} → {targetBranch} +
+ )} + {!title && ( + + {getStatusText()} + + )} +
+
+ + {getStatusText()} + + {getStatusIcon()} + {url && status === "done" && ( + + + + )} + {shouldShowToggle() && ( + + )} +
+
+ + {expanded && description && status === "done" && ( +
+

+ Description +

+
+ {description} +
+
+ )} +
+ ); +} diff --git a/apps/web/src/components/thread/messages/ai.tsx b/apps/web/src/components/thread/messages/ai.tsx index 5465e27c..9ee3db61 100644 --- a/apps/web/src/components/thread/messages/ai.tsx +++ b/apps/web/src/components/thread/messages/ai.tsx @@ -20,12 +20,14 @@ import { useQueryState, parseAsBoolean } from "nuqs"; import { Interrupt } from "./interrupt"; import { ActionStep, ActionItemProps } from "@/components/gen-ui/action-step"; import { TaskSummary } from "@/components/gen-ui/task-summary"; +import { PullRequestOpened } from "@/components/gen-ui/pull-request-opened"; import { ToolCall } from "@langchain/core/messages/tool"; import { createApplyPatchToolFields, createShellToolFields, createSetTaskStatusToolFields, createRgToolFields, + createOpenPrToolFields, createInstallDependenciesToolFields, } from "@open-swe/shared/open-swe/tools"; import { z } from "zod"; @@ -42,6 +44,8 @@ const setTaskStatusTool = createSetTaskStatusToolFields(); type SetTaskStatusToolArgs = z.infer; const rgTool = createRgToolFields(dummyRepo); type RgToolArgs = z.infer; +const openPrTool = createOpenPrToolFields(); +type OpenPrToolArgs = z.infer; const installDependenciesTool = createInstallDependenciesToolFields(dummyRepo); type InstallDependenciesToolArgs = z.infer< typeof installDependenciesTool.schema @@ -235,6 +239,10 @@ export function AssistantMessage({ ? aiToolCalls.find((tc) => tc.name === setTaskStatusTool.name) : undefined; + const openPrToolCall = message + ? aiToolCalls.find((tc) => tc.name === openPrTool.name) + : undefined; + // We can be sure that if the task status tool call is present, it will be the // only tool call/result we need to render for this message. if (taskStatusToolCall) { @@ -257,6 +265,59 @@ export function AssistantMessage({ ); } + // Same for PR tool. If this is present, it's the only tool call we need to render. + if (openPrToolCall) { + let branch: string | undefined; + let targetBranch: string | undefined = "main"; + + if (message && isAIMessageSDK(message)) { + branch = message.additional_kwargs?.branch as string | undefined; + targetBranch = + (message.additional_kwargs?.targetBranch as string | undefined) || + "main"; + } + + const args = openPrToolCall.args as OpenPrToolArgs; + const correspondingToolResult = toolResults.find( + (tr) => tr && tr.tool_call_id === openPrToolCall.id, + ); + + const status = correspondingToolResult ? "done" : "generating"; + + // 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(); + } + } + + // Extract PR number from URL if available + let prNumber: number | undefined = undefined; + if (prUrl) { + const match = prUrl.match(/\/pull\/(\d+)/); + if (match && match[1]) { + prNumber = parseInt(match[1], 10); + } + } + + return ( +
+ +
+ ); + } + if (actionableToolCalls.length > 0) { const actionItems = actionableToolCalls.map((toolCall): ActionItemProps => { const correspondingToolResult = toolResults.find( diff --git a/packages/shared/src/open-swe/planner/types.ts b/packages/shared/src/open-swe/planner/types.ts index ecd21d36..8a152681 100644 --- a/packages/shared/src/open-swe/planner/types.ts +++ b/packages/shared/src/open-swe/planner/types.ts @@ -67,6 +67,13 @@ export const PlannerGraphStateObj = MessagesZodState.extend({ fn: (_state, update) => update, }, }), + proposedPlanTitle: withLangGraph(z.custom(), { + reducer: { + schema: z.custom(), + fn: (_state, update) => update, + }, + default: () => "", + }), }); export type PlannerGraphState = z.infer; diff --git a/packages/shared/src/open-swe/tasks.ts b/packages/shared/src/open-swe/tasks.ts index 030d1ca1..7e53e920 100644 --- a/packages/shared/src/open-swe/tasks.ts +++ b/packages/shared/src/open-swe/tasks.ts @@ -7,16 +7,22 @@ import { PlanItem, Task, TaskPlan, PlanRevision } from "./types.js"; * * @param request The original user request text that initiated this task * @param planItems The plan items to include in the new task - * @param existingTaskPlan Optional existing TaskPlan to add the new task to - * @param parentTaskId Optional ID of a parent task if this task is derived from another + * @param options Optional existing TaskPlan to add the new task to + * @param options.parentTaskId Optional ID of a parent task if this task is derived from another + * @param options.existingTaskPlan Optional existing TaskPlan to add the new task to * @returns The updated TaskPlan with the new task added */ export function createNewTask( request: string, + title: string, planItems: PlanItem[], - existingTaskPlan?: TaskPlan, - parentTaskId?: string, + options?: { + existingTaskPlan?: TaskPlan; + parentTaskId?: string; + }, ): TaskPlan { + const { existingTaskPlan, parentTaskId } = options ?? {}; + // Create the initial plan revision const initialRevision: PlanRevision = { revisionIndex: 0, @@ -30,6 +36,7 @@ export function createNewTask( id: uuidv4(), taskIndex: existingTaskPlan ? existingTaskPlan.tasks.length : 0, request, + title, createdAt: Date.now(), completed: false, planRevisions: [initialRevision], diff --git a/packages/shared/src/open-swe/tools.ts b/packages/shared/src/open-swe/tools.ts index 7d5f8a46..6463ef83 100644 --- a/packages/shared/src/open-swe/tools.ts +++ b/packages/shared/src/open-swe/tools.ts @@ -41,6 +41,11 @@ export function createRequestHumanHelpToolFields() { export function createSessionPlanToolFields() { const sessionPlanSchema = z.object({ + title: z + .string() + .describe( + "The title of the plan. Should be a short, one sentence description of the user's request/plan generated to fulfill it.", + ), plan: z .array(z.string()) .describe("The plan to address the user's request."), @@ -204,3 +209,25 @@ export function createInstallDependenciesToolFields( schema: installDependenciesToolSchema, }; } + +export function createOpenPrToolFields() { + const openPrToolSchema = z.object({ + title: z + .string() + .describe( + "The title of the pull request. Ensure this is a concise and thoughtful title. You should follow conventional commit title format (e.g. prefixing with '[open-swe] fix:', '[open-swe] feat:', '[open-swe] chore:', etc.). Remember to include the '[open-swe]' prefix in the title.", + ), + body: z + .string() + .optional() + .describe( + "The body of the pull request. This should provide a concise description what the PR changes. Do not over-explain, or add technical details unless they're the absolute minimum needed. The user should be able to quickly read your description, and understand what the PR does. Remember: if they want the technical details they can read the changed files, so you don't need to go into great detail here.", + ), + }); + + return { + name: "open_pr", + schema: openPrToolSchema, + description: "Use this tool to open a pull request.", + }; +} diff --git a/packages/shared/src/open-swe/types.ts b/packages/shared/src/open-swe/types.ts index 259098c3..722e311f 100644 --- a/packages/shared/src/open-swe/types.ts +++ b/packages/shared/src/open-swe/types.ts @@ -73,6 +73,10 @@ export type Task = { * The original user request that created this task */ request: string; + /** + * The title of the task. Generated by the LLM. + */ + title: string; /** * When the task was created */