diff --git a/crates/libsy-llm-client/tests/observability.rs b/crates/libsy-llm-client/tests/observability.rs index 820c762c6..8f5071ccd 100644 --- a/crates/libsy-llm-client/tests/observability.rs +++ b/crates/libsy-llm-client/tests/observability.rs @@ -34,9 +34,10 @@ use tracing_subscriber::layer::{Context as LayerContext, SubscriberExt}; use tracing_subscriber::registry::LookupSpan; use switchyard_libsy::{ - Algorithm, ClassifierContractConfig, ClassifyTrigger, DeescalationConfig, Driver, - EscalationJudgeConfig, LibsyError, LlmClassifierConfig, LlmTaskClassifier, PickerMode, - RoutingOutcome, RuntimeModels, StageRouter, StageRouterConfig, Step, TaskClassifierConfig, + Algorithm, CapabilityJudgeConfig, ClassifierContractConfig, ClassifyTrigger, + DeescalationConfig, Driver, EscalationJudgeConfig, LibsyError, LlmCapabilityConfig, + LlmClassifierConfig, LlmTaskClassifier, PickerMode, RoutingOutcome, RuntimeModels, StageRouter, + StageRouterConfig, Step, TaskClassifierConfig, }; use switchyard_llm_client::{ClientRouter, RunObservation, RunObserver}; use switchyard_protocol::{Category, ModelId}; @@ -606,7 +607,10 @@ fn classifier_router() -> switchyard_libsy::Result> { Ok(Arc::new(LlmTaskClassifier::new( LlmClassifierConfig::Capability { config: TaskClassifierConfig { - base_threshold: 0.5, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + ..LlmCapabilityConfig::default() + }), fail_open: false, ..TaskClassifierConfig::default() }, @@ -869,7 +873,10 @@ async fn affinity_keeps_the_algorithm_selection_after_client_fallback() }); let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability { config: TaskClassifierConfig { - base_threshold: 0.5, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + ..LlmCapabilityConfig::default() + }), classify_trigger: ClassifyTrigger::NewSession, ..TaskClassifierConfig::default() }, diff --git a/crates/libsy/src/algorithms/composite.rs b/crates/libsy/src/algorithms/composite.rs index 3afe61486..12ddfb39a 100644 --- a/crates/libsy/src/algorithms/composite.rs +++ b/crates/libsy/src/algorithms/composite.rs @@ -172,6 +172,7 @@ impl Algorithm for CompositeRouter { #[cfg(test)] mod tests { + use crate::{CapabilityJudgeConfig, LlmCapabilityConfig}; use std::collections::HashMap; use std::sync::Arc; @@ -215,7 +216,10 @@ mod tests { fn hash_keyed_router() -> Result> { Ok(Arc::new(CompositeRouter::new(CompositeRouterConfig { judge: TaskClassifierConfig { - base_threshold: 0.5, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + ..LlmCapabilityConfig::default() + }), classify_trigger: ClassifyTrigger::UserTurn, message_hash_fallback: true, ..Default::default() @@ -227,7 +231,10 @@ mod tests { fn router() -> Result> { Ok(Arc::new(CompositeRouter::new(CompositeRouterConfig { judge: TaskClassifierConfig { - base_threshold: 0.5, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + ..LlmCapabilityConfig::default() + }), classify_trigger: ClassifyTrigger::UserTurn, ..Default::default() }, diff --git a/crates/libsy/src/algorithms/llm_class.rs b/crates/libsy/src/algorithms/llm_class.rs index 40d9ad3d9..74ab7e55c 100644 --- a/crates/libsy/src/algorithms/llm_class.rs +++ b/crates/libsy/src/algorithms/llm_class.rs @@ -186,8 +186,9 @@ struct TaskInput { recent_turn_window: Option, } -impl ClassifierInput for TaskInput { - fn build_messages(&self, _state: &State, request: &Request) -> Vec { +impl TaskInput { + /// Selects conversation content without adding judge instructions. + fn messages(&self, request: &Request) -> Vec { // The default preserves the whole-task anchor and latest user update. A // configured window widens that to the surrounding conversation. let mut messages = match self.recent_turn_window { @@ -204,6 +205,13 @@ impl ClassifierInput for TaskInput { .retain(|block| !matches!(block, ContentBlock::Reasoning { .. })); } messages.retain(|message| !message.content.is_empty()); + messages + } +} + +impl ClassifierInput for TaskInput { + fn build_messages(&self, _state: &State, request: &Request) -> Vec { + let mut messages = self.messages(request); // Only the windowed path carries assistant turns and tool traffic for the judge // to be distracted by. The default path is user task messages only — the anchor // and the latest follow-up — so there is nothing there to outrank. @@ -223,7 +231,7 @@ struct TaskClassifierPolicy { } impl TaskClassifierPolicy { - fn new(config: &TaskClassifierConfig) -> Self { + fn new(config: &LlmCapabilityConfig) -> Self { Self { base_threshold: config.base_threshold, threshold_step: config.threshold_step, @@ -297,35 +305,60 @@ fn capability_evidence( })) } +/// Settings shared by capability judges and their surrounding route. #[derive(Clone, Debug)] -/// Settings that control capability classifier prompting and routing. pub struct TaskClassifierConfig { + /// The judge's prediction contract and routing policy. + pub judge: CapabilityJudgeConfig, /// Routes to the capable tier on judge client failures and deadlines. Defaults to true. pub fail_open: bool, - /// Lowest solve probability that routes a supported task to the efficient target. - pub base_threshold: f64, - /// Amount added per capability-boundary step. - /// - /// Supported verdicts use `base_threshold`, uncertain and unmatched verdicts use one - /// step, and unsupported verdicts use two steps. - pub threshold_step: f64, /// How often the classifier re-decides this session's target. pub classify_trigger: ClassifyTrigger, /// Uses the first user message as the SessionKey for sticky routing when session metadata is unavailable. pub message_hash_fallback: bool, - /// Trailing conversation turns the judge sees on top of the client - /// instructions and the opening task. - /// - /// `None` (the default) judges the opening task and latest user follow-up. - /// `Some(n)` widens that to the client instructions, the opening task, and - /// the last `n` turns after it. + /// `None` judges the opening task and latest user follow-up. `Some(n)` also + /// includes the surrounding conversation using the last `n` turns. pub recent_turn_window: Option, +} + +/// Selects how a capability judge obtains and interprets its prediction. +#[derive(Clone, Debug)] +pub enum CapabilityJudgeConfig { + /// A structured LLM verdict with a solve probability and capability boundary. + Llm(LlmCapabilityConfig), +} + +impl Default for CapabilityJudgeConfig { + fn default() -> Self { + Self::Llm(LlmCapabilityConfig::default()) + } +} + +/// Prompt, output, and solve-probability settings for an LLM capability judge. +#[derive(Clone, Debug)] +pub struct LlmCapabilityConfig { + /// Lowest solve probability that routes a supported task to the efficient target. + pub base_threshold: f64, + /// Amount added per capability-boundary step: zero for supported, one for + /// uncertain or unmatched, and two for unsupported. + pub threshold_step: f64, /// Prompt and verdict contract settings for the classifier judge. pub contract: ClassifierContractConfig, /// Maximum completion tokens available to the classifier verdict. pub max_output_tokens: u64, } +impl Default for LlmCapabilityConfig { + fn default() -> Self { + Self { + base_threshold: 0.0, + threshold_step: 0.0, + contract: ClassifierContractConfig::default(), + max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS, + } + } +} + /// Flat serialized shape that maps prompt settings into the runtime contract. #[derive(Deserialize)] #[serde(deny_unknown_fields)] @@ -361,14 +394,16 @@ impl<'de> Deserialize<'de> for TaskClassifierConfig { } contract = contract.with_response_format_type(wire.response_format_type); Ok(Self { - base_threshold: wire.base_threshold, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: wire.base_threshold, + threshold_step: wire.threshold_step, + contract, + max_output_tokens: wire.max_output_tokens, + }), fail_open: wire.fail_open, - threshold_step: wire.threshold_step, classify_trigger: wire.classify_trigger, message_hash_fallback: wire.message_hash_fallback, recent_turn_window: wire.recent_turn_window, - contract, - max_output_tokens: wire.max_output_tokens, }) } } @@ -384,20 +419,16 @@ const fn default_judge_max_output_tokens() -> u64 { impl Default for TaskClassifierConfig { fn default() -> Self { Self { - base_threshold: 0.0, + judge: CapabilityJudgeConfig::default(), fail_open: default_fail_open(), - threshold_step: 0.0, classify_trigger: ClassifyTrigger::default(), message_hash_fallback: false, recent_turn_window: None, - contract: ClassifierContractConfig::default(), - max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS, } } } -impl TaskClassifierConfig { - /// Validates routing thresholds before the classifier is constructed. +impl LlmCapabilityConfig { fn validate(&self) -> Result<()> { if !(0.0..=1.0).contains(&self.base_threshold) { return Err(LibsyError::AlgorithmError { @@ -428,6 +459,15 @@ impl TaskClassifierConfig { message: "max_output_tokens must be at least 1".to_string(), }); } + Ok(()) + } +} + +impl TaskClassifierConfig { + fn validate(&self) -> Result<()> { + match &self.judge { + CapabilityJudgeConfig::Llm(config) => config.validate()?, + } // Only `every_request` is rejected: it retains no target, so a fallback identity has nothing to key on. Both retaining triggers can key the retained target on a message hash when the caller sends no session id. if self.message_hash_fallback && self.classify_trigger == ClassifyTrigger::EveryRequest { return Err(LibsyError::AlgorithmError { @@ -639,7 +679,8 @@ impl LlmTaskClassifier { fn build_capability(config: TaskClassifierConfig) -> Result { config.validate()?; - let contract = Self::load_capability_contract(&config.contract)?; + let CapabilityJudgeConfig::Llm(judge) = &config.judge; + let contract = Self::load_capability_contract(&judge.contract)?; let classify_trigger = config.classify_trigger; let message_hash_fallback = config.message_hash_fallback; let classifier: Arc> = Arc::new( @@ -650,9 +691,9 @@ impl LlmTaskClassifier { }, contract, SerdeDecoder::new(), - JudgeRuntimeConfig::new(config.max_output_tokens)?, + JudgeRuntimeConfig::new(judge.max_output_tokens)?, ), - TaskClassifierPolicy::new(&config), + TaskClassifierPolicy::new(judge), ) .with_error_recovery(config.fail_open) .with_evidence(capability_evidence), @@ -806,15 +847,22 @@ mod tests { type CapabilityJudge = StructuredJudge>; + fn llm_config(base_threshold: f64) -> LlmCapabilityConfig { + LlmCapabilityConfig { + base_threshold, + ..LlmCapabilityConfig::default() + } + } + fn test_config(base_threshold: f64) -> TaskClassifierConfig { TaskClassifierConfig { - base_threshold, + judge: CapabilityJudgeConfig::Llm(llm_config(base_threshold)), ..TaskClassifierConfig::default() } } fn policy() -> TaskClassifierPolicy { - TaskClassifierPolicy::new(&test_config(TEST_THRESHOLD)) + TaskClassifierPolicy::new(&llm_config(TEST_THRESHOLD)) } fn runtime_models() -> HashMap> { @@ -1045,7 +1093,10 @@ mod tests { let recorder = Arc::new(Recorder::default()); let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability { config: TaskClassifierConfig { - max_output_tokens: 512, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + max_output_tokens: 512, + ..llm_config(TEST_THRESHOLD) + }), ..test_config(TEST_THRESHOLD) }, })?); @@ -1067,8 +1118,11 @@ mod tests { let recorder = Arc::new(Recorder::default()); let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability { config: TaskClassifierConfig { - contract: ClassifierContractConfig::default() - .with_prompt("Custom capability rubric."), + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + contract: ClassifierContractConfig::default() + .with_prompt("Custom capability rubric."), + ..llm_config(TEST_THRESHOLD) + }), ..test_config(TEST_THRESHOLD) }, })?); @@ -1229,8 +1283,8 @@ mod tests { #[test] fn the_threshold_moves_the_routing_boundary() -> Result<()> { let borderline = verdict(0.5, "supported", "SUP-1"); - let strict = TaskClassifierPolicy::new(&test_config(0.9)); - let lenient = TaskClassifierPolicy::new(&test_config(0.1)); + let strict = TaskClassifierPolicy::new(&llm_config(0.9)); + let lenient = TaskClassifierPolicy::new(&llm_config(0.1)); assert_eq!(selected(&strict, Some(&borderline))?, "capable"); assert_eq!(selected(&lenient, Some(&borderline))?, "efficient"); Ok(()) @@ -1265,23 +1319,35 @@ mod tests { } for config in [ TaskClassifierConfig { - base_threshold: 0.5, - threshold_step: -0.1, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + threshold_step: -0.1, + ..LlmCapabilityConfig::default() + }), ..TaskClassifierConfig::default() }, TaskClassifierConfig { - base_threshold: 0.8, - threshold_step: 0.11, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.8, + threshold_step: 0.11, + ..LlmCapabilityConfig::default() + }), ..TaskClassifierConfig::default() }, TaskClassifierConfig { - base_threshold: 0.5, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + ..LlmCapabilityConfig::default() + }), message_hash_fallback: true, ..TaskClassifierConfig::default() }, TaskClassifierConfig { - base_threshold: 0.5, - max_output_tokens: 0, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + max_output_tokens: 0, + ..LlmCapabilityConfig::default() + }), ..TaskClassifierConfig::default() }, ] { @@ -1302,7 +1368,10 @@ mod tests { // retained target on a message hash when the caller sends no session id. for trigger in [ClassifyTrigger::NewSession, ClassifyTrigger::UserTurn] { let config = TaskClassifierConfig { - base_threshold: 0.5, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + ..LlmCapabilityConfig::default() + }), classify_trigger: trigger, message_hash_fallback: true, ..TaskClassifierConfig::default() @@ -1314,7 +1383,10 @@ mod tests { )?; } let every_request = TaskClassifierConfig { - base_threshold: 0.5, + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: 0.5, + ..LlmCapabilityConfig::default() + }), classify_trigger: ClassifyTrigger::EveryRequest, message_hash_fallback: true, ..TaskClassifierConfig::default() @@ -1357,9 +1429,9 @@ mod tests { #[test] fn capability_boundaries_apply_monotonic_threshold_steps() -> Result<()> { - let policy = TaskClassifierPolicy::new(&TaskClassifierConfig { + let policy = TaskClassifierPolicy::new(&LlmCapabilityConfig { threshold_step: 0.1, - ..test_config(0.4) + ..llm_config(0.4) }); assert_eq!( diff --git a/crates/libsy/src/algorithms/stage.rs b/crates/libsy/src/algorithms/stage.rs index c2f576287..cc565f61f 100644 --- a/crates/libsy/src/algorithms/stage.rs +++ b/crates/libsy/src/algorithms/stage.rs @@ -259,6 +259,7 @@ pub(crate) fn build_stage_route(config: StageRouterConfig) -> Result PyResult { Ok(Self { inner: TaskClassifierConfig { + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold, + threshold_step, + contract: classifier_contract(prompt, response_format_type)?, + max_output_tokens, + }), fail_open: true, - base_threshold, - threshold_step, classify_trigger: classify_trigger(session_affinity), message_hash_fallback, recent_turn_window, - contract: classifier_contract(prompt, response_format_type)?, - max_output_tokens, }, }) } diff --git a/crates/switchyard-runner/src/algorithm.rs b/crates/switchyard-runner/src/algorithm.rs index acb2644c8..5cc838209 100644 --- a/crates/switchyard-runner/src/algorithm.rs +++ b/crates/switchyard-runner/src/algorithm.rs @@ -10,12 +10,12 @@ use std::path::PathBuf; use std::sync::Arc; use libsy::{ - AdvisorGate, AdvisorGateConfig, Algorithm, ClassifierContractConfig, ClassifierResponseFormat, - ClassifyTrigger, CompositeRouter, CompositeRouterConfig, CustomClassifierConfig, - CustomClassifierPolicy, EscalationJudgeConfig, GateTrigger, HandoffNoteConfig, - LlmClassifierConfig, LlmFallback, LlmTaskClassifier, Noop, Passthrough, PickerMode, - PlanExecute, PlanExecuteConfig, Random, StageRouter, StageRouterConfig, SubagentRouter, - SubagentRouterConfig, TaskClassifierConfig, ToolSemantics, + 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, }; use serde::Deserialize; use switchyard_protocol::{Category, ModelId}; @@ -526,15 +526,17 @@ pub struct StageTierConfig { impl StageClassifierConfig { fn task_classifier_config(&self) -> TaskClassifierConfig { TaskClassifierConfig { + judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig { + base_threshold: self.base_threshold, + threshold_step: self.threshold_step, + contract: classifier_contract(self.prompt.as_deref()) + .with_response_format_type(self.response_format_type), + max_output_tokens: self.max_output_tokens, + }), fail_open: true, - base_threshold: self.base_threshold, - threshold_step: self.threshold_step, classify_trigger: self.classify_trigger, message_hash_fallback: self.message_hash_fallback, recent_turn_window: self.recent_turn_window, - contract: classifier_contract(self.prompt.as_deref()) - .with_response_format_type(self.response_format_type), - max_output_tokens: self.max_output_tokens, } } } @@ -1242,15 +1244,17 @@ fn build_algorithm( let algorithm = match mode { LlmClassifierModeConfig::Capability(config) => { 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, + }), fail_open: config.fail_open, - base_threshold: config.base_threshold, - threshold_step: config.threshold_step, classify_trigger: config.classify_trigger, message_hash_fallback: config.message_hash_fallback, recent_turn_window: config.recent_turn_window, - contract: classifier_contract(config.prompt.as_deref()) - .with_response_format_type(config.response_format_type), - max_output_tokens: config.max_output_tokens, }; LlmTaskClassifier::new(LlmClassifierConfig::Capability { config: classifier_config, @@ -1494,7 +1498,7 @@ fn classifier_contract(prompt: Option<&str>) -> ClassifierContractConfig { } fn default_classifier_max_output_tokens() -> u64 { - TaskClassifierConfig::default().max_output_tokens + LlmCapabilityConfig::default().max_output_tokens } fn warn_single_target_classifier(route_name: &str, models: &CategoryModelConfig) {