mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 09:13:14 +00:00
refactor: Move tool schemas to shared (#170)
* refactor: Move tool schemas to shared * cr
This commit is contained in:
parent
eadf05ed76
commit
b25a29f126
15 changed files with 138 additions and 128 deletions
|
|
@ -7,10 +7,9 @@ import { loadModel, Task } from "../../utils/load-model.js";
|
|||
import {
|
||||
createShellTool,
|
||||
createApplyPatchTool,
|
||||
requestHumanHelpTool,
|
||||
updatePlanTool,
|
||||
createRequestHumanHelpToolFields,
|
||||
createUpdatePlanToolFields,
|
||||
} from "../../tools/index.js";
|
||||
import { getRepoAbsolutePath } from "../../utils/git.js";
|
||||
import { formatPlanPrompt } from "../../utils/plan-prompt.js";
|
||||
import { stopSandbox } from "../../utils/sandbox.js";
|
||||
import { createLogger, LogLevel } from "../../utils/logger.js";
|
||||
|
|
@ -18,6 +17,7 @@ import { getCurrentPlanItem } from "../../utils/current-task.js";
|
|||
import { getMessageContentString } from "@open-swe/shared/messages";
|
||||
import { getActivePlanItems } from "@open-swe/shared/open-swe/tasks";
|
||||
import { SYSTEM_PROMPT } from "./prompt.js";
|
||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "GenerateMessageNode");
|
||||
|
||||
|
|
@ -58,8 +58,8 @@ export async function generateAction(
|
|||
const tools = [
|
||||
createShellTool(state),
|
||||
createApplyPatchTool(state),
|
||||
requestHumanHelpTool,
|
||||
updatePlanTool,
|
||||
createRequestHumanHelpToolFields(),
|
||||
createUpdatePlanToolFields(),
|
||||
];
|
||||
const modelWithTools = model.bindTools(tools, {
|
||||
tool_choice: "auto",
|
||||
|
|
|
|||
|
|
@ -9,13 +9,13 @@ import {
|
|||
cloneRepo,
|
||||
configureGitUserInRepo,
|
||||
getBranchName,
|
||||
getRepoAbsolutePath,
|
||||
pullLatestChanges,
|
||||
} from "../utils/git.js";
|
||||
import { daytonaClient } from "../utils/sandbox.js";
|
||||
import { SNAPSHOT_NAME } from "@open-swe/shared/constants";
|
||||
import { getGitHubTokensFromConfig } from "../utils/github-tokens.js";
|
||||
import { getCodebaseTree } from "../utils/tree.js";
|
||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "Initialize");
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import {
|
|||
createPullRequest,
|
||||
getBranchName,
|
||||
getChangedFilesStatus,
|
||||
getRepoAbsolutePath,
|
||||
} from "../utils/git.js";
|
||||
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||
import { z } from "zod";
|
||||
|
|
@ -20,6 +19,7 @@ import { ToolMessage } from "@langchain/core/messages";
|
|||
import { daytonaClient, deleteSandbox } from "../utils/sandbox.js";
|
||||
import { getGitHubTokensFromConfig } from "../utils/github-tokens.js";
|
||||
import { getActivePlanItems } from "@open-swe/shared/open-swe/tasks";
|
||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "Open PR");
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import { GraphState, GraphConfig } from "@open-swe/shared/open-swe/types";
|
|||
import {
|
||||
checkoutBranchAndCommit,
|
||||
getChangedFilesStatus,
|
||||
getRepoAbsolutePath,
|
||||
} from "../utils/git.js";
|
||||
import {
|
||||
formatBadArgsError,
|
||||
|
|
@ -19,6 +18,7 @@ import { Command } from "@langchain/langgraph";
|
|||
import { truncateOutput } from "../utils/truncate-outputs.js";
|
||||
import { daytonaClient } from "../utils/sandbox.js";
|
||||
import { getCodebaseTree } from "../utils/tree.js";
|
||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
||||
|
||||
|
|
|
|||
|
|
@ -7,8 +7,8 @@ import { getMessageContentString } from "@open-swe/shared/messages";
|
|||
import { getUserRequest } from "../../../../utils/user-request.js";
|
||||
import { isHumanMessage } from "@langchain/core/messages";
|
||||
import { formatFollowupMessagePrompt } from "../../utils/followup-prompt.js";
|
||||
import { getRepoAbsolutePath } from "../../../../utils/git.js";
|
||||
import { SYSTEM_PROMPT } from "./prompt.js";
|
||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "GeneratePlanningMessageNode");
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import {
|
|||
isHumanMessage,
|
||||
ToolMessage,
|
||||
} from "@langchain/core/messages";
|
||||
import { sessionPlanTool } from "../../../tools/index.js";
|
||||
import { createSessionPlanToolFields } from "../../../tools/index.js";
|
||||
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
||||
import { loadModel, Task } from "../../../utils/load-model.js";
|
||||
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
|
||||
|
|
@ -52,6 +52,7 @@ export async function generatePlan(
|
|||
config: GraphConfig,
|
||||
): Promise<PlannerGraphUpdate> {
|
||||
const model = await loadModel(config, Task.PLANNER);
|
||||
const sessionPlanTool = createSessionPlanToolFields();
|
||||
const modelWithTools = model.bindTools([sessionPlanTool], {
|
||||
tool_choice: sessionPlanTool.name,
|
||||
parallel_tool_calls: false,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
import { tool } from "@langchain/core/tools";
|
||||
import { z } from "zod";
|
||||
import { applyPatch } from "diff";
|
||||
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||
import { readFile, writeFile } from "../utils/read-write.js";
|
||||
|
|
@ -7,27 +6,11 @@ import { getCurrentTaskInput } from "@langchain/langgraph";
|
|||
import { fixGitPatch } from "../utils/diff.js";
|
||||
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||
import { daytonaClient } from "../utils/sandbox.js";
|
||||
import { getRepoAbsolutePath } from "../utils/git.js";
|
||||
import { createApplyPatchToolFields } from "@open-swe/shared/open-swe/tools";
|
||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "ApplyPatchTool");
|
||||
|
||||
const createApplyPatchToolDescription = (state: GraphState) => {
|
||||
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
||||
return (
|
||||
"Applies a diff to a file given a file path and diff content." +
|
||||
`The working directory this diff will be applied to is \`${repoRoot}\`. Ensure the file paths you provide are relative to this directory.`
|
||||
);
|
||||
};
|
||||
|
||||
const applyPatchToolSchema = z.object({
|
||||
diff: z
|
||||
.string()
|
||||
.describe(
|
||||
`The diff to apply. Use a standard diff format. Ensure this field is ALWAYS provided.`,
|
||||
),
|
||||
file_path: z.string().describe("The file path to apply the diff to."),
|
||||
});
|
||||
|
||||
export function createApplyPatchTool(state: GraphState) {
|
||||
const applyPatchTool = tool(
|
||||
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
||||
|
|
@ -124,11 +107,7 @@ export function createApplyPatchTool(state: GraphState) {
|
|||
status: "success",
|
||||
};
|
||||
},
|
||||
{
|
||||
name: "apply_patch",
|
||||
description: createApplyPatchToolDescription(state),
|
||||
schema: applyPatchToolSchema,
|
||||
},
|
||||
createApplyPatchToolFields(state.targetRepository),
|
||||
);
|
||||
return applyPatchTool;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
export * from "./apply-patch.js";
|
||||
export * from "./shell.js";
|
||||
export * from "./session-plan.js";
|
||||
export * from "./request-human-help.js";
|
||||
export * from "./update-plan.js";
|
||||
export {
|
||||
createUpdatePlanToolFields,
|
||||
createSessionPlanToolFields,
|
||||
createRequestHumanHelpToolFields,
|
||||
} from "@open-swe/shared/open-swe/tools";
|
||||
|
|
|
|||
|
|
@ -1,16 +0,0 @@
|
|||
import { z } from "zod";
|
||||
|
||||
const requestHumanHelpSchema = z.object({
|
||||
help_request: z
|
||||
.string()
|
||||
.describe(
|
||||
"The help request to send to the human. Should be concise, but descriptive.",
|
||||
),
|
||||
});
|
||||
|
||||
export const requestHumanHelpTool = {
|
||||
name: "request_human_help",
|
||||
schema: requestHumanHelpSchema,
|
||||
description:
|
||||
"Use this tool to request help from the human. This should only be called if you are stuck, and you are unable to continue. This will pause your execution until the user responds. You will not be able to go back and fourth with the user, so ensure the help request contains all of the necessary information and context the user might need to respond to your request.",
|
||||
};
|
||||
|
|
@ -1,11 +0,0 @@
|
|||
import { z } from "zod";
|
||||
|
||||
const sessionPlanSchema = z.object({
|
||||
plan: z.array(z.string()).describe("The plan to address the user's request."),
|
||||
});
|
||||
|
||||
export const sessionPlanTool = {
|
||||
name: "session_plan",
|
||||
description: "Call this tool when you are ready to generate a plan.",
|
||||
schema: sessionPlanSchema,
|
||||
};
|
||||
|
|
@ -1,5 +1,4 @@
|
|||
import { tool } from "@langchain/core/tools";
|
||||
import { z } from "zod";
|
||||
import { Sandbox } from "@daytonaio/sdk";
|
||||
import { GraphState } from "@open-swe/shared/open-swe/types";
|
||||
import { getCurrentTaskInput } from "@langchain/langgraph";
|
||||
|
|
@ -7,7 +6,7 @@ import { getSandboxErrorFields } from "../utils/sandbox-error-fields.js";
|
|||
import { createLogger, LogLevel } from "../utils/logger.js";
|
||||
import { daytonaClient } from "../utils/sandbox.js";
|
||||
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
||||
import { getRepoAbsolutePath } from "../utils/git.js";
|
||||
import { createShellToolFields } from "@open-swe/shared/open-swe/tools";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "ShellTool");
|
||||
|
||||
|
|
@ -16,31 +15,6 @@ const DEFAULT_ENV = {
|
|||
COREPACK_ENABLE_DOWNLOAD_PROMPT: "0",
|
||||
};
|
||||
|
||||
const createShellToolSchema = (state: GraphState) => {
|
||||
const repoRoot = getRepoAbsolutePath(state.targetRepository);
|
||||
const shellToolSchema = z.object({
|
||||
command: z
|
||||
.array(z.string())
|
||||
.describe(
|
||||
"The command to run. Ensure the command is properly formatted, with arguments in the correct order, and including any wrapping strings, quotes, etc. By default, this command will be executed in the root of the repository, unless a custom workdir is specified.",
|
||||
),
|
||||
workdir: z
|
||||
.string()
|
||||
.default(repoRoot)
|
||||
.describe(
|
||||
`The working directory for the command. Defaults to the root of the repository (${repoRoot}). You should only specify this if the command you're running can not be executed from the root of the repository.`,
|
||||
),
|
||||
timeout: z
|
||||
.number()
|
||||
.optional()
|
||||
.default(TIMEOUT_SEC)
|
||||
.describe(
|
||||
"The maximum time to wait for the command to complete in seconds.",
|
||||
),
|
||||
});
|
||||
return shellToolSchema;
|
||||
};
|
||||
|
||||
export function createShellTool(state: GraphState) {
|
||||
const shellTool = tool(
|
||||
async (input): Promise<{ result: string; status: "success" | "error" }> => {
|
||||
|
|
@ -110,11 +84,7 @@ export function createShellTool(state: GraphState) {
|
|||
);
|
||||
}
|
||||
},
|
||||
{
|
||||
name: "shell",
|
||||
description: "Runs a shell command, and returns its output.",
|
||||
schema: createShellToolSchema(state),
|
||||
},
|
||||
createShellToolFields(state.targetRepository),
|
||||
);
|
||||
|
||||
return shellTool;
|
||||
|
|
|
|||
|
|
@ -1,19 +0,0 @@
|
|||
import { z } from "zod";
|
||||
|
||||
const updatePlanSchema = z.object({
|
||||
update_plan_reasoning: z
|
||||
.string()
|
||||
.describe(
|
||||
"The reasoning for why you are updating the plan. This should include context which will be useful when actually updating the plan, such as what plan items to update, edit, or remove, along with any other context that would be useful when updating the plan.",
|
||||
),
|
||||
});
|
||||
|
||||
export const updatePlanTool = {
|
||||
name: "update_plan",
|
||||
schema: updatePlanSchema,
|
||||
description:
|
||||
"Call this tool to update the current plan. This should ONLY be called if you want to remove, edit, or add plan items to the current plan." +
|
||||
"\nDo NOT call this tool to mark a plan item as completed, or add a summary." +
|
||||
"\nYou can not edit/remove completed plan items. This tool can only be used to update/add/remove plan items from the remaining and current plan items." +
|
||||
"\nThe reasoning you pass to this tool will be used in the step that actually updates the plan, so ensure it is useful and concise.",
|
||||
};
|
||||
|
|
@ -2,24 +2,13 @@ import { Octokit } from "@octokit/rest";
|
|||
import { Sandbox } from "@daytonaio/sdk";
|
||||
import { createLogger, LogLevel } from "./logger.js";
|
||||
import { GraphConfig, TargetRepository } from "@open-swe/shared/open-swe/types";
|
||||
import { TIMEOUT_SEC, SANDBOX_ROOT_DIR } from "@open-swe/shared/constants";
|
||||
import { TIMEOUT_SEC } from "@open-swe/shared/constants";
|
||||
import { getSandboxErrorFields } from "./sandbox-error-fields.js";
|
||||
import { ExecuteResponse } from "@daytonaio/sdk/dist/types/ExecuteResponse.js";
|
||||
import path from "node:path";
|
||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||
|
||||
const logger = createLogger(LogLevel.INFO, "GitUtil");
|
||||
|
||||
export function getRepoAbsolutePath(
|
||||
targetRepository: TargetRepository,
|
||||
): string {
|
||||
const repoName = targetRepository.repo;
|
||||
if (!repoName) {
|
||||
throw new Error("No repository name provided");
|
||||
}
|
||||
|
||||
return path.join(SANDBOX_ROOT_DIR, repoName);
|
||||
}
|
||||
|
||||
export function getBranchName(config: GraphConfig): string {
|
||||
const threadId = config.configurable?.thread_id;
|
||||
if (!threadId) {
|
||||
|
|
|
|||
13
packages/shared/src/git.ts
Normal file
13
packages/shared/src/git.ts
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
import { SANDBOX_ROOT_DIR } from "./constants.js";
|
||||
import { TargetRepository } from "./open-swe/types.js";
|
||||
|
||||
export function getRepoAbsolutePath(
|
||||
targetRepository: TargetRepository,
|
||||
): string {
|
||||
const repoName = targetRepository.repo;
|
||||
if (!repoName) {
|
||||
throw new Error("No repository name provided");
|
||||
}
|
||||
|
||||
return `${SANDBOX_ROOT_DIR}/${repoName}`;
|
||||
}
|
||||
102
packages/shared/src/open-swe/tools.ts
Normal file
102
packages/shared/src/open-swe/tools.ts
Normal file
|
|
@ -0,0 +1,102 @@
|
|||
import { z } from "zod";
|
||||
import { TargetRepository } from "./types.js";
|
||||
import { getRepoAbsolutePath } from "../git.js";
|
||||
import { TIMEOUT_SEC } from "../constants.js";
|
||||
|
||||
export function createApplyPatchToolFields(targetRepository: TargetRepository) {
|
||||
const repoRoot = getRepoAbsolutePath(targetRepository);
|
||||
const applyPatchToolSchema = z.object({
|
||||
diff: z
|
||||
.string()
|
||||
.describe(
|
||||
`The diff to apply. Use a standard diff format. Ensure this field is ALWAYS provided.`,
|
||||
),
|
||||
file_path: z.string().describe("The file path to apply the diff to."),
|
||||
});
|
||||
|
||||
return {
|
||||
name: "apply_patch",
|
||||
description:
|
||||
"Applies a diff to a file given a file path and diff content." +
|
||||
`The working directory this diff will be applied to is \`${repoRoot}\`. Ensure the file paths you provide are relative to this directory.`,
|
||||
schema: applyPatchToolSchema,
|
||||
};
|
||||
}
|
||||
|
||||
export function createRequestHumanHelpToolFields() {
|
||||
const requestHumanHelpSchema = z.object({
|
||||
help_request: z
|
||||
.string()
|
||||
.describe(
|
||||
"The help request to send to the human. Should be concise, but descriptive.",
|
||||
),
|
||||
});
|
||||
return {
|
||||
name: "request_human_help",
|
||||
schema: requestHumanHelpSchema,
|
||||
description:
|
||||
"Use this tool to request help from the human. This should only be called if you are stuck, and you are unable to continue. This will pause your execution until the user responds. You will not be able to go back and fourth with the user, so ensure the help request contains all of the necessary information and context the user might need to respond to your request.",
|
||||
};
|
||||
}
|
||||
|
||||
export function createSessionPlanToolFields() {
|
||||
const sessionPlanSchema = z.object({
|
||||
plan: z
|
||||
.array(z.string())
|
||||
.describe("The plan to address the user's request."),
|
||||
});
|
||||
return {
|
||||
name: "session_plan",
|
||||
description: "Call this tool when you are ready to generate a plan.",
|
||||
schema: sessionPlanSchema,
|
||||
};
|
||||
}
|
||||
|
||||
export function createShellToolFields(targetRepository: TargetRepository) {
|
||||
const repoRoot = getRepoAbsolutePath(targetRepository);
|
||||
const shellToolSchema = z.object({
|
||||
command: z
|
||||
.array(z.string())
|
||||
.describe(
|
||||
"The command to run. Ensure the command is properly formatted, with arguments in the correct order, and including any wrapping strings, quotes, etc. By default, this command will be executed in the root of the repository, unless a custom workdir is specified.",
|
||||
),
|
||||
workdir: z
|
||||
.string()
|
||||
.default(repoRoot)
|
||||
.describe(
|
||||
`The working directory for the command. Defaults to the root of the repository (${repoRoot}). You should only specify this if the command you're running can not be executed from the root of the repository.`,
|
||||
),
|
||||
timeout: z
|
||||
.number()
|
||||
.optional()
|
||||
.default(TIMEOUT_SEC)
|
||||
.describe(
|
||||
"The maximum time to wait for the command to complete in seconds.",
|
||||
),
|
||||
});
|
||||
return {
|
||||
name: "shell",
|
||||
description: "Runs a shell command, and returns its output.",
|
||||
schema: shellToolSchema,
|
||||
};
|
||||
}
|
||||
|
||||
export function createUpdatePlanToolFields() {
|
||||
const updatePlanSchema = z.object({
|
||||
update_plan_reasoning: z
|
||||
.string()
|
||||
.describe(
|
||||
"The reasoning for why you are updating the plan. This should include context which will be useful when actually updating the plan, such as what plan items to update, edit, or remove, along with any other context that would be useful when updating the plan.",
|
||||
),
|
||||
});
|
||||
|
||||
return {
|
||||
name: "update_plan",
|
||||
schema: updatePlanSchema,
|
||||
description:
|
||||
"Call this tool to update the current plan. This should ONLY be called if you want to remove, edit, or add plan items to the current plan." +
|
||||
"\nDo NOT call this tool to mark a plan item as completed, or add a summary." +
|
||||
"\nYou can not edit/remove completed plan items. This tool can only be used to update/add/remove plan items from the remaining and current plan items." +
|
||||
"\nThe reasoning you pass to this tool will be used in the step that actually updates the plan, so ensure it is useful and concise.",
|
||||
};
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue