From 0baea2350ef040debbc97abea8b01ca71cd0676b Mon Sep 17 00:00:00 2001 From: Brace Sproul Date: Tue, 15 Jul 2025 11:21:05 -0700 Subject: [PATCH] fix: Simplify configurable model groups (#415) * fix: Simplify configurable model groups * cr --- .../manager/nodes/classify-message/index.ts | 2 +- .../manager/utils/generate-issue-fields.ts | 2 +- .../planner/nodes/determine-needs-context.ts | 2 +- .../planner/nodes/generate-message/index.ts | 2 +- .../planner/nodes/generate-plan/index.ts | 2 +- .../src/graphs/planner/nodes/rewrite-plan.ts | 2 +- .../nodes/generate-message/index.ts | 2 +- .../src/graphs/programmer/nodes/open-pr.ts | 3 +- .../programmer/nodes/progress-plan-step.ts | 2 +- .../graphs/programmer/nodes/update-plan.ts | 2 +- .../src/graphs/reviewer/nodes/final-review.ts | 2 +- .../nodes/generate-review-actions/index.ts | 2 +- apps/open-swe/src/utils/load-model.ts | 38 ++- packages/shared/src/open-swe/types.ts | 288 +++++------------- 14 files changed, 111 insertions(+), 240 deletions(-) diff --git a/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts b/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts index b8b3cb88..b53cf73f 100644 --- a/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts +++ b/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts @@ -87,7 +87,7 @@ export async function classifyMessage( description: "Respond to the user's message and determine how to route it.", schema, }; - const model = await loadModel(config, Task.CLASSIFICATION); + const model = await loadModel(config, Task.ROUTER); const modelWithTools = model.bindTools([respondAndRouteTool], { tool_choice: respondAndRouteTool.name, parallel_tool_calls: false, diff --git a/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts b/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts index bd8f7e2f..3a5e9c4d 100644 --- a/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts +++ b/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts @@ -8,7 +8,7 @@ export async function createIssueFieldsFromMessages( messages: BaseMessage[], configurable: GraphConfig["configurable"], ): Promise<{ title: string; body: string }> { - const model = await loadModel({ configurable }, Task.ACTION_GENERATOR); + const model = await loadModel({ configurable }, Task.SUMMARIZER); const githubIssueTool = { name: "create_github_issue", description: "Create a new GitHub issue with the given title and body.", diff --git a/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts b/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts index 5a1e7eb1..51e7b0a9 100644 --- a/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts +++ b/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts @@ -113,7 +113,7 @@ export async function determineNeedsContext( ): Promise { const [missingMessages, model] = await Promise.all([ getMissingMessages(state, config), - loadModel(config, Task.CLASSIFICATION), + loadModel(config, Task.ROUTER), ]); if (!missingMessages.length) { throw new Error( diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts index f31fa725..f5ebeabd 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts @@ -50,7 +50,7 @@ export async function generateAction( state: PlannerGraphState, config: GraphConfig, ): Promise { - const model = await loadModel(config, Task.ACTION_GENERATOR); + const model = await loadModel(config, Task.PROGRAMMER); const mcpTools = await getMcpTools(config); const tools = [ diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts b/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts index 3a2172c1..96e950c9 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts @@ -49,7 +49,7 @@ export async function generatePlan( state: PlannerGraphState, config: GraphConfig, ): Promise { - const model = await loadModel(config, Task.PLANNER); + const model = await loadModel(config, Task.PROGRAMMER); const sessionPlanTool = createSessionPlanToolFields(); const modelWithTools = model.bindTools([sessionPlanTool], { tool_choice: sessionPlanTool.name, diff --git a/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts b/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts index 11edd22c..012bc220 100644 --- a/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts +++ b/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts @@ -240,7 +240,7 @@ export async function rewritePlan( throw new Error("No plan change request found."); } - const model = await loadModel(config, Task.PLANNER); + const model = await loadModel(config, Task.PROGRAMMER); const tasksToModify = await identifyTasksToModify(state, model); const updatedPlanTasks = await updatePlanTasks(state, tasksToModify, model); diff --git a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts index b18c980e..e1e15ff2 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts @@ -87,7 +87,7 @@ export async function generateAction( state: GraphState, config: GraphConfig, ): Promise { - const model = await loadModel(config, Task.ACTION_GENERATOR); + const model = await loadModel(config, Task.PROGRAMMER); const mcpTools = await getMcpTools(config); const tools = [ diff --git a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts index 6fdb5b1a..b2e6f425 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts @@ -88,7 +88,8 @@ export async function openPullRequest( } const openPrTool = createOpenPrToolFields(); - const model = await loadModel(config, Task.SUMMARIZER); + // use the router model since this is a simple task that doesn't need an advanced model + const model = await loadModel(config, Task.ROUTER); const modelWithTool = model.bindTools([openPrTool], { tool_choice: openPrTool.name, parallel_tool_calls: false, diff --git a/apps/open-swe/src/graphs/programmer/nodes/progress-plan-step.ts b/apps/open-swe/src/graphs/programmer/nodes/progress-plan-step.ts index 9a5f371f..18609953 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/progress-plan-step.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/progress-plan-step.ts @@ -65,7 +65,7 @@ export async function progressPlanStep( ): Promise { const markNotCompletedTool = createMarkTaskNotCompletedToolFields(); const markCompletedTool = createMarkTaskCompletedToolFields(); - const model = await loadModel(config, Task.PROGRESS_PLAN_CHECKER); + const model = await loadModel(config, Task.SUMMARIZER); const modelWithTools = model.bindTools( [markNotCompletedTool, markCompletedTool], { diff --git a/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts b/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts index 47f1268f..9f8f6232 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts @@ -117,7 +117,7 @@ export async function updatePlan( ...updatePlanToolCall, }); - const model = await loadModel(config, Task.PLANNER); + const model = await loadModel(config, Task.PROGRAMMER); const modelWithTools = model.bindTools([updatePlanTool], { tool_choice: updatePlanTool.name, parallel_tool_calls: false, diff --git a/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts b/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts index dfdc41d6..fe7c1c23 100644 --- a/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts +++ b/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts @@ -75,7 +75,7 @@ export async function finalReview( const completedTool = createCodeReviewMarkTaskCompletedFields(); const incompleteTool = createCodeReviewMarkTaskNotCompleteFields(); const tools = [completedTool, incompleteTool]; - const model = await loadModel(config, Task.PLANNER); + const model = await loadModel(config, Task.PROGRAMMER); const modelWithTools = model.bindTools(tools, { tool_choice: "any", parallel_tool_calls: false, diff --git a/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts b/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts index 2814a8f1..8878c807 100644 --- a/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts +++ b/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts @@ -73,7 +73,7 @@ export async function generateReviewActions( state: ReviewerGraphState, config: GraphConfig, ): Promise { - const model = await loadModel(config, Task.ACTION_GENERATOR); + const model = await loadModel(config, Task.PROGRAMMER); const tools = [ createSearchTool(state), createShellTool(state), diff --git a/apps/open-swe/src/utils/load-model.ts b/apps/open-swe/src/utils/load-model.ts index 4be8a589..d1451f7e 100644 --- a/apps/open-swe/src/utils/load-model.ts +++ b/apps/open-swe/src/utils/load-model.ts @@ -2,39 +2,37 @@ import { initChatModel } from "langchain/chat_models/universal"; import { GraphConfig } from "@open-swe/shared/open-swe/types"; export enum Task { - PLANNER = "planner", - PLANNER_CONTEXT = "plannerContext", - ACTION_GENERATOR = "actionGenerator", - PROGRESS_PLAN_CHECKER = "progressPlanChecker", + /** + * Used for programmer tasks. This includes: writing code, + * generating plans, taking context gathering actions, etc. + */ + PROGRAMMER = "programmer", + /** + * Used for routing tasks. This includes: initial request + * routing to different agents. + */ + ROUTER = "router", + /** + * Used for summarizing tasks. This includes: summarizing + * the conversation history, summarizing actions taken during + * a task execution, etc. Should be a slightly advanced model. + */ SUMMARIZER = "summarizer", - CLASSIFICATION = "classification", } const TASK_TO_CONFIG_DEFAULTS_MAP = { - [Task.PLANNER]: { + [Task.PROGRAMMER]: { modelName: "anthropic:claude-sonnet-4-0", temperature: 0, }, - [Task.PLANNER_CONTEXT]: { - modelName: "anthropic:claude-sonnet-4-0", - temperature: 0, - }, - [Task.ACTION_GENERATOR]: { - modelName: "anthropic:claude-sonnet-4-0", - temperature: 0, - }, - [Task.PROGRESS_PLAN_CHECKER]: { - modelName: "anthropic:claude-sonnet-4-0", + [Task.ROUTER]: { + modelName: "anthropic:claude-3-5-haiku-latest", temperature: 0, }, [Task.SUMMARIZER]: { modelName: "anthropic:claude-sonnet-4-0", temperature: 0, }, - [Task.CLASSIFICATION]: { - modelName: "anthropic:claude-3-5-haiku-latest", - temperature: 0, - }, }; export async function loadModel(config: GraphConfig, task: Task) { diff --git a/packages/shared/src/open-swe/types.ts b/packages/shared/src/open-swe/types.ts index 304f11fa..40eb4234 100644 --- a/packages/shared/src/open-swe/types.ts +++ b/packages/shared/src/open-swe/types.ts @@ -291,15 +291,16 @@ export const GraphConfigurationMetadata: { description: "Maximum number of review actions during planning", }, }, - plannerModelName: { + programmerModelName: { x_open_swe_ui_config: { type: "select", default: "anthropic:claude-sonnet-4-0", - description: "The model to use for planning", + description: + "The model to use for programming/other advanced technical tasks", options: MODEL_OPTIONS_NO_THINKING, }, }, - plannerTemperature: { + programmerTemperature: { x_open_swe_ui_config: { type: "slider", default: 0, @@ -309,51 +310,15 @@ export const GraphConfigurationMetadata: { description: "Controls randomness (0 = deterministic, 2 = creative)", }, }, - plannerContextModelName: { + routerModelName: { x_open_swe_ui_config: { type: "select", - default: "anthropic:claude-sonnet-4-0", - description: "The model to use for planning", + default: "anthropic:claude-3-5-haiku-latest", + description: "The model to use for routing tasks", options: MODEL_OPTIONS, }, }, - plannerContextTemperature: { - x_open_swe_ui_config: { - type: "slider", - default: 0, - min: 0, - max: 2, - step: 0.1, - description: "Controls randomness (0 = deterministic, 2 = creative)", - }, - }, - actionGeneratorModelName: { - x_open_swe_ui_config: { - type: "select", - default: "anthropic:claude-sonnet-4-0", - description: "The model to use for action generation", - options: MODEL_OPTIONS, - }, - }, - actionGeneratorTemperature: { - x_open_swe_ui_config: { - type: "slider", - default: 0, - min: 0, - max: 2, - step: 0.1, - description: "Controls randomness (0 = deterministic, 2 = creative)", - }, - }, - progressPlanCheckerModelName: { - x_open_swe_ui_config: { - type: "select", - default: "anthropic:claude-sonnet-4-0", - description: "The model to use for progress plan checking", - options: MODEL_OPTIONS_NO_THINKING, - }, - }, - progressPlanCheckerTemperature: { + routerTemperature: { x_open_swe_ui_config: { type: "slider", default: 0, @@ -367,7 +332,8 @@ export const GraphConfigurationMetadata: { x_open_swe_ui_config: { type: "select", default: "anthropic:claude-sonnet-4-0", - description: "The model to use for summarizing the conversation history", + description: + "The model to use for summarizing the conversation history, or extracting key context from large inputs", options: MODEL_OPTIONS_NO_THINKING, }, }, @@ -381,24 +347,6 @@ export const GraphConfigurationMetadata: { description: "Controls randomness (0 = deterministic, 2 = creative)", }, }, - classificationModelName: { - x_open_swe_ui_config: { - type: "select", - default: "anthropic:claude-3-5-haiku-latest", - description: "The model to use for classifying the user's message", - options: MODEL_OPTIONS_NO_THINKING, - }, - }, - classificationTemperature: { - x_open_swe_ui_config: { - type: "slider", - default: 0, - min: 0, - max: 2, - step: 0.1, - description: "Controls randomness (0 = deterministic, 2 = creative)", - }, - }, maxTokens: { x_open_swe_ui_config: { type: "number", @@ -456,200 +404,124 @@ export const GraphConfiguration = z.object({ * Total messages = maxContextActions * 2 + 1 * @default 75 */ - maxContextActions: z - .number() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.maxContextActions), + maxContextActions: withLangGraph(z.number().optional(), { + metadata: GraphConfigurationMetadata.maxContextActions, + }), /** * The maximum number of context gathering actions to take during review. * Each action consists of 2 messages (request & result), plus 1 human message. * Total messages = maxReviewActions * 2 + 1 * @default 30 */ - maxReviewActions: z - .number() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.maxReviewActions), + maxReviewActions: withLangGraph(z.number().optional(), { + metadata: GraphConfigurationMetadata.maxReviewActions, + }), /** - * The model ID to use for the planning step. - * This includes initial planning, and rewriting. + * The model ID to use for programming/other advanced technical tasks. * @default "anthropic:claude-sonnet-4-0" */ - plannerModelName: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.plannerModelName), + programmerModelName: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata.programmerModelName, + }), /** - * The temperature to use for the planning step. - * This includes initial planning, and rewriting. - * If selecting a reasoning model, this will be ignored. + * The temperature to use for programming/other advanced technical tasks. * @default 0 */ - plannerTemperature: z - .number() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.plannerTemperature), + programmerTemperature: withLangGraph(z.number().optional(), { + metadata: GraphConfigurationMetadata.programmerTemperature, + }), /** - * The model ID to use for the planning step. - * This includes initial planning, and rewriting. - * @default "anthropic:claude-sonnet-4-0" - */ - plannerContextModelName: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.plannerContextModelName), - /** - * The temperature to use for the planning step. - * This includes initial planning, and rewriting. - * If selecting a reasoning model, this will be ignored. - * @default 0 - */ - plannerContextTemperature: z - .number() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.plannerContextTemperature), - - /** - * The model ID to use for action generation. - * @default "anthropic:claude-sonnet-4-0" - */ - actionGeneratorModelName: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.actionGeneratorModelName), - /** - * The temperature to use for action generation. - * If selecting a reasoning model, this will be ignored. - * @default 0 - */ - actionGeneratorTemperature: z - .number() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.actionGeneratorTemperature), - - /** - * The model ID to use for progress plan checking. - * @default "anthropic:claude-sonnet-4-0" - */ - progressPlanCheckerModelName: z - .string() - .optional() - .langgraph.metadata( - GraphConfigurationMetadata.progressPlanCheckerModelName, - ), - /** - * The temperature to use for progress plan checking. - * If selecting a reasoning model, this will be ignored. - * @default 0 - */ - progressPlanCheckerTemperature: z - .number() - .optional() - .langgraph.metadata( - GraphConfigurationMetadata.progressPlanCheckerTemperature, - ), - - /** - * The model ID to use for summarizing the conversation history. - * @default "anthropic:claude-sonnet-4-0" - */ - summarizerModelName: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.summarizerModelName), - /** - * The temperature to use for summarizing the conversation history. - * If selecting a reasoning model, this will be ignored. - * @default 0 - */ - summarizerTemperature: z - .number() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.summarizerTemperature), - - /** - * The model ID to use for classifying the user's message. + * The model ID to use for routing tasks. * @default "anthropic:claude-3-5-haiku-latest" */ - classificationModelName: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.classificationModelName), + routerModelName: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata.routerModelName, + }), /** - * The temperature to use for classifying the user's message. - * If selecting a reasoning model, this will be ignored. + * The temperature to use for routing tasks. * @default 0 */ - classificationTemperature: z - .number() - .optional() - .langgraph.metadata(GraphConfigurationMetadata.classificationTemperature), + routerTemperature: withLangGraph(z.number().optional(), { + metadata: GraphConfigurationMetadata.routerTemperature, + }), + /** + * The model ID to use for summarizing the conversation history, or extracting key context from large inputs. + * @default "anthropic:claude-sonnet-4-0" + */ + summarizerModelName: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata.summarizerModelName, + }), + /** + * The temperature to use for summarizing the conversation history, or extracting key context from large inputs. + * @default 0 + */ + summarizerTemperature: withLangGraph(z.number().optional(), { + metadata: GraphConfigurationMetadata.actionGeneratorTemperature, + }), /** * The maximum number of tokens to generate in an individual generation. * @default 10_000 */ - maxTokens: z - .number() - .optional() - .default(() => 10_000) - .langgraph.metadata(GraphConfigurationMetadata.maxTokens), + maxTokens: withLangGraph(z.number(), { + metadata: GraphConfigurationMetadata.maxTokens, + reducer: { + schema: z.number(), + fn: (_state, update) => update, + }, + default: () => 10_000, + }), /** * The user's GitHub access token. To be used in requests to get information about the user. */ - [GITHUB_TOKEN_COOKIE]: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata[GITHUB_TOKEN_COOKIE]), + [GITHUB_TOKEN_COOKIE]: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata[GITHUB_TOKEN_COOKIE], + }), /** * The installation token from the GitHub app. This token allows us to take actions * on the repos the user has granted us access to, but on behalf of the app, not the user. */ - [GITHUB_INSTALLATION_TOKEN_COOKIE]: z - .string() - .optional() - .langgraph.metadata( - GraphConfigurationMetadata[GITHUB_INSTALLATION_TOKEN_COOKIE], - ), + [GITHUB_INSTALLATION_TOKEN_COOKIE]: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata[GITHUB_INSTALLATION_TOKEN_COOKIE], + }), /** * The user's GitHub ID. Required when creating runs triggered by a bot (e.g. GitHub issue) */ - [GITHUB_USER_ID_HEADER]: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata[GITHUB_USER_ID_HEADER]), + [GITHUB_USER_ID_HEADER]: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata[GITHUB_USER_ID_HEADER], + }), /** * The user's GitHub login. Required when creating runs triggered by a bot (e.g. GitHub issue) */ - [GITHUB_USER_LOGIN_HEADER]: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata[GITHUB_USER_LOGIN_HEADER]), + [GITHUB_USER_LOGIN_HEADER]: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata[GITHUB_USER_LOGIN_HEADER], + }), /** * The installation name of the GitHub app. Required when creating runs triggered by a bot (e.g. GitHub issue) */ - [GITHUB_INSTALLATION_NAME]: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata[GITHUB_INSTALLATION_NAME]), + [GITHUB_INSTALLATION_NAME]: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata[GITHUB_INSTALLATION_NAME], + }), /** * GitHub Personal Access Token. Used for simpler authentication in environments like evals * where GitHub App installation tokens are not available or needed. */ - [GITHUB_PAT]: z - .string() - .optional() - .langgraph.metadata(GraphConfigurationMetadata[GITHUB_PAT]), + [GITHUB_PAT]: withLangGraph(z.string().optional(), { + metadata: GraphConfigurationMetadata[GITHUB_PAT], + }), /** * Custom MCP servers configuration as JSON string. Merges with default servers. * @default Default LangGraph docs MCP server */ - mcpServers: z - .string() - .optional() - .default(JSON.stringify(DEFAULT_MCP_SERVERS)) - .langgraph.metadata(GraphConfigurationMetadata.mcpServers), + mcpServers: withLangGraph(z.string(), { + metadata: GraphConfigurationMetadata.mcpServers, + default: () => JSON.stringify(DEFAULT_MCP_SERVERS), + reducer: { + schema: z.string(), + fn: (_state, update) => update, + }, + }), }); export type GraphConfig = LangGraphRunnableConfig<