mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 06:53:14 +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
|
||||
|
||||
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 { z } from "zod";
|
||||
import { tool } from "@langchain/core/tools";
|
||||
import { ConfigurableModel } from "langchain/chat_models/universal";
|
||||
import { traceable } from "langsmith/traceable";
|
||||
import {
|
||||
PlannerGraphState,
|
||||
|
|
@ -19,6 +18,7 @@ import {
|
|||
supportsParallelToolCallsParam,
|
||||
Task,
|
||||
} 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.
|
||||
|
||||
|
|
@ -97,7 +97,7 @@ const formatSysPromptRewritePlan = (
|
|||
|
||||
async function identifyTasksToModifyFunc(
|
||||
state: PlannerGraphState,
|
||||
model: ConfigurableModel,
|
||||
model: FallbackRunnable,
|
||||
supportsParallelToolCallsParam: boolean,
|
||||
): Promise<PlanItem[]> {
|
||||
if (!state.planChangeRequest) {
|
||||
|
|
@ -188,7 +188,7 @@ const identifyTasksToModify = traceable(identifyTasksToModifyFunc, {
|
|||
async function updatePlanTasksFunc(
|
||||
state: PlannerGraphState,
|
||||
tasksToModify: PlanItem[],
|
||||
model: ConfigurableModel,
|
||||
model: FallbackRunnable,
|
||||
supportsParallelToolCallsParam: boolean,
|
||||
): Promise<string[]> {
|
||||
if (!state.planChangeRequest) {
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
import { initChatModel } from "langchain/chat_models/universal";
|
||||
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
||||
import { isAllowedUser } from "./github/allowed-users.js";
|
||||
import { decryptSecret } from "@open-swe/shared/crypto";
|
||||
import { getModelManager } from "./model-manager.js";
|
||||
import { FallbackRunnable } from "./runtime-fallback.js";
|
||||
|
||||
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) {
|
||||
const modelStr =
|
||||
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 modelManager = getModelManager();
|
||||
|
||||
const [modelProvider, ...modelNameParts] = modelStr.split(":");
|
||||
|
||||
let thinkingModel = false;
|
||||
if (modelNameParts[0] === "extended-thinking") {
|
||||
// Using a thinking model. Remove it from the model name.
|
||||
modelNameParts.shift();
|
||||
thinkingModel = true;
|
||||
const model = await modelManager.loadModel(config, task);
|
||||
if (!model) {
|
||||
throw new Error(`Model loading returned undefined for task: ${task}`);
|
||||
}
|
||||
|
||||
const modelName = modelNameParts.join(":");
|
||||
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 fallbackModel = new FallbackRunnable(model, config, task, modelManager);
|
||||
return fallbackModel;
|
||||
}
|
||||
|
||||
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