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 { 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<typeof condenseContextToolSchema>).context,
additional_kwargs: {
...response.additional_kwargs,
summary_message: true,
},
});
const newMessagesStateUpdate = [...removedMessages, messageWithoutToolCall];
if (!allTasksCompleted) {
return new Command({

View file

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

View file

@ -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<typeof condenseContextToolSchema>).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],
};
}

View file

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

View file

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

View file

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

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