feat: Add opened PR tool (#292)

* feat: Add opened PR tool

* cr
This commit is contained in:
Brace Sproul 2025-06-22 16:35:25 -07:00 • committed by GitHub
parent 08a0945a0e
commit 2ee04861be
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 270 additions and 30 deletions

View file

@ -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 }),
};
}

View file

@ -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);
}

View file

@ -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<typeof openPrToolSchema>;
const { title, body } = toolCall.args as z.infer<typeof openPrTool.schema>;
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

View file

@ -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 (
<div className="h-3.5 w-3.5 rounded-full border border-gray-300" />
);
case "generating":
return <Loader2 className="h-3.5 w-3.5 animate-spin text-gray-500" />;
case "done":
return <CheckCircle className="h-3.5 w-3.5 text-green-500" />;
}
};
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 (
<div className="overflow-hidden rounded-md border border-gray-200">
<div className="flex items-center border-b border-gray-200 bg-gray-50 p-2">
<GitPullRequest className="mr-2 h-3.5 w-3.5 text-gray-500" />
<div className="flex-1">
{title && status === "done" && (
<div className="mb-0.5 text-xs font-normal text-gray-800">
{title}
</div>
)}
{branch && status === "done" && (
<div className="text-xs font-normal text-gray-500">
{branch} → {targetBranch}
</div>
)}
{!title && (
<span className="text-xs font-normal text-gray-800">
{getStatusText()}
</span>
)}
</div>
<div className="flex items-center gap-2">
<span className="text-xs font-normal text-gray-500">
{getStatusText()}
</span>
{getStatusIcon()}
{url && status === "done" && (
<a
href={url}
target="_blank"
rel="noopener noreferrer"
className="text-gray-500 hover:text-gray-700"
title="Open pull request"
>
<ExternalLink className="h-3.5 w-3.5" />
</a>
)}
{shouldShowToggle() && (
<button
onClick={() => setExpanded(!expanded)}
className="text-gray-500 hover:text-gray-700"
>
{expanded ? (
<ChevronUp className="h-3.5 w-3.5" />
) : (
<ChevronDown className="h-3.5 w-3.5" />
)}
</button>
)}
</div>
</div>
{expanded && description && status === "done" && (
<div className="border-t border-gray-200 p-2">
<h3 className="mb-1 text-xs font-normal text-gray-500">
Description
</h3>
<div className="text-xs font-normal whitespace-pre-wrap text-gray-800">
{description}
</div>
</div>
)}
</div>
);
}

View file

@ -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<typeof setTaskStatusTool.schema>;
const rgTool = createRgToolFields(dummyRepo);
type RgToolArgs = z.infer<typeof rgTool.schema>;
const openPrTool = createOpenPrToolFields();
type OpenPrToolArgs = z.infer<typeof openPrTool.schema>;
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 (
<div className="flex flex-col gap-4">
<PullRequestOpened
status={status}
title={args.title}
description={args.body}
url={prUrl}
prNumber={prNumber}
branch={branch}
targetBranch={targetBranch}
/>
</div>
);
}
if (actionableToolCalls.length > 0) {
const actionItems = actionableToolCalls.map((toolCall): ActionItemProps => {
const correspondingToolResult = toolResults.find(

View file

@ -67,6 +67,13 @@ export const PlannerGraphStateObj = MessagesZodState.extend({
fn: (_state, update) => update,
},
}),
proposedPlanTitle: withLangGraph(z.custom<string>(), {
reducer: {
schema: z.custom<string>(),
fn: (_state, update) => update,
},
default: () => "",
}),
});
export type PlannerGraphState = z.infer<typeof PlannerGraphStateObj>;

View file

@ -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],

View file

@ -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.",
};
}

View file

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