mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 15:03:16 +00:00
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:
parent
1380069dbf
commit
88d1b8e734
7 changed files with 323 additions and 113 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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) && {
|
||||
|
|
|
|||
|
|
@ -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}`);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
),
|
||||
});
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue