diff --git a/crates/switchyard-runner/src/algorithm.rs b/crates/switchyard-runner/src/algorithm.rs index 5cc838209..a7e40e751 100644 --- a/crates/switchyard-runner/src/algorithm.rs +++ b/crates/switchyard-runner/src/algorithm.rs @@ -12,12 +12,14 @@ use std::sync::Arc; use libsy::{ AdvisorGate, AdvisorGateConfig, Algorithm, CapabilityJudgeConfig, ClassifierContractConfig, ClassifierResponseFormat, ClassifyTrigger, CompositeRouter, CompositeRouterConfig, - CustomClassifierConfig, CustomClassifierPolicy, EscalationJudgeConfig, GateTrigger, - HandoffNoteConfig, LlmCapabilityConfig, LlmClassifierConfig, LlmFallback, LlmTaskClassifier, - Noop, Passthrough, PickerMode, PlanExecute, PlanExecuteConfig, Random, StageRouter, - StageRouterConfig, SubagentRouter, SubagentRouterConfig, TaskClassifierConfig, ToolSemantics, + CustomClassifierConfig, CustomClassifierPolicy, DecisionJudgeConfig, EscalationJudgeConfig, + GateTrigger, HandoffNoteConfig, LlmCapabilityConfig, LlmClassifierConfig, LlmFallback, + LlmTaskClassifier, Noop, Passthrough, PickerMode, PlanExecute, PlanExecuteConfig, Random, + StageRouter, StageRouterConfig, SubagentRouter, SubagentRouterConfig, TaskClassifierConfig, + ToolSemantics, }; use serde::Deserialize; +use serde_json::Value; use switchyard_protocol::{Category, ModelId}; /// Error returned when an algorithm description cannot be constructed. @@ -102,14 +104,16 @@ struct CapabilityClassifierRouteConfig { fail_open: bool, strong_target: String, weak_target: String, - base_threshold: f64, - threshold_step: f64, + judge: CapabilityJudgeRouteConfig, classify_trigger: ClassifyTrigger, message_hash_fallback: bool, recent_turn_window: Option, - prompt: Option, - response_format_type: ClassifierResponseFormat, - max_output_tokens: u64, +} + +#[derive(Clone, Debug)] +enum CapabilityJudgeRouteConfig { + Llm(LlmCapabilityConfig), + Decision(DecisionJudgeRouteConfig), } #[derive(Clone, Debug)] @@ -219,6 +223,20 @@ impl CategoryModelConfig { } } +/// Relative-advantage judgment using a decision model instead of an LLM prompt. +#[derive(Clone, Debug, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct DecisionJudgeRouteConfig { + /// Route capable only when its advantage score is strictly above this cutoff. + pub cutoff: f64, + /// Replaces the default relative-advantage instructions. + pub instructions: Option, + /// Anonymous evidence labels mapped to configured LLM target names. + pub candidates: BTreeMap, + /// Candidate descriptions and reference outcomes, passed unchanged to the judge. + pub evidence: Value, +} + /// Settings for an `llm_classifier` route. Which fields are required depends on /// the [`ClassifierMode`]; using a field from the wrong mode is an error. #[derive(Clone, Debug, Default, Deserialize)] @@ -260,6 +278,8 @@ pub struct LlmClassifierRouteConfig { /// Escalation mode: how many escalate verdicts latch the session, and how /// much of the transcript the judge sees. pub escalation: Option, + /// Capability mode: use a decision target as the judge, with relative-advantage scoring. + pub decision: Option, /// Custom mode: runtime model groups. pub models: Option, /// Custom mode: category used when the judge fails or its verdict cannot be routed. @@ -542,6 +562,16 @@ impl StageClassifierConfig { } impl AlgorithmSpec { + pub(crate) fn decision_judge(&self) -> Option<(&str, &DecisionJudgeRouteConfig)> { + match self { + Self::LlmClassifier { config, .. } => config + .decision + .as_ref() + .map(|decision| (config.classifier_target.as_str(), decision)), + _ => None, + } + } + /// Completion targets in algorithm order; judge-only targets are excluded. pub fn routing_target_names(&self) -> Vec<&str> { match self { @@ -893,6 +923,7 @@ impl LlmClassifierRouteConfig { response_format_type, max_output_tokens, escalation, + decision, models, default_target, response_schema, @@ -905,6 +936,12 @@ impl LlmClassifierRouteConfig { (None, false) => ClassifierMode::Capability, }; + if decision.is_some() && !matches!(selected_mode, ClassifierMode::Capability) { + return Err(AlgorithmConfigError::new(format!( + "llm_classifier route {route_name}: decision requires capability mode" + ))); + } + if !matches!(selected_mode, ClassifierMode::Capability) && *fail_open == Some(true) { return Err(AlgorithmConfigError::new(format!( "llm_classifier route {route_name}: fail_open = true requires capability mode" @@ -928,6 +965,31 @@ impl LlmClassifierRouteConfig { response_schema, policy, )?; + let judge = if let Some(decision) = decision { + if base_threshold.is_some() + || threshold_step.is_some() + || prompt.is_some() + || *response_format_type != ClassifierResponseFormat::JsonSchema + || *max_output_tokens != default_classifier_max_output_tokens() + { + return Err(AlgorithmConfigError::new(format!( + "llm_classifier route {route_name}: decision cannot use LLM judge settings; use decision.cutoff and decision.instructions" + ))); + } + CapabilityJudgeRouteConfig::Decision(decision.clone()) + } else { + CapabilityJudgeRouteConfig::Llm(LlmCapabilityConfig { + base_threshold: required_classifier_field( + route_name, + "base_threshold", + base_threshold, + )?, + threshold_step: threshold_step.unwrap_or_default(), + contract: classifier_contract(prompt.as_deref()) + .with_response_format_type(*response_format_type), + max_output_tokens: *max_output_tokens, + }) + }; Ok(LlmClassifierModeConfig::Capability( CapabilityClassifierRouteConfig { classifier_target: classifier_target.clone(), @@ -942,18 +1004,10 @@ impl LlmClassifierRouteConfig { "weak_target", weak_target, )?, - base_threshold: required_classifier_field( - route_name, - "base_threshold", - base_threshold, - )?, - threshold_step: threshold_step.unwrap_or_default(), + judge, classify_trigger: *classify_trigger, message_hash_fallback: *message_hash_fallback, recent_turn_window: *recent_turn_window, - prompt: prompt.clone(), - response_format_type: *response_format_type, - max_output_tokens: *max_output_tokens, }, )) } @@ -1243,14 +1297,34 @@ fn build_algorithm( let mode = classifier_config.validated_classifier_mode(route_name)?; let algorithm = match mode { LlmClassifierModeConfig::Capability(config) => { + let judge = match config.judge { + CapabilityJudgeRouteConfig::Llm(judge) => CapabilityJudgeConfig::Llm(judge), + CapabilityJudgeRouteConfig::Decision(judge) => { + for target in [&config.strong_target, &config.weak_target] { + if !judge.candidates.values().any(|name| name == target) { + return Err(AlgorithmConfigError::new(format!( + "llm_classifier route {route_name}: decision.candidates must include target {target}" + ))); + } + } + let candidates = judge + .candidates + .into_iter() + .map(|(label, name)| { + resolve_target_model_id(route_name, &name, targets) + .map(|model| (label, model)) + }) + .collect::>()?; + CapabilityJudgeConfig::Decision(DecisionJudgeConfig { + cutoff: judge.cutoff, + instructions: judge.instructions, + candidates, + evidence: judge.evidence, + }) + } + }; let classifier_config = TaskClassifierConfig { - judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { - base_threshold: config.base_threshold, - threshold_step: config.threshold_step, - contract: classifier_contract(config.prompt.as_deref()) - .with_response_format_type(config.response_format_type), - max_output_tokens: config.max_output_tokens, - }), + judge, fail_open: config.fail_open, classify_trigger: config.classify_trigger, message_hash_fallback: config.message_hash_fallback, diff --git a/crates/switchyard-runner/src/config.rs b/crates/switchyard-runner/src/config.rs index 51105bd42..2c85a97ca 100644 --- a/crates/switchyard-runner/src/config.rs +++ b/crates/switchyard-runner/src/config.rs @@ -15,9 +15,9 @@ use serde::{Deserialize, Deserializer}; use serde_json::Value; use switchyard_llm_client::{ AuxiliaryOperation, Backend, ClientRouter, DEFAULT_MAX_RETRIES, HttpBackendConfig, ModelConfig, - TranslatingLlmClient, + SystemOneClient, TranslatingLlmClient, }; -use switchyard_protocol::{Category, ModelId, RoutedLlmClient, WireFormat}; +use switchyard_protocol::{Category, ModelId, RoutedDecisionClient, RoutedLlmClient, WireFormat}; use crate::{ AlgorithmSpec, AuxiliaryTarget, CallerAuthKind, DecisionTarget, ModelCapabilities, Route, @@ -60,6 +60,10 @@ pub(crate) struct DeploymentConfig { #[serde(default)] llm_clients: BTreeMap, targets: BTreeMap, + #[serde(default)] + decision_clients: BTreeMap, + #[serde(default)] + decision_targets: BTreeMap, routes: BTreeMap, } @@ -211,16 +215,49 @@ impl DeploymentConfig { let mut provider_api_keys = Vec::new(); let clients = self.build_clients(&mut provider_api_keys)?; + let decision_clients = self.build_decision_clients(&mut provider_api_keys)?; + for (name, target) in &self.decision_targets { + validate_value("decision target name", name)?; + validate_value(&format!("decision target {name} id"), &target.id)?; + if self.targets.contains_key(name) { + return Err(RunnerError::configuration(format!( + "target {name} is defined as both an LLM and decision target" + ))); + } + if !decision_clients.contains_key(&target.decision_client) { + return Err(RunnerError::configuration(format!( + "decision target {name} references unknown decision client {}", + target.decision_client + ))); + } + } let targets = self.build_targets(); let fallback_base_url = self.fallback_base_url()?; let mut routes = Vec::with_capacity(self.routes.len()); for (route_name, config) in &self.routes { - for target_name in config.callable_target_names() { - self.targets.get(target_name).ok_or_else(|| { - RunnerError::configuration(format!( - "route references unknown target {target_name}" - )) - })?; + let decision_judge = config.algorithm.decision_judge(); + for name in config.callable_target_names() { + let (exists, kind) = if decision_judge.is_some_and(|(judge, _)| judge == name) { + (self.decision_targets.contains_key(name), "decision") + } else { + (self.targets.contains_key(name), "LLM") + }; + if !exists { + return Err(RunnerError::configuration(format!( + "route references unknown target {name}; route {route_name} requires target kind {kind}" + ))); + } + } + for name in config.routing_target_names().into_iter().chain( + decision_judge + .into_iter() + .flat_map(|(_, judge)| judge.candidates.values().map(String::as_str)), + ) { + if !self.targets.contains_key(name) { + return Err(RunnerError::configuration(format!( + "route {route_name} completion and candidate target {name} must be an LLM target" + ))); + } } let capabilities = config.capabilities(); if capabilities.context_window == Some(0) { @@ -233,7 +270,7 @@ impl DeploymentConfig { .build(route_name, &targets) .map_err(|error| RunnerError::configuration_source(error.to_string(), error))?; let (route_clients, caller_auth) = - self.build_route_clients(route_name, config, &clients)?; + self.build_route_clients(route_name, config, &clients, &decision_clients)?; let anthropic_auxiliary_target = self.build_anthropic_auxiliary_target(config, &clients); let responses_auxiliary_target = @@ -343,10 +380,49 @@ impl DeploymentConfig { Ok(clients) } + fn build_decision_clients( + &self, + provider_api_keys: &mut Vec, + ) -> RunnerResult>> { + self.decision_clients + .iter() + .map(|(name, config)| { + validate_value("decision client name", name)?; + let DecisionClientConfig::SystemOne { + endpoint, + api_key_env, + timeout_ms, + } = config; + if *timeout_ms == 0 { + return Err(RunnerError::configuration(format!( + "decision client {name} timeout_ms must be at least 1" + ))); + } + let api_key = read_api_key(&format!("decision client {name}"), api_key_env)?; + let client = SystemOneClient::new( + endpoint.0.clone(), + api_key.clone(), + Duration::from_millis(*timeout_ms), + ) + .map_err(|error| RunnerError::configuration(error.to_string()))?; + provider_api_keys.push(api_key); + Ok(( + name.clone(), + Arc::new(client) as Arc, + )) + }) + .collect() + } + fn build_targets(&self) -> BTreeMap { self.targets .iter() .map(|(name, config)| (name.clone(), config.id.clone())) + .chain( + self.decision_targets + .iter() + .map(|(name, config)| (name.clone(), config.id.clone())), + ) .collect() } @@ -358,6 +434,7 @@ impl DeploymentConfig { route_name: &str, route: &RouteConfig, clients: &BTreeMap>, + decision_clients: &BTreeMap>, ) -> RunnerResult<(ClientRouter, Option)> { let TargetPromptPolicy { prompts, @@ -368,7 +445,20 @@ impl DeploymentConfig { let mut caller_auth = None; let mut has_mixed_families = false; let mut forwarding_origins = BTreeSet::new(); + let mut decisions_by_model = HashMap::new(); for name in route.callable_target_names() { + if route + .algorithm + .decision_judge() + .is_some_and(|(judge, _)| judge == name) + { + let target = &self.decision_targets[name]; + decisions_by_model.insert( + target.id.clone(), + decision_clients[&target.decision_client].clone(), + ); + continue; + } let target = self.targets.get(name).ok_or_else(|| { RunnerError::configuration(format!("route references unknown target {name}")) })?; @@ -417,8 +507,9 @@ impl DeploymentConfig { .into_iter() .map(|name| self.targets[name].id.clone()) .collect::>(); - let router = ClientRouter::new_with_completion_targets( + let router = ClientRouter::new_with_decision_clients( by_model, + decisions_by_model, prompts, routing_answer_target, &completion_targets, @@ -599,6 +690,23 @@ struct LlmClientConfig { timeout_ms: Option, } +#[derive(Debug, Deserialize)] +#[serde(tag = "format", rename_all = "snake_case", deny_unknown_fields)] +enum DecisionClientConfig { + SystemOne { + endpoint: HttpBaseUrl, + api_key_env: String, + timeout_ms: u64, + }, +} + +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +struct DecisionTargetConfig { + id: ModelId, + decision_client: String, +} + #[derive(Debug, Deserialize)] #[serde(deny_unknown_fields)] struct TargetConfig { @@ -691,24 +799,7 @@ fn build_backend( let api_key = config .api_key_env .as_deref() - .map(|variable| { - if variable.trim().is_empty() { - return Err(RunnerError::configuration(format!( - "llm client {client_name} api_key_env must not be empty" - ))); - } - let api_key = std::env::var(variable).map_err(|error| { - RunnerError::configuration(format!( - "llm client {client_name} could not read api_key_env {variable}: {error}" - )) - })?; - if api_key.trim().is_empty() { - return Err(RunnerError::configuration(format!( - "llm client {client_name} api_key_env {variable} is empty" - ))); - } - Ok(api_key) - }) + .map(|variable| read_api_key(&format!("llm client {client_name}"), variable)) .transpose()?; let http = HttpBackendConfig { base_url: config.base_url.as_str().to_string(), @@ -730,6 +821,25 @@ fn build_backend( Ok(backend) } +fn read_api_key(client: &str, variable: &str) -> RunnerResult { + if variable.trim().is_empty() { + return Err(RunnerError::configuration(format!( + "{client} api_key_env must not be empty" + ))); + } + let api_key = std::env::var(variable).map_err(|error| { + RunnerError::configuration(format!( + "{client} could not read api_key_env {variable}: {error}" + )) + })?; + if api_key.trim().is_empty() { + return Err(RunnerError::configuration(format!( + "{client} api_key_env {variable} is empty" + ))); + } + Ok(api_key) +} + // A function so that serde default can use it. const fn default_max_retries() -> u32 { DEFAULT_MAX_RETRIES diff --git a/crates/switchyard-runner/src/lib.rs b/crates/switchyard-runner/src/lib.rs index a488b31a7..115e8ce83 100644 --- a/crates/switchyard-runner/src/lib.rs +++ b/crates/switchyard-runner/src/lib.rs @@ -12,7 +12,8 @@ mod runner; pub use algorithm::{ AdvisorTriggerConfig, AlgorithmConfigError, AlgorithmSpec, CategoryModelConfig, ClassifierMode, - ClassifierPolicyConfig, LlmClassifierRouteConfig, StageClassifierConfig, SubagentRouteConfig, + ClassifierPolicyConfig, DecisionJudgeRouteConfig, LlmClassifierRouteConfig, + StageClassifierConfig, SubagentRouteConfig, }; pub use failure::{RouteErrorKind, RouteErrorPhase, RouteErrorSummary, stream_error_summary}; // Re-exported because `Route::new` takes it, so a host wiring routes does not need a libsy dep. diff --git a/crates/switchyard-server/CONFIGURATION.md b/crates/switchyard-server/CONFIGURATION.md index f7458bc95..275718ece 100644 --- a/crates/switchyard-server/CONFIGURATION.md +++ b/crates/switchyard-server/CONFIGURATION.md @@ -32,6 +32,47 @@ outside normal assistant `content`. To support another wire format, add its `ClientFormat` variant and explicit construction match in `../switchyard-runner/src/config.rs`. Add a client type only when a second implementation exists. +## Use a decision model as the capability judge + +With `quality` and `economy` already defined as LLM targets: + +```toml +[decision_clients.jev] +format = "system_one" +endpoint = "https://api.typesafe.ai/v1/systemone" +api_key_env = "TYPESAFE_API_KEY" +timeout_ms = 5000 + +[decision_targets.judge] +id = "jev-latest" +decision_client = "jev" + +[routes.capability] +id = "switchyard/capability" +type = "llm_classifier" +mode = "capability" +classifier_target = "judge" +strong_target = "quality" +weak_target = "economy" + +[routes.capability.decision] +cutoff = 0.4 +candidates = { a = "quality", b = "economy" } +evidence = { candidates = { a = "Higher-quality model", b = "Lower-cost model" } } +``` + +- `endpoint` is the full System One URL. Set the named key environment variable + before startup; caller credentials are not forwarded to the decision provider. +- Candidate labels must cover both tiers. Extra candidates add context, not routes. + Supply evidence from your evaluations and choose a cutoff using those results. +- `decision.instructions` can replace the default relative-advantage instructions. + It must preserve the meaning of `advantage` and `no_advantage`. Evidence is passed unchanged. +- The capable tier is selected only when its advantage score exceeds `cutoff`. + Client failures and deadlines use the capable tier by default; set `fail_open = false` + on the route to propagate those errors. +- Decision targets cannot serve final answers. Both normal requests and `/v1/decision` + use this configuration. Omit `decision` to keep the existing LLM judge settings. + ## Add an algorithm 1. Implement and export the algorithm from `libsy`.