mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-03 17:53:20 +00:00
[open-swe] feat: implement robust error handling for Daytona sandbox operations (#372)
* Apply patch * Apply patch * Apply patch * Apply patch * apply everywhere * cr * cr --------- Co-authored-by: open-swe-dev[bot] <open-swe-dev@users.noreply.github.com> Co-authored-by: bracesproul <braceasproul@gmail.com>
This commit is contained in:
parent
77fe0e9a05
commit
4fb0f90f4e
13 changed files with 226 additions and 143 deletions
|
|
@ -12,7 +12,7 @@ import {
|
||||||
HumanInterrupt,
|
HumanInterrupt,
|
||||||
HumanResponse,
|
HumanResponse,
|
||||||
} from "@langchain/langgraph/prebuilt";
|
} from "@langchain/langgraph/prebuilt";
|
||||||
import { startSandbox } from "../../../utils/sandbox.js";
|
import { getSandboxWithErrorHandling } from "../../../utils/sandbox.js";
|
||||||
import { createNewTask } from "@open-swe/shared/open-swe/tasks";
|
import { createNewTask } from "@open-swe/shared/open-swe/tasks";
|
||||||
import { getUserRequest } from "../../../utils/user-request.js";
|
import { getUserRequest } from "../../../utils/user-request.js";
|
||||||
import {
|
import {
|
||||||
|
|
@ -79,7 +79,19 @@ async function startProgrammerRun(input: {
|
||||||
|
|
||||||
const programmerThreadId = uuidv4();
|
const programmerThreadId = uuidv4();
|
||||||
// Restart the sandbox.
|
// Restart the sandbox.
|
||||||
runInput.sandboxSessionId = (await startSandbox(state.sandboxSessionId)).id;
|
const { sandbox, codebaseTree, dependenciesInstalled } =
|
||||||
|
await getSandboxWithErrorHandling(
|
||||||
|
state.sandboxSessionId,
|
||||||
|
state.targetRepository,
|
||||||
|
state.branchName,
|
||||||
|
config,
|
||||||
|
);
|
||||||
|
runInput.sandboxSessionId = sandbox.id;
|
||||||
|
runInput.codebaseTree = codebaseTree ?? runInput.codebaseTree;
|
||||||
|
runInput.dependenciesInstalled =
|
||||||
|
dependenciesInstalled !== null
|
||||||
|
? dependenciesInstalled
|
||||||
|
: runInput.dependenciesInstalled;
|
||||||
|
|
||||||
const run = await langGraphClient.runs.create(
|
const run = await langGraphClient.runs.create(
|
||||||
programmerThreadId,
|
programmerThreadId,
|
||||||
|
|
@ -179,11 +191,6 @@ export async function interruptProposedPlan(
|
||||||
If editing the plan, ensure each step in the plan is separated by "${PLAN_INTERRUPT_DELIMITER}".`,
|
If editing the plan, ensure each step in the plan is separated by "${PLAN_INTERRUPT_DELIMITER}".`,
|
||||||
})[0];
|
})[0];
|
||||||
|
|
||||||
if (!state.sandboxSessionId) {
|
|
||||||
// TODO: This should prob just create a sandbox?
|
|
||||||
throw new Error("No sandbox session ID found.");
|
|
||||||
}
|
|
||||||
|
|
||||||
if (interruptRes.type === "response") {
|
if (interruptRes.type === "response") {
|
||||||
// Plan was responded to, route to the rewrite plan node.
|
// Plan was responded to, route to the rewrite plan node.
|
||||||
throw new Error("RESPONDING TO PLAN NOT IMPLEMENTED.");
|
throw new Error("RESPONDING TO PLAN NOT IMPLEMENTED.");
|
||||||
|
|
|
||||||
|
|
@ -21,9 +21,9 @@ import {
|
||||||
} from "../../../utils/github/git.js";
|
} from "../../../utils/github/git.js";
|
||||||
import { createFindInstancesOfTool } from "../../../tools/find-instances-of.js";
|
import { createFindInstancesOfTool } from "../../../tools/find-instances-of.js";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { daytonaClient } from "../../../utils/sandbox.js";
|
|
||||||
import { createPlannerNotesTool } from "../../../tools/planner-notes.js";
|
import { createPlannerNotesTool } from "../../../tools/planner-notes.js";
|
||||||
import { getMcpTools } from "../../../utils/mcp-client.js";
|
import { getMcpTools } from "../../../utils/mcp-client.js";
|
||||||
|
import { getSandboxWithErrorHandling } from "../../../utils/sandbox.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
||||||
|
|
||||||
|
|
@ -62,6 +62,14 @@ export async function takeActions(
|
||||||
throw new Error("No tool calls found.");
|
throw new Error("No tool calls found.");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const { sandbox, codebaseTree, dependenciesInstalled } =
|
||||||
|
await getSandboxWithErrorHandling(
|
||||||
|
state.sandboxSessionId,
|
||||||
|
state.targetRepository,
|
||||||
|
state.branchName,
|
||||||
|
config,
|
||||||
|
);
|
||||||
|
|
||||||
const toolCallResultsPromise = toolCalls.map(async (toolCall) => {
|
const toolCallResultsPromise = toolCalls.map(async (toolCall) => {
|
||||||
const tool = toolsMap[toolCall.name];
|
const tool = toolsMap[toolCall.name];
|
||||||
if (!tool) {
|
if (!tool) {
|
||||||
|
|
@ -85,7 +93,12 @@ export async function takeActions(
|
||||||
try {
|
try {
|
||||||
const toolResult =
|
const toolResult =
|
||||||
// @ts-expect-error tool.invoke types are weird here...
|
// @ts-expect-error tool.invoke types are weird here...
|
||||||
(await tool.invoke(toolCall.args)) as {
|
(await tool.invoke({
|
||||||
|
...toolCall.args,
|
||||||
|
// Pass in the existing/new sandbox session ID to the tool call.
|
||||||
|
// use `x` prefix to avoid name conflicts with tool args.
|
||||||
|
xSandboxSessionId: sandbox.id,
|
||||||
|
})) as {
|
||||||
result: string;
|
result: string;
|
||||||
status: "success" | "error";
|
status: "success" | "error";
|
||||||
};
|
};
|
||||||
|
|
@ -137,7 +150,6 @@ export async function takeActions(
|
||||||
});
|
});
|
||||||
|
|
||||||
let toolCallResults = await Promise.all(toolCallResultsPromise);
|
let toolCallResults = await Promise.all(toolCallResultsPromise);
|
||||||
const sandbox = await daytonaClient().get(state.sandboxSessionId);
|
|
||||||
const repoPath = getRepoAbsolutePath(state.targetRepository);
|
const repoPath = getRepoAbsolutePath(state.targetRepository);
|
||||||
const changedFiles = await getChangedFilesStatus(repoPath, sandbox);
|
const changedFiles = await getChangedFilesStatus(repoPath, sandbox);
|
||||||
if (changedFiles?.length > 0) {
|
if (changedFiles?.length > 0) {
|
||||||
|
|
@ -174,5 +186,8 @@ ${tc.content}`,
|
||||||
|
|
||||||
return {
|
return {
|
||||||
messages: toolCallResults,
|
messages: toolCallResults,
|
||||||
|
sandboxSessionId: sandbox.id,
|
||||||
|
...(codebaseTree && { codebaseTree }),
|
||||||
|
...(dependenciesInstalled !== null && { dependenciesInstalled }),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -15,7 +15,10 @@ import { loadModel, Task } from "../../../utils/load-model.js";
|
||||||
import { formatPlanPromptWithSummaries } from "../../../utils/plan-prompt.js";
|
import { formatPlanPromptWithSummaries } from "../../../utils/plan-prompt.js";
|
||||||
import { getUserRequest } from "../../../utils/user-request.js";
|
import { getUserRequest } from "../../../utils/user-request.js";
|
||||||
import { AIMessage, ToolMessage } from "@langchain/core/messages";
|
import { AIMessage, ToolMessage } from "@langchain/core/messages";
|
||||||
import { daytonaClient, deleteSandbox } from "../../../utils/sandbox.js";
|
import {
|
||||||
|
deleteSandbox,
|
||||||
|
getSandboxWithErrorHandling,
|
||||||
|
} from "../../../utils/sandbox.js";
|
||||||
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
|
||||||
import { getActivePlanItems } from "@open-swe/shared/open-swe/tasks";
|
import { getActivePlanItems } from "@open-swe/shared/open-swe/tasks";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
|
|
@ -46,15 +49,16 @@ export async function openPullRequest(
|
||||||
state: GraphState,
|
state: GraphState,
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<GraphUpdate> {
|
): Promise<GraphUpdate> {
|
||||||
const sandboxSessionId = state.sandboxSessionId;
|
|
||||||
if (!sandboxSessionId) {
|
|
||||||
throw new Error(
|
|
||||||
"Failed to open pull request: No sandbox session ID found in state.",
|
|
||||||
);
|
|
||||||
}
|
|
||||||
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
|
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
|
||||||
|
|
||||||
const sandbox = await daytonaClient().get(sandboxSessionId);
|
const { sandbox, codebaseTree, dependenciesInstalled } =
|
||||||
|
await getSandboxWithErrorHandling(
|
||||||
|
state.sandboxSessionId,
|
||||||
|
state.targetRepository,
|
||||||
|
state.branchName,
|
||||||
|
config,
|
||||||
|
);
|
||||||
|
const sandboxSessionId = sandbox.id;
|
||||||
|
|
||||||
const { owner, repo } = state.targetRepository;
|
const { owner, repo } = state.targetRepository;
|
||||||
|
|
||||||
|
|
@ -154,5 +158,7 @@ export async function openPullRequest(
|
||||||
sandboxSessionId: undefined,
|
sandboxSessionId: undefined,
|
||||||
dependenciesInstalled: false,
|
dependenciesInstalled: false,
|
||||||
}),
|
}),
|
||||||
|
...(codebaseTree && { codebaseTree }),
|
||||||
|
...(dependenciesInstalled !== null && { dependenciesInstalled }),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,15 @@
|
||||||
import { isAIMessage, ToolMessage } from "@langchain/core/messages";
|
import { isAIMessage, ToolMessage } from "@langchain/core/messages";
|
||||||
import { GraphState, GraphUpdate } from "@open-swe/shared/open-swe/types";
|
import {
|
||||||
|
GraphConfig,
|
||||||
|
GraphState,
|
||||||
|
GraphUpdate,
|
||||||
|
} from "@open-swe/shared/open-swe/types";
|
||||||
import { HumanInterrupt, HumanResponse } from "@langchain/langgraph/prebuilt";
|
import { HumanInterrupt, HumanResponse } from "@langchain/langgraph/prebuilt";
|
||||||
import { END, interrupt, Command } from "@langchain/langgraph";
|
import { END, interrupt, Command } from "@langchain/langgraph";
|
||||||
import { stopSandbox, startSandbox } from "../../../utils/sandbox.js";
|
import {
|
||||||
|
getSandboxWithErrorHandling,
|
||||||
|
stopSandbox,
|
||||||
|
} from "../../../utils/sandbox.js";
|
||||||
|
|
||||||
const constructDescription = (helpRequest: string): string => {
|
const constructDescription = (helpRequest: string): string => {
|
||||||
return `The agent has requested help. Here is the help request:
|
return `The agent has requested help. Here is the help request:
|
||||||
|
|
@ -12,16 +19,18 @@ ${helpRequest}
|
||||||
\`\`\``;
|
\`\`\``;
|
||||||
};
|
};
|
||||||
|
|
||||||
export async function requestHelp(state: GraphState): Promise<Command> {
|
export async function requestHelp(
|
||||||
|
state: GraphState,
|
||||||
|
config: GraphConfig,
|
||||||
|
): Promise<Command> {
|
||||||
const lastMessage = state.internalMessages[state.internalMessages.length - 1];
|
const lastMessage = state.internalMessages[state.internalMessages.length - 1];
|
||||||
if (!isAIMessage(lastMessage) || !lastMessage.tool_calls?.length) {
|
if (!isAIMessage(lastMessage) || !lastMessage.tool_calls?.length) {
|
||||||
throw new Error("Last message is not an AI message with tool calls.");
|
throw new Error("Last message is not an AI message with tool calls.");
|
||||||
}
|
}
|
||||||
const sandboxSessionId = state.sandboxSessionId;
|
const sandboxSessionId = state.sandboxSessionId;
|
||||||
if (!sandboxSessionId) {
|
if (sandboxSessionId) {
|
||||||
throw new Error("Sandbox session ID not found.");
|
await stopSandbox(sandboxSessionId);
|
||||||
}
|
}
|
||||||
await stopSandbox(sandboxSessionId);
|
|
||||||
|
|
||||||
const toolCall = lastMessage.tool_calls[0];
|
const toolCall = lastMessage.tool_calls[0];
|
||||||
|
|
||||||
|
|
@ -52,15 +61,27 @@ export async function requestHelp(state: GraphState): Promise<Command> {
|
||||||
if (typeof interruptRes.args !== "string") {
|
if (typeof interruptRes.args !== "string") {
|
||||||
throw new Error("Interrupt response expected to be a string.");
|
throw new Error("Interrupt response expected to be a string.");
|
||||||
}
|
}
|
||||||
await startSandbox(sandboxSessionId);
|
|
||||||
|
const { sandbox, codebaseTree, dependenciesInstalled } =
|
||||||
|
await getSandboxWithErrorHandling(
|
||||||
|
state.sandboxSessionId,
|
||||||
|
state.targetRepository,
|
||||||
|
state.branchName,
|
||||||
|
config,
|
||||||
|
);
|
||||||
|
|
||||||
const toolMessage = new ToolMessage({
|
const toolMessage = new ToolMessage({
|
||||||
tool_call_id: toolCall.id ?? "",
|
tool_call_id: toolCall.id ?? "",
|
||||||
content: `Human response: ${interruptRes.args}`,
|
content: `Human response: ${interruptRes.args}`,
|
||||||
status: "success",
|
status: "success",
|
||||||
});
|
});
|
||||||
|
|
||||||
const commandUpdate: GraphUpdate = {
|
const commandUpdate: GraphUpdate = {
|
||||||
messages: [toolMessage],
|
messages: [toolMessage],
|
||||||
internalMessages: [toolMessage],
|
internalMessages: [toolMessage],
|
||||||
|
sandboxSessionId: sandbox.id,
|
||||||
|
...(codebaseTree && { codebaseTree }),
|
||||||
|
...(dependenciesInstalled !== null && { dependenciesInstalled }),
|
||||||
};
|
};
|
||||||
return new Command({
|
return new Command({
|
||||||
goto: "generate-action",
|
goto: "generate-action",
|
||||||
|
|
|
||||||
|
|
@ -20,7 +20,7 @@ import {
|
||||||
} from "../../../utils/zod-to-string.js";
|
} from "../../../utils/zod-to-string.js";
|
||||||
import { Command } from "@langchain/langgraph";
|
import { Command } from "@langchain/langgraph";
|
||||||
import { truncateOutput } from "../../../utils/truncate-outputs.js";
|
import { truncateOutput } from "../../../utils/truncate-outputs.js";
|
||||||
import { daytonaClient } from "../../../utils/sandbox.js";
|
import { getSandboxWithErrorHandling } from "../../../utils/sandbox.js";
|
||||||
import { getCodebaseTree } from "../../../utils/tree.js";
|
import { getCodebaseTree } from "../../../utils/tree.js";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { shouldDiagnoseError } from "../utils/tool-message-error.js";
|
import { shouldDiagnoseError } from "../utils/tool-message-error.js";
|
||||||
|
|
@ -40,12 +40,6 @@ export async function takeAction(
|
||||||
throw new Error("Last message is not an AI message with tool calls.");
|
throw new Error("Last message is not an AI message with tool calls.");
|
||||||
}
|
}
|
||||||
|
|
||||||
if (!state.sandboxSessionId) {
|
|
||||||
throw new Error(
|
|
||||||
"Failed to take action: No sandbox session ID found in state.",
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
const applyPatchTool = createApplyPatchTool(state);
|
const applyPatchTool = createApplyPatchTool(state);
|
||||||
const shellTool = createShellTool(state);
|
const shellTool = createShellTool(state);
|
||||||
const rgTool = createRgTool(state);
|
const rgTool = createRgTool(state);
|
||||||
|
|
@ -71,6 +65,13 @@ export async function takeAction(
|
||||||
throw new Error("No tool calls found.");
|
throw new Error("No tool calls found.");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const { sandbox, dependenciesInstalled } = await getSandboxWithErrorHandling(
|
||||||
|
state.sandboxSessionId,
|
||||||
|
state.targetRepository,
|
||||||
|
state.branchName,
|
||||||
|
config,
|
||||||
|
);
|
||||||
|
|
||||||
const toolCallResultsPromise = toolCalls.map(async (toolCall) => {
|
const toolCallResultsPromise = toolCalls.map(async (toolCall) => {
|
||||||
const tool = toolsMap[toolCall.name];
|
const tool = toolsMap[toolCall.name];
|
||||||
|
|
||||||
|
|
@ -90,7 +91,12 @@ export async function takeAction(
|
||||||
try {
|
try {
|
||||||
const toolResult: { result: string; status: "success" | "error" } =
|
const toolResult: { result: string; status: "success" | "error" } =
|
||||||
// @ts-expect-error tool.invoke types are weird here...
|
// @ts-expect-error tool.invoke types are weird here...
|
||||||
await tool.invoke(toolCall.args);
|
await tool.invoke({
|
||||||
|
...toolCall.args,
|
||||||
|
// Pass in the existing/new sandbox session ID to the tool call.
|
||||||
|
// use `x` prefix to avoid name conflicts with tool args.
|
||||||
|
xSandboxSessionId: sandbox.id,
|
||||||
|
});
|
||||||
if (typeof toolResult === "string") {
|
if (typeof toolResult === "string") {
|
||||||
result = toolResult;
|
result = toolResult;
|
||||||
toolCallStatus = "success";
|
toolCallStatus = "success";
|
||||||
|
|
@ -141,7 +147,6 @@ export async function takeAction(
|
||||||
|
|
||||||
// Always check if there are changed files after running a tool.
|
// Always check if there are changed files after running a tool.
|
||||||
// If there are, commit them.
|
// If there are, commit them.
|
||||||
const sandbox = await daytonaClient().get(state.sandboxSessionId);
|
|
||||||
const changedFiles = await getChangedFilesStatus(
|
const changedFiles = await getChangedFilesStatus(
|
||||||
getRepoAbsolutePath(state.targetRepository),
|
getRepoAbsolutePath(state.targetRepository),
|
||||||
sandbox,
|
sandbox,
|
||||||
|
|
@ -169,13 +174,22 @@ export async function takeAction(
|
||||||
|
|
||||||
const codebaseTree = await getCodebaseTree();
|
const codebaseTree = await getCodebaseTree();
|
||||||
|
|
||||||
|
// Prioritize wereDependenciesInstalled over dependenciesInstalled
|
||||||
|
const dependenciesInstalledUpdate =
|
||||||
|
wereDependenciesInstalled !== null
|
||||||
|
? wereDependenciesInstalled
|
||||||
|
: dependenciesInstalled !== null
|
||||||
|
? dependenciesInstalled
|
||||||
|
: null;
|
||||||
|
|
||||||
const commandUpdate: GraphUpdate = {
|
const commandUpdate: GraphUpdate = {
|
||||||
messages: toolCallResults,
|
messages: toolCallResults,
|
||||||
internalMessages: toolCallResults,
|
internalMessages: toolCallResults,
|
||||||
...(branchName && { branchName }),
|
...(branchName && { branchName }),
|
||||||
codebaseTree,
|
codebaseTree,
|
||||||
...(wereDependenciesInstalled !== null && {
|
sandboxSessionId: sandbox.id,
|
||||||
dependenciesInstalled: wereDependenciesInstalled,
|
...(dependenciesInstalledUpdate !== null && {
|
||||||
|
dependenciesInstalled: dependenciesInstalledUpdate,
|
||||||
}),
|
}),
|
||||||
};
|
};
|
||||||
return new Command({
|
return new Command({
|
||||||
|
|
|
||||||
|
|
@ -2,34 +2,22 @@ import { tool } from "@langchain/core/tools";
|
||||||
import { applyPatch } from "diff";
|
import { applyPatch } from "diff";
|
||||||
import { GraphState } from "@open-swe/shared/open-swe/types";
|
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||||
import { readFile, writeFile } from "../utils/read-write.js";
|
import { readFile, writeFile } from "../utils/read-write.js";
|
||||||
import { getCurrentTaskInput } from "@langchain/langgraph";
|
|
||||||
import { fixGitPatch } from "../utils/diff.js";
|
import { fixGitPatch } from "../utils/diff.js";
|
||||||
import { createLogger, LogLevel } from "../utils/logger.js";
|
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||||
import { daytonaClient } from "../utils/sandbox.js";
|
|
||||||
import { createApplyPatchToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createApplyPatchToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
|
import { getSandboxSessionOrThrow } from "./utils/get-sandbox-id.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "ApplyPatchTool");
|
const logger = createLogger(LogLevel.INFO, "ApplyPatchTool");
|
||||||
|
|
||||||
export function createApplyPatchTool(state: GraphState) {
|
export function createApplyPatchTool(state: GraphState) {
|
||||||
const applyPatchTool = tool(
|
const applyPatchTool = tool(
|
||||||
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
||||||
const state = getCurrentTaskInput<GraphState>();
|
const sandbox = await getSandboxSessionOrThrow(input);
|
||||||
const { sandboxSessionId } = state;
|
|
||||||
if (!sandboxSessionId) {
|
|
||||||
logger.error("FAILED TO RUN COMMAND: No sandbox session ID provided", {
|
|
||||||
input,
|
|
||||||
});
|
|
||||||
throw new Error(
|
|
||||||
"FAILED TO RUN COMMAND: No sandbox session ID provided",
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
const { diff, file_path } = input;
|
const { diff, file_path } = input;
|
||||||
|
|
||||||
const workDir = getRepoAbsolutePath(state.targetRepository);
|
const workDir = getRepoAbsolutePath(state.targetRepository);
|
||||||
|
|
||||||
const sandbox = await daytonaClient().get(sandboxSessionId);
|
|
||||||
|
|
||||||
const { success: readFileSuccess, output: readFileOutput } =
|
const { success: readFileSuccess, output: readFileOutput } =
|
||||||
await readFile({
|
await readFile({
|
||||||
sandbox,
|
sandbox,
|
||||||
|
|
|
||||||
|
|
@ -1,15 +1,13 @@
|
||||||
import { tool } from "@langchain/core/tools";
|
import { tool } from "@langchain/core/tools";
|
||||||
import { Sandbox } from "@daytonaio/sdk";
|
|
||||||
import { GraphState } from "@open-swe/shared/open-swe/types";
|
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||||
import { getCurrentTaskInput } from "@langchain/langgraph";
|
|
||||||
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
||||||
import { createLogger, LogLevel } from "../utils/logger.js";
|
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||||
import { daytonaClient } from "../utils/sandbox.js";
|
|
||||||
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
||||||
import { createFindInstancesOfToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createFindInstancesOfToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { z } from "zod";
|
import { z } from "zod";
|
||||||
import { wrapScript } from "../utils/wrap-script.js";
|
import { wrapScript } from "../utils/wrap-script.js";
|
||||||
|
import { getSandboxSessionOrThrow } from "./utils/get-sandbox-id.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "FindInstancesOfTool");
|
const logger = createLogger(LogLevel.INFO, "FindInstancesOfTool");
|
||||||
|
|
||||||
|
|
@ -68,25 +66,10 @@ export function createFindInstancesOfTool(
|
||||||
async (
|
async (
|
||||||
input: z.infer<typeof findInstancesOfFields.schema>,
|
input: z.infer<typeof findInstancesOfFields.schema>,
|
||||||
): Promise<{ result: string; status: "success" | "error" }> => {
|
): Promise<{ result: string; status: "success" | "error" }> => {
|
||||||
let sandbox: Sandbox | undefined;
|
|
||||||
try {
|
try {
|
||||||
const state = getCurrentTaskInput<GraphState>();
|
const sandbox = await getSandboxSessionOrThrow(input);
|
||||||
const { sandboxSessionId } = state;
|
|
||||||
if (!sandboxSessionId) {
|
|
||||||
logger.error(
|
|
||||||
"FAILED TO RUN COMMAND: No sandbox session ID provided",
|
|
||||||
{
|
|
||||||
input,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
throw new Error(
|
|
||||||
"FAILED TO RUN COMMAND: No sandbox session ID provided",
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
||||||
|
|
||||||
sandbox = await daytonaClient().get(sandboxSessionId);
|
|
||||||
const command = formatFindInstancesOfCommand(input);
|
const command = formatFindInstancesOfCommand(input);
|
||||||
logger.info("Running find_instances_of command", {
|
logger.info("Running find_instances_of command", {
|
||||||
command: command.join(" "),
|
command: command.join(" "),
|
||||||
|
|
|
||||||
|
|
@ -1,13 +1,11 @@
|
||||||
import { tool } from "@langchain/core/tools";
|
import { tool } from "@langchain/core/tools";
|
||||||
import { Sandbox } from "@daytonaio/sdk";
|
|
||||||
import { GraphState } from "@open-swe/shared/open-swe/types";
|
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||||
import { getCurrentTaskInput } from "@langchain/langgraph";
|
|
||||||
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
||||||
import { createLogger, LogLevel } from "../utils/logger.js";
|
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||||
import { daytonaClient } from "../utils/sandbox.js";
|
|
||||||
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
||||||
import { createInstallDependenciesToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createInstallDependenciesToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
|
import { getSandboxSessionOrThrow } from "./utils/get-sandbox-id.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "InstallDependenciesTool");
|
const logger = createLogger(LogLevel.INFO, "InstallDependenciesTool");
|
||||||
|
|
||||||
|
|
@ -21,25 +19,10 @@ export function createInstallDependenciesTool(
|
||||||
) {
|
) {
|
||||||
const installDependenciesTool = tool(
|
const installDependenciesTool = tool(
|
||||||
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
||||||
let sandbox: Sandbox | undefined;
|
|
||||||
try {
|
try {
|
||||||
const state = getCurrentTaskInput<GraphState>();
|
const sandbox = await getSandboxSessionOrThrow(input);
|
||||||
const { sandboxSessionId } = state;
|
|
||||||
if (!sandboxSessionId) {
|
|
||||||
logger.error(
|
|
||||||
"FAILED TO INSTALL DEPENDENCIES: No sandbox session ID provided",
|
|
||||||
{
|
|
||||||
input,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
throw new Error(
|
|
||||||
"FAILED TO INSTALL DEPENDENCIES: No sandbox session ID provided",
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
||||||
|
|
||||||
sandbox = await daytonaClient().get(sandboxSessionId);
|
|
||||||
const command = input.command.join(" ");
|
const command = input.command.join(" ");
|
||||||
const workdir = input.workdir || repoRoot;
|
const workdir = input.workdir || repoRoot;
|
||||||
logger.info("Running install dependencies command", {
|
logger.info("Running install dependencies command", {
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,7 @@
|
||||||
import { tool } from "@langchain/core/tools";
|
import { tool } from "@langchain/core/tools";
|
||||||
import { Sandbox } from "@daytonaio/sdk";
|
|
||||||
import { GraphState } from "@open-swe/shared/open-swe/types";
|
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||||
import { getCurrentTaskInput } from "@langchain/langgraph";
|
|
||||||
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
||||||
import { createLogger, LogLevel } from "../utils/logger.js";
|
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||||
import { daytonaClient } from "../utils/sandbox.js";
|
|
||||||
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
||||||
import {
|
import {
|
||||||
createRgToolFields,
|
createRgToolFields,
|
||||||
|
|
@ -12,6 +9,7 @@ import {
|
||||||
} from "@open-swe/shared/open-swe/tools";
|
} from "@open-swe/shared/open-swe/tools";
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { wrapScript } from "../utils/wrap-script.js";
|
import { wrapScript } from "../utils/wrap-script.js";
|
||||||
|
import { getSandboxSessionOrThrow } from "./utils/get-sandbox-id.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "RgTool");
|
const logger = createLogger(LogLevel.INFO, "RgTool");
|
||||||
|
|
||||||
|
|
@ -25,25 +23,10 @@ export function createRgTool(
|
||||||
) {
|
) {
|
||||||
const rgTool = tool(
|
const rgTool = tool(
|
||||||
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
||||||
let sandbox: Sandbox | undefined;
|
|
||||||
try {
|
try {
|
||||||
const state = getCurrentTaskInput<GraphState>();
|
const sandbox = await getSandboxSessionOrThrow(input);
|
||||||
const { sandboxSessionId } = state;
|
|
||||||
if (!sandboxSessionId) {
|
|
||||||
logger.error(
|
|
||||||
"FAILED TO RUN COMMAND: No sandbox session ID provided",
|
|
||||||
{
|
|
||||||
input,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
throw new Error(
|
|
||||||
"FAILED TO RUN COMMAND: No sandbox session ID provided",
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
||||||
|
|
||||||
sandbox = await daytonaClient().get(sandboxSessionId);
|
|
||||||
const command = formatRgCommand({
|
const command = formatRgCommand({
|
||||||
pattern: input.pattern,
|
pattern: input.pattern,
|
||||||
paths: input.paths,
|
paths: input.paths,
|
||||||
|
|
|
||||||
|
|
@ -1,12 +1,10 @@
|
||||||
import { tool } from "@langchain/core/tools";
|
import { tool } from "@langchain/core/tools";
|
||||||
import { Sandbox } from "@daytonaio/sdk";
|
|
||||||
import { GraphState } from "@open-swe/shared/open-swe/types";
|
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||||
import { getCurrentTaskInput } from "@langchain/langgraph";
|
|
||||||
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
||||||
import { createLogger, LogLevel } from "../utils/logger.js";
|
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||||
import { daytonaClient } from "../utils/sandbox.js";
|
|
||||||
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
||||||
import { createShellToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createShellToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
|
import { getSandboxSessionOrThrow } from "./utils/get-sandbox-id.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "ShellTool");
|
const logger = createLogger(LogLevel.INFO, "ShellTool");
|
||||||
|
|
||||||
|
|
@ -20,23 +18,9 @@ export function createShellTool(
|
||||||
) {
|
) {
|
||||||
const shellTool = tool(
|
const shellTool = tool(
|
||||||
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
||||||
let sandbox: Sandbox | undefined;
|
|
||||||
try {
|
try {
|
||||||
const state = getCurrentTaskInput<GraphState>();
|
const sandbox = await getSandboxSessionOrThrow(input);
|
||||||
const { sandboxSessionId } = state;
|
|
||||||
if (!sandboxSessionId) {
|
|
||||||
logger.error(
|
|
||||||
"FAILED TO RUN COMMAND: No sandbox session ID provided",
|
|
||||||
{
|
|
||||||
input,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
throw new Error(
|
|
||||||
"FAILED TO RUN COMMAND: No sandbox session ID provided",
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
sandbox = await daytonaClient().get(sandboxSessionId);
|
|
||||||
const { command, workdir, timeout } = input;
|
const { command, workdir, timeout } = input;
|
||||||
const response = await sandbox.process.executeCommand(
|
const response = await sandbox.process.executeCommand(
|
||||||
command.join(" "),
|
command.join(" "),
|
||||||
|
|
|
||||||
30
apps/open-swe/src/tools/utils/get-sandbox-id.ts
Normal file
30
apps/open-swe/src/tools/utils/get-sandbox-id.ts
Normal file
|
|
@ -0,0 +1,30 @@
|
||||||
|
import { getCurrentTaskInput } from "@langchain/langgraph";
|
||||||
|
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||||
|
import { createLogger, LogLevel } from "../../utils/logger.js";
|
||||||
|
import { daytonaClient } from "../../utils/sandbox.js";
|
||||||
|
import { Sandbox } from "@daytonaio/sdk";
|
||||||
|
|
||||||
|
const logger = createLogger(LogLevel.INFO, "GetSandboxSessionOrThrow");
|
||||||
|
|
||||||
|
export async function getSandboxSessionOrThrow(
|
||||||
|
input: Record<string, unknown>,
|
||||||
|
): Promise<Sandbox> {
|
||||||
|
let sandboxSessionId = "";
|
||||||
|
// Attempt to extract from input.
|
||||||
|
if ("xSandboxSessionId" in input) {
|
||||||
|
sandboxSessionId = input.xSandboxSessionId as string;
|
||||||
|
} else {
|
||||||
|
const state = getCurrentTaskInput<GraphState>();
|
||||||
|
sandboxSessionId = state.sandboxSessionId;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!sandboxSessionId) {
|
||||||
|
logger.error("FAILED TO RUN COMMAND: No sandbox session ID provided", {
|
||||||
|
input,
|
||||||
|
});
|
||||||
|
throw new Error("FAILED TO RUN COMMAND: No sandbox session ID provided");
|
||||||
|
}
|
||||||
|
|
||||||
|
const sandbox = await daytonaClient().get(sandboxSessionId);
|
||||||
|
return sandbox;
|
||||||
|
}
|
||||||
|
|
@ -386,7 +386,7 @@ export async function stashAndClearChanges(
|
||||||
): Promise<ExecuteResponse | false> {
|
): Promise<ExecuteResponse | false> {
|
||||||
try {
|
try {
|
||||||
const gitStashOutput = await sandbox.process.executeCommand(
|
const gitStashOutput = await sandbox.process.executeCommand(
|
||||||
"git stash && git reset --hard",
|
"git add -A && git stash && git reset --hard",
|
||||||
absoluteRepoDir,
|
absoluteRepoDir,
|
||||||
undefined,
|
undefined,
|
||||||
TIMEOUT_SEC,
|
TIMEOUT_SEC,
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,11 @@
|
||||||
import { Daytona, Sandbox, SandboxState } from "@daytonaio/sdk";
|
import { Daytona, Sandbox, SandboxState } from "@daytonaio/sdk";
|
||||||
import { createLogger, LogLevel } from "./logger.js";
|
import { createLogger, LogLevel } from "./logger.js";
|
||||||
|
import { GraphConfig, TargetRepository } from "@open-swe/shared/open-swe/types";
|
||||||
|
import { DEFAULT_SANDBOX_CREATE_PARAMS } from "../constants.js";
|
||||||
|
import { getGitHubTokensFromConfig } from "./github-tokens.js";
|
||||||
|
import { cloneRepo, configureGitUserInRepo } from "./github/git.js";
|
||||||
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
|
import { getCodebaseTree } from "./tree.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "Sandbox");
|
const logger = createLogger(LogLevel.INFO, "Sandbox");
|
||||||
|
|
||||||
|
|
@ -38,22 +44,6 @@ export async function stopSandbox(sandboxSessionId: string): Promise<string> {
|
||||||
return sandbox.id;
|
return sandbox.id;
|
||||||
}
|
}
|
||||||
|
|
||||||
/**
|
|
||||||
* Starts the sandbox.
|
|
||||||
* @param sandboxSessionId The ID of the sandbox to start.
|
|
||||||
* @returns The sandbox client.
|
|
||||||
*/
|
|
||||||
export async function startSandbox(sandboxSessionId: string): Promise<Sandbox> {
|
|
||||||
const sandbox = await daytonaClient().get(sandboxSessionId);
|
|
||||||
if (
|
|
||||||
sandbox.instance.state == SandboxState.STOPPED ||
|
|
||||||
sandbox.instance.state == SandboxState.ARCHIVED
|
|
||||||
) {
|
|
||||||
await daytonaClient().start(sandbox);
|
|
||||||
}
|
|
||||||
return sandbox;
|
|
||||||
}
|
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Deletes the sandbox.
|
* Deletes the sandbox.
|
||||||
* @param sandboxSessionId The ID of the sandbox to delete.
|
* @param sandboxSessionId The ID of the sandbox to delete.
|
||||||
|
|
@ -74,3 +64,82 @@ export async function deleteSandbox(
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export async function getSandboxWithErrorHandling(
|
||||||
|
sandboxSessionId: string | undefined,
|
||||||
|
targetRepository: TargetRepository,
|
||||||
|
branchName: string,
|
||||||
|
config: GraphConfig,
|
||||||
|
): Promise<{
|
||||||
|
sandbox: Sandbox;
|
||||||
|
codebaseTree: string | null;
|
||||||
|
dependenciesInstalled: boolean | null;
|
||||||
|
}> {
|
||||||
|
try {
|
||||||
|
if (!sandboxSessionId) {
|
||||||
|
throw new Error("No sandbox ID provided.");
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.info("Getting sandbox.");
|
||||||
|
// Try to get existing sandbox
|
||||||
|
const sandbox = await daytonaClient().get(sandboxSessionId);
|
||||||
|
|
||||||
|
// Check sandbox state
|
||||||
|
const sandboxInfo = await sandbox.info();
|
||||||
|
const state = sandboxInfo.state;
|
||||||
|
|
||||||
|
if (state === "started") {
|
||||||
|
return {
|
||||||
|
sandbox,
|
||||||
|
codebaseTree: null,
|
||||||
|
dependenciesInstalled: null,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
if (state === "stopped" || state === "archived") {
|
||||||
|
await sandbox.start();
|
||||||
|
return {
|
||||||
|
sandbox,
|
||||||
|
codebaseTree: null,
|
||||||
|
dependenciesInstalled: null,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
// For any other state, recreate sandbox
|
||||||
|
throw new Error(`Sandbox in unrecoverable state: ${state}`);
|
||||||
|
} catch (error) {
|
||||||
|
// Recreate sandbox if any step fails
|
||||||
|
logger.info("Recreating sandbox due to error or unrecoverable state", {
|
||||||
|
error,
|
||||||
|
});
|
||||||
|
|
||||||
|
const sandbox = await daytonaClient().create(DEFAULT_SANDBOX_CREATE_PARAMS);
|
||||||
|
const { githubInstallationToken } = getGitHubTokensFromConfig(config);
|
||||||
|
|
||||||
|
// Clone repository
|
||||||
|
await cloneRepo(sandbox, targetRepository, {
|
||||||
|
githubInstallationToken,
|
||||||
|
stateBranchName: branchName,
|
||||||
|
});
|
||||||
|
|
||||||
|
// Configure git user
|
||||||
|
const absoluteRepoDir = getRepoAbsolutePath(targetRepository);
|
||||||
|
await configureGitUserInRepo(absoluteRepoDir, sandbox, {
|
||||||
|
githubInstallationToken,
|
||||||
|
owner: targetRepository.owner,
|
||||||
|
repo: targetRepository.repo,
|
||||||
|
});
|
||||||
|
|
||||||
|
// Get codebase tree
|
||||||
|
const codebaseTree = await getCodebaseTree(sandbox.id);
|
||||||
|
|
||||||
|
logger.info("Sandbox created successfully", {
|
||||||
|
sandboxId: sandbox.id,
|
||||||
|
});
|
||||||
|
return {
|
||||||
|
sandbox,
|
||||||
|
codebaseTree,
|
||||||
|
dependenciesInstalled: false,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue