From 6e470e6bd88d02b6f70fcaefd77d92f3c672e719 Mon Sep 17 00:00:00 2001 From: Brace Sproul Date: Thu, 22 May 2025 13:19:40 -0700 Subject: [PATCH] fix: Allow different model names depending on task (#9) * fix: Allow different model names depending on task * add gemini 2.5 flash and pro * fix anthropic thinking models --- src/nodes/generate-message.ts | 4 +- src/nodes/generate-plan.ts | 4 +- src/nodes/progress-plan-step.ts | 4 +- src/nodes/rewrite-plan.ts | 4 +- src/types.ts | 185 ++++++++++++++++++++++---------- src/utils/load-model.ts | 48 ++++++++- 6 files changed, 182 insertions(+), 67 deletions(-) diff --git a/src/nodes/generate-message.ts b/src/nodes/generate-message.ts index 7877816d..71d816c3 100644 --- a/src/nodes/generate-message.ts +++ b/src/nodes/generate-message.ts @@ -1,5 +1,5 @@ import { GraphState, GraphConfig, GraphUpdate, PlanItem } from "../types.js"; -import { loadModel } from "../utils/load-model.js"; +import { loadModel, Task } from "../utils/load-model.js"; import { shellTool, applyPatchTool } from "../tools/index.js"; import { formatPlanPrompt } from "../utils/plan-prompt.js"; @@ -58,7 +58,7 @@ export async function generateAction( state: GraphState, config: GraphConfig, ): Promise { - const model = await loadModel(config); + const model = await loadModel(config, Task.ACTION_GENERATOR); const tools = [shellTool, applyPatchTool]; const modelWithTools = model.bindTools(tools, { tool_choice: "auto" }); diff --git a/src/nodes/generate-plan.ts b/src/nodes/generate-plan.ts index 5fb0821a..0a1c7a7a 100644 --- a/src/nodes/generate-plan.ts +++ b/src/nodes/generate-plan.ts @@ -1,6 +1,6 @@ import { sessionPlanTool } from "../tools/index.js"; import { GraphState, GraphConfig, GraphUpdate } from "../types.js"; -import { loadModel } from "../utils/load-model.js"; +import { loadModel, Task } from "../utils/load-model.js"; const systemPrompt = `You are operating as a terminal-based agentic coding assistant built by LangChain. It wraps LLM models to enable natural language interaction with a local codebase. You are expected to be precise, safe, and helpful. @@ -19,7 +19,7 @@ export async function generatePlan( state: GraphState, config: GraphConfig, ): Promise { - const model = await loadModel(config); + const model = await loadModel(config, Task.PLANNER); const modelWithTools = model.bindTools([sessionPlanTool], { tool_choice: "auto", }); diff --git a/src/nodes/progress-plan-step.ts b/src/nodes/progress-plan-step.ts index 665f77fb..9afd6f74 100644 --- a/src/nodes/progress-plan-step.ts +++ b/src/nodes/progress-plan-step.ts @@ -1,6 +1,6 @@ import { z } from "zod"; import { GraphConfig, GraphState, GraphUpdate, PlanItem } from "../types.js"; -import { loadModel } from "../utils/load-model.js"; +import { loadModel, Task } from "../utils/load-model.js"; import { formatPlanPrompt } from "../utils/plan-prompt.js"; const systemPrompt = `You are operating as a terminal-based agentic coding assistant built by LangChain. It wraps LLM models to enable natural language interaction with a local codebase. You are expected to be precise, safe, and helpful. @@ -35,7 +35,7 @@ export async function progressPlanStep( state: GraphState, config: GraphConfig, ): Promise { - const model = await loadModel(config); + const model = await loadModel(config, Task.PROGRESS_PLAN_CHECKER); const modelWithTools = model.bindTools([confirmTaskCompletionTool], { tool_choice: confirmTaskCompletionTool.name, }); diff --git a/src/nodes/rewrite-plan.ts b/src/nodes/rewrite-plan.ts index e4cc3ed7..af2ce714 100644 --- a/src/nodes/rewrite-plan.ts +++ b/src/nodes/rewrite-plan.ts @@ -1,5 +1,5 @@ import { GraphState, GraphConfig, GraphUpdate } from "../types.js"; -import { loadModel } from "../utils/load-model.js"; +import { loadModel, Task } from "../utils/load-model.js"; import { sessionPlanTool } from "../tools/index.js"; const systemPrompt = `You are operating as a terminal-based agentic coding assistant built by LangChain. It wraps LLM models to enable natural language interaction with a local codebase. You are expected to be precise, safe, and helpful. @@ -36,7 +36,7 @@ export async function rewritePlan( throw new Error("No plan change request found."); } - const model = await loadModel(config); + const model = await loadModel(config, Task.PLANNER); const modelWithTools = model.bindTools([sessionPlanTool], { // The model should always call the tool when rewriting the plan. tool_choice: sessionPlanTool.name, diff --git a/src/types.ts b/src/types.ts index ce9de321..1b73203c 100644 --- a/src/types.ts +++ b/src/types.ts @@ -47,16 +47,64 @@ export const GraphAnnotation = Annotation.Root({ export type GraphState = typeof GraphAnnotation.State; export type GraphUpdate = typeof GraphAnnotation.Update; -export const MCPConfig = z.object({ - /** - * The MCP server URL. - */ - url: z.string(), - /** - * The list of tools to provide to the LLM. - */ - tools: z.array(z.string()), -}); +const MODEL_OPTIONS = [ + { + label: "Claude Sonnet 4 (Extended Thinking)", + value: "anthropic:extended-thinking:claude-sonnet-4-0", + }, + { + label: "Claude Opus 4 (Extended Thinking)", + value: "anthropic:extended-thinking:claude-opus-4-0", + }, + { + label: "Claude Sonnet 4", + value: "anthropic:claude-sonnet-4-0", + }, + { + label: "Claude Opus 4", + value: "anthropic:claude-opus-4-0", + }, + { + label: "Claude 3.7 Sonnet", + value: "anthropic:claude-3-7-sonnet-latest", + }, + { + label: "Claude 3.5 Sonnet", + value: "anthropic:claude-3-5-sonnet-latest", + }, + { + label: "o4", + value: "openai:o4", + }, + { + label: "o4 mini", + value: "openai:o4-mini", + }, + { + label: "o3", + value: "openai:o3", + }, + { + label: "o3 mini", + value: "openai:o3-mini", + }, + { + label: "GPT 4o", + value: "openai:gpt-4o", + }, + { + label: "GPT 4.1", + value: "openai:gpt-4.1", + }, + { + label: "Gemini 2.5 Pro Preview", + value: "google-genai:gemini-2.5-pro-preview-05-06", + }, + { + label: "Gemini 2.5 Flash Preview", + value: "google-genai:gemini-2.5-flash-preview-05-20", + }, +]; export const GraphConfiguration = z.object({ /** @@ -85,55 +133,28 @@ export const GraphConfiguration = z.object({ */ sandbox_language: z.enum(["js", "python"]).optional().langgraph.metadata({}), /** - * The model ID to use for the reflection generation. - * Should be in the format `provider:model_name`. - * Defaults to `anthropic:claude-3-7-sonnet-latest`. + * The model ID to use for the planning step. + * This includes initial planning, and rewriting. + * @default "anthropic:extended-thinking:claude-sonnet-4-0" */ - modelName: z + plannerModelName: z .string() .optional() .langgraph.metadata({ x_lg_ui_config: { type: "select", - default: "anthropic:claude-3-7-sonnet-latest", - description: "The model to use in all generations", - options: [ - { - label: "Claude 3.7 Sonnet", - value: "anthropic:claude-3-7-sonnet-latest", - }, - { - label: "Claude 3.5 Sonnet", - value: "anthropic:claude-3-5-sonnet-latest", - }, - { - label: "GPT 4o", - value: "openai:gpt-4o", - }, - { - label: "GPT 4.1", - value: "openai:gpt-4.1", - }, - { - label: "o3", - value: "openai:o3", - }, - { - label: "o3 mini", - value: "openai:o3-mini", - }, - { - label: "o4", - value: "openai:o4", - }, - ], + default: "anthropic:extended-thinking:claude-sonnet-4-0", + description: "The model to use for planning", + options: MODEL_OPTIONS, }, }), /** - * The temperature to use for the reflection generation. - * Defaults to `0.7`. + * 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 */ - temperature: z + plannerTemperature: z .number() .optional() .langgraph.metadata({ @@ -146,18 +167,72 @@ export const GraphConfiguration = z.object({ description: "Controls randomness (0 = deterministic, 2 = creative)", }, }), + /** - * The maximum number of tokens to generate. - * Defaults to `1000`. + * The model ID to use for action generation. + * @default "anthropic:claude-sonnet-4-0" */ - maxTokens: z + actionGeneratorModelName: z + .string() + .optional() + .langgraph.metadata({ + x_lg_ui_config: { + type: "select", + default: "anthropic:claude-sonnet-4-0", + description: "The model to use for action generation", + options: MODEL_OPTIONS, + }, + }), + /** + * The temperature to use for action generation. + * If selecting a reasoning model, this will be ignored. + * @default 0 + */ + actionGeneratorTemperature: z .number() .optional() .langgraph.metadata({ x_lg_ui_config: { - type: "number", - min: 1, - description: "The maximum number of tokens to generate", + type: "slider", + default: 0, + min: 0, + max: 2, + step: 0.1, + description: "Controls randomness (0 = deterministic, 2 = creative)", + }, + }), + + /** + * The model ID to use for progress plan checking. + * @default "anthropic:claude-sonnet-4-0" + */ + progressPlanCheckerModelName: z + .string() + .optional() + .langgraph.metadata({ + x_lg_ui_config: { + type: "select", + default: "anthropic:claude-sonnet-4-0", + description: "The model to use for progress plan checking", + options: MODEL_OPTIONS, + }, + }), + /** + * 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({ + x_lg_ui_config: { + type: "slider", + default: 0, + min: 0, + max: 2, + step: 0.1, + description: "Controls randomness (0 = deterministic, 2 = creative)", }, }), }); diff --git a/src/utils/load-model.ts b/src/utils/load-model.ts index 4b9603fd..dbab9926 100644 --- a/src/utils/load-model.ts +++ b/src/utils/load-model.ts @@ -1,15 +1,55 @@ import { initChatModel } from "langchain/chat_models/universal"; import { GraphConfig } from "../types.js"; -export async function loadModel(config: GraphConfig) { +export enum Task { + PLANNER = "planner", + ACTION_GENERATOR = "actionGenerator", + PROGRESS_PLAN_CHECKER = "progressPlanChecker", +} + +const TASK_TO_CONFIG_DEFAULTS_MAP = { + [Task.PLANNER]: { + modelName: "anthropic:extended-thinking: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", + temperature: 0, + }, +}; + +export async function loadModel(config: GraphConfig, task: Task) { const modelStr = - config.configurable?.modelName ?? "anthropic:claude-3-7-sonnet-latest"; + config.configurable?.[`${task}ModelName`] ?? + TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName; + const temperature = + config.configurable?.[`${task}Temperature`] ?? + TASK_TO_CONFIG_DEFAULTS_MAP[task].temperature; + const [modelProvider, ...modelNameParts] = modelStr.split(":"); + + let thinkingModel = false; + if (modelNameParts[0] === "extended-thinking") { + // Using a thinking model. Remove it from the model name. + modelNameParts.shift(); + thinkingModel = true; + } + const modelName = modelNameParts.join(":"); + if (modelProvider === "openai" && modelName.startsWith("o")) { + thinkingModel = true; + } + const model = await initChatModel(modelName, { modelProvider, - temperature: config.configurable?.temperature ?? 0, - maxTokens: config.configurable?.maxTokens ?? undefined, + temperature: thinkingModel ? undefined : temperature, + ...(thinkingModel && modelProvider === "anthropic" + ? { thinking: { budgetTokens: 5000, type: "enabled" } } + : {}), }); return model;