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 a1c63e2c..2ee2f35a 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 @@ -3,10 +3,12 @@ import { GraphState, GraphConfig, GraphUpdate, + TaskPlan, } from "@open-swe/shared/open-swe/types"; import { getModelManager, loadModel, + Provider, supportsParallelToolCallsParam, Task, } from "../../../../utils/llms/index.js"; @@ -50,7 +52,12 @@ import { trackCachePerformance, } from "../../../../utils/caching.js"; import { createMarkTaskCompletedToolFields } from "@open-swe/shared/open-swe/tools"; -import { HumanMessage } from "@langchain/core/messages"; +import { + BaseMessage, + BaseMessageLike, + HumanMessage, +} from "@langchain/core/messages"; +import { BindToolsInput } from "@langchain/core/language_models/chat_models"; const logger = createLogger(LogLevel.INFO, "GenerateMessageNode"); @@ -91,7 +98,10 @@ const formatStaticInstructionsPrompt = ( const formatCacheablePrompt = ( state: GraphState, - isAnthropicModel: boolean, + args?: { + isAnthropicModel?: boolean; + excludeCacheControl?: boolean; + }, ): CacheablePromptSegment[] => { const codeReview = getCodeReviewFields(state.internalMessages); @@ -99,15 +109,19 @@ const formatCacheablePrompt = ( // Cache Breakpoint 2: Static Instructions { type: "text", - text: formatStaticInstructionsPrompt(state, isAnthropicModel), - cache_control: { type: "ephemeral" }, + text: formatStaticInstructionsPrompt(state, !!args?.isAnthropicModel), + ...(!args?.excludeCacheControl + ? { cache_control: { type: "ephemeral" } } + : {}), }, // Cache Breakpoint 3: Dynamic Context { type: "text", text: formatDynamicContextPrompt(state), - cache_control: { type: "ephemeral" }, + ...(!args?.excludeCacheControl + ? { cache_control: { type: "ephemeral" } } + : {}), }, ]; @@ -119,7 +133,9 @@ const formatCacheablePrompt = ( review: codeReview.review, newActions: codeReview.newActions, }), - cache_control: { type: "ephemeral" }, + ...(!args?.excludeCacheControl + ? { cache_control: { type: "ephemeral" } } + : {}), }); } @@ -147,20 +163,18 @@ const formatSpecificPlanPrompt = (state: GraphState): HumanMessage => { }); }; -export async function generateAction( +async function createToolsAndPrompt( state: GraphState, config: GraphConfig, -): Promise { - const model = await loadModel(config, Task.PROGRAMMER); - const modelManager = getModelManager(); - const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); - const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( - config, - Task.PROGRAMMER, - ); + options: { + latestTaskPlan: TaskPlan | null; + missingMessages: BaseMessage[]; + }, +): Promise<{ + providerTools: Record; + providerMessages: Record; +}> { const mcpTools = await getMcpTools(config); - const markTaskCompletedTool = createMarkTaskCompletedToolFields(); - const isAnthropicModel = modelName.includes("claude-"); const sharedTools = [ createGrepTool(state), createShellTool(state), @@ -168,69 +182,132 @@ export async function generateAction( createUpdatePlanToolFields(), createGetURLContentTool(state), createInstallDependenciesTool(state), - markTaskCompletedTool, + createMarkTaskCompletedToolFields(), createSearchDocumentForTool(state, config), ...mcpTools, ]; - const anthropicModelTools = [ - { - type: "text_editor_20250429", - name: "str_replace_based_edit_tool", - }, - ]; - const nonAnthropicModelTools = [createApplyPatchTool(state)]; - const tools = [ - ...sharedTools, - ...(isAnthropicModel ? anthropicModelTools : nonAnthropicModelTools), - ]; logger.info( `MCP tools added to Programmer: ${mcpTools.map((t) => t.name).join(", ")}`, ); - // Cache Breakpoint 1: Add cache_control marker to the last tool for tools definition caching - tools[tools.length - 1] = { - ...tools[tools.length - 1], - cache_control: { type: "ephemeral" }, - } as any; - const modelWithTools = model.bindTools(tools, { - tool_choice: "auto", - ...(modelSupportsParallelToolCallsParam - ? { - parallel_tool_calls: true, - } - : {}), - }); + const anthropicModelTools = [ + ...sharedTools, + { + type: "text_editor_20250429", + name: "str_replace_based_edit_tool", + cache_control: { type: "ephemeral" }, + }, + ]; + const nonAnthropicModelTools = [ + ...sharedTools, + { ...createApplyPatchTool(state), cache_control: { type: "ephemeral" } }, + ]; + + const inputMessages = filterMessagesWithoutContent([ + ...state.internalMessages, + ...options.missingMessages, + ]); + if (!inputMessages.length) { + throw new Error("No messages to process."); + } + + const anthropicMessages = [ + { + role: "system", + content: formatCacheablePrompt( + { + ...state, + taskPlan: options.latestTaskPlan ?? state.taskPlan, + }, + { + isAnthropicModel: true, + excludeCacheControl: false, + }, + ), + }, + ...convertMessagesToCacheControlledMessages(inputMessages), + formatSpecificPlanPrompt(state), + ]; + + const nonAnthropicMessages = [ + { + role: "system", + content: formatCacheablePrompt( + { + ...state, + taskPlan: options.latestTaskPlan ?? state.taskPlan, + }, + { + isAnthropicModel: false, + excludeCacheControl: true, + }, + ), + }, + ...inputMessages, + formatSpecificPlanPrompt(state), + ]; + + return { + providerTools: { + anthropic: anthropicModelTools, + openai: nonAnthropicModelTools, + "google-genai": nonAnthropicModelTools, + }, + providerMessages: { + anthropic: anthropicMessages, + openai: nonAnthropicMessages, + "google-genai": nonAnthropicMessages, + }, + }; +} + +export async function generateAction( + state: GraphState, + config: GraphConfig, +): Promise { + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); + const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( + config, + Task.PROGRAMMER, + ); + const markTaskCompletedTool = createMarkTaskCompletedToolFields(); + const isAnthropicModel = modelName.includes("claude-"); const [missingMessages, { taskPlan: latestTaskPlan }] = await Promise.all([ getMissingMessages(state, config), getPlansFromIssue(state, config), ]); - const inputMessages = filterMessagesWithoutContent([ - ...state.internalMessages, - ...missingMessages, - ]); - if (!inputMessages.length) { - throw new Error("No messages to process."); - } - - const inputMessagesWithCache = - convertMessagesToCacheControlledMessages(inputMessages); - const response = await modelWithTools.invoke([ + const { providerTools, providerMessages } = await createToolsAndPrompt( + state, + config, { - role: "system", - content: formatCacheablePrompt( - { - ...state, - taskPlan: latestTaskPlan ?? state.taskPlan, - }, - isAnthropicModel, - ), + latestTaskPlan, + missingMessages, }, - ...inputMessagesWithCache, - formatSpecificPlanPrompt(state), - ]); + ); + + const model = await loadModel(config, Task.PROGRAMMER, { + providerTools: providerTools, + providerMessages: providerMessages, + }); + + const modelWithTools = model.bindTools( + isAnthropicModel ? providerTools.anthropic : providerTools.openai, + { + tool_choice: "auto", + ...(modelSupportsParallelToolCallsParam + ? { + parallel_tool_calls: true, + } + : {}), + }, + ); + const response = await modelWithTools.invoke( + isAnthropicModel ? providerMessages.anthropic : providerMessages.openai, + ); const hasToolCalls = !!response.tool_calls?.length; // No tool calls means the graph is going to end. Stop the sandbox. 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 e30d3801..3c1f5093 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 @@ -1,6 +1,7 @@ import { getModelManager, loadModel, + Provider, supportsParallelToolCallsParam, Task, } from "../../../../utils/llms/index.js"; @@ -26,7 +27,7 @@ import { formatCodeReviewPrompt, getCodeReviewFields, } from "../../../../utils/review.js"; -import { BaseMessage } from "@langchain/core/messages"; +import { BaseMessage, BaseMessageLike } from "@langchain/core/messages"; import { getMessageString } from "../../../../utils/message/content.js"; import { CacheablePromptSegment, @@ -35,6 +36,7 @@ import { } from "../../../../utils/caching.js"; import { createScratchpadTool } from "../../../../tools/scratchpad.js"; import { createViewTool } from "../../../../tools/builtin-tools/view.js"; +import { BindToolsInput } from "@langchain/core/language_models/chat_models"; const logger = createLogger(LogLevel.INFO, "GenerateReviewActionsNode"); @@ -66,6 +68,9 @@ function formatSystemPrompt(state: ReviewerGraphState): string { const formatCacheablePrompt = ( state: ReviewerGraphState, + args?: { + excludeCacheControl?: boolean; + }, ): CacheablePromptSegment[] => { const codeReview = getCodeReviewFields(state.internalMessages); @@ -73,7 +78,9 @@ const formatCacheablePrompt = ( { type: "text", text: formatSystemPrompt(state), - cache_control: { type: "ephemeral" }, + ...(!args?.excludeCacheControl + ? { cache_control: { type: "ephemeral" } } + : {}), }, ]; @@ -85,7 +92,9 @@ const formatCacheablePrompt = ( review: codeReview.review, newActions: codeReview.newActions, }), - cache_control: { type: "ephemeral" }, + ...(!args?.excludeCacheControl + ? { cache_control: { type: "ephemeral" } } + : {}), }); } @@ -94,6 +103,9 @@ const formatCacheablePrompt = ( function formatUserConversationHistoryMessage( messages: BaseMessage[], + args?: { + excludeCacheControl?: boolean; + }, ): CacheablePromptSegment[] { return [ { @@ -104,22 +116,17 @@ If the history has been truncated, it is because the conversation was too long. ${messages.map(getMessageString).join("\n")} `, - cache_control: { type: "ephemeral" }, + ...(!args?.excludeCacheControl + ? { cache_control: { type: "ephemeral" } } + : {}), }, ]; } -export async function generateReviewActions( - state: ReviewerGraphState, - config: GraphConfig, -): Promise { - const model = await loadModel(config, Task.REVIEWER); - const modelManager = getModelManager(); - const modelName = modelManager.getModelNameForTask(config, Task.REVIEWER); - const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( - config, - Task.REVIEWER, - ); +function createToolsAndPrompt(state: ReviewerGraphState): { + providerTools: Record; + providerMessages: Record; +} { const tools = [ createGrepTool(state), createShellTool(state), @@ -129,34 +136,87 @@ export async function generateReviewActions( "when generating a final review, after all context gathering and reviewing is complete", ), ]; - tools[tools.length - 1] = { - ...tools[tools.length - 1], + const anthropicTools = tools; + anthropicTools[anthropicTools.length - 1] = { + ...anthropicTools[anthropicTools.length - 1], cache_control: { type: "ephemeral" }, } as any; + const nonAnthropicTools = tools; - const modelWithTools = model.bindTools(tools, { - tool_choice: "auto", - ...(modelSupportsParallelToolCallsParam - ? { - parallel_tool_calls: true, - } - : {}), - }); - - const reviewerMessagesWithCache = convertMessagesToCacheControlledMessages( - state.reviewerMessages, - ); - const response = await modelWithTools.invoke([ + const anthropicMessages = [ { role: "system", - content: formatCacheablePrompt(state), + content: formatCacheablePrompt(state, { excludeCacheControl: false }), }, { role: "user", - content: formatUserConversationHistoryMessage(state.internalMessages), + content: formatUserConversationHistoryMessage(state.internalMessages, { + excludeCacheControl: false, + }), }, - ...reviewerMessagesWithCache, - ]); + ...convertMessagesToCacheControlledMessages(state.reviewerMessages), + ]; + const nonAnthropicMessages = [ + { + role: "system", + content: formatCacheablePrompt(state, { excludeCacheControl: true }), + }, + { + role: "user", + content: formatUserConversationHistoryMessage(state.internalMessages, { + excludeCacheControl: true, + }), + }, + ...state.reviewerMessages, + ]; + + return { + providerTools: { + anthropic: anthropicTools, + openai: nonAnthropicTools, + "google-genai": nonAnthropicTools, + }, + providerMessages: { + anthropic: anthropicMessages, + openai: nonAnthropicMessages, + "google-genai": nonAnthropicMessages, + }, + }; +} + +export async function generateReviewActions( + state: ReviewerGraphState, + config: GraphConfig, +): Promise { + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.REVIEWER); + const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( + config, + Task.REVIEWER, + ); + const isAnthropicModel = modelName.includes("claude-"); + + const { providerTools, providerMessages } = createToolsAndPrompt(state); + + const model = await loadModel(config, Task.REVIEWER, { + providerTools, + providerMessages, + }); + const modelWithTools = model.bindTools( + isAnthropicModel ? providerTools.anthropic : providerTools.openai, + { + tool_choice: "auto", + ...(modelSupportsParallelToolCallsParam + ? { + parallel_tool_calls: true, + } + : {}), + }, + ); + + const response = await modelWithTools.invoke( + isAnthropicModel ? providerMessages.anthropic : providerMessages.openai, + ); logger.info("Generated review actions", { ...(getMessageContentString(response.content) && { diff --git a/apps/open-swe/src/tools/builtin-tools/view.ts b/apps/open-swe/src/tools/builtin-tools/view.ts index 1485af63..1ca6f57a 100644 --- a/apps/open-swe/src/tools/builtin-tools/view.ts +++ b/apps/open-swe/src/tools/builtin-tools/view.ts @@ -26,7 +26,7 @@ export function createViewTool( sandbox, path, workDir, - view_range, + view_range as [number, number] | undefined, ); logger.info(`View command executed successfully on ${path}`); diff --git a/apps/open-swe/src/utils/llms/load-model.ts b/apps/open-swe/src/utils/llms/load-model.ts index 089af472..1d074210 100644 --- a/apps/open-swe/src/utils/llms/load-model.ts +++ b/apps/open-swe/src/utils/llms/load-model.ts @@ -1,16 +1,31 @@ import { GraphConfig } from "@open-swe/shared/open-swe/types"; -import { getModelManager } from "./model-manager.js"; +import { getModelManager, Provider } from "./model-manager.js"; import { FallbackRunnable } from "../runtime-fallback.js"; import { Task, TASK_TO_CONFIG_DEFAULTS_MAP } from "./constants.js"; +import { BindToolsInput } from "@langchain/core/language_models/chat_models"; +import { BaseMessageLike } from "@langchain/core/messages"; -export async function loadModel(config: GraphConfig, task: Task) { +export async function loadModel( + config: GraphConfig, + task: Task, + options?: { + providerTools?: Record; + providerMessages?: Record; + }, +) { const modelManager = getModelManager(); const model = await modelManager.loadModel(config, task); if (!model) { throw new Error(`Model loading returned undefined for task: ${task}`); } - const fallbackModel = new FallbackRunnable(model, config, task, modelManager); + const fallbackModel = new FallbackRunnable( + model, + config, + task, + modelManager, + options, + ); return fallbackModel; } diff --git a/apps/open-swe/src/utils/runtime-fallback.ts b/apps/open-swe/src/utils/runtime-fallback.ts index bb14b2ba..5f5a6b21 100644 --- a/apps/open-swe/src/utils/runtime-fallback.ts +++ b/apps/open-swe/src/utils/runtime-fallback.ts @@ -1,6 +1,6 @@ import { GraphConfig } from "@open-swe/shared/open-swe/types"; import { Task } from "./llms/index.js"; -import { ModelManager } from "./llms/model-manager.js"; +import { ModelManager, Provider } from "./llms/model-manager.js"; import { createLogger, LogLevel } from "./logger.js"; import { Runnable, RunnableConfig } from "@langchain/core/runnables"; import { StructuredToolInterface } from "@langchain/core/tools"; @@ -8,7 +8,11 @@ import { ConfigurableChatModelCallOptions, ConfigurableModel, } from "langchain/chat_models/universal"; -import { AIMessageChunk, BaseMessage } from "@langchain/core/messages"; +import { + AIMessageChunk, + BaseMessage, + BaseMessageLike, +} from "@langchain/core/messages"; import { ChatResult, ChatGeneration } from "@langchain/core/outputs"; import { BaseLanguageModelInput } from "@langchain/core/language_models/base"; import { BindToolsInput } from "@langchain/core/language_models/chat_models"; @@ -17,10 +21,21 @@ import { getMessageContentString } from "@open-swe/shared/messages"; const logger = createLogger(LogLevel.DEBUG, "FallbackRunnable"); interface ExtractedTools { - tools: StructuredToolInterface[]; + tools: BindToolsInput[]; kwargs: Record; } +function useProviderMessages( + initialInput: BaseLanguageModelInput, + providerMessages?: Record, + provider?: Provider, +): BaseLanguageModelInput { + if (!provider || !providerMessages?.[provider]) { + return initialInput; + } + return providerMessages[provider]; +} + export class FallbackRunnable< RunInput extends BaseLanguageModelInput = BaseLanguageModelInput, CallOptions extends @@ -30,12 +45,18 @@ export class FallbackRunnable< private config: GraphConfig; private task: Task; private modelManager: ModelManager; + private providerTools?: Record; + private providerMessages?: Record; constructor( primaryRunnable: any, config: GraphConfig, task: Task, modelManager: ModelManager, + options?: { + providerTools?: Record; + providerMessages?: Record; + }, ) { super({ configurableFields: "any", @@ -47,6 +68,8 @@ export class FallbackRunnable< this.config = config; this.task = task; this.modelManager = modelManager; + this.providerTools = options?.providerTools; + this.providerMessages = options?.providerMessages; } async _generate( @@ -90,11 +113,31 @@ export class FallbackRunnable< let runnableToUse: Runnable = model; - const tools = this.extractBoundTools(); - if (tools && "bindTools" in runnableToUse && runnableToUse.bindTools) { + // Check if provider-specific tools exist for this provider + const providerSpecificTools = + this.providerTools?.[modelConfig.provider]; + let toolsToUse: ExtractedTools | null = null; + + if (providerSpecificTools) { + // Use provider-specific tools if available + const extractedTools = this.extractBoundTools(); + toolsToUse = { + tools: providerSpecificTools, + kwargs: extractedTools?.kwargs || {}, + }; + } else { + // Fall back to extracted bound tools from primary model + toolsToUse = this.extractBoundTools(); + } + + if ( + toolsToUse && + "bindTools" in runnableToUse && + runnableToUse.bindTools + ) { runnableToUse = (runnableToUse as ConfigurableModel).bindTools( - tools.tools, - tools.kwargs, + toolsToUse.tools, + toolsToUse.kwargs, ); } @@ -103,7 +146,14 @@ export class FallbackRunnable< runnableToUse = runnableToUse.withConfig(config); } - const result = await runnableToUse.invoke(input, options); + const result = await runnableToUse.invoke( + useProviderMessages( + input, + this.providerMessages, + modelConfig.provider, + ), + options, + ); this.modelManager.recordSuccess(modelKey); return result; } catch (error) { @@ -131,6 +181,10 @@ export class FallbackRunnable< this.config, this.task, this.modelManager, + { + providerTools: this.providerTools, + providerMessages: this.providerMessages, + }, ) as unknown as ConfigurableModel; } @@ -145,6 +199,10 @@ export class FallbackRunnable< this.config, this.task, this.modelManager, + { + providerTools: this.providerTools, + providerMessages: this.providerMessages, + }, ) as unknown as ConfigurableModel; } diff --git a/apps/web/src/components/thread/messages/ai.tsx b/apps/web/src/components/thread/messages/ai.tsx index 0b650de8..bea37ea7 100644 --- a/apps/web/src/components/thread/messages/ai.tsx +++ b/apps/web/src/components/thread/messages/ai.tsx @@ -315,7 +315,7 @@ export function mapToolMessageToActionStepProps( success, command: args.command || "view", path: args.path || "", - view_range: args.view_range, + view_range: args.view_range as [number, number] | undefined, output, reasoningText, }; diff --git a/packages/shared/src/open-swe/tools.ts b/packages/shared/src/open-swe/tools.ts index 866479ca..e027c14f 100644 --- a/packages/shared/src/open-swe/tools.ts +++ b/packages/shared/src/open-swe/tools.ts @@ -554,10 +554,10 @@ export function createViewToolFields(targetRepository: TargetRepository) { .string() .describe("The path to the file or directory to operate on"), view_range: z - .tuple([z.number(), z.number()]) + .array(z.number()) .optional() .describe( - "Optional array of two integers [start, end] specifying line numbers to view. Line numbers are 1-indexed. Use -1 for end to read to end of file. Only applies to view command.", + "Optional array of two integers [start, end] specifying line numbers to view. Line numbers are 1-indexed. Use -1 for end to read to end of file. Only applies to view command. If this is passed, ensure it is a valid array, containing only two positive integers.", ), });