fix: fallback parallel tool calling fix (#628)

This commit is contained in:
Aliyan Ishfaq 2025-07-31 14:16:43 -07:00 • committed by GitHub
parent bedd2c9b68
commit 65b700e698
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 29 additions and 7 deletions

View file

@ -29,7 +29,7 @@ export async function loadModel(
return fallbackModel;
}
const MODELS_NO_PARALLEL_TOOL_CALLING = ["openai:o3", "openai:o3-mini"];
export const MODELS_NO_PARALLEL_TOOL_CALLING = ["openai:o3", "openai:o3-mini"];
export function supportsParallelToolCallsParam(
config: GraphConfig,

View file

@ -42,9 +42,9 @@ export enum CircuitState {
}
export const PROVIDER_FALLBACK_ORDER = [
"google-genai",
"anthropic",
"openai",
"anthropic",
"google-genai",
] as const;
export type Provider = (typeof PROVIDER_FALLBACK_ORDER)[number];
@ -238,10 +238,21 @@ export class ModelManager {
(!selectedModelConfig ||
fallbackModel.modelName !== selectedModelConfig.modelName)
) {
// Check if fallback model is a thinking model
const isThinkingModel =
(provider === "openai" && fallbackModel.modelName.startsWith("o")) ||
fallbackModel.modelName.includes("extended-thinking");
const fallbackConfig = {
...fallbackModel,
temperature: baseConfig.temperature,
temperature: isThinkingModel ? undefined : baseConfig.temperature,
maxTokens: baseConfig.maxTokens,
...(isThinkingModel
? {
thinkingModel: true,
thinkingBudgetTokens: THINKING_BUDGET_TOKENS,
}
: {}),
};
configs.push(fallbackConfig);
}
@ -348,9 +359,9 @@ export class ModelManager {
[Task.SUMMARIZER]: "gemini-2.5-pro",
},
openai: {
[Task.PLANNER]: "gpt-4.1",
[Task.PLANNER]: "o3",
[Task.PROGRAMMER]: "gpt-4.1",
[Task.REVIEWER]: "gpt-4.1",
[Task.REVIEWER]: "o3",
[Task.ROUTER]: "gpt-4o-mini",
[Task.SUMMARIZER]: "gpt-4.1-mini",
},

View file

@ -18,6 +18,7 @@ import { BaseLanguageModelInput } from "@langchain/core/language_models/base";
import { BindToolsInput } from "@langchain/core/language_models/chat_models";
import { getMessageContentString } from "@open-swe/shared/messages";
import { getConfig } from "@langchain/langgraph";
import { MODELS_NO_PARALLEL_TOOL_CALLING } from "./llms/load-model.js";
const logger = createLogger(LogLevel.DEBUG, "FallbackRunnable");
@ -141,9 +142,19 @@ export class FallbackRunnable<
"bindTools" in runnableToUse &&
runnableToUse.bindTools
) {
const supportsParallelToolCall =
!MODELS_NO_PARALLEL_TOOL_CALLING.some(
(modelName) => modelKey === modelName,
);
const kwargs = { ...toolsToUse.kwargs };
if (!supportsParallelToolCall && "parallel_tool_calls" in kwargs) {
delete kwargs.parallel_tool_calls;
}
runnableToUse = (runnableToUse as ConfigurableModel).bindTools(
toolsToUse.tools,
toolsToUse.kwargs,
kwargs,
);
}