mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-04 04:23:20 +00:00
feat: Added LLM fallback (#433)
* prompt changes and eval fix * prompt deletions reverted * formatting * chore: code cleaning * feat: add LangGraph evaluation support * fix: py lg check logging fix * chore: code cleaning * dataset added * chore: code cleaning * chore: code cleaning * Readme changes * feat: LLM fallback added * feat: add llm fallback * chore: code cleaning * chore: code cleaning * chore: code cleaned * feat: llm fallback in runtime * fix: build errors * refactor: FallbackRunnable now extends ConfigurableModel * Readme update * feat: llm as a judge script added * chore: code formatted * fix: use GA gemini models instead of preview * fix: type improvements * feat: added skeleton for tool py-dev-server * fix: code cleaning & type improvements * fix: code cleaning * cr --------- Co-authored-by: bracesproul <braceasproul@gmail.com>
This commit is contained in:
parent
6a140027f2
commit
f822fda82d
6 changed files with 1244 additions and 674 deletions
|
|
@ -36,4 +36,4 @@ Open SWE can be used in multiple ways:
|
||||||
|
|
||||||
# Documentation
|
# Documentation
|
||||||
|
|
||||||
To get started using Open SWE locally, see the [documentation here](https://docs.langchain.com/labs/swe/).
|
To get started using Open SWE locally, see the [documentation here](https://docs.langchain.com/labs/swe/).
|
||||||
|
|
@ -4,7 +4,6 @@
|
||||||
import { GraphConfig, PlanItem } from "@open-swe/shared/open-swe/types";
|
import { GraphConfig, PlanItem } from "@open-swe/shared/open-swe/types";
|
||||||
import { z } from "zod";
|
import { z } from "zod";
|
||||||
import { tool } from "@langchain/core/tools";
|
import { tool } from "@langchain/core/tools";
|
||||||
import { ConfigurableModel } from "langchain/chat_models/universal";
|
|
||||||
import { traceable } from "langsmith/traceable";
|
import { traceable } from "langsmith/traceable";
|
||||||
import {
|
import {
|
||||||
PlannerGraphState,
|
PlannerGraphState,
|
||||||
|
|
@ -19,6 +18,7 @@ import {
|
||||||
supportsParallelToolCallsParam,
|
supportsParallelToolCallsParam,
|
||||||
Task,
|
Task,
|
||||||
} from "../../../utils/load-model.js";
|
} from "../../../utils/load-model.js";
|
||||||
|
import { FallbackRunnable } from "../../../utils/runtime-fallback.js";
|
||||||
|
|
||||||
const systemPromptIdentifyChanges = `You are operating as an agentic coding assistant built by LangChain. You've previously been given a task to generate a plan of action for, to address the user's initial request.
|
const systemPromptIdentifyChanges = `You are operating as an agentic coding assistant built by LangChain. You've previously been given a task to generate a plan of action for, to address the user's initial request.
|
||||||
|
|
||||||
|
|
@ -97,7 +97,7 @@ const formatSysPromptRewritePlan = (
|
||||||
|
|
||||||
async function identifyTasksToModifyFunc(
|
async function identifyTasksToModifyFunc(
|
||||||
state: PlannerGraphState,
|
state: PlannerGraphState,
|
||||||
model: ConfigurableModel,
|
model: FallbackRunnable,
|
||||||
supportsParallelToolCallsParam: boolean,
|
supportsParallelToolCallsParam: boolean,
|
||||||
): Promise<PlanItem[]> {
|
): Promise<PlanItem[]> {
|
||||||
if (!state.planChangeRequest) {
|
if (!state.planChangeRequest) {
|
||||||
|
|
@ -188,7 +188,7 @@ const identifyTasksToModify = traceable(identifyTasksToModifyFunc, {
|
||||||
async function updatePlanTasksFunc(
|
async function updatePlanTasksFunc(
|
||||||
state: PlannerGraphState,
|
state: PlannerGraphState,
|
||||||
tasksToModify: PlanItem[],
|
tasksToModify: PlanItem[],
|
||||||
model: ConfigurableModel,
|
model: FallbackRunnable,
|
||||||
supportsParallelToolCallsParam: boolean,
|
supportsParallelToolCallsParam: boolean,
|
||||||
): Promise<string[]> {
|
): Promise<string[]> {
|
||||||
if (!state.planChangeRequest) {
|
if (!state.planChangeRequest) {
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,6 @@
|
||||||
import { initChatModel } from "langchain/chat_models/universal";
|
|
||||||
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
||||||
import { isAllowedUser } from "./github/allowed-users.js";
|
import { getModelManager } from "./model-manager.js";
|
||||||
import { decryptSecret } from "@open-swe/shared/crypto";
|
import { FallbackRunnable } from "./runtime-fallback.js";
|
||||||
|
|
||||||
export enum Task {
|
export enum Task {
|
||||||
/**
|
/**
|
||||||
|
|
@ -37,91 +36,15 @@ const TASK_TO_CONFIG_DEFAULTS_MAP = {
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
const providerToApiKey = (
|
|
||||||
providerName: string,
|
|
||||||
apiKeys: Record<string, string>,
|
|
||||||
): string => {
|
|
||||||
switch (providerName) {
|
|
||||||
case "openai":
|
|
||||||
return apiKeys.openaiApiKey;
|
|
||||||
case "anthropic":
|
|
||||||
return apiKeys.anthropicApiKey;
|
|
||||||
case "google-genai":
|
|
||||||
return apiKeys.googleApiKey;
|
|
||||||
default:
|
|
||||||
throw new Error(`Unknown provider: ${providerName}`);
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
export async function loadModel(config: GraphConfig, task: Task) {
|
export async function loadModel(config: GraphConfig, task: Task) {
|
||||||
const modelStr =
|
const modelManager = getModelManager();
|
||||||
config.configurable?.[`${task}ModelName`] ??
|
|
||||||
TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName;
|
|
||||||
const temperature =
|
|
||||||
config.configurable?.[`${task}Temperature`] ??
|
|
||||||
TASK_TO_CONFIG_DEFAULTS_MAP[task].temperature;
|
|
||||||
|
|
||||||
const [modelProvider, ...modelNameParts] = modelStr.split(":");
|
const model = await modelManager.loadModel(config, task);
|
||||||
|
if (!model) {
|
||||||
let thinkingModel = false;
|
throw new Error(`Model loading returned undefined for task: ${task}`);
|
||||||
if (modelNameParts[0] === "extended-thinking") {
|
|
||||||
// Using a thinking model. Remove it from the model name.
|
|
||||||
modelNameParts.shift();
|
|
||||||
thinkingModel = true;
|
|
||||||
}
|
}
|
||||||
|
const fallbackModel = new FallbackRunnable(model, config, task, modelManager);
|
||||||
const modelName = modelNameParts.join(":");
|
return fallbackModel;
|
||||||
if (modelProvider === "openai" && modelName.startsWith("o")) {
|
|
||||||
thinkingModel = true;
|
|
||||||
}
|
|
||||||
|
|
||||||
const thinkingBudgetTokens = 5000;
|
|
||||||
const thinkingMaxTokens = thinkingBudgetTokens * 4;
|
|
||||||
|
|
||||||
let maxTokens = config.configurable?.maxTokens ?? 10_000;
|
|
||||||
if (modelName.includes("claude-3-5-haiku")) {
|
|
||||||
// The max tokens for haiku is 8192
|
|
||||||
maxTokens = maxTokens > 8_192 ? 8_192 : maxTokens;
|
|
||||||
}
|
|
||||||
|
|
||||||
// TODO: Fix types
|
|
||||||
const userLogin = (config.configurable as any)?.langgraph_auth_user
|
|
||||||
?.display_name;
|
|
||||||
const secretsEncryptionKey = process.env.SECRETS_ENCRYPTION_KEY;
|
|
||||||
if (!secretsEncryptionKey) {
|
|
||||||
throw new Error("SECRETS_ENCRYPTION_KEY environment variable is required");
|
|
||||||
}
|
|
||||||
if (!userLogin) {
|
|
||||||
throw new Error("User login not found in config");
|
|
||||||
}
|
|
||||||
const apiKeys = config.configurable?.apiKeys;
|
|
||||||
let apiKey: string | null = null;
|
|
||||||
if (!isAllowedUser(userLogin)) {
|
|
||||||
if (!apiKeys) {
|
|
||||||
throw new Error("API keys not found in config");
|
|
||||||
}
|
|
||||||
apiKey = decryptSecret(
|
|
||||||
providerToApiKey(modelProvider, apiKeys),
|
|
||||||
secretsEncryptionKey,
|
|
||||||
);
|
|
||||||
if (!apiKey) {
|
|
||||||
throw new Error("No API key found for provider: " + modelProvider);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
const model = await initChatModel(modelName, {
|
|
||||||
modelProvider,
|
|
||||||
temperature: thinkingModel ? undefined : temperature,
|
|
||||||
...(apiKey ? { apiKey } : {}),
|
|
||||||
...(thinkingModel && modelProvider === "anthropic"
|
|
||||||
? {
|
|
||||||
thinking: { budget_tokens: thinkingBudgetTokens, type: "enabled" },
|
|
||||||
maxTokens: thinkingMaxTokens,
|
|
||||||
}
|
|
||||||
: { maxTokens }),
|
|
||||||
});
|
|
||||||
|
|
||||||
return model;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
const MODELS_NO_PARALLEL_TOOL_CALLING = ["openai:o3", "openai:o3-mini"];
|
const MODELS_NO_PARALLEL_TOOL_CALLING = ["openai:o3", "openai:o3-mini"];
|
||||||
|
|
|
||||||
429
apps/open-swe/src/utils/model-manager.ts
Normal file
429
apps/open-swe/src/utils/model-manager.ts
Normal file
|
|
@ -0,0 +1,429 @@
|
||||||
|
import {
|
||||||
|
ConfigurableModel,
|
||||||
|
initChatModel,
|
||||||
|
} from "langchain/chat_models/universal";
|
||||||
|
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
||||||
|
import { createLogger, LogLevel } from "./logger.js";
|
||||||
|
import { Task } from "./load-model.js";
|
||||||
|
import { isAllowedUser } from "./github/allowed-users.js";
|
||||||
|
import { decryptSecret } from "@open-swe/shared/crypto";
|
||||||
|
|
||||||
|
const logger = createLogger(LogLevel.INFO, "ModelManager");
|
||||||
|
|
||||||
|
type InitChatModelArgs = Parameters<typeof initChatModel>[1];
|
||||||
|
|
||||||
|
export interface CircuitBreakerState {
|
||||||
|
state: CircuitState;
|
||||||
|
failureCount: number;
|
||||||
|
lastFailureTime: number;
|
||||||
|
openedAt?: number;
|
||||||
|
}
|
||||||
|
|
||||||
|
interface ModelLoadConfig {
|
||||||
|
provider: Provider;
|
||||||
|
modelName: string;
|
||||||
|
temperature?: number;
|
||||||
|
maxTokens?: number;
|
||||||
|
thinkingModel?: boolean;
|
||||||
|
thinkingBudgetTokens?: number;
|
||||||
|
}
|
||||||
|
|
||||||
|
export enum CircuitState {
|
||||||
|
/*
|
||||||
|
* CLOSED: Normal operation
|
||||||
|
*/
|
||||||
|
CLOSED = "CLOSED",
|
||||||
|
/*
|
||||||
|
* OPEN: Failing, use fallback
|
||||||
|
*/
|
||||||
|
OPEN = "OPEN",
|
||||||
|
}
|
||||||
|
|
||||||
|
export const PROVIDER_FALLBACK_ORDER = [
|
||||||
|
"google-genai",
|
||||||
|
"anthropic",
|
||||||
|
"openai",
|
||||||
|
] as const;
|
||||||
|
export type Provider = (typeof PROVIDER_FALLBACK_ORDER)[number];
|
||||||
|
|
||||||
|
export interface ModelManagerConfig {
|
||||||
|
/*
|
||||||
|
* Failures before opening circuit
|
||||||
|
*/
|
||||||
|
circuitBreakerFailureThreshold: number;
|
||||||
|
/*
|
||||||
|
* Time to wait before trying again (ms)
|
||||||
|
*/
|
||||||
|
circuitBreakerTimeoutMs: number;
|
||||||
|
fallbackOrder: Provider[];
|
||||||
|
}
|
||||||
|
|
||||||
|
export const DEFAULT_MODEL_MANAGER_CONFIG: ModelManagerConfig = {
|
||||||
|
circuitBreakerFailureThreshold: 2, // TBD, need to test
|
||||||
|
circuitBreakerTimeoutMs: 180000, // 3 minutes timeout
|
||||||
|
fallbackOrder: [...PROVIDER_FALLBACK_ORDER],
|
||||||
|
};
|
||||||
|
|
||||||
|
const MAX_RETRIES = 3;
|
||||||
|
const THINKING_BUDGET_TOKENS = 5000;
|
||||||
|
|
||||||
|
const providerToApiKey = (
|
||||||
|
providerName: string,
|
||||||
|
apiKeys: Record<string, string>,
|
||||||
|
): string => {
|
||||||
|
switch (providerName) {
|
||||||
|
case "openai":
|
||||||
|
return apiKeys.openaiApiKey;
|
||||||
|
case "anthropic":
|
||||||
|
return apiKeys.anthropicApiKey;
|
||||||
|
case "google-genai":
|
||||||
|
return apiKeys.googleApiKey;
|
||||||
|
default:
|
||||||
|
throw new Error(`Unknown provider: ${providerName}`);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
export class ModelManager {
|
||||||
|
private config: ModelManagerConfig;
|
||||||
|
private circuitBreakers: Map<string, CircuitBreakerState> = new Map();
|
||||||
|
|
||||||
|
constructor(config: Partial<ModelManagerConfig> = {}) {
|
||||||
|
this.config = { ...DEFAULT_MODEL_MANAGER_CONFIG, ...config };
|
||||||
|
|
||||||
|
logger.info("Initialized", {
|
||||||
|
config: this.config,
|
||||||
|
fallbackOrder: this.config.fallbackOrder,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Load a single model (no fallback during loading)
|
||||||
|
*/
|
||||||
|
async loadModel(graphConfig: GraphConfig, task: Task) {
|
||||||
|
const baseConfig = this.getBaseConfigForTask(graphConfig, task);
|
||||||
|
const model = await this.initializeModel(baseConfig, graphConfig);
|
||||||
|
return model;
|
||||||
|
}
|
||||||
|
/**
|
||||||
|
* Initialize the model instance
|
||||||
|
*/
|
||||||
|
public async initializeModel(
|
||||||
|
config: ModelLoadConfig,
|
||||||
|
graphConfig?: GraphConfig,
|
||||||
|
) {
|
||||||
|
const {
|
||||||
|
provider,
|
||||||
|
modelName,
|
||||||
|
temperature,
|
||||||
|
maxTokens,
|
||||||
|
thinkingModel,
|
||||||
|
thinkingBudgetTokens,
|
||||||
|
} = config;
|
||||||
|
|
||||||
|
const thinkingMaxTokens = thinkingBudgetTokens
|
||||||
|
? thinkingBudgetTokens * 4
|
||||||
|
: undefined;
|
||||||
|
|
||||||
|
let finalMaxTokens = maxTokens ?? 10_000;
|
||||||
|
if (modelName.includes("claude-3-5-haiku")) {
|
||||||
|
finalMaxTokens = finalMaxTokens > 8_192 ? 8_192 : finalMaxTokens;
|
||||||
|
}
|
||||||
|
|
||||||
|
let apiKey: string | null = null;
|
||||||
|
if (graphConfig) {
|
||||||
|
const userLogin = (graphConfig.configurable as any)?.langgraph_auth_user
|
||||||
|
?.display_name;
|
||||||
|
const secretsEncryptionKey = process.env.SECRETS_ENCRYPTION_KEY;
|
||||||
|
if (!secretsEncryptionKey) {
|
||||||
|
throw new Error(
|
||||||
|
"SECRETS_ENCRYPTION_KEY environment variable is required",
|
||||||
|
);
|
||||||
|
}
|
||||||
|
if (!userLogin) {
|
||||||
|
throw new Error("User login not found in config");
|
||||||
|
}
|
||||||
|
const apiKeys = graphConfig.configurable?.apiKeys;
|
||||||
|
if (!isAllowedUser(userLogin)) {
|
||||||
|
if (!apiKeys) {
|
||||||
|
throw new Error("API keys not found in config");
|
||||||
|
}
|
||||||
|
apiKey = decryptSecret(
|
||||||
|
providerToApiKey(provider, apiKeys),
|
||||||
|
secretsEncryptionKey,
|
||||||
|
);
|
||||||
|
if (!apiKey) {
|
||||||
|
throw new Error("No API key found for provider: " + provider);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const modelOptions: InitChatModelArgs = {
|
||||||
|
modelProvider: provider,
|
||||||
|
temperature: thinkingModel ? undefined : temperature,
|
||||||
|
max_retries: MAX_RETRIES,
|
||||||
|
...(apiKey ? { apiKey } : {}),
|
||||||
|
...(thinkingModel && provider === "anthropic"
|
||||||
|
? {
|
||||||
|
thinking: { budget_tokens: thinkingBudgetTokens, type: "enabled" },
|
||||||
|
maxTokens: thinkingMaxTokens,
|
||||||
|
}
|
||||||
|
: { maxTokens: finalMaxTokens }),
|
||||||
|
};
|
||||||
|
|
||||||
|
logger.debug("Initializing model", {
|
||||||
|
provider,
|
||||||
|
modelName,
|
||||||
|
});
|
||||||
|
|
||||||
|
return await initChatModel(modelName, modelOptions);
|
||||||
|
}
|
||||||
|
|
||||||
|
public getModelConfigs(
|
||||||
|
config: GraphConfig,
|
||||||
|
task: Task,
|
||||||
|
selectedModel: ConfigurableModel,
|
||||||
|
) {
|
||||||
|
const configs: ModelLoadConfig[] = [];
|
||||||
|
const baseConfig = this.getBaseConfigForTask(config, task);
|
||||||
|
|
||||||
|
const defaultConfig = selectedModel._defaultConfig;
|
||||||
|
let selectedModelConfig: ModelLoadConfig | null = null;
|
||||||
|
|
||||||
|
if (defaultConfig) {
|
||||||
|
const provider = defaultConfig.modelProvider as Provider;
|
||||||
|
const modelName = defaultConfig.model;
|
||||||
|
|
||||||
|
if (provider && modelName) {
|
||||||
|
const isThinkingModel = baseConfig.thinkingModel;
|
||||||
|
selectedModelConfig = {
|
||||||
|
provider,
|
||||||
|
modelName,
|
||||||
|
temperature: defaultConfig.temperature ?? baseConfig.temperature,
|
||||||
|
maxTokens: defaultConfig.maxTokens ?? baseConfig.maxTokens,
|
||||||
|
...(isThinkingModel
|
||||||
|
? {
|
||||||
|
thinkingModel: true,
|
||||||
|
thinkingBudgetTokens: THINKING_BUDGET_TOKENS,
|
||||||
|
}
|
||||||
|
: {}),
|
||||||
|
};
|
||||||
|
configs.push(selectedModelConfig);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add fallback models
|
||||||
|
for (const provider of this.config.fallbackOrder) {
|
||||||
|
const fallbackModel = this.getDefaultModelForProvider(provider, task);
|
||||||
|
if (
|
||||||
|
fallbackModel &&
|
||||||
|
(!selectedModelConfig ||
|
||||||
|
fallbackModel.modelName !== selectedModelConfig.modelName)
|
||||||
|
) {
|
||||||
|
const fallbackConfig = {
|
||||||
|
...fallbackModel,
|
||||||
|
temperature: baseConfig.temperature,
|
||||||
|
maxTokens: baseConfig.maxTokens,
|
||||||
|
};
|
||||||
|
configs.push(fallbackConfig);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return configs;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Get base configuration for a task from GraphConfig
|
||||||
|
*/
|
||||||
|
private getBaseConfigForTask(
|
||||||
|
config: GraphConfig,
|
||||||
|
task: Task,
|
||||||
|
): ModelLoadConfig {
|
||||||
|
const taskMap = {
|
||||||
|
[Task.PROGRAMMER]: {
|
||||||
|
modelName:
|
||||||
|
config.configurable?.[`${task}ModelName`] ??
|
||||||
|
"google-genai:gemini-2.5-pro",
|
||||||
|
temperature: config.configurable?.[`${task}Temperature`] ?? 0,
|
||||||
|
},
|
||||||
|
[Task.ROUTER]: {
|
||||||
|
modelName:
|
||||||
|
config.configurable?.[`${task}ModelName`] ??
|
||||||
|
"anthropic:claude-3-5-haiku-latest",
|
||||||
|
temperature: config.configurable?.[`${task}Temperature`] ?? 0,
|
||||||
|
},
|
||||||
|
[Task.SUMMARIZER]: {
|
||||||
|
modelName:
|
||||||
|
config.configurable?.[`${task}ModelName`] ??
|
||||||
|
"anthropic:claude-sonnet-4-0",
|
||||||
|
temperature: config.configurable?.[`${task}Temperature`] ?? 0,
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
const taskConfig = taskMap[task];
|
||||||
|
const modelStr = taskConfig.modelName;
|
||||||
|
const [modelProvider, ...modelNameParts] = modelStr.split(":");
|
||||||
|
|
||||||
|
let thinkingModel = false;
|
||||||
|
if (modelNameParts[0] === "extended-thinking") {
|
||||||
|
thinkingModel = true;
|
||||||
|
modelNameParts.shift();
|
||||||
|
}
|
||||||
|
|
||||||
|
const modelName = modelNameParts.join(":");
|
||||||
|
if (modelProvider === "openai" && modelName.startsWith("o")) {
|
||||||
|
thinkingModel = true;
|
||||||
|
}
|
||||||
|
|
||||||
|
const thinkingBudgetTokens = THINKING_BUDGET_TOKENS;
|
||||||
|
|
||||||
|
return {
|
||||||
|
modelName,
|
||||||
|
provider: modelProvider as Provider,
|
||||||
|
temperature: taskConfig.temperature,
|
||||||
|
maxTokens: config.configurable?.maxTokens ?? 10_000,
|
||||||
|
thinkingModel,
|
||||||
|
thinkingBudgetTokens,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Get default model for a provider and task
|
||||||
|
*/
|
||||||
|
private getDefaultModelForProvider(
|
||||||
|
provider: Provider,
|
||||||
|
task: Task,
|
||||||
|
): ModelLoadConfig | null {
|
||||||
|
const defaultModels: Record<Provider, Record<Task, string>> = {
|
||||||
|
anthropic: {
|
||||||
|
[Task.PROGRAMMER]: "claude-sonnet-4-0",
|
||||||
|
[Task.ROUTER]: "claude-3-5-haiku-latest",
|
||||||
|
[Task.SUMMARIZER]: "claude-sonnet-4-0",
|
||||||
|
},
|
||||||
|
"google-genai": {
|
||||||
|
[Task.PROGRAMMER]: "gemini-2.5-pro",
|
||||||
|
[Task.ROUTER]: "gemini-2.5-flash",
|
||||||
|
[Task.SUMMARIZER]: "gemini-2.5-pro",
|
||||||
|
},
|
||||||
|
openai: {
|
||||||
|
[Task.PROGRAMMER]: "gpt-4o",
|
||||||
|
[Task.ROUTER]: "gpt-4o-mini",
|
||||||
|
[Task.SUMMARIZER]: "gpt-4o",
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
const modelName = defaultModels[provider][task];
|
||||||
|
if (!modelName) {
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
return { provider, modelName };
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Circuit breaker methods
|
||||||
|
*/
|
||||||
|
public isCircuitClosed(modelKey: string): boolean {
|
||||||
|
const state = this.getCircuitState(modelKey);
|
||||||
|
|
||||||
|
if (state.state === CircuitState.CLOSED) {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (state.state === CircuitState.OPEN && state.openedAt) {
|
||||||
|
const timeElapsed = Date.now() - state.openedAt;
|
||||||
|
if (timeElapsed >= this.config.circuitBreakerTimeoutMs) {
|
||||||
|
state.state = CircuitState.CLOSED;
|
||||||
|
state.failureCount = 0;
|
||||||
|
delete state.openedAt;
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
`${modelKey}: Circuit breaker automatically recovered: OPEN → CLOSED`,
|
||||||
|
{
|
||||||
|
timeElapsed: (timeElapsed / 1000).toFixed(1) + "s",
|
||||||
|
},
|
||||||
|
);
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
private getCircuitState(modelKey: string): CircuitBreakerState {
|
||||||
|
if (!this.circuitBreakers.has(modelKey)) {
|
||||||
|
this.circuitBreakers.set(modelKey, {
|
||||||
|
state: CircuitState.CLOSED,
|
||||||
|
failureCount: 0,
|
||||||
|
lastFailureTime: 0,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
return this.circuitBreakers.get(modelKey)!;
|
||||||
|
}
|
||||||
|
|
||||||
|
public recordSuccess(modelKey: string): void {
|
||||||
|
const circuitState = this.getCircuitState(modelKey);
|
||||||
|
|
||||||
|
circuitState.state = CircuitState.CLOSED;
|
||||||
|
circuitState.failureCount = 0;
|
||||||
|
delete circuitState.openedAt;
|
||||||
|
|
||||||
|
logger.debug(`${modelKey}: Circuit breaker reset after successful request`);
|
||||||
|
}
|
||||||
|
|
||||||
|
public recordFailure(modelKey: string): void {
|
||||||
|
const circuitState = this.getCircuitState(modelKey);
|
||||||
|
const now = Date.now();
|
||||||
|
|
||||||
|
circuitState.lastFailureTime = now;
|
||||||
|
circuitState.failureCount++;
|
||||||
|
|
||||||
|
if (
|
||||||
|
circuitState.failureCount >= this.config.circuitBreakerFailureThreshold
|
||||||
|
) {
|
||||||
|
circuitState.state = CircuitState.OPEN;
|
||||||
|
circuitState.openedAt = now;
|
||||||
|
|
||||||
|
logger.warn(
|
||||||
|
`${modelKey}: Circuit breaker opened after ${circuitState.failureCount} failures`,
|
||||||
|
{
|
||||||
|
timeoutMs: this.config.circuitBreakerTimeoutMs,
|
||||||
|
willRetryAt: new Date(
|
||||||
|
now + this.config.circuitBreakerTimeoutMs,
|
||||||
|
).toISOString(),
|
||||||
|
},
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Monitoring and observability methods
|
||||||
|
*/
|
||||||
|
public getCircuitBreakerStatus(): Map<string, CircuitBreakerState> {
|
||||||
|
return new Map(this.circuitBreakers);
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Cleanup on shutdown
|
||||||
|
*/
|
||||||
|
public shutdown(): void {
|
||||||
|
this.circuitBreakers.clear();
|
||||||
|
logger.info("Shutdown complete");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let globalModelManager: ModelManager | null = null;
|
||||||
|
|
||||||
|
export function getModelManager(
|
||||||
|
config?: Partial<ModelManagerConfig>,
|
||||||
|
): ModelManager {
|
||||||
|
if (!globalModelManager) {
|
||||||
|
globalModelManager = new ModelManager(config);
|
||||||
|
}
|
||||||
|
return globalModelManager;
|
||||||
|
}
|
||||||
|
|
||||||
|
export function resetModelManager(): void {
|
||||||
|
if (globalModelManager) {
|
||||||
|
globalModelManager.shutdown();
|
||||||
|
globalModelManager = null;
|
||||||
|
}
|
||||||
|
}
|
||||||
208
apps/open-swe/src/utils/runtime-fallback.ts
Normal file
208
apps/open-swe/src/utils/runtime-fallback.ts
Normal file
|
|
@ -0,0 +1,208 @@
|
||||||
|
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
||||||
|
import { Task } from "./load-model.js";
|
||||||
|
import { ModelManager } from "./model-manager.js";
|
||||||
|
import { createLogger, LogLevel } from "./logger.js";
|
||||||
|
import { Runnable, RunnableConfig } from "@langchain/core/runnables";
|
||||||
|
import { StructuredToolInterface } from "@langchain/core/tools";
|
||||||
|
import {
|
||||||
|
ConfigurableChatModelCallOptions,
|
||||||
|
ConfigurableModel,
|
||||||
|
} from "langchain/chat_models/universal";
|
||||||
|
import { AIMessageChunk, BaseMessage } 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";
|
||||||
|
import { getMessageContentString } from "@open-swe/shared/messages";
|
||||||
|
|
||||||
|
const logger = createLogger(LogLevel.DEBUG, "FallbackRunnable");
|
||||||
|
|
||||||
|
interface ExtractedTools {
|
||||||
|
tools: StructuredToolInterface[];
|
||||||
|
kwargs: Record<string, any>;
|
||||||
|
}
|
||||||
|
|
||||||
|
export class FallbackRunnable<
|
||||||
|
RunInput extends BaseLanguageModelInput = BaseLanguageModelInput,
|
||||||
|
CallOptions extends
|
||||||
|
ConfigurableChatModelCallOptions = ConfigurableChatModelCallOptions,
|
||||||
|
> extends ConfigurableModel<RunInput, CallOptions> {
|
||||||
|
private primaryRunnable: any;
|
||||||
|
private config: GraphConfig;
|
||||||
|
private task: Task;
|
||||||
|
private modelManager: ModelManager;
|
||||||
|
|
||||||
|
constructor(
|
||||||
|
primaryRunnable: any,
|
||||||
|
config: GraphConfig,
|
||||||
|
task: Task,
|
||||||
|
modelManager: ModelManager,
|
||||||
|
) {
|
||||||
|
super({
|
||||||
|
configurableFields: "any",
|
||||||
|
configPrefix: "fallback",
|
||||||
|
queuedMethodOperations: {},
|
||||||
|
disableStreaming: false,
|
||||||
|
});
|
||||||
|
this.primaryRunnable = primaryRunnable;
|
||||||
|
this.config = config;
|
||||||
|
this.task = task;
|
||||||
|
this.modelManager = modelManager;
|
||||||
|
}
|
||||||
|
|
||||||
|
async _generate(
|
||||||
|
messages: BaseMessage[],
|
||||||
|
options?: Record<string, any>,
|
||||||
|
): Promise<ChatResult> {
|
||||||
|
const result = await this.invoke(messages, options);
|
||||||
|
const generation: ChatGeneration = {
|
||||||
|
message: result,
|
||||||
|
text: result?.content ? getMessageContentString(result.content) : "",
|
||||||
|
};
|
||||||
|
return {
|
||||||
|
generations: [generation],
|
||||||
|
llmOutput: {},
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
async invoke(
|
||||||
|
input: BaseLanguageModelInput,
|
||||||
|
options?: Record<string, any>,
|
||||||
|
): Promise<AIMessageChunk> {
|
||||||
|
const modelConfigs = this.modelManager.getModelConfigs(
|
||||||
|
this.config,
|
||||||
|
this.task,
|
||||||
|
this.getPrimaryModel(),
|
||||||
|
);
|
||||||
|
|
||||||
|
let lastError: Error | undefined;
|
||||||
|
|
||||||
|
for (let i = 0; i < modelConfigs.length; i++) {
|
||||||
|
const modelConfig = modelConfigs[i];
|
||||||
|
const modelKey = `${modelConfig.provider}:${modelConfig.modelName}`;
|
||||||
|
|
||||||
|
if (!this.modelManager.isCircuitClosed(modelKey)) {
|
||||||
|
logger.warn(`Circuit breaker open for ${modelKey}, skipping`);
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
const model = await this.modelManager.initializeModel(modelConfig);
|
||||||
|
let runnableToUse: Runnable<BaseLanguageModelInput, AIMessageChunk> =
|
||||||
|
model;
|
||||||
|
|
||||||
|
const tools = this.extractBoundTools();
|
||||||
|
if (tools && "bindTools" in runnableToUse && runnableToUse.bindTools) {
|
||||||
|
runnableToUse = (runnableToUse as ConfigurableModel).bindTools(
|
||||||
|
tools.tools,
|
||||||
|
tools.kwargs,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
const config = this.extractConfig();
|
||||||
|
if (config) {
|
||||||
|
runnableToUse = runnableToUse.withConfig(config);
|
||||||
|
}
|
||||||
|
|
||||||
|
const result = await runnableToUse.invoke(input, options);
|
||||||
|
this.modelManager.recordSuccess(modelKey);
|
||||||
|
return result;
|
||||||
|
} catch (error) {
|
||||||
|
logger.warn(
|
||||||
|
`${modelKey} failed: ${error instanceof Error ? error.message : String(error)}`,
|
||||||
|
);
|
||||||
|
lastError = error instanceof Error ? error : new Error(String(error));
|
||||||
|
this.modelManager.recordFailure(modelKey);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
throw new Error(
|
||||||
|
`All fallback models exhausted for task ${this.task}. Last error: ${lastError?.message}`,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
bindTools(
|
||||||
|
tools: BindToolsInput[],
|
||||||
|
kwargs?: Record<string, any>,
|
||||||
|
): ConfigurableModel<RunInput, CallOptions> {
|
||||||
|
const boundPrimary =
|
||||||
|
this.primaryRunnable.bindTools?.(tools, kwargs) ?? this.primaryRunnable;
|
||||||
|
return new FallbackRunnable(
|
||||||
|
boundPrimary,
|
||||||
|
this.config,
|
||||||
|
this.task,
|
||||||
|
this.modelManager,
|
||||||
|
) as unknown as ConfigurableModel<RunInput, CallOptions>;
|
||||||
|
}
|
||||||
|
|
||||||
|
// @ts-expect-error - types are hard man :/
|
||||||
|
withConfig(
|
||||||
|
config?: RunnableConfig,
|
||||||
|
): ConfigurableModel<RunInput, CallOptions> {
|
||||||
|
const configuredPrimary =
|
||||||
|
this.primaryRunnable.withConfig?.(config) ?? this.primaryRunnable;
|
||||||
|
return new FallbackRunnable(
|
||||||
|
configuredPrimary,
|
||||||
|
this.config,
|
||||||
|
this.task,
|
||||||
|
this.modelManager,
|
||||||
|
) as unknown as ConfigurableModel<RunInput, CallOptions>;
|
||||||
|
}
|
||||||
|
|
||||||
|
private getPrimaryModel(): ConfigurableModel {
|
||||||
|
let current = this.primaryRunnable;
|
||||||
|
|
||||||
|
// Unwrap any LangChain bindings to get to the actual model
|
||||||
|
while (current?.bound) {
|
||||||
|
current = current.bound;
|
||||||
|
}
|
||||||
|
|
||||||
|
// The unwrapped object should be a chat model with _llmType
|
||||||
|
if (current && typeof current._llmType !== "undefined") {
|
||||||
|
return current;
|
||||||
|
}
|
||||||
|
|
||||||
|
throw new Error(
|
||||||
|
"Could not extract primary model from runnable - no _llmType found",
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
private extractBoundTools(): ExtractedTools | null {
|
||||||
|
let current: any = this.primaryRunnable;
|
||||||
|
|
||||||
|
while (current) {
|
||||||
|
if (current._queuedMethodOperations?.bindTools) {
|
||||||
|
const bindToolsOp = current._queuedMethodOperations.bindTools;
|
||||||
|
|
||||||
|
if (Array.isArray(bindToolsOp) && bindToolsOp.length > 0) {
|
||||||
|
const tools = bindToolsOp[0] as StructuredToolInterface[];
|
||||||
|
const toolOptions = bindToolsOp[1] || {};
|
||||||
|
|
||||||
|
return {
|
||||||
|
tools: tools,
|
||||||
|
kwargs: {
|
||||||
|
tool_choice: (toolOptions as Record<string, any>).tool_choice,
|
||||||
|
parallel_tool_calls: (toolOptions as Record<string, any>)
|
||||||
|
.parallel_tool_calls,
|
||||||
|
},
|
||||||
|
};
|
||||||
|
}
|
||||||
|
}
|
||||||
|
current = current.bound;
|
||||||
|
}
|
||||||
|
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
|
||||||
|
private extractConfig(): Partial<RunnableConfig> | null {
|
||||||
|
let current: any = this.primaryRunnable;
|
||||||
|
|
||||||
|
while (current) {
|
||||||
|
if (current.config) {
|
||||||
|
return current.config;
|
||||||
|
}
|
||||||
|
current = current.bound;
|
||||||
|
}
|
||||||
|
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue