feat: Parallel tool calling (#282)

* feat: Parallel tool calling

* feat: Support parallel actions

* render parallel tool calling in ui

* new error handling logic

* cr
This commit is contained in:
Brace Sproul 2025-06-22 12:14:56 -07:00 • committed by GitHub
parent c348d88309
commit 96298ad22b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 566 additions and 316 deletions

View file

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

View file

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

View file

@ -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.
</context_gathering_guidelines>
<workspace_information>

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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<any>,
): 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<any>) {
// 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,
);
}

View file

@ -43,7 +43,6 @@ type ShellActionProps = BaseActionProps &
errorCode?: number;
};
// Apply patch specific props
type PatchActionProps = BaseActionProps &
Partial<ApplyPatchToolArgs> & {
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 (
<div className="border-border h-3.5 w-3.5 rounded-full border" />
);
return <div className="border-border size-3.5 rounded-full border" />;
case "generating":
return (
<Loader2 className="text-muted-foreground h-3.5 w-3.5 animate-spin" />
<Loader2 className="text-muted-foreground size-3.5 animate-spin" />
);
case "done":
return props.success ? (
<CheckCircle className="h-3.5 w-3.5 text-green-500" />
<CheckCircle className="size-3.5 text-green-500" />
) : (
<XCircle className="h-3.5 w-3.5 text-red-500" />
<XCircle className="size-3.5 text-red-500" />
);
}
};
@ -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 <Loader2 className="text-muted-foreground mr-2 h-3.5 w-3.5" />;
return <Loader2 className="text-muted-foreground mr-2 size-3.5" />;
}
return props.actionType === "shell" ? (
<Terminal className="text-muted-foreground mr-2 h-3.5 w-3.5" />
<Terminal className="text-muted-foreground mr-2 size-3.5" />
) : (
<FileCode className="text-muted-foreground mr-2 h-3.5 w-3.5" />
<FileCode className="text-muted-foreground mr-2 size-3.5" />
);
};
@ -231,25 +231,8 @@ export function ActionStep(props: ActionStepProps) {
};
return (
<div className="border-border overflow-hidden rounded-md border">
{props.reasoningText && (
<div className="border-b border-blue-300 bg-blue-100/50 p-2 dark:border-blue-800 dark:bg-blue-900/50">
<button
onClick={() => setShowReasoning(!showReasoning)}
className="flex items-center gap-1 text-xs font-normal text-blue-600 hover:text-blue-700 dark:text-blue-400 dark:hover:text-blue-300"
>
<MessageSquare className="h-3 w-3" />
{showReasoning ? "Hide reasoning" : "Show reasoning"}
</button>
{showReasoning && (
<p className="mt-1 text-xs font-normal text-blue-700 dark:text-blue-300">
{props.reasoningText}
</p>
)}
</div>
)}
<div className="border-border bg-card flex items-center border-b p-2 dark:bg-gray-800">
<div className="border-border mb-2 overflow-hidden rounded-md border last:mb-0">
<div className="border-border flex items-center border-b bg-gray-50 p-2 dark:bg-gray-800">
{renderHeaderIcon()}
{renderHeaderContent()}
<div className="flex items-center gap-2">
@ -263,9 +246,9 @@ export function ActionStep(props: ActionStepProps) {
className="text-muted-foreground hover:text-foreground"
>
{expanded ? (
<ChevronUp className="h-3.5 w-3.5" />
<ChevronUp className="size-3.5" />
) : (
<ChevronDown className="h-3.5 w-3.5" />
<ChevronDown className="size-3.5" />
)}
</button>
)}
@ -273,19 +256,60 @@ export function ActionStep(props: ActionStepProps) {
</div>
{renderContent()}
</div>
);
}
{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 (
<div className="border-border overflow-hidden rounded-md border">
<div className="border-b border-blue-300 bg-blue-100/50 p-2 dark:border-blue-800 dark:bg-blue-900/50">
<button
onClick={() => setShowReasoning(!showReasoning)}
className="flex cursor-pointer items-center gap-1 text-xs font-normal text-blue-600 hover:text-blue-700 dark:text-blue-400 dark:hover:text-blue-300"
>
<MessageSquare className="h-3 w-3" />
{showReasoning ? "Hide reasoning" : "Show reasoning"}
</button>
{showReasoning && (
<p className="mt-1 text-xs font-normal text-blue-700 dark:text-blue-300">
{reasoningText || "No reasoning provided."}
</p>
)}
</div>
<div className="p-2">
{props.actions.map((action: ActionItemProps, index: number) => (
<ActionItem
key={index}
{...action}
/>
))}
</div>
{summaryText && anyActionDone && (
<div className="border-t border-green-300 bg-green-100/50 p-2 dark:border-green-800 dark:bg-green-900/50">
<button
onClick={() => setShowSummary(!showSummary)}
className="flex items-center gap-1 text-xs font-normal text-green-600 hover:text-green-700 dark:text-green-400 dark:hover:text-green-300"
className="flex cursor-pointer items-center gap-1 text-xs font-normal text-green-600 hover:text-green-700 dark:text-green-400 dark:hover:text-green-300"
>
<FileText className="h-3 w-3" />
{showSummary ? "Hide summary" : "Show summary"}
</button>
{showSummary && (
<p className="mt-1 text-xs font-normal text-green-700 dark:text-green-300">
{props.summaryText}
{summaryText}
</p>
)}
</div>

View file

@ -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 (
<ActionStep {...mapToolMessageToActionStepProps(toolResult, thread)} />
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 (
<ActionStep
actionType={aiToolCallName === shellTool.name ? "shell" : "apply-patch"}
status="generating"
command={
aiToolCallName === shellTool.name ? aiToolCallArgs?.command || [] : []
}
workdir={
aiToolCallName === shellTool.name ? aiToolCallArgs?.workdir : ""
}
file_path={
aiToolCallName === applyPatchTool.name
? aiToolCallArgs?.file_path || ""
: ""
}
diff={
aiToolCallName === applyPatchTool.name ? aiToolCallArgs?.diff : ""
}
reasoningText={contentString}
/>
<div className="flex flex-col gap-4">
<ActionStep
actions={actionItems}
reasoningText={contentString}
/>
</div>
);
}
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({
</div>
)}
{!hideToolCalls && (
{!hideToolCalls && aiToolCalls.length > 0 && (
<span>
{(hasToolCalls && toolCallsHaveContents && (
<ToolCalls toolCalls={message.tool_calls} />
)) ||
(hasAnthropicToolCalls && (
<ToolCalls toolCalls={anthropicStreamedToolCalls} />
)) ||
(hasToolCalls && (
<ToolCalls toolCalls={message.tool_calls} />
))}
<ToolCalls toolCalls={aiToolCalls} />
</span>
)}