mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 13:53:15 +00:00
fix: Use plain messages instead of tool calls for context messages (#26)
This commit is contained in:
parent
e6d113e65a
commit
c6ba3ee52e
7 changed files with 149 additions and 64 deletions
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
});
|
||||
|
|
|
|||
50
apps/open-swe/src/utils/zod-to-string.ts
Normal file
50
apps/open-swe/src/utils/zod-to-string.ts
Normal 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";
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue