fix: Use plain messages instead of tool calls for context messages (#26)

This commit is contained in:
Brace Sproul 2025-05-26 17:38:28 -07:00 • committed by GitHub
parent e6d113e65a
commit c6ba3ee52e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 149 additions and 64 deletions

View file

@ -1,11 +1,7 @@
import { z } from "zod"; import { z } from "zod";
import { GraphConfig, GraphState, PlanItem } from "../types.js"; import { GraphConfig, GraphState, PlanItem } from "../types.js";
import { loadModel, Task } from "../utils/load-model.js"; import { loadModel, Task } from "../utils/load-model.js";
import { import { AIMessage, isHumanMessage } from "@langchain/core/messages";
AIMessage,
isHumanMessage,
ToolMessage,
} from "@langchain/core/messages";
import { formatPlanPrompt } from "../utils/plan-prompt.js"; import { formatPlanPrompt } from "../utils/plan-prompt.js";
import { createLogger, LogLevel } from "../utils/logger.js"; import { createLogger, LogLevel } from "../utils/logger.js";
import { getMessageString } from "../utils/message/content.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"); 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); const removedMessages = removeLastTaskMessages(state.messages);
logger.info(`Removing ${removedMessages.length} message(s) from state.`); logger.info(`Removing ${removedMessages.length} message(s) from state.`);
const allTasksCompleted = state.plan.every((p) => p.completed); const allTasksCompleted = state.plan.every((p) => p.completed);
const newMessagesStateUpdate = [ // Ensure all tool calls are removed from the message.
...removedMessages, delete response.tool_call_chunks;
new AIMessage({ delete response.tool_calls;
...response, delete response.invalid_tool_calls;
additional_kwargs: {
...response.additional_kwargs, const messageWithoutToolCall = new AIMessage({
summary_message: true, ...response,
}, content:
}), "Condensed Task Context:\n\n" +
toolMessage, (toolCall.args as z.infer<typeof condenseContextToolSchema>).context,
]; additional_kwargs: {
...response.additional_kwargs,
summary_message: true,
},
});
const newMessagesStateUpdate = [...removedMessages, messageWithoutToolCall];
if (!allTasksCompleted) { if (!allTasksCompleted) {
return new Command({ return new Command({

View file

@ -8,9 +8,17 @@ import {
getRepoAbsolutePath, getRepoAbsolutePath,
} from "../utils/git/index.js"; } from "../utils/git/index.js";
import { Sandbox } from "@e2b/code-interpreter"; import { Sandbox } from "@e2b/code-interpreter";
import { zodSchemaToString } from "../utils/zod-to-string.js";
import { z } from "zod";
const logger = createLogger(LogLevel.INFO, "TakeAction"); 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( export async function takeAction(
state: GraphState, state: GraphState,
config: GraphConfig, config: GraphConfig,
@ -48,13 +56,24 @@ export async function takeAction(
// @ts-expect-error tool.invoke types are weird here... // @ts-expect-error tool.invoke types are weird here...
result = await tool.invoke(toolCall.args); result = await tool.invoke(toolCall.args);
} catch (e) { } catch (e) {
logger.error("Failed to call tool", { if (
...(e instanceof Error e instanceof Error &&
? { name: e.name, message: e.message, stack: e.stack } e.message === "Received tool input did not match expected schema"
: { error: e }), ) {
}); logger.error("Received tool input did not match expected schema", {
const errMessage = e instanceof Error ? e.message : "Unknown error"; toolCall,
result = `FAILED TO CALL TOOL: "${toolCall.name}"\n\nError: ${errMessage}`; 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({ const toolMessage = new ToolMessage({

View file

@ -2,11 +2,7 @@ import { z } from "zod";
import { GraphConfig } from "../../../types.js"; import { GraphConfig } from "../../../types.js";
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js"; import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
import { loadModel, Task } from "../../../utils/load-model.js"; import { loadModel, Task } from "../../../utils/load-model.js";
import { import { AIMessage, isHumanMessage } from "@langchain/core/messages";
AIMessage,
isHumanMessage,
ToolMessage,
} from "@langchain/core/messages";
import { import {
getMessageContentString, getMessageContentString,
getMessageString, getMessageString,
@ -83,25 +79,22 @@ ${state.plannerMessages.map(getMessageString).join("\n")}`;
throw new Error("Failed to generate plan"); throw new Error("Failed to generate plan");
} }
const toolMessage = new ToolMessage({ delete response.tool_call_chunks;
tool_call_id: toolCall.id ?? "", delete response.tool_calls;
name: toolCall.name, delete response.invalid_tool_calls;
content: `Successfully summarized planning context.`,
const messageWithoutToolCall = new AIMessage({
...response,
content:
"Condensed Planning Context:\n\n" +
(toolCall.args as z.infer<typeof condenseContextToolSchema>).context,
additional_kwargs: { additional_kwargs: {
...response.additional_kwargs,
summary_message: true, summary_message: true,
}, },
}); });
return { return {
messages: [ messages: [messageWithoutToolCall],
new AIMessage({
...response,
additional_kwargs: {
...response.additional_kwargs,
summary_message: true,
},
}),
toolMessage,
],
}; };
} }

View file

@ -13,6 +13,12 @@ const logger = createLogger(LogLevel.INFO, "ApplyPatchTool");
const applyPatchToolSchema = z.object({ const applyPatchToolSchema = z.object({
diff: z.string().describe("The diff to apply. Use a standard diff format."), 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."), 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( export const applyPatchTool = tool(
@ -26,13 +32,16 @@ export const applyPatchTool = tool(
throw new Error("FAILED TO RUN COMMAND: No sandbox session ID provided"); 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 sandbox = await Sandbox.connect(sandboxSessionId);
const { success: readFileSuccess, output: readFileOutput } = await readFile( const { success: readFileSuccess, output: readFileOutput } = await readFile(
sandbox, sandbox,
file_path, file_path,
{
workDir: workdir,
},
); );
if (!readFileSuccess) { if (!readFileSuccess) {
logger.error("Failed to read file", readFileOutput); logger.error("Failed to read file", readFileOutput);
@ -64,7 +73,9 @@ export const applyPatchTool = tool(
} }
const { success: writeFileSuccess, output: writeFileOutput } = const { success: writeFileSuccess, output: writeFileOutput } =
await writeFile(sandbox, file_path, patchedContent); await writeFile(sandbox, file_path, patchedContent, {
workDir: workdir,
});
if (!writeFileSuccess) { if (!writeFileSuccess) {
logger.error("Failed to write file", { logger.error("Failed to write file", {
writeFileOutput, writeFileOutput,

View file

@ -12,13 +12,13 @@ const logger = createLogger(LogLevel.INFO, "ShellTool");
const DEFAULT_COMMAND_TIMEOUT = 60_000; // 1 minute const DEFAULT_COMMAND_TIMEOUT = 60_000; // 1 minute
const shellToolSchema = z.object({ const shellToolSchema = z.object({
command: z.array(z.string()).describe("The command to run"),
workdir: z workdir: z
.string() .string()
.optional() .default("/home/user")
.describe( .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 timeout: z
.number() .number()
.optional() .optional()

View file

@ -2,18 +2,24 @@ import { Sandbox } from "@e2b/code-interpreter";
import { createLogger, LogLevel } from "./logger.js"; import { createLogger, LogLevel } from "./logger.js";
import { TIMEOUT_MS } from "../constants.js"; import { TIMEOUT_MS } from "../constants.js";
import { getSandboxErrorFields } from "./sandbox-error-fields.js"; import { getSandboxErrorFields } from "./sandbox-error-fields.js";
import { traceable } from "langsmith/traceable";
const logger = createLogger(LogLevel.INFO, "ReadWriteUtil"); const logger = createLogger(LogLevel.INFO, "ReadWriteUtil");
export async function readFile( async function readFileFunc(
sandbox: Sandbox, sandbox: Sandbox,
filePath: string, filePath: string,
args?: {
workDir?: string;
},
): Promise<{ ): Promise<{
success: boolean; success: boolean;
output: string; output: string;
}> { }> {
try { 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. // Add an extra 5 min timeout to the sandbox.
await sandbox.setTimeout(TIMEOUT_MS); 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, sandbox: Sandbox,
filePath: string, filePath: string,
content: string, content: string,
args?: {
workDir?: string;
},
): Promise<{ ): Promise<{
success: boolean; success: boolean;
output: string; output: string;
@ -72,12 +85,14 @@ export async function writeFile(
const writeCommand = `cat > "${filePath}" << '${delimiter}' const writeCommand = `cat > "${filePath}" << '${delimiter}'
${content} ${content}
${delimiter}`; ${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. // Add an extra 5 min timeout to the sandbox.
await sandbox.setTimeout(TIMEOUT_MS); await sandbox.setTimeout(TIMEOUT_MS);
if (writeOutput.exitCode !== 0) { 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, writeOutput,
}); });
return { return {
@ -87,16 +102,16 @@ ${delimiter}`;
} }
if (writeOutput.stderr) { if (writeOutput.stderr) {
logger.warn( 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 { return {
success: true, success: true,
output: `Successfully wrote file '${filePath}' to sandbox via printf.`, output: `Successfully wrote file '${filePath}' to sandbox via cat.`,
}; };
} catch (e: any) { } catch (e: any) {
logger.error( 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 ...(e instanceof Error
? { name: e.name, message: e.message, stack: e.stack } ? { name: e.name, message: e.message, stack: e.stack }
@ -118,3 +133,7 @@ ${delimiter}`;
}; };
} }
} }
export const writeFile = traceable(writeFileFunc, {
name: "write_file",
});

View file

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