feat: Add provider-specific tool support for FallbackRunnable (#560)

* Apply patch

* Apply patch

* Apply patch

* Apply patch

* Apply patch

* Apply patch

* support custom system prompts

* cr

* cr

* cr

---------

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
Co-authored-by: bracesproul <braceasproul@gmail.com>
This commit is contained in:
open-swe[bot] 2025-07-28 23:14:41 +00:00 • committed by GitHub
parent 1380069dbf
commit 88d1b8e734
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 323 additions and 113 deletions

View file

@ -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<GraphUpdate> {
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<Provider, BindToolsInput[]>;
providerMessages: Record<Provider, BaseMessageLike[]>;
}> {
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<GraphUpdate> {
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.

View file

@ -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.
<conversation_history>
${messages.map(getMessageString).join("\n")}
</conversation_history>`,
cache_control: { type: "ephemeral" },
...(!args?.excludeCacheControl
? { cache_control: { type: "ephemeral" } }
: {}),
},
];
}
export async function generateReviewActions(
state: ReviewerGraphState,
config: GraphConfig,
): Promise<ReviewerGraphUpdate> {
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<Provider, BindToolsInput[]>;
providerMessages: Record<Provider, BaseMessageLike[]>;
} {
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<ReviewerGraphUpdate> {
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) && {

View file

@ -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}`);

View file

@ -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<Provider, BindToolsInput[]>;
providerMessages?: Record<Provider, BaseMessageLike[]>;
},
) {
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;
}

View file

@ -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<string, any>;
}
function useProviderMessages(
initialInput: BaseLanguageModelInput,
providerMessages?: Record<Provider, BaseMessageLike[]>,
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<Provider, BindToolsInput[]>;
private providerMessages?: Record<Provider, BaseMessageLike[]>;
constructor(
primaryRunnable: any,
config: GraphConfig,
task: Task,
modelManager: ModelManager,
options?: {
providerTools?: Record<Provider, BindToolsInput[]>;
providerMessages?: Record<Provider, BaseMessageLike[]>;
},
) {
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<BaseLanguageModelInput, AIMessageChunk> =
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<RunInput, CallOptions>;
}
@ -145,6 +199,10 @@ export class FallbackRunnable<
this.config,
this.task,
this.modelManager,
{
providerTools: this.providerTools,
providerMessages: this.providerMessages,
},
) as unknown as ConfigurableModel<RunInput, CallOptions>;
}

View file

@ -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,
};

View file

@ -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.",
),
});