diff --git a/apps/open-swe/src/nodes/summarize-task-steps.ts b/apps/open-swe/src/nodes/summarize-task-steps.ts index d6bf6274..32d0fb5b 100644 --- a/apps/open-swe/src/nodes/summarize-task-steps.ts +++ b/apps/open-swe/src/nodes/summarize-task-steps.ts @@ -1,11 +1,7 @@ import { z } from "zod"; import { GraphConfig, GraphState, PlanItem } from "../types.js"; import { loadModel, Task } from "../utils/load-model.js"; -import { - AIMessage, - isHumanMessage, - ToolMessage, -} from "@langchain/core/messages"; +import { AIMessage, isHumanMessage } from "@langchain/core/messages"; import { formatPlanPrompt } from "../utils/plan-prompt.js"; import { createLogger, LogLevel } from "../utils/logger.js"; import { getMessageString } from "../utils/message/content.js"; @@ -95,31 +91,28 @@ Given this full conversation history please generate a concise, and useful summa throw new Error("Failed to generate plan"); } - const toolMessage = new ToolMessage({ - tool_call_id: toolCall.id ?? "", - name: toolCall.name, - content: `Successfully summarized planning context.`, - additional_kwargs: { - summary_message: true, - }, - }); - const removedMessages = removeLastTaskMessages(state.messages); logger.info(`Removing ${removedMessages.length} message(s) from state.`); const allTasksCompleted = state.plan.every((p) => p.completed); - const newMessagesStateUpdate = [ - ...removedMessages, - new AIMessage({ - ...response, - additional_kwargs: { - ...response.additional_kwargs, - summary_message: true, - }, - }), - toolMessage, - ]; + // Ensure all tool calls are removed from the message. + delete response.tool_call_chunks; + delete response.tool_calls; + delete response.invalid_tool_calls; + + const messageWithoutToolCall = new AIMessage({ + ...response, + content: + "Condensed Task Context:\n\n" + + (toolCall.args as z.infer).context, + additional_kwargs: { + ...response.additional_kwargs, + summary_message: true, + }, + }); + + const newMessagesStateUpdate = [...removedMessages, messageWithoutToolCall]; if (!allTasksCompleted) { return new Command({ diff --git a/apps/open-swe/src/nodes/take-action.ts b/apps/open-swe/src/nodes/take-action.ts index 7e49a81e..772a6769 100644 --- a/apps/open-swe/src/nodes/take-action.ts +++ b/apps/open-swe/src/nodes/take-action.ts @@ -8,9 +8,17 @@ import { getRepoAbsolutePath, } from "../utils/git/index.js"; import { Sandbox } from "@e2b/code-interpreter"; +import { zodSchemaToString } from "../utils/zod-to-string.js"; +import { z } from "zod"; const logger = createLogger(LogLevel.INFO, "TakeAction"); +function formatBadArgsError(schema: z.ZodTypeAny, args: any) { + return `Invalid arguments for tool call. Expected:\n${zodSchemaToString( + schema, + )}.\nGot:\n${JSON.stringify(args)}`; +} + export async function takeAction( state: GraphState, config: GraphConfig, @@ -48,13 +56,24 @@ export async function takeAction( // @ts-expect-error tool.invoke types are weird here... result = await tool.invoke(toolCall.args); } catch (e) { - 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 ( + 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({ diff --git a/apps/open-swe/src/subgraphs/planner/nodes/summarizer.ts b/apps/open-swe/src/subgraphs/planner/nodes/summarizer.ts index f2b2f96f..b769a433 100644 --- a/apps/open-swe/src/subgraphs/planner/nodes/summarizer.ts +++ b/apps/open-swe/src/subgraphs/planner/nodes/summarizer.ts @@ -2,11 +2,7 @@ import { z } from "zod"; import { GraphConfig } from "../../../types.js"; import { PlannerGraphState, PlannerGraphUpdate } from "../types.js"; import { loadModel, Task } from "../../../utils/load-model.js"; -import { - AIMessage, - isHumanMessage, - ToolMessage, -} from "@langchain/core/messages"; +import { AIMessage, isHumanMessage } from "@langchain/core/messages"; import { getMessageContentString, getMessageString, @@ -83,25 +79,22 @@ ${state.plannerMessages.map(getMessageString).join("\n")}`; throw new Error("Failed to generate plan"); } - const toolMessage = new ToolMessage({ - tool_call_id: toolCall.id ?? "", - name: toolCall.name, - content: `Successfully summarized planning context.`, + delete response.tool_call_chunks; + delete response.tool_calls; + delete response.invalid_tool_calls; + + const messageWithoutToolCall = new AIMessage({ + ...response, + content: + "Condensed Planning Context:\n\n" + + (toolCall.args as z.infer).context, additional_kwargs: { + ...response.additional_kwargs, summary_message: true, }, }); return { - messages: [ - new AIMessage({ - ...response, - additional_kwargs: { - ...response.additional_kwargs, - summary_message: true, - }, - }), - toolMessage, - ], + messages: [messageWithoutToolCall], }; } diff --git a/apps/open-swe/src/tools/apply-patch.ts b/apps/open-swe/src/tools/apply-patch.ts index 64ac2b0a..965e7d70 100644 --- a/apps/open-swe/src/tools/apply-patch.ts +++ b/apps/open-swe/src/tools/apply-patch.ts @@ -13,6 +13,12 @@ const logger = createLogger(LogLevel.INFO, "ApplyPatchTool"); const applyPatchToolSchema = z.object({ diff: z.string().describe("The diff to apply. Use a standard diff format."), file_path: z.string().describe("The file path to apply the diff to."), + workdir: z + .string() + .default("/home/user") + .describe( + "The working directory for the command. Ensure this path is NOT included in any command arguments, as it will be added automatically. Defaults to '/home/user' as this is the root directory of the sandbox.", + ), }); export const applyPatchTool = tool( @@ -26,13 +32,16 @@ export const applyPatchTool = tool( throw new Error("FAILED TO RUN COMMAND: No sandbox session ID provided"); } - const { diff, file_path } = input; + const { diff, file_path, workdir } = input; const sandbox = await Sandbox.connect(sandboxSessionId); const { success: readFileSuccess, output: readFileOutput } = await readFile( sandbox, file_path, + { + workDir: workdir, + }, ); if (!readFileSuccess) { logger.error("Failed to read file", readFileOutput); @@ -64,7 +73,9 @@ export const applyPatchTool = tool( } const { success: writeFileSuccess, output: writeFileOutput } = - await writeFile(sandbox, file_path, patchedContent); + await writeFile(sandbox, file_path, patchedContent, { + workDir: workdir, + }); if (!writeFileSuccess) { logger.error("Failed to write file", { writeFileOutput, diff --git a/apps/open-swe/src/tools/shell.ts b/apps/open-swe/src/tools/shell.ts index 1cf41d10..5283af19 100644 --- a/apps/open-swe/src/tools/shell.ts +++ b/apps/open-swe/src/tools/shell.ts @@ -12,13 +12,13 @@ const logger = createLogger(LogLevel.INFO, "ShellTool"); const DEFAULT_COMMAND_TIMEOUT = 60_000; // 1 minute const shellToolSchema = z.object({ + command: z.array(z.string()).describe("The command to run"), workdir: z .string() - .optional() + .default("/home/user") .describe( - "The working directory for the command. Ensure this path is NOT included in any command arguments, as it will be added automatically.", + "The working directory for the command. Ensure this path is NOT included in any command arguments, as it will be added automatically. Defaults to '/home/user' as this is the root directory of the sandbox.", ), - command: z.array(z.string()).describe("The command to run"), timeout: z .number() .optional() diff --git a/apps/open-swe/src/utils/read-write.ts b/apps/open-swe/src/utils/read-write.ts index 5613bca5..1728aa5a 100644 --- a/apps/open-swe/src/utils/read-write.ts +++ b/apps/open-swe/src/utils/read-write.ts @@ -2,18 +2,24 @@ import { Sandbox } from "@e2b/code-interpreter"; import { createLogger, LogLevel } from "./logger.js"; import { TIMEOUT_MS } from "../constants.js"; import { getSandboxErrorFields } from "./sandbox-error-fields.js"; +import { traceable } from "langsmith/traceable"; const logger = createLogger(LogLevel.INFO, "ReadWriteUtil"); -export async function readFile( +async function readFileFunc( sandbox: Sandbox, filePath: string, + args?: { + workDir?: string; + }, ): Promise<{ success: boolean; output: string; }> { try { - const readOutput = await sandbox.commands.run(`cat "${filePath}"`); + const readOutput = await sandbox.commands.run(`cat "${filePath}"`, { + cwd: args?.workDir, + }); // Add an extra 5 min timeout to the sandbox. await sandbox.setTimeout(TIMEOUT_MS); @@ -59,10 +65,17 @@ export async function readFile( } } -export async function writeFile( +export const readFile = traceable(readFileFunc, { + name: "read_file", +}); + +async function writeFileFunc( sandbox: Sandbox, filePath: string, content: string, + args?: { + workDir?: string; + }, ): Promise<{ success: boolean; output: string; @@ -72,12 +85,14 @@ export async function writeFile( const writeCommand = `cat > "${filePath}" << '${delimiter}' ${content} ${delimiter}`; - const writeOutput = await sandbox.commands.run(writeCommand); + const writeOutput = await sandbox.commands.run(writeCommand, { + cwd: args?.workDir, + }); // Add an extra 5 min timeout to the sandbox. await sandbox.setTimeout(TIMEOUT_MS); if (writeOutput.exitCode !== 0) { - logger.error(`Error writing file '${filePath}' to sandbox via printf:`, { + logger.error(`Error writing file '${filePath}' to sandbox via cat:`, { writeOutput, }); return { @@ -87,16 +102,16 @@ ${delimiter}`; } if (writeOutput.stderr) { logger.warn( - `Stderr while writing file '${filePath}' to sandbox via printf: ${writeOutput.stderr}`, + `Stderr while writing file '${filePath}' to sandbox via cat: ${writeOutput.stderr}`, ); } return { success: true, - output: `Successfully wrote file '${filePath}' to sandbox via printf.`, + output: `Successfully wrote file '${filePath}' to sandbox via cat.`, }; } catch (e: any) { logger.error( - `Exception while trying to write file '${filePath}' to sandbox via printf:`, + `Exception while trying to write file '${filePath}' to sandbox via cat:`, { ...(e instanceof Error ? { name: e.name, message: e.message, stack: e.stack } @@ -118,3 +133,7 @@ ${delimiter}`; }; } } + +export const writeFile = traceable(writeFileFunc, { + name: "write_file", +}); diff --git a/apps/open-swe/src/utils/zod-to-string.ts b/apps/open-swe/src/utils/zod-to-string.ts new file mode 100644 index 00000000..132572f4 --- /dev/null +++ b/apps/open-swe/src/utils/zod-to-string.ts @@ -0,0 +1,50 @@ +import { z } from "zod"; + +export function zodSchemaToString(schema: z.ZodTypeAny, indent = 0): string { + const spaces = " ".repeat(indent); + + if (schema instanceof z.ZodObject) { + const shape = schema._def.shape(); + const lines: string[] = [`${spaces}{`]; + + for (const [key, value] of Object.entries(shape)) { + const fieldSchema = value as z.ZodTypeAny; + const description = fieldSchema._def.description + ? ` // ${fieldSchema._def.description}` + : ""; + + if (fieldSchema instanceof z.ZodObject) { + lines.push( + `${spaces} ${key}: ${zodSchemaToString(fieldSchema, indent + 2)}${description}`, + ); + } else { + const type = getZodType(fieldSchema); + lines.push(`${spaces} ${key}: ${type}${description}`); + } + } + + lines.push(`${spaces}}`); + return lines.join("\n"); + } + + return getZodType(schema); +} + +function getZodType(schema: z.ZodTypeAny): string { + const def = schema._def; + + if (schema instanceof z.ZodString) return "string"; + if (schema instanceof z.ZodNumber) return "number"; + if (schema instanceof z.ZodBoolean) return "boolean"; + if (schema instanceof z.ZodArray) return `${getZodType(def.type)}[]`; + if (schema instanceof z.ZodOptional) + return `${getZodType(def.innerType)} | undefined`; + if (schema instanceof z.ZodNullable) + return `${getZodType(def.innerType)} | null`; + if (schema instanceof z.ZodUnion) + return def.options.map(getZodType).join(" | "); + if (schema instanceof z.ZodEnum) + return def.values.map((v: any) => `"${v}"`).join(" | "); + + return def.typeName || "unknown"; +}