import { readFileSync } from 'fs'; import { join, dirname } from 'path'; import { fileURLToPath } from 'url'; import yaml from 'js-yaml'; import { logger } from '../observability/logger.js'; const __dirname = dirname(fileURLToPath(import.meta.url)); const CONFIG_DIR = join(__dirname, '../config'); export interface RoutingRule { model: string; tier: 'fast' | 'medium' | 'large'; prompt_template: string; temperature: number; max_tokens: number; output_format: 'text' | 'json'; requires_fact_check: boolean; validators: string[]; callers: string[]; } export interface ModelConfig { tier: 'fast' | 'medium' | 'large'; context_length: number; strengths: string[]; max_tokens_default: number; } export interface ModelsYaml { ollama_base_url: string; tiers: Record; models: Record; fallback_chains: Record; tier_fallback: Record; } export interface RoutingRulesYaml { routing_rules: Record; validators: Record>; } export interface RouterDecision { model: string; fallback_chain: string[]; tier: 'fast' | 'medium' | 'large'; prompt_template: string; temperature: number; max_tokens: number; output_format: 'text' | 'json'; requires_fact_check: boolean; validators: string[]; ollama_base_url: string; timeout_ms: number; } let modelsConfig: ModelsYaml | null = null; let routingConfig: RoutingRulesYaml | null = null; function loadModels(): ModelsYaml { if (modelsConfig) return modelsConfig; try { const raw = readFileSync(join(CONFIG_DIR, 'models.yaml'), 'utf-8'); modelsConfig = yaml.load(raw) as ModelsYaml; return modelsConfig; } catch (err) { logger.error({ err }, 'Failed to load models.yaml'); throw new Error('Could not load models configuration'); } } function loadRoutingRules(): RoutingRulesYaml { if (routingConfig) return routingConfig; try { const raw = readFileSync(join(CONFIG_DIR, 'routing-rules.yaml'), 'utf-8'); routingConfig = yaml.load(raw) as RoutingRulesYaml; return routingConfig; } catch (err) { logger.error({ err }, 'Failed to load routing-rules.yaml'); throw new Error('Could not load routing rules configuration'); } } export function reloadConfigs(): void { modelsConfig = null; routingConfig = null; loadModels(); loadRoutingRules(); } function isCallerAllowed(rule: RoutingRule, caller: string): boolean { return rule.callers.includes('all') || rule.callers.includes(caller); } function buildFallbackChain( primaryModel: string, tier: string, models: ModelsYaml, ): string[] { const chain = models.fallback_chains[tier] ?? []; // Put primary first, then other fallbacks excluding primary return [primaryModel, ...chain.filter((m) => m !== primaryModel)]; } export function route( taskType: string, caller: string, overrides?: { model?: string; temperature?: number; max_tokens?: number; }, ): RouterDecision { const models = loadModels(); const rules = loadRoutingRules(); const rule = rules.routing_rules[taskType]; if (!rule) { // Fall back to generic_qa const fallbackRule = rules.routing_rules['generic_qa']; if (!fallbackRule) { throw new Error(`No routing rule for task_type: ${taskType}`); } logger.warn({ taskType, caller }, 'Unknown task_type, falling back to generic_qa'); return buildDecision('generic_qa', fallbackRule, caller, models, overrides); } if (!isCallerAllowed(rule, caller)) { throw new Error(`Caller "${caller}" is not allowed to use task_type "${taskType}"`); } return buildDecision(taskType, rule, caller, models, overrides); } function buildDecision( _taskType: string, rule: RoutingRule, _caller: string, models: ModelsYaml, overrides?: { model?: string; temperature?: number; max_tokens?: number }, ): RouterDecision { const selectedModel = overrides?.model ?? rule.model; const tier = rule.tier; const tierConfig = models.tiers[tier]; if (!tierConfig) { throw new Error(`Unknown model tier: ${tier}`); } return { model: selectedModel, fallback_chain: buildFallbackChain(selectedModel, tier, models), tier, prompt_template: rule.prompt_template, temperature: overrides?.temperature ?? rule.temperature, max_tokens: overrides?.max_tokens ?? rule.max_tokens, output_format: rule.output_format, requires_fact_check: rule.requires_fact_check, validators: rule.validators, ollama_base_url: models.ollama_base_url, timeout_ms: tierConfig.timeout_ms, }; } export function getModelTier(model: string): 'fast' | 'medium' | 'large' { const models = loadModels(); const config = models.models[model]; return config?.tier ?? 'medium'; } export function getOllamaBaseUrl(): string { const models = loadModels(); return models.ollama_base_url; }