mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-04 10:12:10 +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,
|
GraphState,
|
||||||
GraphConfig,
|
GraphConfig,
|
||||||
GraphUpdate,
|
GraphUpdate,
|
||||||
|
TaskPlan,
|
||||||
} from "@open-swe/shared/open-swe/types";
|
} from "@open-swe/shared/open-swe/types";
|
||||||
import {
|
import {
|
||||||
getModelManager,
|
getModelManager,
|
||||||
loadModel,
|
loadModel,
|
||||||
|
Provider,
|
||||||
supportsParallelToolCallsParam,
|
supportsParallelToolCallsParam,
|
||||||
Task,
|
Task,
|
||||||
} from "../../../../utils/llms/index.js";
|
} from "../../../../utils/llms/index.js";
|
||||||
|
|
@ -50,7 +52,12 @@ import {
|
||||||
trackCachePerformance,
|
trackCachePerformance,
|
||||||
} from "../../../../utils/caching.js";
|
} from "../../../../utils/caching.js";
|
||||||
import { createMarkTaskCompletedToolFields } from "@open-swe/shared/open-swe/tools";
|
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");
|
const logger = createLogger(LogLevel.INFO, "GenerateMessageNode");
|
||||||
|
|
||||||
|
|
@ -91,7 +98,10 @@ const formatStaticInstructionsPrompt = (
|
||||||
|
|
||||||
const formatCacheablePrompt = (
|
const formatCacheablePrompt = (
|
||||||
state: GraphState,
|
state: GraphState,
|
||||||
isAnthropicModel: boolean,
|
args?: {
|
||||||
|
isAnthropicModel?: boolean;
|
||||||
|
excludeCacheControl?: boolean;
|
||||||
|
},
|
||||||
): CacheablePromptSegment[] => {
|
): CacheablePromptSegment[] => {
|
||||||
const codeReview = getCodeReviewFields(state.internalMessages);
|
const codeReview = getCodeReviewFields(state.internalMessages);
|
||||||
|
|
||||||
|
|
@ -99,15 +109,19 @@ const formatCacheablePrompt = (
|
||||||
// Cache Breakpoint 2: Static Instructions
|
// Cache Breakpoint 2: Static Instructions
|
||||||
{
|
{
|
||||||
type: "text",
|
type: "text",
|
||||||
text: formatStaticInstructionsPrompt(state, isAnthropicModel),
|
text: formatStaticInstructionsPrompt(state, !!args?.isAnthropicModel),
|
||||||
cache_control: { type: "ephemeral" },
|
...(!args?.excludeCacheControl
|
||||||
|
? { cache_control: { type: "ephemeral" } }
|
||||||
|
: {}),
|
||||||
},
|
},
|
||||||
|
|
||||||
// Cache Breakpoint 3: Dynamic Context
|
// Cache Breakpoint 3: Dynamic Context
|
||||||
{
|
{
|
||||||
type: "text",
|
type: "text",
|
||||||
text: formatDynamicContextPrompt(state),
|
text: formatDynamicContextPrompt(state),
|
||||||
cache_control: { type: "ephemeral" },
|
...(!args?.excludeCacheControl
|
||||||
|
? { cache_control: { type: "ephemeral" } }
|
||||||
|
: {}),
|
||||||
},
|
},
|
||||||
];
|
];
|
||||||
|
|
||||||
|
|
@ -119,7 +133,9 @@ const formatCacheablePrompt = (
|
||||||
review: codeReview.review,
|
review: codeReview.review,
|
||||||
newActions: codeReview.newActions,
|
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,
|
state: GraphState,
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<GraphUpdate> {
|
options: {
|
||||||
const model = await loadModel(config, Task.PROGRAMMER);
|
latestTaskPlan: TaskPlan | null;
|
||||||
const modelManager = getModelManager();
|
missingMessages: BaseMessage[];
|
||||||
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER);
|
},
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
): Promise<{
|
||||||
config,
|
providerTools: Record<Provider, BindToolsInput[]>;
|
||||||
Task.PROGRAMMER,
|
providerMessages: Record<Provider, BaseMessageLike[]>;
|
||||||
);
|
}> {
|
||||||
const mcpTools = await getMcpTools(config);
|
const mcpTools = await getMcpTools(config);
|
||||||
const markTaskCompletedTool = createMarkTaskCompletedToolFields();
|
|
||||||
const isAnthropicModel = modelName.includes("claude-");
|
|
||||||
const sharedTools = [
|
const sharedTools = [
|
||||||
createGrepTool(state),
|
createGrepTool(state),
|
||||||
createShellTool(state),
|
createShellTool(state),
|
||||||
|
|
@ -168,69 +182,132 @@ export async function generateAction(
|
||||||
createUpdatePlanToolFields(),
|
createUpdatePlanToolFields(),
|
||||||
createGetURLContentTool(state),
|
createGetURLContentTool(state),
|
||||||
createInstallDependenciesTool(state),
|
createInstallDependenciesTool(state),
|
||||||
markTaskCompletedTool,
|
createMarkTaskCompletedToolFields(),
|
||||||
createSearchDocumentForTool(state, config),
|
createSearchDocumentForTool(state, config),
|
||||||
...mcpTools,
|
...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(
|
logger.info(
|
||||||
`MCP tools added to Programmer: ${mcpTools.map((t) => t.name).join(", ")}`,
|
`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, {
|
const anthropicModelTools = [
|
||||||
tool_choice: "auto",
|
...sharedTools,
|
||||||
...(modelSupportsParallelToolCallsParam
|
{
|
||||||
? {
|
type: "text_editor_20250429",
|
||||||
parallel_tool_calls: true,
|
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([
|
const [missingMessages, { taskPlan: latestTaskPlan }] = await Promise.all([
|
||||||
getMissingMessages(state, config),
|
getMissingMessages(state, config),
|
||||||
getPlansFromIssue(state, config),
|
getPlansFromIssue(state, config),
|
||||||
]);
|
]);
|
||||||
|
|
||||||
const inputMessages = filterMessagesWithoutContent([
|
const { providerTools, providerMessages } = await createToolsAndPrompt(
|
||||||
...state.internalMessages,
|
state,
|
||||||
...missingMessages,
|
config,
|
||||||
]);
|
|
||||||
if (!inputMessages.length) {
|
|
||||||
throw new Error("No messages to process.");
|
|
||||||
}
|
|
||||||
|
|
||||||
const inputMessagesWithCache =
|
|
||||||
convertMessagesToCacheControlledMessages(inputMessages);
|
|
||||||
const response = await modelWithTools.invoke([
|
|
||||||
{
|
{
|
||||||
role: "system",
|
latestTaskPlan,
|
||||||
content: formatCacheablePrompt(
|
missingMessages,
|
||||||
{
|
|
||||||
...state,
|
|
||||||
taskPlan: latestTaskPlan ?? state.taskPlan,
|
|
||||||
},
|
|
||||||
isAnthropicModel,
|
|
||||||
),
|
|
||||||
},
|
},
|
||||||
...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;
|
const hasToolCalls = !!response.tool_calls?.length;
|
||||||
// No tool calls means the graph is going to end. Stop the sandbox.
|
// No tool calls means the graph is going to end. Stop the sandbox.
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
import {
|
import {
|
||||||
getModelManager,
|
getModelManager,
|
||||||
loadModel,
|
loadModel,
|
||||||
|
Provider,
|
||||||
supportsParallelToolCallsParam,
|
supportsParallelToolCallsParam,
|
||||||
Task,
|
Task,
|
||||||
} from "../../../../utils/llms/index.js";
|
} from "../../../../utils/llms/index.js";
|
||||||
|
|
@ -26,7 +27,7 @@ import {
|
||||||
formatCodeReviewPrompt,
|
formatCodeReviewPrompt,
|
||||||
getCodeReviewFields,
|
getCodeReviewFields,
|
||||||
} from "../../../../utils/review.js";
|
} 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 { getMessageString } from "../../../../utils/message/content.js";
|
||||||
import {
|
import {
|
||||||
CacheablePromptSegment,
|
CacheablePromptSegment,
|
||||||
|
|
@ -35,6 +36,7 @@ import {
|
||||||
} from "../../../../utils/caching.js";
|
} from "../../../../utils/caching.js";
|
||||||
import { createScratchpadTool } from "../../../../tools/scratchpad.js";
|
import { createScratchpadTool } from "../../../../tools/scratchpad.js";
|
||||||
import { createViewTool } from "../../../../tools/builtin-tools/view.js";
|
import { createViewTool } from "../../../../tools/builtin-tools/view.js";
|
||||||
|
import { BindToolsInput } from "@langchain/core/language_models/chat_models";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "GenerateReviewActionsNode");
|
const logger = createLogger(LogLevel.INFO, "GenerateReviewActionsNode");
|
||||||
|
|
||||||
|
|
@ -66,6 +68,9 @@ function formatSystemPrompt(state: ReviewerGraphState): string {
|
||||||
|
|
||||||
const formatCacheablePrompt = (
|
const formatCacheablePrompt = (
|
||||||
state: ReviewerGraphState,
|
state: ReviewerGraphState,
|
||||||
|
args?: {
|
||||||
|
excludeCacheControl?: boolean;
|
||||||
|
},
|
||||||
): CacheablePromptSegment[] => {
|
): CacheablePromptSegment[] => {
|
||||||
const codeReview = getCodeReviewFields(state.internalMessages);
|
const codeReview = getCodeReviewFields(state.internalMessages);
|
||||||
|
|
||||||
|
|
@ -73,7 +78,9 @@ const formatCacheablePrompt = (
|
||||||
{
|
{
|
||||||
type: "text",
|
type: "text",
|
||||||
text: formatSystemPrompt(state),
|
text: formatSystemPrompt(state),
|
||||||
cache_control: { type: "ephemeral" },
|
...(!args?.excludeCacheControl
|
||||||
|
? { cache_control: { type: "ephemeral" } }
|
||||||
|
: {}),
|
||||||
},
|
},
|
||||||
];
|
];
|
||||||
|
|
||||||
|
|
@ -85,7 +92,9 @@ const formatCacheablePrompt = (
|
||||||
review: codeReview.review,
|
review: codeReview.review,
|
||||||
newActions: codeReview.newActions,
|
newActions: codeReview.newActions,
|
||||||
}),
|
}),
|
||||||
cache_control: { type: "ephemeral" },
|
...(!args?.excludeCacheControl
|
||||||
|
? { cache_control: { type: "ephemeral" } }
|
||||||
|
: {}),
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -94,6 +103,9 @@ const formatCacheablePrompt = (
|
||||||
|
|
||||||
function formatUserConversationHistoryMessage(
|
function formatUserConversationHistoryMessage(
|
||||||
messages: BaseMessage[],
|
messages: BaseMessage[],
|
||||||
|
args?: {
|
||||||
|
excludeCacheControl?: boolean;
|
||||||
|
},
|
||||||
): CacheablePromptSegment[] {
|
): CacheablePromptSegment[] {
|
||||||
return [
|
return [
|
||||||
{
|
{
|
||||||
|
|
@ -104,22 +116,17 @@ If the history has been truncated, it is because the conversation was too long.
|
||||||
<conversation_history>
|
<conversation_history>
|
||||||
${messages.map(getMessageString).join("\n")}
|
${messages.map(getMessageString).join("\n")}
|
||||||
</conversation_history>`,
|
</conversation_history>`,
|
||||||
cache_control: { type: "ephemeral" },
|
...(!args?.excludeCacheControl
|
||||||
|
? { cache_control: { type: "ephemeral" } }
|
||||||
|
: {}),
|
||||||
},
|
},
|
||||||
];
|
];
|
||||||
}
|
}
|
||||||
|
|
||||||
export async function generateReviewActions(
|
function createToolsAndPrompt(state: ReviewerGraphState): {
|
||||||
state: ReviewerGraphState,
|
providerTools: Record<Provider, BindToolsInput[]>;
|
||||||
config: GraphConfig,
|
providerMessages: Record<Provider, BaseMessageLike[]>;
|
||||||
): 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,
|
|
||||||
);
|
|
||||||
const tools = [
|
const tools = [
|
||||||
createGrepTool(state),
|
createGrepTool(state),
|
||||||
createShellTool(state),
|
createShellTool(state),
|
||||||
|
|
@ -129,34 +136,87 @@ export async function generateReviewActions(
|
||||||
"when generating a final review, after all context gathering and reviewing is complete",
|
"when generating a final review, after all context gathering and reviewing is complete",
|
||||||
),
|
),
|
||||||
];
|
];
|
||||||
tools[tools.length - 1] = {
|
const anthropicTools = tools;
|
||||||
...tools[tools.length - 1],
|
anthropicTools[anthropicTools.length - 1] = {
|
||||||
|
...anthropicTools[anthropicTools.length - 1],
|
||||||
cache_control: { type: "ephemeral" },
|
cache_control: { type: "ephemeral" },
|
||||||
} as any;
|
} as any;
|
||||||
|
const nonAnthropicTools = tools;
|
||||||
|
|
||||||
const modelWithTools = model.bindTools(tools, {
|
const anthropicMessages = [
|
||||||
tool_choice: "auto",
|
|
||||||
...(modelSupportsParallelToolCallsParam
|
|
||||||
? {
|
|
||||||
parallel_tool_calls: true,
|
|
||||||
}
|
|
||||||
: {}),
|
|
||||||
});
|
|
||||||
|
|
||||||
const reviewerMessagesWithCache = convertMessagesToCacheControlledMessages(
|
|
||||||
state.reviewerMessages,
|
|
||||||
);
|
|
||||||
const response = await modelWithTools.invoke([
|
|
||||||
{
|
{
|
||||||
role: "system",
|
role: "system",
|
||||||
content: formatCacheablePrompt(state),
|
content: formatCacheablePrompt(state, { excludeCacheControl: false }),
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
role: "user",
|
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", {
|
logger.info("Generated review actions", {
|
||||||
...(getMessageContentString(response.content) && {
|
...(getMessageContentString(response.content) && {
|
||||||
|
|
|
||||||
|
|
@ -26,7 +26,7 @@ export function createViewTool(
|
||||||
sandbox,
|
sandbox,
|
||||||
path,
|
path,
|
||||||
workDir,
|
workDir,
|
||||||
view_range,
|
view_range as [number, number] | undefined,
|
||||||
);
|
);
|
||||||
|
|
||||||
logger.info(`View command executed successfully on ${path}`);
|
logger.info(`View command executed successfully on ${path}`);
|
||||||
|
|
|
||||||
|
|
@ -1,16 +1,31 @@
|
||||||
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
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 { FallbackRunnable } from "../runtime-fallback.js";
|
||||||
import { Task, TASK_TO_CONFIG_DEFAULTS_MAP } from "./constants.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 modelManager = getModelManager();
|
||||||
|
|
||||||
const model = await modelManager.loadModel(config, task);
|
const model = await modelManager.loadModel(config, task);
|
||||||
if (!model) {
|
if (!model) {
|
||||||
throw new Error(`Model loading returned undefined for task: ${task}`);
|
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;
|
return fallbackModel;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
||||||
import { Task } from "./llms/index.js";
|
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 { createLogger, LogLevel } from "./logger.js";
|
||||||
import { Runnable, RunnableConfig } from "@langchain/core/runnables";
|
import { Runnable, RunnableConfig } from "@langchain/core/runnables";
|
||||||
import { StructuredToolInterface } from "@langchain/core/tools";
|
import { StructuredToolInterface } from "@langchain/core/tools";
|
||||||
|
|
@ -8,7 +8,11 @@ import {
|
||||||
ConfigurableChatModelCallOptions,
|
ConfigurableChatModelCallOptions,
|
||||||
ConfigurableModel,
|
ConfigurableModel,
|
||||||
} from "langchain/chat_models/universal";
|
} 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 { ChatResult, ChatGeneration } from "@langchain/core/outputs";
|
||||||
import { BaseLanguageModelInput } from "@langchain/core/language_models/base";
|
import { BaseLanguageModelInput } from "@langchain/core/language_models/base";
|
||||||
import { BindToolsInput } from "@langchain/core/language_models/chat_models";
|
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");
|
const logger = createLogger(LogLevel.DEBUG, "FallbackRunnable");
|
||||||
|
|
||||||
interface ExtractedTools {
|
interface ExtractedTools {
|
||||||
tools: StructuredToolInterface[];
|
tools: BindToolsInput[];
|
||||||
kwargs: Record<string, any>;
|
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<
|
export class FallbackRunnable<
|
||||||
RunInput extends BaseLanguageModelInput = BaseLanguageModelInput,
|
RunInput extends BaseLanguageModelInput = BaseLanguageModelInput,
|
||||||
CallOptions extends
|
CallOptions extends
|
||||||
|
|
@ -30,12 +45,18 @@ export class FallbackRunnable<
|
||||||
private config: GraphConfig;
|
private config: GraphConfig;
|
||||||
private task: Task;
|
private task: Task;
|
||||||
private modelManager: ModelManager;
|
private modelManager: ModelManager;
|
||||||
|
private providerTools?: Record<Provider, BindToolsInput[]>;
|
||||||
|
private providerMessages?: Record<Provider, BaseMessageLike[]>;
|
||||||
|
|
||||||
constructor(
|
constructor(
|
||||||
primaryRunnable: any,
|
primaryRunnable: any,
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
task: Task,
|
task: Task,
|
||||||
modelManager: ModelManager,
|
modelManager: ModelManager,
|
||||||
|
options?: {
|
||||||
|
providerTools?: Record<Provider, BindToolsInput[]>;
|
||||||
|
providerMessages?: Record<Provider, BaseMessageLike[]>;
|
||||||
|
},
|
||||||
) {
|
) {
|
||||||
super({
|
super({
|
||||||
configurableFields: "any",
|
configurableFields: "any",
|
||||||
|
|
@ -47,6 +68,8 @@ export class FallbackRunnable<
|
||||||
this.config = config;
|
this.config = config;
|
||||||
this.task = task;
|
this.task = task;
|
||||||
this.modelManager = modelManager;
|
this.modelManager = modelManager;
|
||||||
|
this.providerTools = options?.providerTools;
|
||||||
|
this.providerMessages = options?.providerMessages;
|
||||||
}
|
}
|
||||||
|
|
||||||
async _generate(
|
async _generate(
|
||||||
|
|
@ -90,11 +113,31 @@ export class FallbackRunnable<
|
||||||
let runnableToUse: Runnable<BaseLanguageModelInput, AIMessageChunk> =
|
let runnableToUse: Runnable<BaseLanguageModelInput, AIMessageChunk> =
|
||||||
model;
|
model;
|
||||||
|
|
||||||
const tools = this.extractBoundTools();
|
// Check if provider-specific tools exist for this provider
|
||||||
if (tools && "bindTools" in runnableToUse && runnableToUse.bindTools) {
|
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(
|
runnableToUse = (runnableToUse as ConfigurableModel).bindTools(
|
||||||
tools.tools,
|
toolsToUse.tools,
|
||||||
tools.kwargs,
|
toolsToUse.kwargs,
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -103,7 +146,14 @@ export class FallbackRunnable<
|
||||||
runnableToUse = runnableToUse.withConfig(config);
|
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);
|
this.modelManager.recordSuccess(modelKey);
|
||||||
return result;
|
return result;
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
|
|
@ -131,6 +181,10 @@ export class FallbackRunnable<
|
||||||
this.config,
|
this.config,
|
||||||
this.task,
|
this.task,
|
||||||
this.modelManager,
|
this.modelManager,
|
||||||
|
{
|
||||||
|
providerTools: this.providerTools,
|
||||||
|
providerMessages: this.providerMessages,
|
||||||
|
},
|
||||||
) as unknown as ConfigurableModel<RunInput, CallOptions>;
|
) as unknown as ConfigurableModel<RunInput, CallOptions>;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -145,6 +199,10 @@ export class FallbackRunnable<
|
||||||
this.config,
|
this.config,
|
||||||
this.task,
|
this.task,
|
||||||
this.modelManager,
|
this.modelManager,
|
||||||
|
{
|
||||||
|
providerTools: this.providerTools,
|
||||||
|
providerMessages: this.providerMessages,
|
||||||
|
},
|
||||||
) as unknown as ConfigurableModel<RunInput, CallOptions>;
|
) as unknown as ConfigurableModel<RunInput, CallOptions>;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -315,7 +315,7 @@ export function mapToolMessageToActionStepProps(
|
||||||
success,
|
success,
|
||||||
command: args.command || "view",
|
command: args.command || "view",
|
||||||
path: args.path || "",
|
path: args.path || "",
|
||||||
view_range: args.view_range,
|
view_range: args.view_range as [number, number] | undefined,
|
||||||
output,
|
output,
|
||||||
reasoningText,
|
reasoningText,
|
||||||
};
|
};
|
||||||
|
|
|
||||||
|
|
@ -554,10 +554,10 @@ export function createViewToolFields(targetRepository: TargetRepository) {
|
||||||
.string()
|
.string()
|
||||||
.describe("The path to the file or directory to operate on"),
|
.describe("The path to the file or directory to operate on"),
|
||||||
view_range: z
|
view_range: z
|
||||||
.tuple([z.number(), z.number()])
|
.array(z.number())
|
||||||
.optional()
|
.optional()
|
||||||
.describe(
|
.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