mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-05 03:42:13 +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 { 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({
|
||||||
|
|
|
||||||
|
|
@ -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({
|
||||||
|
|
|
||||||
|
|
@ -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,
|
|
||||||
],
|
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
});
|
||||||
|
|
|
||||||
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