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
1412 .respond(std::future::ready(Err(LibsyError::client_call(
1413 "judge",
1414 LlmClientError::General("private provider body".into()),
1415 ))))
1416 .await;
1417 }
1418 if name == "dropped reply" {
1419 drop(call);
1420 return Ok(());
1421 }
1422 let value = if name == "wrong type" {
1423 DecisionValue::Boolean(BooleanEstimate::Value(true))
1424 } else {
1425 DecisionValue::Choice {
1426 selected: "no_advantage".into(),
1427 probabilities: (name != "no distribution").then(|| {
1428 BTreeMap::from([
1429 ("advantage".into(), Probability(score)),
1430 ("no_advantage".into(), Probability(1.0 - score)),
1431 ])
1432 }),
1433 }
1434 };
1435 let answers = if name == "missing answer" {
1436 BTreeMap::new()
1437 } else {
1438 BTreeMap::from([(
1439 "route".into(),
1440 DecisionAnswer {
1441 value,
1442 provider_confidence: Some(ProviderConfidence(0.99)),
1443 },
1444 )])
1445 };
1446 call.respond(std::future::ready(Ok(DecisionResponse {
1447 id: None,
1448 model: Some("provider-judge".into()),
1449 answers,
1450 usage: Default::default(),
1451 })))
1452 .await
1453 }
1454 };
1455 let models = Arc::new(RuntimeModels::new(runtime_models()));
1456 let result = drive(router.clone(), request.clone(), models.clone(), &serve).await;
1457 assert_eq!(
1458 calls.load(Ordering::SeqCst),
1459 usize::from(name != "missing candidate"),
1460 "{name}"
1461 );
1462 let Some(expected) = expected else {
1463 if name == "missing candidate" {
1464 assert!(
1465 matches!(result, Err(LibsyError::AlgorithmError { message }) if message.contains("candidate is missing"))
1466 );
1467 } else {
1468 assert!(
1469 matches!(result, Err(LibsyError::ClientCall { .. })),
1470 "{name}"
1471 );
1472 }
1473 continue;
1474 };
1475 let outcome = result?;
1476 assert_eq!(outcome.selected_model_id()?, expected, "{name}");
1477 assert!(outcome.response.is_none());
1478 assert_eq!(outcome.request.llm_request.messages, original_messages);
1479 let evidence = outcome
1480 .metadata
1481 .and_then(|metadata| metadata.evidence)
1482 .expect("routing evidence");
1483 if matches!(name, "above" | "equal" | "below") {
1484 assert_eq!(evidence["source"], "decision_classifier");
1485 assert_eq!(evidence["verdict"], "relative_advantage");
1486 assert_eq!(evidence["threshold"], settings.cutoff);
1487 assert_eq!(evidence["score"], score);
1488 } else {
1489 let reason = match name {
1490 "provider error" => "client_error",
1491 "dropped reply" => "call_error",
1492 _ => "invalid_verdict",
1493 };
1494 assert_eq!(
1495 evidence,
1496 json!({"source": "fail_open", "reason_code": reason})
1497 );
1498 }
1499 let retained = drive(router, request.clone(), models, &serve).await?;
1500 assert_eq!(retained.selected_model_id()?, expected);
1501 assert_eq!(
1502 calls.load(Ordering::SeqCst),
1503 1,
1504 "affinity should skip the judge: {name}"
1505 );
1506 }
1507 for cutoff in [-0.1, 1.1, f64::NAN] {
1508 let mut config = settings.clone();
1509 config.cutoff = cutoff;
1510 assert!(
1511 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1512 config: TaskClassifierConfig {
1513 judge: CapabilityJudgeConfig::Decision(config),
1514 ..TaskClassifierConfig::default()
1515 },
1516 })
1517 .is_err()
1518 );
1519 }
1520 let mut duplicate = settings.clone();
1521 duplicate
1522 .candidates
1523 .insert("duplicate-a".into(), "capable".into());
1524 assert!(duplicate.validate().is_err());
1525 Ok(())
1526 }
1527
1528 #[test]
1529 fn the_threshold_boundary_is_inclusive() -> Result<()> {
1530 let policy = policy();
1531 let at_threshold = verdict(0.5, "supported", "SUP-1");
1532 let below_threshold = verdict(0.49, "supported", "SUP-1");
1533 assert_eq!(selected(&policy, Some(&at_threshold))?, "efficient");
1534 assert_eq!(selected(&policy, Some(&below_threshold))?, "capable");
1535 Ok(())
1536 }
1537
1538 #[test]
1539 fn the_threshold_moves_the_routing_boundary() -> Result<()> {
1540 let borderline = verdict(0.5, "supported", "SUP-1");
1541 let strict = TaskClassifierPolicy::new(&llm_config(0.9));
1542 let lenient = TaskClassifierPolicy::new(&llm_config(0.1));
1543 assert_eq!(selected(&strict, Some(&borderline))?, "capable");
1544 assert_eq!(selected(&lenient, Some(&borderline))?, "efficient");
1545 Ok(())
1546 }
1547
1548 #[test]
1549 fn classifier_config_rejects_unknown_fields() {
1550 let error = serde_json::from_value::<TaskClassifierConfig>(serde_json::json!({
1551 "base_threshold": 0.5,
1552 "classifier_magic": true,
1553 }))
1554 .expect_err("unknown classifier fields must be rejected");
1555
1556 assert!(
1557 error
1558 .to_string()
1559 .contains("unknown field `classifier_magic`"),
1560 "{error}"
1561 );
1562 }
1563
1564 #[test]
1565 fn invalid_classifier_config_is_rejected() -> Result<()> {
1566 for bad in [1.5, -0.1, f64::NAN, f64::INFINITY] {
1567 assert!(
1568 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1569 config: test_config(bad),
1570 })
1571 .is_err(),
1572 "base threshold {bad} should be rejected"
1573 );
1574 }
1575 for config in [
1576 TaskClassifierConfig {
1577 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1578 base_threshold: 0.5,
1579 threshold_step: -0.1,
1580 ..LlmCapabilityConfig::default()
1581 }),
1582 ..TaskClassifierConfig::default()
1583 },
1584 TaskClassifierConfig {
1585 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1586 base_threshold: 0.8,
1587 threshold_step: 0.11,
1588 ..LlmCapabilityConfig::default()
1589 }),
1590 ..TaskClassifierConfig::default()
1591 },
1592 TaskClassifierConfig {
1593 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1594 base_threshold: 0.5,
1595 ..LlmCapabilityConfig::default()
1596 }),
1597 message_hash_fallback: true,
1598 ..TaskClassifierConfig::default()
1599 },
1600 TaskClassifierConfig {
1601 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1602 base_threshold: 0.5,
1603 max_output_tokens: 0,
1604 ..LlmCapabilityConfig::default()
1605 }),
1606 ..TaskClassifierConfig::default()
1607 },
1608 ] {
1609 assert!(LlmTaskClassifier::new(LlmClassifierConfig::Capability { config }).is_err());
1610 }
1611 for base_threshold in [0.0, 1.0] {
1612 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1613 config: test_config(base_threshold),
1614 })?;
1615 }
1616 Ok(())
1617 }
1618
1619 #[test]
1620 fn message_hash_fallback_accepts_retaining_triggers() -> Result<()> {
1621 for trigger in [ClassifyTrigger::NewSession, ClassifyTrigger::UserTurn] {
1625 let config = TaskClassifierConfig {
1626 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1627 base_threshold: 0.5,
1628 ..LlmCapabilityConfig::default()
1629 }),
1630 classify_trigger: trigger,
1631 message_hash_fallback: true,
1632 ..TaskClassifierConfig::default()
1633 };
1634 LlmTaskClassifier::new(LlmClassifierConfig::Capability { config }).map_err(
1635 |error| LibsyError::AlgorithmError {
1636 message: format!("{trigger:?} with message_hash_fallback rejected: {error}"),
1637 },
1638 )?;
1639 }
1640 let every_request = TaskClassifierConfig {
1641 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
1642 base_threshold: 0.5,
1643 ..LlmCapabilityConfig::default()
1644 }),
1645 classify_trigger: ClassifyTrigger::EveryRequest,
1646 message_hash_fallback: true,
1647 ..TaskClassifierConfig::default()
1648 };
1649 assert!(
1650 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1651 config: every_request
1652 })
1653 .is_err(),
1654 "every_request with message_hash_fallback should stay rejected"
1655 );
1656 Ok(())
1657 }
1658
1659 #[test]
1660 fn an_unusable_verdict_is_ambiguous() -> Result<()> {
1661 let policy = policy();
1662 let inconsistent_rule = TaskClassifierVerdict {
1663 capability_boundary: "uncertain".to_string(),
1664 ..verdict(1.0, "supported", "SUP-1")
1665 };
1666 let empty_crux = TaskClassifierVerdict {
1667 crux: " ".to_string(),
1668 ..verdict(1.0, "supported", "SUP-1")
1669 };
1670 let unusable = [
1671 Some(verdict(1.1, "supported", "SUP-1")),
1672 Some(inconsistent_rule),
1673 Some(empty_crux),
1674 None,
1675 ];
1676 for verdict in unusable {
1677 let classification = policy.to_classification(verdict.as_ref(), &policy_driver())?;
1678 assert!(matches!(classification, Classification::Ambiguous(_)));
1679 assert!(classification.argmax(false)?.is_none());
1680 assert!(classification.argmax(true)?.is_none());
1681 }
1682 Ok(())
1683 }
1684
1685 #[test]
1686 fn capability_boundaries_apply_monotonic_threshold_steps() -> Result<()> {
1687 let policy = TaskClassifierPolicy::new(&LlmCapabilityConfig {
1688 threshold_step: 0.1,
1689 ..llm_config(0.4)
1690 });
1691
1692 assert_eq!(
1693 selected(&policy, Some(&verdict(0.4, "supported", "SUP-2")))?,
1694 "efficient"
1695 );
1696 assert_eq!(
1697 selected(&policy, Some(&verdict(0.49, "uncertain", "UNC-1")))?,
1698 "capable"
1699 );
1700 assert_eq!(
1701 selected(&policy, Some(&verdict(0.5, "uncertain", "UNC-1")))?,
1702 "efficient"
1703 );
1704 assert_eq!(
1705 selected(&policy, Some(&verdict(0.5, "unmatched", "none")))?,
1706 "efficient"
1707 );
1708 assert_eq!(
1709 selected(&policy, Some(&verdict(0.59, "unsupported", "LIM-1")))?,
1710 "capable"
1711 );
1712 assert_eq!(
1713 selected(&policy, Some(&verdict(0.6, "unsupported", "LIM-1")))?,
1714 "efficient"
1715 );
1716 Ok(())
1717 }
1718
1719 fn capability_judge(recent_turn_window: Option<usize>) -> Result<CapabilityJudge> {
1722 Ok(StructuredJudge::new(
1723 TaskInput { recent_turn_window },
1724 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?,
1725 SerdeDecoder::new(),
1726 JudgeRuntimeConfig::new(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?,
1727 ))
1728 }
1729
1730 fn judged_contents(recent_turn_window: usize) -> Result<Vec<String>> {
1731 let judge = capability_judge(Some(recent_turn_window))?;
1732 let request = Request {
1733 llm_request: LlmRequest {
1734 messages: vec![
1735 Message::text(Role::System, "client instructions"),
1736 Message::text(Role::User, "initial task"),
1737 Message::text(Role::Assistant, "old response"),
1738 Message::text(Role::User, "old follow-up"),
1739 Message::text(Role::Assistant, "recent 1"),
1740 Message::text(Role::User, "recent 2"),
1741 ],
1742 ..LlmRequest::default()
1743 },
1744 raw_request: None,
1745 metadata: None,
1746 };
1747 Ok(judge
1748 .build_request(&State::default(), &request)
1749 .llm_request
1750 .messages
1751 .iter()
1752 .filter_map(|message| message.text_content("\n"))
1753 .collect())
1754 }
1755
1756 #[test]
1757 fn a_window_widens_the_judge_to_the_surrounding_conversation() -> Result<()> {
1758 let contents = judged_contents(2)?;
1760 assert!(contents.contains(&"client instructions".to_string()));
1761 assert!(contents.contains(&"initial task".to_string()));
1762 assert!(contents.contains(&"recent 1".to_string()));
1763 assert!(contents.contains(&"recent 2".to_string()));
1764 assert!(!contents.contains(&"old response".to_string()));
1765 Ok(())
1766 }
1767
1768 #[test]
1769 fn a_zero_window_keeps_only_the_instructions_and_the_task() -> Result<()> {
1770 let contents = judged_contents(0)?;
1771 assert!(contents.contains(&"client instructions".to_string()));
1772 assert!(contents.contains(&"initial task".to_string()));
1773 assert!(!contents.contains(&"recent 2".to_string()));
1774 Ok(())
1775 }
1776
1777 fn tool_call(id: &str) -> Message {
1778 Message {
1779 role: Role::Assistant,
1780 content: vec![ContentBlock::ToolCall(ToolCall {
1781 id: id.to_string(),
1782 name: "search".to_string(),
1783 arguments: Value::Null,
1784 })],
1785 }
1786 }
1787
1788 fn tool_result(id: &str) -> Message {
1789 Message {
1790 role: Role::Tool,
1791 content: vec![ContentBlock::ToolResult(ToolResult {
1792 tool_call_id: id.to_string(),
1793 content: vec![ContentBlock::Text {
1794 text: "tool output".to_string(),
1795 }],
1796 is_error: None,
1797 })],
1798 }
1799 }
1800
1801 #[test]
1802 fn default_task_input_keeps_user_content_around_tool_results() {
1803 let mut result = tool_result("call-1");
1804 result.role = Role::User;
1805 let mut mixed = result.clone();
1806 mixed.content.push(ContentBlock::Text {
1807 text: "latest follow-up".to_string(),
1808 });
1809 let input = TaskInput {
1810 recent_turn_window: None,
1811 };
1812 let mut request = Request {
1813 llm_request: LlmRequest {
1814 messages: vec![
1815 result.clone(),
1816 Message::text(Role::User, "initial task"),
1817 tool_call("call-1"),
1818 mixed,
1819 result.clone(),
1820 ],
1821 ..LlmRequest::default()
1822 },
1823 ..Request::default()
1824 };
1825 assert_eq!(
1826 input.build_messages(&State::default(), &request),
1827 vec![
1828 Message::text(Role::User, "initial task"),
1829 Message::text(Role::User, "latest follow-up"),
1830 ]
1831 );
1832 request.llm_request.messages = vec![result];
1833 assert!(input.build_messages(&State::default(), &request).is_empty());
1834 }
1835
1836 #[test]
1839 fn trimming_keeps_the_call_that_introduced_a_kept_tool_result() {
1840 let messages = vec![
1841 Message::text(Role::System, "client instructions"),
1842 Message::text(Role::User, "initial task"),
1843 Message::text(Role::Assistant, "old response"),
1844 tool_call("call-1"),
1845 tool_result("call-1"),
1846 Message::text(Role::Assistant, "recent 1"),
1847 Message::text(Role::User, "recent 2"),
1848 Message::text(Role::Assistant, "recent 3"),
1849 Message::text(Role::User, "recent 4"),
1850 ];
1851
1852 let kept = trim_messages(&messages, 5);
1854
1855 assert_eq!(
1856 kept,
1857 vec![
1858 Message::text(Role::System, "client instructions"),
1859 Message::text(Role::User, "initial task"),
1860 tool_call("call-1"),
1861 tool_result("call-1"),
1862 Message::text(Role::Assistant, "recent 1"),
1863 Message::text(Role::User, "recent 2"),
1864 Message::text(Role::Assistant, "recent 3"),
1865 Message::text(Role::User, "recent 4"),
1866 ]
1867 );
1868 }
1869
1870 #[test]
1873 fn trimming_pairs_a_repeated_id_with_the_call_that_precedes_it() {
1874 let messages = vec![
1875 Message::text(Role::System, "client instructions"),
1876 Message::text(Role::User, "initial task"),
1877 tool_call("x"),
1878 tool_result("x"),
1879 Message::text(Role::Assistant, "later"),
1880 tool_call("x"),
1881 tool_result("x"),
1882 ];
1883
1884 let kept = trim_messages(&messages, 4);
1886
1887 assert_eq!(
1888 kept,
1889 vec![
1890 Message::text(Role::System, "client instructions"),
1891 Message::text(Role::User, "initial task"),
1892 tool_call("x"),
1893 tool_result("x"),
1894 Message::text(Role::Assistant, "later"),
1895 tool_call("x"),
1896 tool_result("x"),
1897 ]
1898 );
1899 }
1900
1901 #[test]
1904 fn trimming_keeps_the_counted_window_when_a_result_cannot_be_paired() {
1905 let messages = vec![
1906 Message::text(Role::System, "client instructions"),
1907 tool_call("orphan"),
1908 Message::text(Role::User, "initial task"),
1909 Message::text(Role::Assistant, "old response"),
1910 tool_result("orphan"),
1911 Message::text(Role::Assistant, "recent 1"),
1912 Message::text(Role::User, "recent 2"),
1913 ];
1914
1915 let kept = trim_messages(&messages, 3);
1916
1917 assert_eq!(
1918 kept,
1919 vec![
1920 Message::text(Role::System, "client instructions"),
1921 Message::text(Role::User, "initial task"),
1922 tool_result("orphan"),
1923 Message::text(Role::Assistant, "recent 1"),
1924 Message::text(Role::User, "recent 2"),
1925 ]
1926 );
1927 }
1928
1929 #[test]
1930 fn a_window_restates_the_routing_instruction_last() -> Result<()> {
1931 let contents = judged_contents(2)?;
1932
1933 assert_eq!(
1936 contents.last().map(String::as_str),
1937 Some(TRAILING_ROUTING_INSTRUCTION)
1938 );
1939 Ok(())
1940 }
1941
1942 #[test]
1947 fn a_window_drops_reasoning_but_keeps_visible_text_and_tool_pairs() {
1948 let messages = vec![
1949 Message::text(Role::System, "client instructions"),
1950 Message::text(Role::User, "initial task"),
1951 Message {
1952 role: Role::Assistant,
1953 content: vec![
1954 ContentBlock::Reasoning {
1955 text: "private chain of thought".to_string(),
1956 signature: None,
1957 details: Vec::new(),
1958 },
1959 ContentBlock::Text {
1960 text: "visible answer".to_string(),
1961 },
1962 ],
1963 },
1964 tool_call("call-1"),
1965 tool_result("call-1"),
1966 Message {
1967 role: Role::Assistant,
1968 content: vec![ContentBlock::Reasoning {
1969 text: "reasoning-only turn".to_string(),
1970 signature: None,
1971 details: Vec::new(),
1972 }],
1973 },
1974 Message::text(Role::User, "follow-up"),
1975 ];
1976 let request = Request {
1977 llm_request: LlmRequest {
1978 messages,
1979 ..LlmRequest::default()
1980 },
1981 raw_request: None,
1982 metadata: None,
1983 };
1984
1985 let built = TaskInput {
1986 recent_turn_window: Some(10),
1987 }
1988 .build_messages(&State::default(), &request);
1989
1990 assert!(
1991 !built
1992 .iter()
1993 .flat_map(|message| &message.content)
1994 .any(|block| matches!(block, ContentBlock::Reasoning { .. })),
1995 "{built:?}"
1996 );
1997 assert!(
1998 built
1999 .iter()
2000 .any(|message| message.text_content("\n").as_deref() == Some("visible answer"))
2001 );
2002 assert!(built.contains(&tool_call("call-1")));
2003 assert!(built.contains(&tool_result("call-1")));
2004 assert_eq!(built.len(), 7);
2007 assert!(built.iter().all(|message| !message.content.is_empty()));
2008 }
2009
2010 #[test]
2011 fn the_default_path_is_left_unchanged() -> Result<()> {
2012 let judge = capability_judge(None)?;
2015 let request = Request {
2016 llm_request: LlmRequest {
2017 messages: vec![Message::text(Role::User, "the task")],
2018 ..LlmRequest::default()
2019 },
2020 raw_request: None,
2021 metadata: None,
2022 };
2023
2024 let built = judge.build_request(&State::default(), &request);
2025
2026 assert_eq!(built.llm_request.messages.len(), 1);
2029 assert!(!built.llm_request.instructions.is_empty());
2030 assert!(
2031 !built
2032 .llm_request
2033 .messages
2034 .iter()
2035 .filter_map(|message| message.text_content("\n"))
2036 .any(|text| text.contains(TRAILING_ROUTING_INSTRUCTION))
2037 );
2038 Ok(())
2039 }
2040
2041 #[test]
2042 fn capability_judge_builds_a_structured_request() -> Result<()> {
2043 let judge = capability_judge(None)?;
2044 let request = Request {
2045 llm_request: LlmRequest {
2046 model: Some("inbound".to_string()),
2047 messages: vec![
2048 Message::text(Role::System, "client instructions"),
2049 Message::text(Role::Developer, "client developer instructions"),
2050 Message::text(Role::User, "initial task"),
2051 Message::text(Role::Assistant, "old response"),
2052 Message::text(Role::User, "old follow-up"),
2053 Message::text(Role::Assistant, "recent 1"),
2054 Message::text(Role::User, "recent 2"),
2055 Message::text(Role::Assistant, "recent 3"),
2056 Message::text(Role::User, "recent 4"),
2057 Message::text(Role::Assistant, "recent 5"),
2058 ],
2059 ..LlmRequest::default()
2060 },
2061 raw_request: None,
2062 metadata: None,
2063 };
2064 let judge_request = judge.build_request(&State::default(), &request);
2065
2066 assert_eq!(judge_request.llm_request.model, request.llm_request.model);
2067 assert_eq!(judge_request.llm_request.instructions.len(), 1);
2068 assert_eq!(judge_request.llm_request.instructions[0].role, Role::System);
2069 assert_eq!(
2070 judge_request.llm_request.instructions[0].content,
2071 InstructionBlock {
2072 role: Role::System,
2073 content: Message::text(Role::System, judge.contract().system_prompt()).content,
2074 }
2075 .content,
2076 );
2077 assert_eq!(judge_request.llm_request.messages.len(), 2);
2078 let contents = judge_request
2079 .llm_request
2080 .messages
2081 .iter()
2082 .filter_map(|message| message.text_content("\n"))
2083 .collect::<Vec<_>>();
2084 assert!(contents.contains(&"recent 4".to_string()));
2085 assert!(contents.contains(&"initial task".to_string()));
2086 assert!(!contents.contains(&"recent 5".to_string()));
2087 assert!(!contents.contains(&"client instructions".to_string()));
2088 assert_eq!(
2089 judge_request.llm_request.output.response_format,
2090 Some(judge.contract().response_format().clone())
2091 );
2092 assert_eq!(
2093 judge_request.llm_request.output.max_output_tokens,
2094 Some(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)
2095 );
2096 Ok(())
2097 }
2098
2099 fn sample_value(spec: &Value) -> Value {
2100 if let Some(first) = spec
2101 .get("enum")
2102 .and_then(Value::as_array)
2103 .and_then(|values| values.first())
2104 {
2105 return first.clone();
2106 }
2107 match spec.get("type").and_then(Value::as_str) {
2108 Some("number") => serde_json::json!(0.5),
2109 Some("boolean") => serde_json::json!(false),
2110 _ => serde_json::json!("sample"),
2111 }
2112 }
2113
2114 fn schema_shaped_verdict(schema: &Value) -> Result<String> {
2115 let properties = schema
2116 .pointer("/json_schema/schema/properties")
2117 .and_then(Value::as_object)
2118 .ok_or_else(|| LibsyError::AlgorithmError {
2119 message: "packaged schema declares no properties".to_string(),
2120 })?;
2121 Ok(Value::Object(
2122 properties
2123 .iter()
2124 .map(|(name, spec)| (name.clone(), sample_value(spec)))
2125 .collect(),
2126 )
2127 .to_string())
2128 }
2129
2130 #[test]
2133 fn every_schema_property_round_trips_through_the_judge_parser() -> Result<()> {
2134 let contract =
2135 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?;
2136 let schema = contract.response_format();
2137 let reply = schema_shaped_verdict(schema)?;
2138 let judge: CapabilityJudge = StructuredJudge::new(
2139 TaskInput {
2140 recent_turn_window: None,
2141 },
2142 contract,
2143 SerdeDecoder::new(),
2144 JudgeRuntimeConfig::new(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?,
2145 );
2146
2147 let verdict = judge.parse(&text_response(None, reply))?;
2148
2149 assert!(verdict.is_valid());
2150 assert!((0.0..=1.0).contains(&verdict.p_solve));
2151 Ok(())
2152 }
2153
2154 #[test]
2155 fn packaged_prompt_keeps_the_schema_in_the_structured_request() -> Result<()> {
2156 let contract =
2157 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?;
2158 let prompt = contract.system_prompt();
2159 let schema_name = contract
2160 .response_format()
2161 .pointer("/json_schema/name")
2162 .and_then(Value::as_str)
2163 .ok_or_else(|| LibsyError::AlgorithmError {
2164 message: "packaged response schema has no name".to_string(),
2165 })?;
2166 assert_eq!(schema_name, "CapabilityClassifierDecision");
2167 assert!(prompt.contains("SUP-1 [supported]"));
2168 assert!(prompt.contains("SUP-5 [supported]"));
2169 assert!(!prompt.contains("{{RESPONSE_SCHEMA}}"));
2170 assert!(!prompt.contains("\"type\": \"object\""));
2171 assert!(!prompt.contains("\"json_schema\""));
2172 assert!(!prompt.contains(schema_name));
2173 let rule_values = contract
2174 .response_format()
2175 .pointer("/json_schema/schema/properties/primary_rule/enum")
2176 .and_then(Value::as_array)
2177 .ok_or_else(|| LibsyError::AlgorithmError {
2178 message: "rendered response schema has no primary rule enum".to_string(),
2179 })?;
2180 assert!(
2181 rule_values
2182 .iter()
2183 .any(|value| value.as_str() == Some("SUP-1"))
2184 );
2185 assert!(
2186 rule_values
2187 .iter()
2188 .any(|value| value.as_str() == Some("none"))
2189 );
2190 Ok(())
2191 }
2192}