From 96298ad22bde70710b2390b481f0380e40e7e04c Mon Sep 17 00:00:00 2001 From: Brace Sproul Date: Sun, 22 Jun 2025 12:14:56 -0700 Subject: [PATCH] feat: Parallel tool calling (#282) * feat: Parallel tool calling * feat: Support parallel actions * render parallel tool calling in ui * new error handling logic * cr --- apps/open-swe/src/graphs/planner/index.ts | 12 +- .../planner/nodes/generate-message/index.ts | 10 +- .../planner/nodes/generate-message/prompt.ts | 2 + .../src/graphs/planner/nodes/take-action.ts | 126 +++++------ .../programmer/__tests__/take-action.test.ts | 179 ++++++++++++++++ .../nodes/generate-message/index.ts | 10 +- .../nodes/generate-message/prompt.ts | 1 + .../graphs/programmer/nodes/take-action.ts | 157 ++++++-------- .../programmer/utils/tool-message-error.ts | 88 ++++++++ .../web/src/components/gen-ui/action-step.tsx | 102 +++++---- .../web/src/components/thread/messages/ai.tsx | 195 ++++++++---------- 11 files changed, 566 insertions(+), 316 deletions(-) create mode 100644 apps/open-swe/src/graphs/programmer/__tests__/take-action.test.ts create mode 100644 apps/open-swe/src/graphs/programmer/utils/tool-message-error.ts diff --git a/apps/open-swe/src/graphs/planner/index.ts b/apps/open-swe/src/graphs/planner/index.ts index cffc150b..524b27f5 100644 --- a/apps/open-swe/src/graphs/planner/index.ts +++ b/apps/open-swe/src/graphs/planner/index.ts @@ -13,7 +13,7 @@ import { interruptProposedPlan, prepareGraphState, notetaker, - takeAction, + takeActions, } from "./nodes/index.js"; import { isAIMessage } from "@langchain/core/messages"; import { initializeSandbox } from "../shared/initialize-sandbox.js"; @@ -21,7 +21,7 @@ import { initializeSandbox } from "../shared/initialize-sandbox.js"; function takeActionOrGeneratePlan( state: PlannerGraphState, config: GraphConfig, -): "take-plan-action" | "generate-plan" { +): "take-plan-actions" | "generate-plan" { const { messages } = state; const lastMessage = messages[messages.length - 1]; // If the last message is a tool call, and we have executed less than 75 actions, take action. @@ -34,7 +34,7 @@ function takeActionOrGeneratePlan( lastMessage.tool_calls?.length && messages.length < maxActionsCount ) { - return "take-plan-action"; + return "take-plan-actions"; } // If the last message does not have tool calls, continue to generate plan without modifications. @@ -47,7 +47,7 @@ const workflow = new StateGraph(PlannerGraphStateObj, GraphConfiguration) }) .addNode("initialize-sandbox", initializeSandbox) .addNode("generate-plan-context-action", generateAction) - .addNode("take-plan-action", takeAction) + .addNode("take-plan-actions", takeActions) .addNode("generate-plan", generatePlan) .addNode("notetaker", notetaker) .addNode("interrupt-proposed-plan", interruptProposedPlan) @@ -56,9 +56,9 @@ const workflow = new StateGraph(PlannerGraphStateObj, GraphConfiguration) .addConditionalEdges( "generate-plan-context-action", takeActionOrGeneratePlan, - ["take-plan-action", "generate-plan"], + ["take-plan-actions", "generate-plan"], ) - .addEdge("take-plan-action", "generate-plan-context-action") + .addEdge("take-plan-actions", "generate-plan-context-action") .addEdge("generate-plan", "notetaker") .addEdge("notetaker", "interrupt-proposed-plan") .addEdge("interrupt-proposed-plan", END); diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts index c1adf9dd..b3d57031 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts @@ -46,7 +46,7 @@ export async function generateAction( const tools = [createShellTool(state)]; const modelWithTools = model.bindTools(tools, { tool_choice: "auto", - parallel_tool_calls: false, + parallel_tool_calls: true, }); const [missingMessages, latestTaskPlan] = await Promise.all([ @@ -71,10 +71,10 @@ export async function generateAction( ...(getMessageContentString(response.content) && { content: getMessageContentString(response.content), }), - ...(response.tool_calls?.[0] && { - name: response.tool_calls?.[0].name, - args: response.tool_calls?.[0].args, - }), + ...response.tool_calls?.map((tc) => ({ + name: tc.name, + args: tc.args, + })), }); return { diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-message/prompt.ts b/apps/open-swe/src/graphs/planner/nodes/generate-message/prompt.ts index 7678e29b..35c7cda0 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-message/prompt.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-message/prompt.ts @@ -53,6 +53,8 @@ Your sole objective in this phase is to gather comprehensive context about the c 5. **Format shell commands precisely**: Ensure all shell commands include proper quoting and escaping. Well-formatted commands prevent errors and provide reliable results. 6. **Signal completion clearly**: When you have gathered sufficient context, respond with exactly 'done' without any tool calls. This indicates readiness to proceed to the planning phase. + +7. **Parallel tool calling**: It is highly recommended that you use parallel tool calling to gather context as quickly and efficiently as possible. When you know ahead of time there are multiple commands you want to run to gather context, of which they are independent and can be run in parallel, you should use parallel tool calling. diff --git a/apps/open-swe/src/graphs/planner/nodes/take-action.ts b/apps/open-swe/src/graphs/planner/nodes/take-action.ts index 6bf343c2..15c83c70 100644 --- a/apps/open-swe/src/graphs/planner/nodes/take-action.ts +++ b/apps/open-swe/src/graphs/planner/nodes/take-action.ts @@ -12,7 +12,7 @@ import { truncateOutput } from "../../../utils/truncate-outputs.js"; const logger = createLogger(LogLevel.INFO, "TakeAction"); -export async function takeAction( +export async function takeActions( state: PlannerGraphState, _config: GraphConfig, ): Promise { @@ -28,75 +28,81 @@ export async function takeAction( [shellTool.name]: shellTool, }; - const toolCall = lastMessage.tool_calls[0]; - if (!toolCall) { - throw new Error("No tool call found."); + const toolCalls = lastMessage.tool_calls; + if (!toolCalls?.length) { + throw new Error("No tool calls found."); } - const tool = toolsMap[toolCall.name]; - if (!tool) { - logger.error(`Unknown tool: ${toolCall.name}`); - const toolMessage = new ToolMessage({ - tool_call_id: toolCall.id ?? "", - content: `Unknown tool: ${toolCall.name}`, - name: toolCall.name, - status: "error", + const toolCallResultsPromise = toolCalls.map(async (toolCall) => { + const tool = toolsMap[toolCall.name]; + if (!tool) { + logger.error(`Unknown tool: ${toolCall.name}`); + const toolMessage = new ToolMessage({ + tool_call_id: toolCall.id ?? "", + content: `Unknown tool: ${toolCall.name}`, + name: toolCall.name, + status: "error", + }); + + return toolMessage; + } + + logger.info("Executing planner tool action", { + ...toolCall, }); - return { - messages: [toolMessage], - }; - } - - logger.info("Executing planner tool action", { - ...toolCall, - }); - - let result = ""; - let toolCallStatus: "success" | "error" = "success"; - try { - const toolResult = - // @ts-expect-error tool.invoke types are weird here... - (await tool.invoke(toolCall.args)) as { - result: string; - status: "success" | "error"; - }; - result = toolResult.result; - toolCallStatus = toolResult.status; - } catch (e) { - toolCallStatus = "error"; - if ( - e instanceof Error && - e.message === "Received tool input did not match expected schema" - ) { - logger.error("Received tool input did not match expected schema", { - toolCall, - expectedSchema: zodSchemaToString(tool.schema), - }); - result = formatBadArgsError(tool.schema, toolCall.args); - } else { - logger.error("Failed to call tool", { - ...(e instanceof Error - ? { name: e.name, message: e.message, stack: e.stack } - : { error: e }), - }); - const errMessage = e instanceof Error ? e.message : "Unknown error"; - result = `FAILED TO CALL TOOL: "${toolCall.name}"\n\nError: ${errMessage}`; + let result = ""; + let toolCallStatus: "success" | "error" = "success"; + try { + const toolResult = + // @ts-expect-error tool.invoke types are weird here... + (await tool.invoke(toolCall.args)) as { + result: string; + status: "success" | "error"; + }; + result = toolResult.result; + toolCallStatus = toolResult.status; + } catch (e) { + toolCallStatus = "error"; + if ( + e instanceof Error && + e.message === "Received tool input did not match expected schema" + ) { + logger.error("Received tool input did not match expected schema", { + toolCall, + expectedSchema: zodSchemaToString(tool.schema), + }); + result = formatBadArgsError(tool.schema, toolCall.args); + } else { + logger.error("Failed to call tool", { + ...(e instanceof Error + ? { name: e.name, message: e.message, stack: e.stack } + : { error: e }), + }); + const errMessage = e instanceof Error ? e.message : "Unknown error"; + result = `FAILED TO CALL TOOL: "${toolCall.name}"\n\nError: ${errMessage}`; + } } - } - const toolMessage = new ToolMessage({ - tool_call_id: toolCall.id ?? "", - content: truncateOutput(result), - name: toolCall.name, - status: toolCallStatus, + const toolMessage = new ToolMessage({ + tool_call_id: toolCall.id ?? "", + content: truncateOutput(result), + name: toolCall.name, + status: toolCallStatus, + }); + return toolMessage; }); + const toolCallResults = await Promise.all(toolCallResultsPromise); + logger.info("Completed planner tool action", { - tool_call_id: toolCall.id, - status: toolCallStatus, + ...toolCallResults.map((tc) => ({ + tool_call_id: tc.tool_call_id, + status: tc.status, + })), }); + return { - messages: [toolMessage], + messages: toolCallResults, }; } diff --git a/apps/open-swe/src/graphs/programmer/__tests__/take-action.test.ts b/apps/open-swe/src/graphs/programmer/__tests__/take-action.test.ts new file mode 100644 index 00000000..902875f2 --- /dev/null +++ b/apps/open-swe/src/graphs/programmer/__tests__/take-action.test.ts @@ -0,0 +1,179 @@ +import { AIMessage, ToolMessage, HumanMessage } from "@langchain/core/messages"; +import { describe, expect, test } from "@jest/globals"; +import { + calculateErrorRate, + groupToolMessagesByAIMessage, + shouldDiagnoseError, +} from "../utils/tool-message-error.js"; + +// Helper function to create a tool message with the specified parameters +function createToolMessage( + tool_call_id: string, + name: string, + status: "success" | "error", + is_diagnosis: boolean = false, +): ToolMessage { + const message = new ToolMessage({ + tool_call_id, + content: `Result of ${name}`, + name, + status, + ...(is_diagnosis ? { additional_kwargs: { is_diagnosis: true } } : {}), + }); + + return message; +} + +describe("Error diagnosis logic", () => { + describe("groupToolMessagesByAIMessage", () => { + test("should group tool messages by their parent AI message", () => { + const messages = [ + new HumanMessage({ content: "Human response" }), + new AIMessage({ content: "AI message 1" }), + createToolMessage("1", "tool1", "success"), + createToolMessage("2", "tool2", "error"), + new HumanMessage({ content: "Human response" }), + new AIMessage({ content: "AI message 2" }), + createToolMessage("3", "tool3", "success"), + createToolMessage("4", "tool4", "success"), + createToolMessage("5", "tool5", "error"), + ]; + + const groups = groupToolMessagesByAIMessage(messages); + + expect(groups.length).toBe(2); + expect(groups[0].length).toBe(2); // First group has 2 tool messages + expect(groups[1].length).toBe(3); // Second group has 3 tool messages + }); + + test("should filter out diagnostic tool messages", () => { + const messages = [ + new HumanMessage({ content: "Human response" }), + new AIMessage({ content: "AI message" }), + createToolMessage("1", "tool1", "success"), + createToolMessage("2", "tool2", "error", true), + createToolMessage("3", "tool3", "error"), + ]; + + const groups = groupToolMessagesByAIMessage(messages); + + expect(groups.length).toBe(1); + expect(groups[0].length).toBe(2); // Only non-diagnostic tools + expect(groups[0][0].tool_call_id).toBe("1"); + expect(groups[0][1].tool_call_id).toBe("3"); + }); + }); + + describe("calculateErrorRate", () => { + test("should return 0 for empty group", () => { + expect(calculateErrorRate([])).toBe(0); + }); + + test("should calculate correct error rate", () => { + const group = [ + createToolMessage("1", "tool1", "success"), + createToolMessage("2", "tool2", "error"), + createToolMessage("3", "tool3", "error"), + createToolMessage("4", "tool4", "success"), + ]; + + expect(calculateErrorRate(group)).toBe(0.5); // 2 errors out of 4 = 50% + }); + + test("should return 1 for all errors", () => { + const group = [ + createToolMessage("1", "tool1", "error"), + createToolMessage("2", "tool2", "error"), + ]; + + expect(calculateErrorRate(group)).toBe(1); // 100% errors + }); + }); + + describe("shouldDiagnoseError", () => { + test("should return false if less than 3 groups", () => { + const messages = [ + new HumanMessage({ content: "Human response" }), + new AIMessage({ content: "AI message 1" }), + createToolMessage("1", "tool1", "error"), + createToolMessage("2", "tool2", "error"), + new AIMessage({ content: "AI message 2" }), // AI message 2 + createToolMessage("3", "tool3", "error"), + createToolMessage("4", "tool4", "error"), + ]; + + expect(shouldDiagnoseError(messages)).toBe(false); + }); + + test("should return true if last three groups all have >= 75% error rate", () => { + const messages = [ + new HumanMessage({ content: "Human response" }), + new AIMessage({ content: "AI message 1" }), // AI message 1 (not part of last 3) + createToolMessage("1", "tool1", "success"), + createToolMessage("2", "tool2", "success"), + + new AIMessage({ content: "AI message 2" }), // AI message 2 (part of last 3) + createToolMessage("3", "tool3", "error"), + createToolMessage("4", "tool4", "error"), + createToolMessage("5", "tool5", "error"), + createToolMessage("6", "tool6", "success"), // 75% error rate + + new AIMessage({ content: "AI message 3" }), // AI message 3 (part of last 3) + createToolMessage("7", "tool7", "error"), + createToolMessage("8", "tool8", "error"), + createToolMessage("9", "tool9", "error"), // 100% error rate + + new AIMessage({ content: "AI message 4" }), // AI message 4 (part of last 3) + createToolMessage("10", "tool10", "error"), + createToolMessage("11", "tool11", "error"), + createToolMessage("12", "tool12", "success"), + createToolMessage("13", "tool13", "error"), // 75% error rate + ]; + + expect(shouldDiagnoseError(messages)).toBe(true); + }); + + test("should return false if any of the last three groups has < 75% error rate", () => { + const messages = [ + new AIMessage({ content: "AI message 1" }), // AI message 1 + createToolMessage("1", "tool1", "error"), + createToolMessage("2", "tool2", "error"), + + new AIMessage({ content: "AI message 2" }), // AI message 2 + createToolMessage("3", "tool3", "error"), + createToolMessage("4", "tool4", "error"), + createToolMessage("5", "tool5", "error"), + + new AIMessage({ content: "AI message 3" }), // AI message 3 + createToolMessage("6", "tool6", "success"), + createToolMessage("7", "tool7", "success"), + createToolMessage("8", "tool8", "error"), // 33% error rate (below threshold) + + new AIMessage({ content: "AI message 4" }), // AI message 4 + createToolMessage("9", "tool9", "error"), + createToolMessage("10", "tool10", "error"), + ]; + + expect(shouldDiagnoseError(messages)).toBe(false); + }); + + test("should ignore diagnostic tool messages", () => { + const messages = [ + new AIMessage({ content: "AI message 1" }), // AI message 1 + createToolMessage("1", "tool1", "error"), + createToolMessage("2", "tool2", "error", true), // Diagnostic (ignored) + + new AIMessage({ content: "AI message 2" }), // AI message 2 + createToolMessage("3", "tool3", "error"), + createToolMessage("4", "tool4", "error"), + + new AIMessage({ content: "AI message 3" }), // AI message 3 + createToolMessage("5", "tool5", "error"), + createToolMessage("6", "tool6", "error"), + createToolMessage("7", "tool7", "error", true), // Diagnostic (ignored) + ]; + + expect(shouldDiagnoseError(messages)).toBe(true); // All 3 groups have 100% error rate + }); + }); +}); diff --git a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts index bba92bc5..5e108d2e 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts @@ -65,7 +65,7 @@ export async function generateAction( ]; const modelWithTools = model.bindTools(tools, { tool_choice: "auto", - parallel_tool_calls: false, + parallel_tool_calls: true, }); const [missingMessages, latestTaskPlan] = await Promise.all([ @@ -98,10 +98,10 @@ export async function generateAction( ...(getMessageContentString(response.content) && { content: getMessageContentString(response.content), }), - ...(response.tool_calls?.[0] && { - name: response.tool_calls?.[0].name, - args: response.tool_calls?.[0].args, - }), + ...(response.tool_calls?.map((tc) => ({ + name: tc.name, + args: tc.args, + })) || []), }); const newMessagesList = [...missingMessages, response]; diff --git a/apps/open-swe/src/graphs/programmer/nodes/generate-message/prompt.ts b/apps/open-swe/src/graphs/programmer/nodes/generate-message/prompt.ts index 972255e9..ebbe515c 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/generate-message/prompt.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/generate-message/prompt.ts @@ -133,6 +133,7 @@ You are currently executing a specific task from a pre-generated plan. You have * **Dependencies**: Use the correct package manager; skip if installation fails * **Pre-commit**: Run \`pre-commit run --files ...\` if .pre-commit-config.yaml exists * **History**: Use \`git log\` and \`git blame\` for additional context when needed +* **Parallel Tool Calling**: You're allowed, and encouraged to call multiple tools at once, as long as they do not conflict, or depend on each other. ### Coding Standards 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 fab60d76..2a59c2b8 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/take-action.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/take-action.ts @@ -1,8 +1,4 @@ -import { - isAIMessage, - isToolMessage, - ToolMessage, -} from "@langchain/core/messages"; +import { isAIMessage, ToolMessage } from "@langchain/core/messages"; import { createLogger, LogLevel } from "../../../utils/logger.js"; import { createApplyPatchTool, createShellTool } from "../../../tools/index.js"; import { @@ -23,30 +19,10 @@ import { truncateOutput } from "../../../utils/truncate-outputs.js"; import { daytonaClient } from "../../../utils/sandbox.js"; import { getCodebaseTree } from "../../../utils/tree.js"; import { getRepoAbsolutePath } from "@open-swe/shared/git"; +import { shouldDiagnoseError } from "../utils/tool-message-error.js"; const logger = createLogger(LogLevel.INFO, "TakeAction"); -/** - * Whether or not to route to the diagnose error step. This is true if: - * - the last two tool messages are of an error status - * - two of the last three messages are an error status, including the last tool message - * @param toolMessages The tool messages to check the status of. - */ -function shouldDiagnoseError(toolMessages: ToolMessage[]) { - if ( - toolMessages[toolMessages.length - 1].status !== "error" || - toolMessages.length < 2 - ) { - // Last message is not an error, then neither of the below two conditions should be true. - return false; - } - return ( - // Two of the three last tool calls are errors, return true - // (this is either the last two, or the 3rd, and last since the check above ensures the last is an error) - toolMessages.slice(-3).filter((m) => m.status === "error").length >= 2 - ); -} - export async function takeAction( state: GraphState, config: GraphConfig, @@ -57,6 +33,12 @@ export async function takeAction( throw new Error("Last message is not an AI message with tool calls."); } + if (!state.sandboxSessionId) { + throw new Error( + "Failed to take action: No sandbox session ID found in state.", + ); + } + const applyPatchTool = createApplyPatchTool(state); const shellTool = createShellTool(state); const toolsMap = { @@ -64,74 +46,67 @@ export async function takeAction( [shellTool.name]: shellTool, }; - const toolCall = lastMessage.tool_calls[0]; - - if (!toolCall) { - throw new Error("No tool call found."); + const toolCalls = lastMessage.tool_calls; + if (!toolCalls?.length) { + throw new Error("No tool calls found."); } - const tool = toolsMap[toolCall.name]; + const toolCallResultsPromise = toolCalls.map(async (toolCall) => { + const tool = toolsMap[toolCall.name]; + + if (!tool) { + logger.error(`Unknown tool: ${toolCall.name}`); + const toolMessage = new ToolMessage({ + tool_call_id: toolCall.id ?? "", + content: `Unknown tool: ${toolCall.name}`, + name: toolCall.name, + status: "error", + }); + return toolMessage; + } + + let result = ""; + let toolCallStatus: "success" | "error" = "success"; + try { + const toolResult: { result: string; status: "success" | "error" } = + // @ts-expect-error tool.invoke types are weird here... + await tool.invoke(toolCall.args); + result = toolResult.result; + toolCallStatus = toolResult.status; + } catch (e) { + toolCallStatus = "error"; + if ( + e instanceof Error && + e.message === "Received tool input did not match expected schema" + ) { + logger.error("Received tool input did not match expected schema", { + toolCall, + expectedSchema: zodSchemaToString(tool.schema), + }); + result = formatBadArgsError(tool.schema, toolCall.args); + } else { + logger.error("Failed to call tool", { + ...(e instanceof Error + ? { name: e.name, message: e.message, stack: e.stack } + : { error: e }), + }); + const errMessage = e instanceof Error ? e.message : "Unknown error"; + result = `FAILED TO CALL TOOL: "${toolCall.name}"\n\nError: ${errMessage}`; + } + } - if (!tool) { - logger.error(`Unknown tool: ${toolCall.name}`); const toolMessage = new ToolMessage({ tool_call_id: toolCall.id ?? "", - content: `Unknown tool: ${toolCall.name}`, + content: truncateOutput(result), name: toolCall.name, - status: "error", + status: toolCallStatus, }); - return new Command({ - goto: "progress-plan-step", - update: { - messages: [toolMessage], - internalMessages: [toolMessage], - }, - }); - } - if (!state.sandboxSessionId) { - throw new Error( - "Failed to take action: No sandbox session ID found in state.", - ); - } - - let result = ""; - let toolCallStatus: "success" | "error" = "success"; - try { - const toolResult: { result: string; status: "success" | "error" } = - // @ts-expect-error tool.invoke types are weird here... - await tool.invoke(toolCall.args); - result = toolResult.result; - toolCallStatus = toolResult.status; - } catch (e) { - toolCallStatus = "error"; - if ( - e instanceof Error && - e.message === "Received tool input did not match expected schema" - ) { - logger.error("Received tool input did not match expected schema", { - toolCall, - expectedSchema: zodSchemaToString(tool.schema), - }); - result = formatBadArgsError(tool.schema, toolCall.args); - } else { - logger.error("Failed to call tool", { - ...(e instanceof Error - ? { name: e.name, message: e.message, stack: e.stack } - : { error: e }), - }); - const errMessage = e instanceof Error ? e.message : "Unknown error"; - result = `FAILED TO CALL TOOL: "${toolCall.name}"\n\nError: ${errMessage}`; - } - } - - const toolMessage = new ToolMessage({ - tool_call_id: toolCall.id ?? "", - content: truncateOutput(result), - name: toolCall.name, - status: toolCallStatus, + return toolMessage; }); + const toolCallResults = await Promise.all(toolCallResultsPromise); + // Always check if there are changed files after running a tool. // If there are, commit them. const sandbox = await daytonaClient().get(state.sandboxSessionId); @@ -155,18 +130,16 @@ export async function takeAction( ); } - const shouldRouteDiagnoseNode = shouldDiagnoseError( - [...state.internalMessages, toolMessage].filter( - (m): m is ToolMessage => - isToolMessage(m) && !m.additional_kwargs?.is_diagnosis, - ), - ); + const shouldRouteDiagnoseNode = shouldDiagnoseError([ + ...state.internalMessages, + ...toolCallResults, + ]); const codebaseTree = await getCodebaseTree(); const commandUpdate: GraphUpdate = { - messages: [toolMessage], - internalMessages: [toolMessage], + messages: toolCallResults, + internalMessages: toolCallResults, ...(branchName && { branchName }), codebaseTree, }; diff --git a/apps/open-swe/src/graphs/programmer/utils/tool-message-error.ts b/apps/open-swe/src/graphs/programmer/utils/tool-message-error.ts new file mode 100644 index 00000000..367a25ca --- /dev/null +++ b/apps/open-swe/src/graphs/programmer/utils/tool-message-error.ts @@ -0,0 +1,88 @@ +import { + isAIMessage, + isToolMessage, + ToolMessage, +} from "@langchain/core/messages"; + +/** + * Group tool messages by their parent AI message + * @param messages Array of messages to process + * @returns Array of tool message groups, where each group contains tool messages tied to the same AI message + */ +export function groupToolMessagesByAIMessage( + messages: Array, +): ToolMessage[][] { + const groups: ToolMessage[][] = []; + let currentGroup: ToolMessage[] = []; + let processingToolsForAI = false; + + for (let i = 0; i < messages.length; i++) { + const message = messages[i]; + + if (isAIMessage(message)) { + // If we were already processing tools for a previous AI message, save that group + if (currentGroup.length > 0) { + groups.push([...currentGroup]); + currentGroup = []; + } + processingToolsForAI = true; + } else if ( + isToolMessage(message) && + processingToolsForAI && + !message.additional_kwargs?.is_diagnosis + ) { + currentGroup.push(message); + } else if (!isToolMessage(message) && processingToolsForAI) { + // We've encountered a non-tool message after an AI message, end the current group + if (currentGroup.length > 0) { + groups.push([...currentGroup]); + currentGroup = []; + } + processingToolsForAI = false; + } + } + + // Add the last group if it exists + if (currentGroup.length > 0) { + groups.push(currentGroup); + } + + return groups; +} + +/** + * Calculate the error rate for a group of tool messages + * @param group Array of tool messages + * @returns Error rate as a number between 0 and 1 + */ +export function calculateErrorRate(group: ToolMessage[]): number { + if (group.length === 0) return 0; + const errorCount = group.filter((m) => m.status === "error").length; + return errorCount / group.length; +} + +/** + * Whether or not to route to the diagnose error step. This is true if: + * - the last three tool call groups all have >= 75% error rates + * + * TBD: Should this be checking that each of the last 3 have >= 75% error rates, + * or >= 75% error rate of all tool messages from the last 3 groups? + * + * @param messages All messages to analyze + */ +export function shouldDiagnoseError(messages: Array) { + // Group tool messages by their parent AI message + const toolGroups = groupToolMessagesByAIMessage(messages); + + // If we don't have at least 3 groups, we can't make a determination + if (toolGroups.length < 3) return false; + + // Get the last three groups + const lastThreeGroups = toolGroups.slice(-3); + + // Check if all of the last three groups have an error rate >= 75% + const ERROR_THRESHOLD = 0.75; // 75% + return lastThreeGroups.every( + (group) => calculateErrorRate(group) >= ERROR_THRESHOLD, + ); +} diff --git a/apps/web/src/components/gen-ui/action-step.tsx b/apps/web/src/components/gen-ui/action-step.tsx index 095fd739..5b3100ae 100644 --- a/apps/web/src/components/gen-ui/action-step.tsx +++ b/apps/web/src/components/gen-ui/action-step.tsx @@ -43,7 +43,6 @@ type ShellActionProps = BaseActionProps & errorCode?: number; }; -// Apply patch specific props type PatchActionProps = BaseActionProps & Partial & { actionType: "apply-patch"; @@ -51,32 +50,33 @@ type PatchActionProps = BaseActionProps & fixedDiff?: string; }; -// Union type for all possible action props -export type ActionStepProps = +export type ActionItemProps = | (BaseActionProps & { status: "loading" }) | ShellActionProps | PatchActionProps; -export function ActionStep(props: ActionStepProps) { +export type ActionStepProps = { + actions: ActionItemProps[]; + reasoningText?: string; + summaryText?: string; +}; + +function ActionItem(props: ActionItemProps) { const [expanded, setExpanded] = useState(false); - const [showReasoning, setShowReasoning] = useState(false); - const [showSummary, setShowSummary] = useState(false); const getStatusIcon = () => { switch (props.status) { case "loading": - return ( -
- ); + return
; case "generating": return ( - + ); case "done": return props.success ? ( - + ) : ( - + ); } }; @@ -120,13 +120,13 @@ export function ActionStep(props: ActionStepProps) { const renderHeaderIcon = () => { if (props.status === "loading" || !("actionType" in props)) { // In loading state, we don't know the type yet, use a generic icon - return ; + return ; } return props.actionType === "shell" ? ( - + ) : ( - + ); }; @@ -231,25 +231,8 @@ export function ActionStep(props: ActionStepProps) { }; return ( -
- {props.reasoningText && ( -
- - {showReasoning && ( -

- {props.reasoningText} -

- )} -
- )} - -
+
+
{renderHeaderIcon()} {renderHeaderContent()}
@@ -263,9 +246,9 @@ export function ActionStep(props: ActionStepProps) { className="text-muted-foreground hover:text-foreground" > {expanded ? ( - + ) : ( - + )} )} @@ -273,19 +256,60 @@ export function ActionStep(props: ActionStepProps) {
{renderContent()} +
+ ); +} - {props.summaryText && props.status === "done" && ( +export function ActionStep(props: ActionStepProps) { + const [showReasoning, setShowReasoning] = useState(false); + const [showSummary, setShowSummary] = useState(false); + + const reasoningText = + "reasoningText" in props ? props.reasoningText : undefined; + const summaryText = "summaryText" in props ? props.summaryText : undefined; + + const anyActionDone = props.actions.some( + (action: ActionItemProps) => action.status === "done", + ); + + return ( +
+
+ + {showReasoning && ( +

+ {reasoningText || "No reasoning provided."} +

+ )} +
+ +
+ {props.actions.map((action: ActionItemProps, index: number) => ( + + ))} +
+ + {summaryText && anyActionDone && (
{showSummary && (

- {props.summaryText} + {summaryText}

)}
diff --git a/apps/web/src/components/thread/messages/ai.tsx b/apps/web/src/components/thread/messages/ai.tsx index 8c107c04..710bd0f3 100644 --- a/apps/web/src/components/thread/messages/ai.tsx +++ b/apps/web/src/components/thread/messages/ai.tsx @@ -18,10 +18,7 @@ import { MessageContentComplex } from "@langchain/core/messages"; import { Fragment } from "react/jsx-runtime"; import { useQueryState, parseAsBoolean } from "nuqs"; import { Interrupt } from "./interrupt"; -import { - ActionStep, - type ActionStepProps, -} from "@/components/gen-ui/action-step"; +import { ActionStep, ActionItemProps } from "@/components/gen-ui/action-step"; import { ToolCall } from "@langchain/core/messages/tool"; import { createApplyPatchToolFields, @@ -95,7 +92,7 @@ function parseAnthropicStreamedToolCalls( export function mapToolMessageToActionStepProps( message: ToolMessage, thread: { messages: Message[] }, -): ActionStepProps { +): ActionItemProps { const toolCall: ToolCall | undefined = thread.messages .filter(isAIMessageSDK) .flatMap((m) => m.tool_calls ?? []) @@ -108,7 +105,7 @@ export function mapToolMessageToActionStepProps( ? getContentString(aiMessage.content) : undefined; - const status: ActionStepProps["status"] = "done"; + const status: ActionItemProps["status"] = "done"; const success = message.status === "success"; if (toolCall?.name === shellTool.name) { @@ -162,7 +159,6 @@ export function AssistantMessage({ const messages = thread.messages; const idx = message ? messages.findIndex((m) => m.id === message.id) : -1; - const nextMessage = idx >= 0 ? messages[idx + 1] : undefined; const meta = message ? thread.getMessagesMetadata(message) : undefined; const threadInterrupt = thread.interrupt; @@ -171,97 +167,98 @@ export function AssistantMessage({ ? parseAnthropicStreamedToolCalls(content) : undefined; - // Helper: get tool call name from AI message (OpenAI or Anthropic) - const aiToolCallName = (() => { + const aiToolCalls: ToolCall[] = (() => { if (message && isAIMessageSDK(message)) { - return message.tool_calls?.[0]?.name; + return message.tool_calls || []; } if (anthropicStreamedToolCalls?.length) { - return anthropicStreamedToolCalls[anthropicStreamedToolCalls.length - 1] - .name; + return anthropicStreamedToolCalls; } - return undefined; + return []; })(); - const aiToolCallArgs = (() => { - if (message && isAIMessageSDK(message)) { - return message.tool_calls?.[0]?.args; - } - if (anthropicStreamedToolCalls?.length) { - return anthropicStreamedToolCalls[anthropicStreamedToolCalls.length - 1] - .args; - } - return undefined; - })(); - - const toolResult = - nextMessage && - isToolMessageSDK(nextMessage) && - aiToolCallName && - nextMessage.tool_call_id === - (message && isAIMessageSDK(message) - ? message.tool_calls?.[0]?.id - : anthropicStreamedToolCalls?.length - ? anthropicStreamedToolCalls[anthropicStreamedToolCalls.length - 1].id - : undefined) - ? nextMessage - : undefined; - - if ( - message && - (aiToolCallName === shellTool.name || - aiToolCallName === applyPatchTool.name) - ) { - if (toolResult) { - return ( - + const toolResults = aiToolCalls + .map((toolCall) => { + const matchingToolMessage = messages.find( + (m) => isToolMessageSDK(m) && m.tool_call_id === toolCall.id, ); - } + + return matchingToolMessage as ToolMessage | undefined; + }) + .filter((m): m is ToolMessage => !!m); + + const shellOrPatchToolCalls = message + ? aiToolCalls.filter( + (tc) => tc.name === shellTool.name || tc.name === applyPatchTool.name, + ) + : []; + + if (shellOrPatchToolCalls.length > 0) { + const actionItems = shellOrPatchToolCalls.map((toolCall) => { + const correspondingToolResult = toolResults.find( + (tr) => tr && tr.tool_call_id === toolCall.id, + ); + + const isShellTool = toolCall.name === shellTool.name; + + if (correspondingToolResult) { + // If we have a tool result, map it to action props + return mapToolMessageToActionStepProps(correspondingToolResult, thread); + } else { + if (isShellTool) { + const args = toolCall.args as ShellToolArgs; + return { + actionType: "shell", + status: "generating", + command: args?.command || [], + workdir: args?.workdir, + timeout: args?.timeout, + } as ActionItemProps; + } else { + const args = toolCall.args as ApplyPatchToolArgs; + return { + actionType: "apply-patch", + status: "generating", + file_path: args?.file_path || "", + diff: args?.diff || "", + } as ActionItemProps; + } + } + }); + return ( - +
+ +
); } - if ( - message?.type === "tool" && - (message.name === shellTool.name || message.name === applyPatchTool.name) && - idx > 0 && - messages[idx - 1] && - ((messages[idx - 1] && - isAIMessageSDK(messages[idx - 1]) && - (messages[idx - 1] as AIMessage).tool_calls?.some( - (tc) => - tc.id === (message as ToolMessage).tool_call_id && - (tc.name === shellTool.name || tc.name === applyPatchTool.name), - )) || - (Array.isArray(messages[idx - 1].content) && - parseAnthropicStreamedToolCalls( - messages[idx - 1].content as MessageContentComplex[], - )?.some( - (tc) => - tc.id === (message as ToolMessage).tool_call_id && - (tc.name === shellTool.name || tc.name === applyPatchTool.name), - ))) - ) { - return null; + if (message?.type === "tool" && idx > 0) { + const isPreviousToolCall = messages.slice(0, idx).some((prevMessage) => { + if (isAIMessageSDK(prevMessage) && prevMessage.tool_calls) { + return prevMessage.tool_calls.some( + (tc) => tc.id === (message as ToolMessage).tool_call_id, + ); + } + + if (Array.isArray(prevMessage.content)) { + const toolCalls = parseAnthropicStreamedToolCalls( + prevMessage.content as MessageContentComplex[], + ); + return toolCalls?.some( + (tc) => tc.id === (message as ToolMessage).tool_call_id, + ); + } + + return false; + }); + + if (isPreviousToolCall) { + return null; + } } const isLastMessage = @@ -269,18 +266,6 @@ export function AssistantMessage({ const hasNoAIOrToolMessages = !thread.messages.find( (m) => m.type === "ai" || m.type === "tool", ); - - const hasToolCalls = - message && - "tool_calls" in message && - message.tool_calls && - message.tool_calls.length > 0; - const toolCallsHaveContents = - hasToolCalls && - message.tool_calls?.some( - (tc) => tc.args && Object.keys(tc.args).length > 0, - ); - const hasAnthropicToolCalls = !!anthropicStreamedToolCalls?.length; const isToolResult = message?.type === "tool"; if (isToolResult && hideToolCalls) { @@ -309,17 +294,9 @@ export function AssistantMessage({
)} - {!hideToolCalls && ( + {!hideToolCalls && aiToolCalls.length > 0 && ( - {(hasToolCalls && toolCallsHaveContents && ( - - )) || - (hasAnthropicToolCalls && ( - - )) || - (hasToolCalls && ( - - ))} + )}