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