1use std::collections::HashSet;
7use std::sync::Arc;
8
9use async_trait::async_trait;
10use serde::{Deserialize, Deserializer};
11use serde_json::Value;
12use switchyard_protocol::{Category, ContentBlock, Message, Role};
13
14mod decision;
15use decision::DecisionClassifier;
16pub use decision::DecisionJudgeConfig;
17
18use super::escalation;
19use super::fall_through::FallThrough;
20use super::util::DEFAULT_JUDGE_MAX_OUTPUT_TOKENS;
21use super::util::affinity::{AffinityRouter, ClassifyTrigger};
22use super::util::classifier_contract::{
23 ClassifierContract, ClassifierContractConfig, ClassifierResponseFormat,
24};
25use super::util::escalation::EscalationJudgeConfig;
26use super::util::llm_judge::{
27 ClassifierInput, JsonSchemaDecoder, JudgeClassifier, JudgePolicy, JudgeRuntimeConfig,
28 SerdeDecoder, StructuredJudge,
29};
30use super::util::target_selector::TargetSelectorPolicy;
31use crate::core::algorithm::{Algorithm, Driver};
32use crate::core::classifier::{Classification, Classifier, Score};
33use crate::core::state::State;
34use crate::{LibsyError, Result};
35use switchyard_protocol::{Request, Response};
36
37const PROMPT_TEMPLATE: &str = include_str!("../prompts/capability-classifier/prompt.md");
38const SCHEMA_TEMPLATE: &str = include_str!("../prompts/capability-classifier/schema.json");
39const ALGORITHM_NAME: &str = "llm_task_classifier";
41
42const TRAILING_ROUTING_INSTRUCTION: &str =
51 "Route the conversation above. Output ONLY the routing JSON object, nothing else.";
52
53#[derive(Deserialize)]
54#[serde(deny_unknown_fields)]
55struct TaskClassifierVerdict {
56 crux: String,
57 primary_rule: String,
58 capability_boundary: String,
59 p_solve: f64,
60}
61
62impl TaskClassifierVerdict {
63 fn is_valid(&self) -> bool {
65 (0.0..=1.0).contains(&self.p_solve)
66 && !self.crux.trim().is_empty()
67 && matches!(
68 (
69 self.primary_rule.as_str(),
70 self.capability_boundary.as_str()
71 ),
72 ("SUP-1" | "SUP-2" | "SUP-3" | "SUP-4" | "SUP-5", "supported")
73 | ("UNC-1" | "UNC-2", "uncertain")
74 | ("LIM-1" | "LIM-2", "unsupported")
75 | ("none", "unmatched")
76 )
77 }
78
79 fn boundary_steps(&self) -> Option<u8> {
81 match self.capability_boundary.as_str() {
82 "supported" => Some(0),
83 "uncertain" | "unmatched" => Some(1),
84 "unsupported" => Some(2),
85 _ => None,
86 }
87 }
88}
89
90fn trim_messages(messages: &[Message], recent_turn_window: usize) -> Vec<Message> {
100 let is_instruction = |message: &Message| matches!(message.role, Role::System | Role::Developer);
101 let mut kept: Vec<&Message> = messages.iter().filter(|m| is_instruction(m)).collect();
102 let Some(task) = messages.iter().position(|m| m.role == Role::User) else {
103 return kept.into_iter().cloned().collect();
104 };
105 kept.push(&messages[task]);
106
107 let tail: Vec<&Message> = messages[task + 1..]
108 .iter()
109 .filter(|m| !is_instruction(m))
110 .collect();
111 kept.extend(&tail[window_start(&tail, recent_turn_window)..]);
112 kept.into_iter().cloned().collect()
113}
114
115fn window_start(tail: &[&Message], recent_turn_window: usize) -> usize {
128 let counted = tail.len().saturating_sub(recent_turn_window);
129 if counted == tail.len() {
131 return counted;
132 }
133 let mut unpaired: HashSet<&str> = HashSet::new();
134 for (start, message) in tail.iter().enumerate().rev() {
135 for block in message.content.iter().rev() {
138 match block {
139 ContentBlock::ToolResult(result) => {
140 unpaired.insert(result.tool_call_id.as_str());
141 }
142 ContentBlock::ToolCall(call) => {
143 unpaired.remove(call.id.as_str());
144 }
145 _ => {}
146 }
147 }
148 if start <= counted && unpaired.is_empty() {
149 return start;
150 }
151 }
152 counted
153}
154
155fn task_messages(messages: &[Message]) -> Vec<Message> {
157 let is_task_content = |block: &ContentBlock| {
160 !matches!(
161 block,
162 ContentBlock::ToolCall(_)
163 | ContentBlock::ToolResult(_)
164 | ContentBlock::Reasoning { .. }
165 )
166 };
167 let mut user_messages = messages.iter().filter(|message| {
168 message.role == Role::User && message.content.iter().any(is_task_content)
169 });
170 let Some(opening_task) = user_messages.next() else {
171 return Vec::new();
172 };
173 [Some(opening_task), user_messages.next_back()]
174 .into_iter()
175 .flatten()
176 .map(|message| Message {
177 role: Role::User,
178 content: message
179 .content
180 .iter()
181 .filter(|block| is_task_content(block))
182 .cloned()
183 .collect(),
184 })
185 .collect()
186}
187
188struct TaskInput {
190 recent_turn_window: Option<usize>,
191}
192
193impl TaskInput {
194 fn messages(&self, request: &Request) -> Vec<Message> {
196 let mut messages = match self.recent_turn_window {
199 Some(window) => trim_messages(&request.llm_request.messages, window),
200 None => task_messages(&request.llm_request.messages),
201 };
202 for message in &mut messages {
207 message
208 .content
209 .retain(|block| !matches!(block, ContentBlock::Reasoning { .. }));
210 }
211 messages.retain(|message| !message.content.is_empty());
212 messages
213 }
214}
215
216impl ClassifierInput for TaskInput {
217 fn build_messages(&self, _state: &State, request: &Request) -> Vec<Message> {
218 let mut messages = self.messages(request);
219 if self.recent_turn_window.is_some() {
223 messages.push(Message::text(
224 Role::User,
225 TRAILING_ROUTING_INSTRUCTION.to_string(),
226 ));
227 }
228 messages
229 }
230}
231
232struct TaskClassifierPolicy {
233 base_threshold: f64,
234 threshold_step: f64,
235}
236
237impl TaskClassifierPolicy {
238 fn new(config: &LlmCapabilityConfig) -> Self {
239 Self {
240 base_threshold: config.base_threshold,
241 threshold_step: config.threshold_step,
242 }
243 }
244
245 fn threshold(&self, verdict: &TaskClassifierVerdict) -> Option<f64> {
247 Some(self.base_threshold + f64::from(verdict.boundary_steps()?) * self.threshold_step)
248 }
249}
250
251impl JudgePolicy for TaskClassifierPolicy {
252 type Verdict = TaskClassifierVerdict;
253
254 fn to_classification(
255 &self,
256 verdict: Option<&Self::Verdict>,
257 driver: &Driver,
258 ) -> Result<Classification> {
259 let Some(verdict) = verdict.filter(|verdict| verdict.is_valid()) else {
262 return Ok(Classification::Ambiguous(vec![]));
263 };
264 let Some(threshold) = self.threshold(verdict) else {
267 return Ok(Classification::Ambiguous(vec![]));
268 };
269 let category = if verdict.p_solve >= threshold
270 || (threshold - verdict.p_solve).abs() <= f64::EPSILON
271 {
272 Category::Efficient
273 } else {
274 Category::Capable
275 };
276 let Some(target) = driver.models_for(&category).first().cloned() else {
279 return Ok(Classification::Ambiguous(vec![]));
280 };
281 Ok(Classification::Scores(vec![Score {
282 target,
283 confidence: 1.0,
284 category: Some(category),
285 }]))
286 }
287}
288
289fn capability_evidence(
291 policy: &TaskClassifierPolicy,
292 verdict: Option<&TaskClassifierVerdict>,
293) -> Option<Value> {
294 let verdict = verdict?;
295 let Some(threshold) = verdict
296 .is_valid()
297 .then(|| policy.threshold(verdict))
298 .flatten()
299 else {
300 return Some(serde_json::json!({
301 "source": "fail_open",
302 "reason_code": "invalid_verdict",
303 }));
304 };
305 Some(serde_json::json!({
306 "source": "llm-classifier",
307 "score": verdict.p_solve,
308 "threshold": threshold,
309 }))
310}
311
312#[derive(Clone, Debug)]
314pub struct TaskClassifierConfig {
315 pub judge: CapabilityJudgeConfig,
317 pub fail_open: bool,
319 pub classify_trigger: ClassifyTrigger,
321 pub message_hash_fallback: bool,
323 pub recent_turn_window: Option<usize>,
326}
327
328#[derive(Clone, Debug)]
330pub enum CapabilityJudgeConfig {
331 Llm(LlmCapabilityConfig),
333 Decision(DecisionJudgeConfig),
335}
336
337impl Default for CapabilityJudgeConfig {
338 fn default() -> Self {
339 Self::Llm(LlmCapabilityConfig::default())
340 }
341}
342
343#[derive(Clone, Debug)]
345pub struct LlmCapabilityConfig {
346 pub base_threshold: f64,
348 pub threshold_step: f64,
351 pub contract: ClassifierContractConfig,
353 pub max_output_tokens: u64,
355}
356
357impl Default for LlmCapabilityConfig {
358 fn default() -> Self {
359 Self {
360 base_threshold: 0.0,
361 threshold_step: 0.0,
362 contract: ClassifierContractConfig::default(),
363 max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
364 }
365 }
366}
367
368#[derive(Deserialize)]
370#[serde(deny_unknown_fields)]
371struct TaskClassifierConfigWire {
372 base_threshold: f64,
373 #[serde(default = "default_fail_open")]
374 fail_open: bool,
375 #[serde(default)]
376 threshold_step: f64,
377 #[serde(default)]
378 classify_trigger: ClassifyTrigger,
379 #[serde(default)]
380 message_hash_fallback: bool,
381 #[serde(default)]
382 recent_turn_window: Option<usize>,
383 #[serde(default)]
384 prompt: Option<String>,
385 #[serde(default)]
386 response_format_type: ClassifierResponseFormat,
387 #[serde(default = "default_judge_max_output_tokens")]
388 max_output_tokens: u64,
389}
390
391impl<'de> Deserialize<'de> for TaskClassifierConfig {
392 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
393 where
394 D: Deserializer<'de>,
395 {
396 let wire = TaskClassifierConfigWire::deserialize(deserializer)?;
397 let mut contract = ClassifierContractConfig::default();
398 if let Some(prompt) = wire.prompt {
399 contract = contract.with_prompt(prompt);
400 }
401 contract = contract.with_response_format_type(wire.response_format_type);
402 Ok(Self {
403 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
404 base_threshold: wire.base_threshold,
405 threshold_step: wire.threshold_step,
406 contract,
407 max_output_tokens: wire.max_output_tokens,
408 }),
409 fail_open: wire.fail_open,
410 classify_trigger: wire.classify_trigger,
411 message_hash_fallback: wire.message_hash_fallback,
412 recent_turn_window: wire.recent_turn_window,
413 })
414 }
415}
416
417const fn default_fail_open() -> bool {
418 true
419}
420
421const fn default_judge_max_output_tokens() -> u64 {
422 DEFAULT_JUDGE_MAX_OUTPUT_TOKENS
423}
424
425impl Default for TaskClassifierConfig {
426 fn default() -> Self {
427 Self {
428 judge: CapabilityJudgeConfig::default(),
429 fail_open: default_fail_open(),
430 classify_trigger: ClassifyTrigger::default(),
431 message_hash_fallback: false,
432 recent_turn_window: None,
433 }
434 }
435}
436
437impl LlmCapabilityConfig {
438 fn validate(&self) -> Result<()> {
439 if !(0.0..=1.0).contains(&self.base_threshold) {
440 return Err(LibsyError::AlgorithmError {
441 message: format!(
442 "base_threshold must be between 0 and 1, got {}",
443 self.base_threshold
444 ),
445 });
446 }
447 if !self.threshold_step.is_finite() || self.threshold_step < 0.0 {
448 return Err(LibsyError::AlgorithmError {
449 message: format!(
450 "threshold_step must be finite and greater than or equal to 0, got {}",
451 self.threshold_step
452 ),
453 });
454 }
455 let unsupported_threshold = self.base_threshold + 2.0 * self.threshold_step;
456 if unsupported_threshold > 1.0 && unsupported_threshold - 1.0 > f64::EPSILON {
457 return Err(LibsyError::AlgorithmError {
458 message: format!(
459 "base_threshold + 2 * threshold_step must be at most 1, got {unsupported_threshold}"
460 ),
461 });
462 }
463 if self.max_output_tokens == 0 {
464 return Err(LibsyError::AlgorithmError {
465 message: "max_output_tokens must be at least 1".to_string(),
466 });
467 }
468 Ok(())
469 }
470}
471
472impl TaskClassifierConfig {
473 fn validate(&self) -> Result<()> {
474 match &self.judge {
475 CapabilityJudgeConfig::Llm(config) => config.validate()?,
476 CapabilityJudgeConfig::Decision(config) => config.validate()?,
477 }
478 if self.message_hash_fallback && self.classify_trigger == ClassifyTrigger::EveryRequest {
480 return Err(LibsyError::AlgorithmError {
481 message:
482 "message_hash_fallback requires classify_trigger = new_session or user_turn"
483 .to_string(),
484 });
485 }
486 Ok(())
487 }
488}
489
490#[derive(Clone, Debug)]
492pub enum CustomClassifierPolicy {
493 TargetSelector {
495 selector: String,
497 },
498}
499
500impl CustomClassifierPolicy {
501 pub fn target_selector(selector: impl Into<String>) -> Self {
503 Self::TargetSelector {
504 selector: selector.into(),
505 }
506 }
507}
508
509#[derive(Clone, Debug)]
511pub struct CustomClassifierConfig {
512 pub prompt: String,
514 pub response_schema: Value,
516 pub policy: CustomClassifierPolicy,
518 pub classify_trigger: ClassifyTrigger,
520 pub message_hash_fallback: bool,
522 pub recent_turn_window: Option<usize>,
524 pub max_output_tokens: u64,
526}
527
528impl CustomClassifierConfig {
529 pub fn new(
531 prompt: impl Into<String>,
532 response_schema: Value,
533 policy: CustomClassifierPolicy,
534 ) -> Self {
535 Self {
536 prompt: prompt.into(),
537 response_schema,
538 policy,
539 classify_trigger: ClassifyTrigger::default(),
540 message_hash_fallback: false,
541 recent_turn_window: None,
542 max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
543 }
544 }
545
546 fn validate(&self) -> Result<()> {
547 if self.max_output_tokens == 0 {
548 return Err(LibsyError::AlgorithmError {
549 message: "max_output_tokens must be at least 1".to_string(),
550 });
551 }
552 if self.message_hash_fallback && self.classify_trigger == ClassifyTrigger::EveryRequest {
554 return Err(LibsyError::AlgorithmError {
555 message:
556 "message_hash_fallback requires classify_trigger = new_session or user_turn"
557 .to_string(),
558 });
559 }
560 Ok(())
561 }
562}
563
564enum CustomPolicyRuntime {
565 TargetSelector(TargetSelectorPolicy),
566}
567
568impl JudgePolicy for CustomPolicyRuntime {
569 type Verdict = Value;
570
571 fn to_classification(
572 &self,
573 verdict: Option<&Self::Verdict>,
574 driver: &Driver,
575 ) -> Result<Classification> {
576 match self {
577 Self::TargetSelector(policy) => policy.to_classification(verdict, driver),
578 }
579 }
580}
581
582fn affinity_router(
584 trigger: ClassifyTrigger,
585 message_hash_fallback: bool,
586) -> Option<Arc<AffinityRouter>> {
587 let router = match trigger {
588 ClassifyTrigger::EveryRequest => return None,
589 ClassifyTrigger::NewSession => AffinityRouter::new(),
590 ClassifyTrigger::UserTurn => AffinityRouter::new().with_release_on_user_turn(),
591 };
592 let router = if message_hash_fallback {
593 router.with_message_hash_fallback()
594 } else {
595 router
596 };
597 Some(Arc::new(router))
598}
599
600pub struct LlmTaskClassifier {
602 route: FallThrough<State>,
603 inner: Arc<dyn Classifier<State>>,
605}
606
607struct ClassifierRouteConfig {
608 default_target: Category,
609 classify_trigger: ClassifyTrigger,
610 message_hash_fallback: bool,
611}
612
613pub struct DefaultCategoryClassifier(pub Category);
616
617#[async_trait]
618impl<S: Send> Classifier<S> for DefaultCategoryClassifier {
619 async fn score(
620 &self,
621 _state: &mut S,
622 _request: &mut Request,
623 driver: &Driver,
624 ) -> Result<(Classification, Option<Response>)> {
625 let target = driver.first_model_for(&self.0)?;
626 driver.set_evidence_if_empty(serde_json::json!({"source": "fall_open"}));
627 Ok((
628 Classification::Scores(vec![Score {
629 target: target.clone(),
630 confidence: 0.0,
631 category: Some(self.0.clone()),
632 }]),
633 None,
634 ))
635 }
636}
637
638#[derive(Clone)]
640#[non_exhaustive]
641pub enum LlmClassifierConfig {
642 Capability {
644 config: TaskClassifierConfig,
646 },
647 Escalation {
649 contract: ClassifierContractConfig,
651 config: EscalationJudgeConfig,
653 max_output_tokens: u64,
655 },
656 Custom {
658 default_target: Category,
660 config: CustomClassifierConfig,
662 },
663}
664
665impl LlmTaskClassifier {
666 pub fn new(config: LlmClassifierConfig) -> Result<Self> {
673 match config {
674 LlmClassifierConfig::Capability { config } => Self::build_capability(config),
675 LlmClassifierConfig::Escalation {
676 contract,
677 config,
678 max_output_tokens,
679 } => Self::build_escalation(contract, config, max_output_tokens),
680 LlmClassifierConfig::Custom {
681 default_target,
682 config,
683 } => Self::build_custom(default_target, config),
684 }
685 }
686
687 fn build_capability(config: TaskClassifierConfig) -> Result<Self> {
688 config.validate()?;
689 let classify_trigger = config.classify_trigger;
690 let message_hash_fallback = config.message_hash_fallback;
691 let input = TaskInput {
692 recent_turn_window: config.recent_turn_window,
693 };
694 let classifier: Arc<dyn Classifier<State>> = match config.judge {
695 CapabilityJudgeConfig::Llm(judge) => Arc::new(
696 JudgeClassifier::new(
697 StructuredJudge::new(
698 input,
699 Self::load_capability_contract(&judge.contract)?,
700 SerdeDecoder::new(),
701 JudgeRuntimeConfig::new(judge.max_output_tokens)?,
702 ),
703 TaskClassifierPolicy::new(&judge),
704 )
705 .with_error_recovery(config.fail_open)
706 .with_evidence(capability_evidence),
707 ),
708 CapabilityJudgeConfig::Decision(judge) => {
709 Arc::new(DecisionClassifier::new(judge, input, config.fail_open)?)
710 }
711 };
712 Self::from_classifier(
713 classifier,
714 ClassifierRouteConfig {
715 default_target: Category::Capable,
716 classify_trigger,
717 message_hash_fallback,
718 },
719 )
720 }
721
722 fn build_custom(default_target: Category, config: CustomClassifierConfig) -> Result<Self> {
723 config.validate()?;
724 let CustomClassifierConfig {
725 prompt,
726 response_schema,
727 policy,
728 classify_trigger,
729 message_hash_fallback,
730 recent_turn_window,
731 max_output_tokens,
732 } = config;
733 let contract = ClassifierContract::from_inner_schema(&prompt, response_schema)?;
734 let policy = match policy {
735 CustomClassifierPolicy::TargetSelector { selector } => {
736 CustomPolicyRuntime::TargetSelector(TargetSelectorPolicy::new(selector)?)
737 }
738 };
739 let classifier: Arc<dyn Classifier<State>> = Arc::new(JudgeClassifier::new(
740 StructuredJudge::new(
741 TaskInput { recent_turn_window },
742 contract,
743 JsonSchemaDecoder::new(),
744 JudgeRuntimeConfig::new(max_output_tokens)?,
745 ),
746 policy,
747 ));
748
749 Self::from_classifier(
750 classifier,
751 ClassifierRouteConfig {
752 default_target,
753 classify_trigger,
754 message_hash_fallback,
755 },
756 )
757 }
758
759 fn build_escalation(
760 contract_config: ClassifierContractConfig,
761 config: EscalationJudgeConfig,
762 max_output_tokens: u64,
763 ) -> Result<Self> {
764 let inner = escalation::build_classifier(contract_config, config, max_output_tokens)?;
765 Ok(Self {
766 route: FallThrough::<State>::new_with_state()
767 .with_name(ALGORITHM_NAME)
768 .with_classifier(Arc::clone(&inner)),
769 inner,
770 })
771 }
772
773 fn load_capability_contract(config: &ClassifierContractConfig) -> Result<ClassifierContract> {
775 ClassifierContract::from_config(config, PROMPT_TEMPLATE, SCHEMA_TEMPLATE)
776 }
777
778 fn from_classifier(
780 inner: Arc<dyn Classifier<State>>,
781 config: ClassifierRouteConfig,
782 ) -> Result<Self> {
783 if config.message_hash_fallback && config.classify_trigger == ClassifyTrigger::EveryRequest
785 {
786 return Err(LibsyError::AlgorithmError {
787 message:
788 "message_hash_fallback requires classify_trigger = new_session or user_turn"
789 .to_string(),
790 });
791 }
792 let mut route = FallThrough::<State>::new_with_state().with_name(ALGORITHM_NAME);
794 if let Some(affinity) =
795 affinity_router(config.classify_trigger, config.message_hash_fallback).as_ref()
796 {
797 route = route
799 .with_processor(affinity.clone())
800 .with_classifier(affinity.clone());
801 }
802 let fallback = DefaultCategoryClassifier(config.default_target);
803 Ok(Self {
804 route: route
805 .with_classifier(inner.clone())
806 .with_classifier(Arc::new(fallback)),
807 inner,
808 })
809 }
810}
811
812#[async_trait]
813impl Classifier<State> for LlmTaskClassifier {
814 async fn score(
815 &self,
816 state: &mut State,
817 request: &mut Request,
818 driver: &Driver,
819 ) -> Result<(Classification, Option<Response>)> {
820 self.inner.score(state, request, driver).await
821 }
822}
823
824#[async_trait]
825impl Algorithm for LlmTaskClassifier {
826 fn name(&self) -> &str {
827 "llm_task_classifier"
828 }
829
830 async fn route(
831 self: Arc<Self>,
832 driver: Driver,
833 request: Request,
834 ) -> Result<crate::RoutingOutcome> {
835 self.route.execute(driver, request).await
836 }
837}
838
839#[cfg(test)]
840mod tests {
841 use std::collections::HashMap;
842 use std::sync::Arc;
843
844 use parking_lot::Mutex;
845 use serde_json::Value;
846
847 use super::*;
848 use switchyard_protocol::{
849 ContentBlock, InstructionBlock, LlmClientError, LlmRequest, Metadata, ModelId, ToolCall,
850 ToolResult, completion_text, text_request, text_response,
851 };
852
853 use crate::algorithms::util::llm_judge::Judge;
854 use crate::core::testing::{Serve, test_drive_with_models};
855 use switchyard_protocol::{LlmResponse, Response};
856
857 const TEST_THRESHOLD: f64 = 0.5;
858
859 type CapabilityJudge = StructuredJudge<TaskInput, SerdeDecoder<TaskClassifierVerdict>>;
860
861 fn llm_config(base_threshold: f64) -> LlmCapabilityConfig {
862 LlmCapabilityConfig {
863 base_threshold,
864 ..LlmCapabilityConfig::default()
865 }
866 }
867
868 fn test_config(base_threshold: f64) -> TaskClassifierConfig {
869 TaskClassifierConfig {
870 judge: CapabilityJudgeConfig::Llm(llm_config(base_threshold)),
871 ..TaskClassifierConfig::default()
872 }
873 }
874
875 fn policy() -> TaskClassifierPolicy {
876 TaskClassifierPolicy::new(&llm_config(TEST_THRESHOLD))
877 }
878
879 fn runtime_models() -> HashMap<Category, Vec<ModelId>> {
880 [
881 (Category::Judge, vec![ModelId::from("judge")]),
882 (Category::Efficient, vec![ModelId::from("efficient")]),
883 (Category::Capable, vec![ModelId::from("capable")]),
884 (
885 Category::Any,
886 vec![ModelId::from("efficient"), ModelId::from("capable")],
887 ),
888 ]
889 .into()
890 }
891
892 fn policy_driver() -> Driver {
893 Driver::new("test", Arc::new(runtime_models().into())).0
894 }
895
896 fn verdict(
897 p_solve: f64,
898 capability_boundary: &str,
899 primary_rule: &str,
900 ) -> TaskClassifierVerdict {
901 TaskClassifierVerdict {
902 crux: "test crux".to_string(),
903 primary_rule: primary_rule.to_string(),
904 capability_boundary: capability_boundary.to_string(),
905 p_solve,
906 }
907 }
908
909 fn selected(
910 policy: &TaskClassifierPolicy,
911 verdict: Option<&TaskClassifierVerdict>,
912 ) -> Result<ModelId> {
913 policy
914 .to_classification(verdict, &policy_driver())?
915 .argmax(false)?
916 .map(|score| score.target)
917 .ok_or_else(|| LibsyError::AlgorithmError {
918 message: "policy abstained".to_string(),
919 })
920 }
921
922 #[derive(Default)]
925 struct Recorder {
926 calls: Mutex<Vec<String>>,
927 call_roles: Mutex<Vec<(String, bool)>>,
928 judge_max_output_tokens: Mutex<Vec<Option<u64>>>,
929 judge_system_prompts: Mutex<Vec<String>>,
930 }
931
932 impl Recorder {
933 fn calls(&self) -> Vec<String> {
934 self.calls.lock().clone()
935 }
936
937 fn call_roles(&self) -> Vec<(String, bool)> {
938 self.call_roles.lock().clone()
939 }
940
941 fn judge_max_output_tokens(&self) -> Vec<Option<u64>> {
942 self.judge_max_output_tokens.lock().clone()
943 }
944
945 fn judge_system_prompts(&self) -> Vec<String> {
946 self.judge_system_prompts.lock().clone()
947 }
948
949 fn serve(self: &Arc<Self>) -> impl Serve {
950 let recorder = Arc::clone(self);
951 move |model: ModelId, request: Request| {
952 let recorder = Arc::clone(&recorder);
953 async move {
954 let model = model.to_string();
955 recorder.calls.lock().push(model.clone());
956 recorder
957 .call_roles
958 .lock()
959 .push((model.clone(), model != "judge"));
960 let completion = if model == "judge" {
961 recorder
962 .judge_max_output_tokens
963 .lock()
964 .push(request.llm_request.output.max_output_tokens);
965 recorder.judge_system_prompts.lock().extend(
966 request
967 .llm_request
968 .instructions
969 .first()
970 .and_then(|instruction| {
971 instruction.content.iter().find_map(|b| {
972 if let ContentBlock::Text { text } = b {
973 Some(text.clone())
974 } else {
975 None
976 }
977 })
978 }),
979 );
980 r#"{"crux":"bounded task","primary_rule":"SUP-1","capability_boundary":"supported","p_solve":0.9}"#.to_string()
981 } else {
982 format!("answer from {model}")
983 };
984 Ok(Response {
985 llm_response: LlmResponse::Agg(text_response(None, completion)),
986 metadata: request.metadata,
987 upstream_headers: http::HeaderMap::new(),
988 })
989 }
990 }
991 }
992 }
993
994 fn unreachable_judge() -> impl Serve {
996 |model: ModelId, request: Request| async move {
997 let model = model.to_string();
998 if model == "judge" {
999 return Err(LlmClientError::Timeout {
1000 source: Box::new(std::io::Error::other("judge unreachable")),
1001 });
1002 }
1003 Ok(Response {
1004 llm_response: LlmResponse::Agg(text_response(None, format!("answer from {model}"))),
1005 metadata: request.metadata,
1006 upstream_headers: http::HeaderMap::new(),
1007 })
1008 }
1009 }
1010
1011 fn router() -> Result<Arc<LlmTaskClassifier>> {
1012 Ok(Arc::new(LlmTaskClassifier::new(
1013 LlmClassifierConfig::Capability {
1014 config: test_config(TEST_THRESHOLD),
1015 },
1016 )?))
1017 }
1018
1019 fn classify_request() -> Request {
1020 Request {
1021 llm_request: text_request(Some("auto".to_string()), "classify this task"),
1022 raw_request: None,
1023 metadata: None,
1024 }
1025 }
1026
1027 fn classify_session_request() -> Request {
1028 Request {
1029 metadata: Some(Metadata {
1030 session_id: Some("session-1".to_string()),
1031 ..Metadata::default()
1032 }),
1033 ..classify_request()
1034 }
1035 }
1036
1037 fn classify_follow_up_request() -> Request {
1038 let mut request = classify_request();
1039 request
1040 .llm_request
1041 .messages
1042 .push(Message::text(Role::Assistant, "I will add the test."));
1043 request.llm_request.messages.push(Message::text(
1044 Role::User,
1045 "Now run the test suite and report the result.",
1046 ));
1047 request
1048 }
1049
1050 #[tokio::test]
1051 async fn an_unreachable_judge_routes_capable_instead_of_failing_the_request() -> Result<()> {
1052 let router = router()?;
1053
1054 let (selected_model, response) = test_drive_with_models(
1055 router,
1056 classify_request(),
1057 runtime_models(),
1058 unreachable_judge(),
1059 )
1060 .await?;
1061
1062 assert_eq!(selected_model, "capable");
1063 assert_eq!(
1064 response.llm_response.as_agg().map(completion_text),
1065 Some("answer from capable".to_string())
1066 );
1067 Ok(())
1068 }
1069
1070 #[tokio::test]
1071 async fn classifier_judges_each_request_without_affinity() -> Result<()> {
1072 let recorder = Arc::new(Recorder::default());
1073 let router = router()?;
1074 let request = classify_request();
1075 let models = runtime_models();
1076
1077 test_drive_with_models(
1078 router.clone(),
1079 request.clone(),
1080 models.clone(),
1081 recorder.serve(),
1082 )
1083 .await?;
1084 test_drive_with_models(router, request, models, recorder.serve()).await?;
1085
1086 assert_eq!(
1087 recorder.calls(),
1088 vec!["judge", "efficient", "judge", "efficient"]
1089 );
1090 assert_eq!(
1091 recorder.call_roles(),
1092 vec![
1093 ("judge".to_string(), false),
1094 ("efficient".to_string(), true),
1095 ("judge".to_string(), false),
1096 ("efficient".to_string(), true),
1097 ]
1098 );
1099 Ok(())
1100 }
1101
1102 #[tokio::test]
1103 async fn classifier_config_sets_the_judge_completion_cap() -> Result<()> {
1104 let recorder = Arc::new(Recorder::default());
1105 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1106 config: TaskClassifierConfig {
1107 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1108 max_output_tokens: 512,
1109 ..llm_config(TEST_THRESHOLD)
1110 }),
1111 ..test_config(TEST_THRESHOLD)
1112 },
1113 })?);
1114
1115 test_drive_with_models(
1116 router,
1117 classify_request(),
1118 runtime_models(),
1119 recorder.serve(),
1120 )
1121 .await?;
1122
1123 assert_eq!(recorder.judge_max_output_tokens(), vec![Some(512)]);
1124 Ok(())
1125 }
1126
1127 #[tokio::test]
1128 async fn classifier_config_overrides_the_packaged_prompt() -> Result<()> {
1129 let recorder = Arc::new(Recorder::default());
1130 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1131 config: TaskClassifierConfig {
1132 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1133 contract: ClassifierContractConfig::default()
1134 .with_prompt("Custom capability rubric."),
1135 ..llm_config(TEST_THRESHOLD)
1136 }),
1137 ..test_config(TEST_THRESHOLD)
1138 },
1139 })?);
1140
1141 test_drive_with_models(
1142 router,
1143 classify_request(),
1144 runtime_models(),
1145 recorder.serve(),
1146 )
1147 .await?;
1148
1149 let prompts = recorder.judge_system_prompts();
1150 assert_eq!(prompts.len(), 1);
1151 assert_eq!(prompts[0], "Custom capability rubric.");
1152 Ok(())
1153 }
1154
1155 #[tokio::test]
1156 async fn classifier_config_enables_new_session_trigger() -> Result<()> {
1157 let recorder = Arc::new(Recorder::default());
1158 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1159 config: TaskClassifierConfig {
1160 classify_trigger: ClassifyTrigger::NewSession,
1161 ..test_config(TEST_THRESHOLD)
1162 },
1163 })?);
1164
1165 let request = classify_session_request();
1166 let models = runtime_models();
1167 test_drive_with_models(
1168 router.clone(),
1169 request.clone(),
1170 models.clone(),
1171 recorder.serve(),
1172 )
1173 .await?;
1174 test_drive_with_models(router, request, models, recorder.serve()).await?;
1175
1176 assert_eq!(recorder.calls(), vec!["judge", "efficient", "efficient"]);
1177 Ok(())
1178 }
1179
1180 #[tokio::test]
1181 async fn classifier_config_reuses_message_hash_affinity_for_a_follow_up() -> Result<()> {
1182 let recorder = Arc::new(Recorder::default());
1183 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1184 config: TaskClassifierConfig {
1185 classify_trigger: ClassifyTrigger::NewSession,
1186 message_hash_fallback: true,
1187 recent_turn_window: None,
1188 ..test_config(TEST_THRESHOLD)
1189 },
1190 })?);
1191
1192 let models = runtime_models();
1193 test_drive_with_models(
1194 router.clone(),
1195 classify_request(),
1196 models.clone(),
1197 recorder.serve(),
1198 )
1199 .await?;
1200 test_drive_with_models(
1201 router,
1202 classify_follow_up_request(),
1203 models,
1204 recorder.serve(),
1205 )
1206 .await?;
1207
1208 assert_eq!(recorder.calls(), vec!["judge", "efficient", "efficient"]);
1209 Ok(())
1210 }
1211
1212 #[tokio::test]
1213 async fn one_classifier_uses_each_requests_runtime_models() -> Result<()> {
1214 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1215 config: TaskClassifierConfig {
1216 classify_trigger: ClassifyTrigger::NewSession,
1217 ..test_config(TEST_THRESHOLD)
1218 },
1219 })?);
1220 let calls = Arc::new(Mutex::new(Vec::new()));
1221 let serve = |calls: Arc<Mutex<Vec<String>>>| {
1222 move |model: ModelId, _request: Request| {
1223 let calls = Arc::clone(&calls);
1224 async move {
1225 calls.lock().push(model.to_string());
1226 let text = if model.as_str().starts_with("judge-") {
1227 r#"{"crux":"bounded task","primary_rule":"SUP-1","capability_boundary":"supported","p_solve":0.9}"#.to_string()
1228 } else {
1229 model.to_string()
1230 };
1231 Ok(Response {
1232 llm_response: LlmResponse::Agg(text_response(None, text)),
1233 metadata: None,
1234 upstream_headers: Default::default(),
1235 })
1236 }
1237 }
1238 };
1239 let models = |suffix: &str| -> HashMap<Category, Vec<ModelId>> {
1240 [
1241 (
1242 Category::Judge,
1243 vec![ModelId::from(format!("judge-{suffix}"))],
1244 ),
1245 (
1246 Category::Efficient,
1247 vec![ModelId::from(format!("efficient-{suffix}"))],
1248 ),
1249 (
1250 Category::Capable,
1251 vec![ModelId::from(format!("capable-{suffix}"))],
1252 ),
1253 (
1254 Category::Any,
1255 vec![
1256 ModelId::from(format!("efficient-{suffix}")),
1257 ModelId::from(format!("capable-{suffix}")),
1258 ],
1259 ),
1260 ]
1261 .into()
1262 };
1263 let request = classify_session_request();
1264
1265 let (first, _) = test_drive_with_models(
1266 router.clone(),
1267 request.clone(),
1268 models("a"),
1269 serve(Arc::clone(&calls)),
1270 )
1271 .await?;
1272 let (second, _) =
1273 test_drive_with_models(router, request, models("b"), serve(Arc::clone(&calls))).await?;
1274
1275 assert_eq!(first, "efficient-a");
1276 assert_eq!(second, "efficient-b");
1277 assert_eq!(
1278 &*calls.lock(),
1279 &["judge-a", "efficient-a", "judge-b", "efficient-b"]
1280 );
1281 Ok(())
1282 }
1283
1284 #[tokio::test]
1285 async fn decision_judge_routes_on_relative_advantage() -> Result<()> {
1286 use crate::{Call, RuntimeModels, drive};
1287 use serde_json::json;
1288 use std::collections::BTreeMap;
1289 use std::sync::atomic::{AtomicUsize, Ordering};
1290 use switchyard_protocol::{
1291 BooleanEstimate, DecisionAnswer, DecisionResponse, DecisionValue, Probability,
1292 ProviderConfidence,
1293 };
1294
1295 let settings = DecisionJudgeConfig {
1296 cutoff: 0.4,
1297 instructions: None,
1298 candidates: BTreeMap::from([
1299 ("a".into(), "capable".into()),
1300 ("b".into(), "efficient".into()),
1301 ("c".into(), "third".into()),
1302 ]),
1303 evidence: json!({
1304 "candidate_descriptions": {"a": "Profile A", "b": "Profile B", "c": "Profile C"},
1305 "reference_cases": [{"task": "Independent example", "observed_solved": {"a": true, "b": null, "c": false}}],
1306 "summaries": [],
1307 }),
1308 };
1309 let mut request = classify_session_request();
1310 request.llm_request.messages.push(Message {
1311 role: Role::Assistant,
1312 content: vec![
1313 ContentBlock::Reasoning {
1314 text: "private reasoning".into(),
1315 signature: None,
1316 details: vec![],
1317 },
1318 ContentBlock::Text {
1319 text: "visible progress".into(),
1320 },
1321 ],
1322 });
1323 request
1324 .llm_request
1325 .messages
1326 .push(Message::text(Role::User, "latest follow-up"));
1327 let original_messages = request.llm_request.messages.clone();
1328 let expected_task = json!([
1329 original_messages[0].clone(),
1330 Message::text(Role::Assistant, "visible progress"),
1331 Message::text(Role::User, "latest follow-up"),
1332 ]);
1333
1334 for (name, score, fail_open, expected) in [
1335 ("above", 0.45, true, Some("capable")),
1336 ("equal", 0.4, true, Some("efficient")),
1337 ("below", 0.39, false, Some("efficient")),
1338 ("missing answer", 0.0, true, Some("capable")),
1339 ("no distribution", 0.0, true, Some("capable")),
1340 ("wrong type", 0.0, true, Some("capable")),
1341 ("out of range", 1.1, true, Some("capable")),
1342 ("nonfinite", f64::NAN, false, Some("capable")),
1343 ("provider error", 0.0, true, Some("capable")),
1344 ("provider error", 0.0, false, None),
1345 ("dropped reply", 0.0, true, Some("capable")),
1346 ("missing candidate", 0.0, true, None),
1347 ] {
1348 let mut config = settings.clone();
1349 if name == "missing candidate" {
1350 config.candidates.remove("a");
1351 }
1352 if name == "below" {
1353 config.instructions = Some(
1354 json!({"policy": "Judge capable-only success using the supplied evidence."}),
1355 );
1356 }
1357 let override_instructions = config.instructions.clone();
1358 let router: Arc<dyn Algorithm> =
1359 Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1360 config: TaskClassifierConfig {
1361 judge: CapabilityJudgeConfig::Decision(config),
1362 fail_open,
1363 recent_turn_window: Some(3),
1364 classify_trigger: ClassifyTrigger::NewSession,
1365 ..TaskClassifierConfig::default()
1366 },
1367 })?);
1368 let calls = AtomicUsize::new(0);
1369 let serve = |call| {
1370 let calls = &calls;
1371 let expected_task = &expected_task;
1372 let settings = &settings;
1373 let override_instructions = &override_instructions;
1374 async move {
1375 let Call::Decision(call) = call else {
1376 panic!("judge must use the decision step")
1377 };
1378 calls.fetch_add(1, Ordering::SeqCst);
1379 assert_eq!(call.model, "judge");
1380 assert_eq!(call.request.model, Some("judge".into()));
1381 assert_eq!(
1382 call.request.context,
1383 json!({
1384 "task": expected_task,
1385 "candidates": ["a", "b", "c"],
1386 "comparison": {"capable": "a", "efficient": "b"},
1387 "evidence": settings.evidence,
1388 })
1389 );
1390 assert_eq!(call.request.questions.len(), 1);
1391 let question = &call.request.questions["route"];
1392 let switchyard_protocol::DecisionKind::Choice { options } = &question.kind
1393 else {
1394 panic!("expected a choice")
1395 };
1396 assert_eq!(
1397 options
1398 .iter()
1399 .map(|option| option.id.as_str())
1400 .collect::<Vec<_>>(),
1401 ["advantage", "no_advantage"]
1402 );
1403 if let Some(instructions) = override_instructions {
1404 assert_eq!(&question.instructions, instructions);
1405 } else {
1406 for key in ["setting", "evidence", "comparison", "boundary", "policy"] {
1407 assert!(question.instructions[key].as_str().is_some());
1408 }
1409 }
1410 if name == "provider error" {
1411 return call.respond(Err(LibsyError::client_call(
1412 "judge",
1413 LlmClientError::General("private provider body".into()),
1414 )));
1415 }
1416 if name == "dropped reply" {
1417 drop(call);
1418 return Ok(());
1419 }
1420 let value = if name == "wrong type" {
1421 DecisionValue::Boolean(BooleanEstimate::Value(true))
1422 } else {
1423 DecisionValue::Choice {
1424 selected: "no_advantage".into(),
1425 probabilities: (name != "no distribution").then(|| {
1426 BTreeMap::from([
1427 ("advantage".into(), Probability(score)),
1428 ("no_advantage".into(), Probability(1.0 - score)),
1429 ])
1430 }),
1431 }
1432 };
1433 let answers = if name == "missing answer" {
1434 BTreeMap::new()
1435 } else {
1436 BTreeMap::from([(
1437 "route".into(),
1438 DecisionAnswer {
1439 value,
1440 provider_confidence: Some(ProviderConfidence(0.99)),
1441 },
1442 )])
1443 };
1444 call.respond(Ok(DecisionResponse {
1445 id: None,
1446 model: Some("provider-judge".into()),
1447 answers,
1448 usage: Default::default(),
1449 }))
1450 }
1451 };
1452 let models = Arc::new(RuntimeModels::new(runtime_models()));
1453 let result = drive(router.clone(), request.clone(), models.clone(), &serve).await;
1454 assert_eq!(
1455 calls.load(Ordering::SeqCst),
1456 usize::from(name != "missing candidate"),
1457 "{name}"
1458 );
1459 let Some(expected) = expected else {
1460 if name == "missing candidate" {
1461 assert!(
1462 matches!(result, Err(LibsyError::AlgorithmError { message }) if message.contains("candidate is missing"))
1463 );
1464 } else {
1465 assert!(
1466 matches!(result, Err(LibsyError::ClientCall { .. })),
1467 "{name}"
1468 );
1469 }
1470 continue;
1471 };
1472 let outcome = result?;
1473 assert_eq!(outcome.selected_model_id()?, expected, "{name}");
1474 assert!(outcome.response.is_none());
1475 assert_eq!(outcome.request.llm_request.messages, original_messages);
1476 let evidence = outcome
1477 .metadata
1478 .and_then(|metadata| metadata.evidence)
1479 .expect("routing evidence");
1480 if matches!(name, "above" | "equal" | "below") {
1481 assert_eq!(evidence["source"], "decision_classifier");
1482 assert_eq!(evidence["verdict"], "relative_advantage");
1483 assert_eq!(evidence["threshold"], settings.cutoff);
1484 assert_eq!(evidence["score"], score);
1485 } else {
1486 let reason = match name {
1487 "provider error" => "client_error",
1488 "dropped reply" => "call_error",
1489 _ => "invalid_verdict",
1490 };
1491 assert_eq!(
1492 evidence,
1493 json!({"source": "fail_open", "reason_code": reason})
1494 );
1495 }
1496 let retained = drive(router, request.clone(), models, &serve).await?;
1497 assert_eq!(retained.selected_model_id()?, expected);
1498 assert_eq!(
1499 calls.load(Ordering::SeqCst),
1500 1,
1501 "affinity should skip the judge: {name}"
1502 );
1503 }
1504 for cutoff in [-0.1, 1.1, f64::NAN] {
1505 let mut config = settings.clone();
1506 config.cutoff = cutoff;
1507 assert!(
1508 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1509 config: TaskClassifierConfig {
1510 judge: CapabilityJudgeConfig::Decision(config),
1511 ..TaskClassifierConfig::default()
1512 },
1513 })
1514 .is_err()
1515 );
1516 }
1517 let mut duplicate = settings.clone();
1518 duplicate
1519 .candidates
1520 .insert("duplicate-a".into(), "capable".into());
1521 assert!(duplicate.validate().is_err());
1522 Ok(())
1523 }
1524
1525 #[test]
1526 fn the_threshold_boundary_is_inclusive() -> Result<()> {
1527 let policy = policy();
1528 let at_threshold = verdict(0.5, "supported", "SUP-1");
1529 let below_threshold = verdict(0.49, "supported", "SUP-1");
1530 assert_eq!(selected(&policy, Some(&at_threshold))?, "efficient");
1531 assert_eq!(selected(&policy, Some(&below_threshold))?, "capable");
1532 Ok(())
1533 }
1534
1535 #[test]
1536 fn the_threshold_moves_the_routing_boundary() -> Result<()> {
1537 let borderline = verdict(0.5, "supported", "SUP-1");
1538 let strict = TaskClassifierPolicy::new(&llm_config(0.9));
1539 let lenient = TaskClassifierPolicy::new(&llm_config(0.1));
1540 assert_eq!(selected(&strict, Some(&borderline))?, "capable");
1541 assert_eq!(selected(&lenient, Some(&borderline))?, "efficient");
1542 Ok(())
1543 }
1544
1545 #[test]
1546 fn classifier_config_rejects_unknown_fields() {
1547 let error = serde_json::from_value::<TaskClassifierConfig>(serde_json::json!({
1548 "base_threshold": 0.5,
1549 "classifier_magic": true,
1550 }))
1551 .expect_err("unknown classifier fields must be rejected");
1552
1553 assert!(
1554 error
1555 .to_string()
1556 .contains("unknown field `classifier_magic`"),
1557 "{error}"
1558 );
1559 }
1560
1561 #[test]
1562 fn invalid_classifier_config_is_rejected() -> Result<()> {
1563 for bad in [1.5, -0.1, f64::NAN, f64::INFINITY] {
1564 assert!(
1565 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1566 config: test_config(bad),
1567 })
1568 .is_err(),
1569 "base threshold {bad} should be rejected"
1570 );
1571 }
1572 for config in [
1573 TaskClassifierConfig {
1574 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1575 base_threshold: 0.5,
1576 threshold_step: -0.1,
1577 ..LlmCapabilityConfig::default()
1578 }),
1579 ..TaskClassifierConfig::default()
1580 },
1581 TaskClassifierConfig {
1582 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1583 base_threshold: 0.8,
1584 threshold_step: 0.11,
1585 ..LlmCapabilityConfig::default()
1586 }),
1587 ..TaskClassifierConfig::default()
1588 },
1589 TaskClassifierConfig {
1590 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1591 base_threshold: 0.5,
1592 ..LlmCapabilityConfig::default()
1593 }),
1594 message_hash_fallback: true,
1595 ..TaskClassifierConfig::default()
1596 },
1597 TaskClassifierConfig {
1598 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1599 base_threshold: 0.5,
1600 max_output_tokens: 0,
1601 ..LlmCapabilityConfig::default()
1602 }),
1603 ..TaskClassifierConfig::default()
1604 },
1605 ] {
1606 assert!(LlmTaskClassifier::new(LlmClassifierConfig::Capability { config }).is_err());
1607 }
1608 for base_threshold in [0.0, 1.0] {
1609 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1610 config: test_config(base_threshold),
1611 })?;
1612 }
1613 Ok(())
1614 }
1615
1616 #[test]
1617 fn message_hash_fallback_accepts_retaining_triggers() -> Result<()> {
1618 for trigger in [ClassifyTrigger::NewSession, ClassifyTrigger::UserTurn] {
1622 let config = TaskClassifierConfig {
1623 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1624 base_threshold: 0.5,
1625 ..LlmCapabilityConfig::default()
1626 }),
1627 classify_trigger: trigger,
1628 message_hash_fallback: true,
1629 ..TaskClassifierConfig::default()
1630 };
1631 LlmTaskClassifier::new(LlmClassifierConfig::Capability { config }).map_err(
1632 |error| LibsyError::AlgorithmError {
1633 message: format!("{trigger:?} with message_hash_fallback rejected: {error}"),
1634 },
1635 )?;
1636 }
1637 let every_request = TaskClassifierConfig {
1638 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1639 base_threshold: 0.5,
1640 ..LlmCapabilityConfig::default()
1641 }),
1642 classify_trigger: ClassifyTrigger::EveryRequest,
1643 message_hash_fallback: true,
1644 ..TaskClassifierConfig::default()
1645 };
1646 assert!(
1647 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1648 config: every_request
1649 })
1650 .is_err(),
1651 "every_request with message_hash_fallback should stay rejected"
1652 );
1653 Ok(())
1654 }
1655
1656 #[test]
1657 fn an_unusable_verdict_is_ambiguous() -> Result<()> {
1658 let policy = policy();
1659 let inconsistent_rule = TaskClassifierVerdict {
1660 capability_boundary: "uncertain".to_string(),
1661 ..verdict(1.0, "supported", "SUP-1")
1662 };
1663 let empty_crux = TaskClassifierVerdict {
1664 crux: " ".to_string(),
1665 ..verdict(1.0, "supported", "SUP-1")
1666 };
1667 let unusable = [
1668 Some(verdict(1.1, "supported", "SUP-1")),
1669 Some(inconsistent_rule),
1670 Some(empty_crux),
1671 None,
1672 ];
1673 for verdict in unusable {
1674 let classification = policy.to_classification(verdict.as_ref(), &policy_driver())?;
1675 assert!(matches!(classification, Classification::Ambiguous(_)));
1676 assert!(classification.argmax(false)?.is_none());
1677 assert!(classification.argmax(true)?.is_none());
1678 }
1679 Ok(())
1680 }
1681
1682 #[test]
1683 fn capability_boundaries_apply_monotonic_threshold_steps() -> Result<()> {
1684 let policy = TaskClassifierPolicy::new(&LlmCapabilityConfig {
1685 threshold_step: 0.1,
1686 ..llm_config(0.4)
1687 });
1688
1689 assert_eq!(
1690 selected(&policy, Some(&verdict(0.4, "supported", "SUP-2")))?,
1691 "efficient"
1692 );
1693 assert_eq!(
1694 selected(&policy, Some(&verdict(0.49, "uncertain", "UNC-1")))?,
1695 "capable"
1696 );
1697 assert_eq!(
1698 selected(&policy, Some(&verdict(0.5, "uncertain", "UNC-1")))?,
1699 "efficient"
1700 );
1701 assert_eq!(
1702 selected(&policy, Some(&verdict(0.5, "unmatched", "none")))?,
1703 "efficient"
1704 );
1705 assert_eq!(
1706 selected(&policy, Some(&verdict(0.59, "unsupported", "LIM-1")))?,
1707 "capable"
1708 );
1709 assert_eq!(
1710 selected(&policy, Some(&verdict(0.6, "unsupported", "LIM-1")))?,
1711 "efficient"
1712 );
1713 Ok(())
1714 }
1715
1716 fn capability_judge(recent_turn_window: Option<usize>) -> Result<CapabilityJudge> {
1719 Ok(StructuredJudge::new(
1720 TaskInput { recent_turn_window },
1721 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?,
1722 SerdeDecoder::new(),
1723 JudgeRuntimeConfig::new(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?,
1724 ))
1725 }
1726
1727 fn judged_contents(recent_turn_window: usize) -> Result<Vec<String>> {
1728 let judge = capability_judge(Some(recent_turn_window))?;
1729 let request = Request {
1730 llm_request: LlmRequest {
1731 messages: vec![
1732 Message::text(Role::System, "client instructions"),
1733 Message::text(Role::User, "initial task"),
1734 Message::text(Role::Assistant, "old response"),
1735 Message::text(Role::User, "old follow-up"),
1736 Message::text(Role::Assistant, "recent 1"),
1737 Message::text(Role::User, "recent 2"),
1738 ],
1739 ..LlmRequest::default()
1740 },
1741 raw_request: None,
1742 metadata: None,
1743 };
1744 Ok(judge
1745 .build_request(&State::default(), &request)
1746 .llm_request
1747 .messages
1748 .iter()
1749 .filter_map(|message| message.text_content("\n"))
1750 .collect())
1751 }
1752
1753 #[test]
1754 fn a_window_widens_the_judge_to_the_surrounding_conversation() -> Result<()> {
1755 let contents = judged_contents(2)?;
1757 assert!(contents.contains(&"client instructions".to_string()));
1758 assert!(contents.contains(&"initial task".to_string()));
1759 assert!(contents.contains(&"recent 1".to_string()));
1760 assert!(contents.contains(&"recent 2".to_string()));
1761 assert!(!contents.contains(&"old response".to_string()));
1762 Ok(())
1763 }
1764
1765 #[test]
1766 fn a_zero_window_keeps_only_the_instructions_and_the_task() -> Result<()> {
1767 let contents = judged_contents(0)?;
1768 assert!(contents.contains(&"client instructions".to_string()));
1769 assert!(contents.contains(&"initial task".to_string()));
1770 assert!(!contents.contains(&"recent 2".to_string()));
1771 Ok(())
1772 }
1773
1774 fn tool_call(id: &str) -> Message {
1775 Message {
1776 role: Role::Assistant,
1777 content: vec![ContentBlock::ToolCall(ToolCall {
1778 id: id.to_string(),
1779 name: "search".to_string(),
1780 arguments: Value::Null,
1781 })],
1782 }
1783 }
1784
1785 fn tool_result(id: &str) -> Message {
1786 Message {
1787 role: Role::Tool,
1788 content: vec![ContentBlock::ToolResult(ToolResult {
1789 tool_call_id: id.to_string(),
1790 content: vec![ContentBlock::Text {
1791 text: "tool output".to_string(),
1792 }],
1793 is_error: None,
1794 })],
1795 }
1796 }
1797
1798 #[test]
1799 fn default_task_input_keeps_user_content_around_tool_results() {
1800 let mut result = tool_result("call-1");
1801 result.role = Role::User;
1802 let mut mixed = result.clone();
1803 mixed.content.push(ContentBlock::Text {
1804 text: "latest follow-up".to_string(),
1805 });
1806 let input = TaskInput {
1807 recent_turn_window: None,
1808 };
1809 let mut request = Request {
1810 llm_request: LlmRequest {
1811 messages: vec![
1812 result.clone(),
1813 Message::text(Role::User, "initial task"),
1814 tool_call("call-1"),
1815 mixed,
1816 result.clone(),
1817 ],
1818 ..LlmRequest::default()
1819 },
1820 ..Request::default()
1821 };
1822 assert_eq!(
1823 input.build_messages(&State::default(), &request),
1824 vec![
1825 Message::text(Role::User, "initial task"),
1826 Message::text(Role::User, "latest follow-up"),
1827 ]
1828 );
1829 request.llm_request.messages = vec![result];
1830 assert!(input.build_messages(&State::default(), &request).is_empty());
1831 }
1832
1833 #[test]
1836 fn trimming_keeps_the_call_that_introduced_a_kept_tool_result() {
1837 let messages = vec![
1838 Message::text(Role::System, "client instructions"),
1839 Message::text(Role::User, "initial task"),
1840 Message::text(Role::Assistant, "old response"),
1841 tool_call("call-1"),
1842 tool_result("call-1"),
1843 Message::text(Role::Assistant, "recent 1"),
1844 Message::text(Role::User, "recent 2"),
1845 Message::text(Role::Assistant, "recent 3"),
1846 Message::text(Role::User, "recent 4"),
1847 ];
1848
1849 let kept = trim_messages(&messages, 5);
1851
1852 assert_eq!(
1853 kept,
1854 vec![
1855 Message::text(Role::System, "client instructions"),
1856 Message::text(Role::User, "initial task"),
1857 tool_call("call-1"),
1858 tool_result("call-1"),
1859 Message::text(Role::Assistant, "recent 1"),
1860 Message::text(Role::User, "recent 2"),
1861 Message::text(Role::Assistant, "recent 3"),
1862 Message::text(Role::User, "recent 4"),
1863 ]
1864 );
1865 }
1866
1867 #[test]
1870 fn trimming_pairs_a_repeated_id_with_the_call_that_precedes_it() {
1871 let messages = vec![
1872 Message::text(Role::System, "client instructions"),
1873 Message::text(Role::User, "initial task"),
1874 tool_call("x"),
1875 tool_result("x"),
1876 Message::text(Role::Assistant, "later"),
1877 tool_call("x"),
1878 tool_result("x"),
1879 ];
1880
1881 let kept = trim_messages(&messages, 4);
1883
1884 assert_eq!(
1885 kept,
1886 vec![
1887 Message::text(Role::System, "client instructions"),
1888 Message::text(Role::User, "initial task"),
1889 tool_call("x"),
1890 tool_result("x"),
1891 Message::text(Role::Assistant, "later"),
1892 tool_call("x"),
1893 tool_result("x"),
1894 ]
1895 );
1896 }
1897
1898 #[test]
1901 fn trimming_keeps_the_counted_window_when_a_result_cannot_be_paired() {
1902 let messages = vec![
1903 Message::text(Role::System, "client instructions"),
1904 tool_call("orphan"),
1905 Message::text(Role::User, "initial task"),
1906 Message::text(Role::Assistant, "old response"),
1907 tool_result("orphan"),
1908 Message::text(Role::Assistant, "recent 1"),
1909 Message::text(Role::User, "recent 2"),
1910 ];
1911
1912 let kept = trim_messages(&messages, 3);
1913
1914 assert_eq!(
1915 kept,
1916 vec![
1917 Message::text(Role::System, "client instructions"),
1918 Message::text(Role::User, "initial task"),
1919 tool_result("orphan"),
1920 Message::text(Role::Assistant, "recent 1"),
1921 Message::text(Role::User, "recent 2"),
1922 ]
1923 );
1924 }
1925
1926 #[test]
1927 fn a_window_restates_the_routing_instruction_last() -> Result<()> {
1928 let contents = judged_contents(2)?;
1929
1930 assert_eq!(
1933 contents.last().map(String::as_str),
1934 Some(TRAILING_ROUTING_INSTRUCTION)
1935 );
1936 Ok(())
1937 }
1938
1939 #[test]
1944 fn a_window_drops_reasoning_but_keeps_visible_text_and_tool_pairs() {
1945 let messages = vec![
1946 Message::text(Role::System, "client instructions"),
1947 Message::text(Role::User, "initial task"),
1948 Message {
1949 role: Role::Assistant,
1950 content: vec![
1951 ContentBlock::Reasoning {
1952 text: "private chain of thought".to_string(),
1953 signature: None,
1954 details: Vec::new(),
1955 },
1956 ContentBlock::Text {
1957 text: "visible answer".to_string(),
1958 },
1959 ],
1960 },
1961 tool_call("call-1"),
1962 tool_result("call-1"),
1963 Message {
1964 role: Role::Assistant,
1965 content: vec![ContentBlock::Reasoning {
1966 text: "reasoning-only turn".to_string(),
1967 signature: None,
1968 details: Vec::new(),
1969 }],
1970 },
1971 Message::text(Role::User, "follow-up"),
1972 ];
1973 let request = Request {
1974 llm_request: LlmRequest {
1975 messages,
1976 ..LlmRequest::default()
1977 },
1978 raw_request: None,
1979 metadata: None,
1980 };
1981
1982 let built = TaskInput {
1983 recent_turn_window: Some(10),
1984 }
1985 .build_messages(&State::default(), &request);
1986
1987 assert!(
1988 !built
1989 .iter()
1990 .flat_map(|message| &message.content)
1991 .any(|block| matches!(block, ContentBlock::Reasoning { .. })),
1992 "{built:?}"
1993 );
1994 assert!(
1995 built
1996 .iter()
1997 .any(|message| message.text_content("\n").as_deref() == Some("visible answer"))
1998 );
1999 assert!(built.contains(&tool_call("call-1")));
2000 assert!(built.contains(&tool_result("call-1")));
2001 assert_eq!(built.len(), 7);
2004 assert!(built.iter().all(|message| !message.content.is_empty()));
2005 }
2006
2007 #[test]
2008 fn the_default_path_is_left_unchanged() -> Result<()> {
2009 let judge = capability_judge(None)?;
2012 let request = Request {
2013 llm_request: LlmRequest {
2014 messages: vec![Message::text(Role::User, "the task")],
2015 ..LlmRequest::default()
2016 },
2017 raw_request: None,
2018 metadata: None,
2019 };
2020
2021 let built = judge.build_request(&State::default(), &request);
2022
2023 assert_eq!(built.llm_request.messages.len(), 1);
2026 assert!(!built.llm_request.instructions.is_empty());
2027 assert!(
2028 !built
2029 .llm_request
2030 .messages
2031 .iter()
2032 .filter_map(|message| message.text_content("\n"))
2033 .any(|text| text.contains(TRAILING_ROUTING_INSTRUCTION))
2034 );
2035 Ok(())
2036 }
2037
2038 #[test]
2039 fn capability_judge_builds_a_structured_request() -> Result<()> {
2040 let judge = capability_judge(None)?;
2041 let request = Request {
2042 llm_request: LlmRequest {
2043 model: Some("inbound".to_string()),
2044 messages: vec![
2045 Message::text(Role::System, "client instructions"),
2046 Message::text(Role::Developer, "client developer instructions"),
2047 Message::text(Role::User, "initial task"),
2048 Message::text(Role::Assistant, "old response"),
2049 Message::text(Role::User, "old follow-up"),
2050 Message::text(Role::Assistant, "recent 1"),
2051 Message::text(Role::User, "recent 2"),
2052 Message::text(Role::Assistant, "recent 3"),
2053 Message::text(Role::User, "recent 4"),
2054 Message::text(Role::Assistant, "recent 5"),
2055 ],
2056 ..LlmRequest::default()
2057 },
2058 raw_request: None,
2059 metadata: None,
2060 };
2061 let judge_request = judge.build_request(&State::default(), &request);
2062
2063 assert_eq!(judge_request.llm_request.model, request.llm_request.model);
2064 assert_eq!(judge_request.llm_request.instructions.len(), 1);
2065 assert_eq!(judge_request.llm_request.instructions[0].role, Role::System);
2066 assert_eq!(
2067 judge_request.llm_request.instructions[0].content,
2068 InstructionBlock {
2069 role: Role::System,
2070 content: Message::text(Role::System, judge.contract().system_prompt()).content,
2071 }
2072 .content,
2073 );
2074 assert_eq!(judge_request.llm_request.messages.len(), 2);
2075 let contents = judge_request
2076 .llm_request
2077 .messages
2078 .iter()
2079 .filter_map(|message| message.text_content("\n"))
2080 .collect::<Vec<_>>();
2081 assert!(contents.contains(&"recent 4".to_string()));
2082 assert!(contents.contains(&"initial task".to_string()));
2083 assert!(!contents.contains(&"recent 5".to_string()));
2084 assert!(!contents.contains(&"client instructions".to_string()));
2085 assert_eq!(
2086 judge_request.llm_request.output.response_format,
2087 Some(judge.contract().response_format().clone())
2088 );
2089 assert_eq!(
2090 judge_request.llm_request.output.max_output_tokens,
2091 Some(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)
2092 );
2093 Ok(())
2094 }
2095
2096 fn sample_value(spec: &Value) -> Value {
2097 if let Some(first) = spec
2098 .get("enum")
2099 .and_then(Value::as_array)
2100 .and_then(|values| values.first())
2101 {
2102 return first.clone();
2103 }
2104 match spec.get("type").and_then(Value::as_str) {
2105 Some("number") => serde_json::json!(0.5),
2106 Some("boolean") => serde_json::json!(false),
2107 _ => serde_json::json!("sample"),
2108 }
2109 }
2110
2111 fn schema_shaped_verdict(schema: &Value) -> Result<String> {
2112 let properties = schema
2113 .pointer("/json_schema/schema/properties")
2114 .and_then(Value::as_object)
2115 .ok_or_else(|| LibsyError::AlgorithmError {
2116 message: "packaged schema declares no properties".to_string(),
2117 })?;
2118 Ok(Value::Object(
2119 properties
2120 .iter()
2121 .map(|(name, spec)| (name.clone(), sample_value(spec)))
2122 .collect(),
2123 )
2124 .to_string())
2125 }
2126
2127 #[test]
2130 fn every_schema_property_round_trips_through_the_judge_parser() -> Result<()> {
2131 let contract =
2132 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?;
2133 let schema = contract.response_format();
2134 let reply = schema_shaped_verdict(schema)?;
2135 let judge: CapabilityJudge = StructuredJudge::new(
2136 TaskInput {
2137 recent_turn_window: None,
2138 },
2139 contract,
2140 SerdeDecoder::new(),
2141 JudgeRuntimeConfig::new(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?,
2142 );
2143
2144 let verdict = judge.parse(&text_response(None, reply))?;
2145
2146 assert!(verdict.is_valid());
2147 assert!((0.0..=1.0).contains(&verdict.p_solve));
2148 Ok(())
2149 }
2150
2151 #[test]
2152 fn packaged_prompt_keeps_the_schema_in_the_structured_request() -> Result<()> {
2153 let contract =
2154 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?;
2155 let prompt = contract.system_prompt();
2156 let schema_name = contract
2157 .response_format()
2158 .pointer("/json_schema/name")
2159 .and_then(Value::as_str)
2160 .ok_or_else(|| LibsyError::AlgorithmError {
2161 message: "packaged response schema has no name".to_string(),
2162 })?;
2163 assert_eq!(schema_name, "CapabilityClassifierDecision");
2164 assert!(prompt.contains("SUP-1 [supported]"));
2165 assert!(prompt.contains("SUP-5 [supported]"));
2166 assert!(!prompt.contains("{{RESPONSE_SCHEMA}}"));
2167 assert!(!prompt.contains("\"type\": \"object\""));
2168 assert!(!prompt.contains("\"json_schema\""));
2169 assert!(!prompt.contains(schema_name));
2170 let rule_values = contract
2171 .response_format()
2172 .pointer("/json_schema/schema/properties/primary_rule/enum")
2173 .and_then(Value::as_array)
2174 .ok_or_else(|| LibsyError::AlgorithmError {
2175 message: "rendered response schema has no primary rule enum".to_string(),
2176 })?;
2177 assert!(
2178 rule_values
2179 .iter()
2180 .any(|value| value.as_str() == Some("SUP-1"))
2181 );
2182 assert!(
2183 rule_values
2184 .iter()
2185 .any(|value| value.as_str() == Some("none"))
2186 );
2187 Ok(())
2188 }
2189}