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: &TaskClassifierConfig) -> 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)]
309pub struct TaskClassifierConfig {
311 pub fail_open: bool,
313 pub base_threshold: f64,
315 pub threshold_step: f64,
320 pub classify_trigger: ClassifyTrigger,
322 pub message_hash_fallback: bool,
324 pub recent_turn_window: Option<usize>,
331 pub contract: ClassifierContractConfig,
333 pub max_output_tokens: u64,
335}
336
337#[derive(Deserialize)]
339#[serde(deny_unknown_fields)]
340struct TaskClassifierConfigWire {
341 base_threshold: f64,
342 #[serde(default = "default_fail_open")]
343 fail_open: bool,
344 #[serde(default)]
345 threshold_step: f64,
346 #[serde(default)]
347 classify_trigger: ClassifyTrigger,
348 #[serde(default)]
349 message_hash_fallback: bool,
350 #[serde(default)]
351 recent_turn_window: Option<usize>,
352 #[serde(default)]
353 prompt: Option<String>,
354 #[serde(default)]
355 response_format_type: ClassifierResponseFormat,
356 #[serde(default = "default_judge_max_output_tokens")]
357 max_output_tokens: u64,
358}
359
360impl<'de> Deserialize<'de> for TaskClassifierConfig {
361 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
362 where
363 D: Deserializer<'de>,
364 {
365 let wire = TaskClassifierConfigWire::deserialize(deserializer)?;
366 let mut contract = ClassifierContractConfig::default();
367 if let Some(prompt) = wire.prompt {
368 contract = contract.with_prompt(prompt);
369 }
370 contract = contract.with_response_format_type(wire.response_format_type);
371 Ok(Self {
372 base_threshold: wire.base_threshold,
373 fail_open: wire.fail_open,
374 threshold_step: wire.threshold_step,
375 classify_trigger: wire.classify_trigger,
376 message_hash_fallback: wire.message_hash_fallback,
377 recent_turn_window: wire.recent_turn_window,
378 contract,
379 max_output_tokens: wire.max_output_tokens,
380 })
381 }
382}
383
384const fn default_fail_open() -> bool {
385 true
386}
387
388const fn default_judge_max_output_tokens() -> u64 {
389 DEFAULT_JUDGE_MAX_OUTPUT_TOKENS
390}
391
392impl Default for TaskClassifierConfig {
393 fn default() -> Self {
394 Self {
395 base_threshold: 0.0,
396 fail_open: default_fail_open(),
397 threshold_step: 0.0,
398 classify_trigger: ClassifyTrigger::default(),
399 message_hash_fallback: false,
400 recent_turn_window: None,
401 contract: ClassifierContractConfig::default(),
402 max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
403 }
404 }
405}
406
407impl TaskClassifierConfig {
408 fn validate(&self) -> Result<()> {
410 if !(0.0..=1.0).contains(&self.base_threshold) {
411 return Err(LibsyError::AlgorithmError {
412 message: format!(
413 "base_threshold must be between 0 and 1, got {}",
414 self.base_threshold
415 ),
416 });
417 }
418 if !self.threshold_step.is_finite() || self.threshold_step < 0.0 {
419 return Err(LibsyError::AlgorithmError {
420 message: format!(
421 "threshold_step must be finite and greater than or equal to 0, got {}",
422 self.threshold_step
423 ),
424 });
425 }
426 let unsupported_threshold = self.base_threshold + 2.0 * self.threshold_step;
427 if unsupported_threshold > 1.0 && unsupported_threshold - 1.0 > f64::EPSILON {
428 return Err(LibsyError::AlgorithmError {
429 message: format!(
430 "base_threshold + 2 * threshold_step must be at most 1, got {unsupported_threshold}"
431 ),
432 });
433 }
434 if self.max_output_tokens == 0 {
435 return Err(LibsyError::AlgorithmError {
436 message: "max_output_tokens must be at least 1".to_string(),
437 });
438 }
439 if self.message_hash_fallback && self.classify_trigger == ClassifyTrigger::EveryRequest {
441 return Err(LibsyError::AlgorithmError {
442 message:
443 "message_hash_fallback requires classify_trigger = new_session or user_turn"
444 .to_string(),
445 });
446 }
447 Ok(())
448 }
449}
450
451#[derive(Clone, Debug)]
453pub enum CustomClassifierPolicy {
454 TargetSelector {
456 selector: String,
458 },
459}
460
461impl CustomClassifierPolicy {
462 pub fn target_selector(selector: impl Into<String>) -> Self {
464 Self::TargetSelector {
465 selector: selector.into(),
466 }
467 }
468}
469
470#[derive(Clone, Debug)]
472pub struct CustomClassifierConfig {
473 pub prompt: String,
475 pub response_schema: Value,
477 pub policy: CustomClassifierPolicy,
479 pub classify_trigger: ClassifyTrigger,
481 pub message_hash_fallback: bool,
483 pub recent_turn_window: Option<usize>,
485 pub max_output_tokens: u64,
487}
488
489impl CustomClassifierConfig {
490 pub fn new(
492 prompt: impl Into<String>,
493 response_schema: Value,
494 policy: CustomClassifierPolicy,
495 ) -> Self {
496 Self {
497 prompt: prompt.into(),
498 response_schema,
499 policy,
500 classify_trigger: ClassifyTrigger::default(),
501 message_hash_fallback: false,
502 recent_turn_window: None,
503 max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
504 }
505 }
506
507 fn validate(&self) -> Result<()> {
508 if self.max_output_tokens == 0 {
509 return Err(LibsyError::AlgorithmError {
510 message: "max_output_tokens must be at least 1".to_string(),
511 });
512 }
513 if self.message_hash_fallback && self.classify_trigger == ClassifyTrigger::EveryRequest {
515 return Err(LibsyError::AlgorithmError {
516 message:
517 "message_hash_fallback requires classify_trigger = new_session or user_turn"
518 .to_string(),
519 });
520 }
521 Ok(())
522 }
523}
524
525enum CustomPolicyRuntime {
526 TargetSelector(TargetSelectorPolicy),
527}
528
529impl JudgePolicy for CustomPolicyRuntime {
530 type Verdict = Value;
531
532 fn to_classification(
533 &self,
534 verdict: Option<&Self::Verdict>,
535 driver: &Driver,
536 ) -> Result<Classification> {
537 match self {
538 Self::TargetSelector(policy) => policy.to_classification(verdict, driver),
539 }
540 }
541}
542
543fn affinity_router(
545 trigger: ClassifyTrigger,
546 message_hash_fallback: bool,
547) -> Option<Arc<AffinityRouter>> {
548 let router = match trigger {
549 ClassifyTrigger::EveryRequest => return None,
550 ClassifyTrigger::NewSession => AffinityRouter::new(),
551 ClassifyTrigger::UserTurn => AffinityRouter::new().with_release_on_user_turn(),
552 };
553 let router = if message_hash_fallback {
554 router.with_message_hash_fallback()
555 } else {
556 router
557 };
558 Some(Arc::new(router))
559}
560
561pub struct LlmTaskClassifier {
563 route: FallThrough<State>,
564 inner: Arc<dyn Classifier<State>>,
566}
567
568struct ClassifierRouteConfig {
569 default_target: Category,
570 classify_trigger: ClassifyTrigger,
571 message_hash_fallback: bool,
572}
573
574pub struct DefaultCategoryClassifier(pub Category);
577
578#[async_trait]
579impl<S: Send> Classifier<S> for DefaultCategoryClassifier {
580 async fn score(
581 &self,
582 _state: &mut S,
583 _request: &mut Request,
584 driver: &Driver,
585 ) -> Result<(Classification, Option<Response>)> {
586 let target = driver.first_model_for(&self.0)?;
587 driver.set_evidence_if_empty(serde_json::json!({"source": "fall_open"}));
588 Ok((
589 Classification::Scores(vec![Score {
590 target: target.clone(),
591 confidence: 0.0,
592 category: Some(self.0.clone()),
593 }]),
594 None,
595 ))
596 }
597}
598
599#[derive(Clone)]
601#[non_exhaustive]
602pub enum LlmClassifierConfig {
603 Capability {
605 config: TaskClassifierConfig,
607 },
608 Escalation {
610 contract: ClassifierContractConfig,
612 config: EscalationJudgeConfig,
614 max_output_tokens: u64,
616 },
617 Custom {
619 default_target: Category,
621 config: CustomClassifierConfig,
623 },
624}
625
626impl LlmTaskClassifier {
627 pub fn new(config: LlmClassifierConfig) -> Result<Self> {
634 match config {
635 LlmClassifierConfig::Capability { config } => Self::build_capability(config),
636 LlmClassifierConfig::Escalation {
637 contract,
638 config,
639 max_output_tokens,
640 } => Self::build_escalation(contract, config, max_output_tokens),
641 LlmClassifierConfig::Custom {
642 default_target,
643 config,
644 } => Self::build_custom(default_target, config),
645 }
646 }
647
648 fn build_capability(config: TaskClassifierConfig) -> Result<Self> {
649 config.validate()?;
650 let contract = Self::load_capability_contract(&config.contract)?;
651 let classify_trigger = config.classify_trigger;
652 let message_hash_fallback = config.message_hash_fallback;
653 let classifier: Arc<dyn Classifier<State>> = Arc::new(
654 JudgeClassifier::new(
655 StructuredJudge::new(
656 TaskInput {
657 recent_turn_window: config.recent_turn_window,
658 },
659 contract,
660 SerdeDecoder::new(),
661 JudgeRuntimeConfig::new(config.max_output_tokens)?,
662 ),
663 TaskClassifierPolicy::new(&config),
664 )
665 .with_error_recovery(config.fail_open)
666 .with_evidence(capability_evidence),
667 );
668 Self::from_classifier(
669 classifier,
670 ClassifierRouteConfig {
671 default_target: Category::Capable,
672 classify_trigger,
673 message_hash_fallback,
674 },
675 )
676 }
677
678 fn build_custom(default_target: Category, config: CustomClassifierConfig) -> Result<Self> {
679 config.validate()?;
680 let CustomClassifierConfig {
681 prompt,
682 response_schema,
683 policy,
684 classify_trigger,
685 message_hash_fallback,
686 recent_turn_window,
687 max_output_tokens,
688 } = config;
689 let contract = ClassifierContract::from_inner_schema(&prompt, response_schema)?;
690 let policy = match policy {
691 CustomClassifierPolicy::TargetSelector { selector } => {
692 CustomPolicyRuntime::TargetSelector(TargetSelectorPolicy::new(selector)?)
693 }
694 };
695 let classifier: Arc<dyn Classifier<State>> = Arc::new(JudgeClassifier::new(
696 StructuredJudge::new(
697 TaskInput { recent_turn_window },
698 contract,
699 JsonSchemaDecoder::new(),
700 JudgeRuntimeConfig::new(max_output_tokens)?,
701 ),
702 policy,
703 ));
704
705 Self::from_classifier(
706 classifier,
707 ClassifierRouteConfig {
708 default_target,
709 classify_trigger,
710 message_hash_fallback,
711 },
712 )
713 }
714
715 fn build_escalation(
716 contract_config: ClassifierContractConfig,
717 config: EscalationJudgeConfig,
718 max_output_tokens: u64,
719 ) -> Result<Self> {
720 let inner = escalation::build_classifier(contract_config, config, max_output_tokens)?;
721 Ok(Self {
722 route: FallThrough::<State>::new_with_state()
723 .with_name(ALGORITHM_NAME)
724 .with_classifier(Arc::clone(&inner)),
725 inner,
726 })
727 }
728
729 fn load_capability_contract(config: &ClassifierContractConfig) -> Result<ClassifierContract> {
731 ClassifierContract::from_config(config, PROMPT_TEMPLATE, SCHEMA_TEMPLATE)
732 }
733
734 fn from_classifier(
736 inner: Arc<dyn Classifier<State>>,
737 config: ClassifierRouteConfig,
738 ) -> Result<Self> {
739 if config.message_hash_fallback && config.classify_trigger == ClassifyTrigger::EveryRequest
741 {
742 return Err(LibsyError::AlgorithmError {
743 message:
744 "message_hash_fallback requires classify_trigger = new_session or user_turn"
745 .to_string(),
746 });
747 }
748 let mut route = FallThrough::<State>::new_with_state().with_name(ALGORITHM_NAME);
750 if let Some(affinity) =
751 affinity_router(config.classify_trigger, config.message_hash_fallback).as_ref()
752 {
753 route = route
755 .with_processor(affinity.clone())
756 .with_classifier(affinity.clone());
757 }
758 let fallback = DefaultCategoryClassifier(config.default_target);
759 Ok(Self {
760 route: route
761 .with_classifier(inner.clone())
762 .with_classifier(Arc::new(fallback)),
763 inner,
764 })
765 }
766}
767
768#[async_trait]
769impl Classifier<State> for LlmTaskClassifier {
770 async fn score(
771 &self,
772 state: &mut State,
773 request: &mut Request,
774 driver: &Driver,
775 ) -> Result<(Classification, Option<Response>)> {
776 self.inner.score(state, request, driver).await
777 }
778}
779
780#[async_trait]
781impl Algorithm for LlmTaskClassifier {
782 fn name(&self) -> &str {
783 "llm_task_classifier"
784 }
785
786 async fn route(
787 self: Arc<Self>,
788 driver: Driver,
789 request: Request,
790 ) -> Result<crate::RoutingOutcome> {
791 self.route.execute(driver, request).await
792 }
793}
794
795#[cfg(test)]
796mod tests {
797 use std::collections::HashMap;
798 use std::sync::Arc;
799
800 use parking_lot::Mutex;
801 use serde_json::Value;
802
803 use super::*;
804 use switchyard_protocol::{
805 ContentBlock, InstructionBlock, LlmClientError, LlmRequest, Metadata, ModelId, ToolCall,
806 ToolResult, completion_text, text_request, text_response,
807 };
808
809 use crate::algorithms::util::llm_judge::Judge;
810 use crate::core::testing::{Serve, test_drive_with_models};
811 use switchyard_protocol::{LlmResponse, Response};
812
813 const TEST_THRESHOLD: f64 = 0.5;
814
815 type CapabilityJudge = StructuredJudge<TaskInput, SerdeDecoder<TaskClassifierVerdict>>;
816
817 fn test_config(base_threshold: f64) -> TaskClassifierConfig {
818 TaskClassifierConfig {
819 base_threshold,
820 ..TaskClassifierConfig::default()
821 }
822 }
823
824 fn policy() -> TaskClassifierPolicy {
825 TaskClassifierPolicy::new(&test_config(TEST_THRESHOLD))
826 }
827
828 fn runtime_models() -> HashMap<Category, Vec<ModelId>> {
829 [
830 (Category::Judge, vec![ModelId::from("judge")]),
831 (Category::Efficient, vec![ModelId::from("efficient")]),
832 (Category::Capable, vec![ModelId::from("capable")]),
833 (
834 Category::Any,
835 vec![ModelId::from("efficient"), ModelId::from("capable")],
836 ),
837 ]
838 .into()
839 }
840
841 fn policy_driver() -> Driver {
842 Driver::new("test", Arc::new(runtime_models().into())).0
843 }
844
845 fn verdict(
846 p_solve: f64,
847 capability_boundary: &str,
848 primary_rule: &str,
849 ) -> TaskClassifierVerdict {
850 TaskClassifierVerdict {
851 crux: "test crux".to_string(),
852 primary_rule: primary_rule.to_string(),
853 capability_boundary: capability_boundary.to_string(),
854 p_solve,
855 }
856 }
857
858 fn selected(
859 policy: &TaskClassifierPolicy,
860 verdict: Option<&TaskClassifierVerdict>,
861 ) -> Result<ModelId> {
862 policy
863 .to_classification(verdict, &policy_driver())?
864 .argmax(false)?
865 .map(|score| score.target)
866 .ok_or_else(|| LibsyError::AlgorithmError {
867 message: "policy abstained".to_string(),
868 })
869 }
870
871 #[derive(Default)]
874 struct Recorder {
875 calls: Mutex<Vec<String>>,
876 call_roles: Mutex<Vec<(String, bool)>>,
877 judge_max_output_tokens: Mutex<Vec<Option<u64>>>,
878 judge_system_prompts: Mutex<Vec<String>>,
879 }
880
881 impl Recorder {
882 fn calls(&self) -> Vec<String> {
883 self.calls.lock().clone()
884 }
885
886 fn call_roles(&self) -> Vec<(String, bool)> {
887 self.call_roles.lock().clone()
888 }
889
890 fn judge_max_output_tokens(&self) -> Vec<Option<u64>> {
891 self.judge_max_output_tokens.lock().clone()
892 }
893
894 fn judge_system_prompts(&self) -> Vec<String> {
895 self.judge_system_prompts.lock().clone()
896 }
897
898 fn serve(self: &Arc<Self>) -> impl Serve {
899 let recorder = Arc::clone(self);
900 move |model: ModelId, request: Request| {
901 let recorder = Arc::clone(&recorder);
902 async move {
903 let model = model.to_string();
904 recorder.calls.lock().push(model.clone());
905 recorder
906 .call_roles
907 .lock()
908 .push((model.clone(), model != "judge"));
909 let completion = if model == "judge" {
910 recorder
911 .judge_max_output_tokens
912 .lock()
913 .push(request.llm_request.output.max_output_tokens);
914 recorder.judge_system_prompts.lock().extend(
915 request
916 .llm_request
917 .instructions
918 .first()
919 .and_then(|instruction| {
920 instruction.content.iter().find_map(|b| {
921 if let ContentBlock::Text { text } = b {
922 Some(text.clone())
923 } else {
924 None
925 }
926 })
927 }),
928 );
929 r#"{"crux":"bounded task","primary_rule":"SUP-1","capability_boundary":"supported","p_solve":0.9}"#.to_string()
930 } else {
931 format!("answer from {model}")
932 };
933 Ok(Response {
934 llm_response: LlmResponse::Agg(text_response(None, completion)),
935 metadata: request.metadata,
936 upstream_headers: http::HeaderMap::new(),
937 })
938 }
939 }
940 }
941 }
942
943 fn unreachable_judge() -> impl Serve {
945 |model: ModelId, request: Request| async move {
946 let model = model.to_string();
947 if model == "judge" {
948 return Err(LlmClientError::Timeout {
949 source: Box::new(std::io::Error::other("judge unreachable")),
950 });
951 }
952 Ok(Response {
953 llm_response: LlmResponse::Agg(text_response(None, format!("answer from {model}"))),
954 metadata: request.metadata,
955 upstream_headers: http::HeaderMap::new(),
956 })
957 }
958 }
959
960 fn router() -> Result<Arc<LlmTaskClassifier>> {
961 Ok(Arc::new(LlmTaskClassifier::new(
962 LlmClassifierConfig::Capability {
963 config: test_config(TEST_THRESHOLD),
964 },
965 )?))
966 }
967
968 fn classify_request() -> Request {
969 Request {
970 llm_request: text_request(Some("auto".to_string()), "classify this task"),
971 raw_request: None,
972 metadata: None,
973 }
974 }
975
976 fn classify_session_request() -> Request {
977 Request {
978 metadata: Some(Metadata {
979 session_id: Some("session-1".to_string()),
980 ..Metadata::default()
981 }),
982 ..classify_request()
983 }
984 }
985
986 fn classify_follow_up_request() -> Request {
987 let mut request = classify_request();
988 request
989 .llm_request
990 .messages
991 .push(Message::text(Role::Assistant, "I will add the test."));
992 request.llm_request.messages.push(Message::text(
993 Role::User,
994 "Now run the test suite and report the result.",
995 ));
996 request
997 }
998
999 #[tokio::test]
1000 async fn an_unreachable_judge_routes_capable_instead_of_failing_the_request() -> Result<()> {
1001 let router = router()?;
1002
1003 let (selected_model, response) = test_drive_with_models(
1004 router,
1005 classify_request(),
1006 runtime_models(),
1007 unreachable_judge(),
1008 )
1009 .await?;
1010
1011 assert_eq!(selected_model, "capable");
1012 assert_eq!(
1013 response.llm_response.as_agg().map(completion_text),
1014 Some("answer from capable".to_string())
1015 );
1016 Ok(())
1017 }
1018
1019 #[tokio::test]
1020 async fn classifier_judges_each_request_without_affinity() -> Result<()> {
1021 let recorder = Arc::new(Recorder::default());
1022 let router = router()?;
1023 let request = classify_request();
1024 let models = runtime_models();
1025
1026 test_drive_with_models(
1027 router.clone(),
1028 request.clone(),
1029 models.clone(),
1030 recorder.serve(),
1031 )
1032 .await?;
1033 test_drive_with_models(router, request, models, recorder.serve()).await?;
1034
1035 assert_eq!(
1036 recorder.calls(),
1037 vec!["judge", "efficient", "judge", "efficient"]
1038 );
1039 assert_eq!(
1040 recorder.call_roles(),
1041 vec![
1042 ("judge".to_string(), false),
1043 ("efficient".to_string(), true),
1044 ("judge".to_string(), false),
1045 ("efficient".to_string(), true),
1046 ]
1047 );
1048 Ok(())
1049 }
1050
1051 #[tokio::test]
1052 async fn classifier_config_sets_the_judge_completion_cap() -> Result<()> {
1053 let recorder = Arc::new(Recorder::default());
1054 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1055 config: TaskClassifierConfig {
1056 max_output_tokens: 512,
1057 ..test_config(TEST_THRESHOLD)
1058 },
1059 })?);
1060
1061 test_drive_with_models(
1062 router,
1063 classify_request(),
1064 runtime_models(),
1065 recorder.serve(),
1066 )
1067 .await?;
1068
1069 assert_eq!(recorder.judge_max_output_tokens(), vec![Some(512)]);
1070 Ok(())
1071 }
1072
1073 #[tokio::test]
1074 async fn classifier_config_overrides_the_packaged_prompt() -> Result<()> {
1075 let recorder = Arc::new(Recorder::default());
1076 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1077 config: TaskClassifierConfig {
1078 contract: ClassifierContractConfig::default()
1079 .with_prompt("Custom capability rubric."),
1080 ..test_config(TEST_THRESHOLD)
1081 },
1082 })?);
1083
1084 test_drive_with_models(
1085 router,
1086 classify_request(),
1087 runtime_models(),
1088 recorder.serve(),
1089 )
1090 .await?;
1091
1092 let prompts = recorder.judge_system_prompts();
1093 assert_eq!(prompts.len(), 1);
1094 assert_eq!(prompts[0], "Custom capability rubric.");
1095 Ok(())
1096 }
1097
1098 #[tokio::test]
1099 async fn classifier_config_enables_new_session_trigger() -> Result<()> {
1100 let recorder = Arc::new(Recorder::default());
1101 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1102 config: TaskClassifierConfig {
1103 classify_trigger: ClassifyTrigger::NewSession,
1104 ..test_config(TEST_THRESHOLD)
1105 },
1106 })?);
1107
1108 let request = classify_session_request();
1109 let models = runtime_models();
1110 test_drive_with_models(
1111 router.clone(),
1112 request.clone(),
1113 models.clone(),
1114 recorder.serve(),
1115 )
1116 .await?;
1117 test_drive_with_models(router, request, models, recorder.serve()).await?;
1118
1119 assert_eq!(recorder.calls(), vec!["judge", "efficient", "efficient"]);
1120 Ok(())
1121 }
1122
1123 #[tokio::test]
1124 async fn classifier_config_reuses_message_hash_affinity_for_a_follow_up() -> Result<()> {
1125 let recorder = Arc::new(Recorder::default());
1126 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1127 config: TaskClassifierConfig {
1128 classify_trigger: ClassifyTrigger::NewSession,
1129 message_hash_fallback: true,
1130 recent_turn_window: None,
1131 ..test_config(TEST_THRESHOLD)
1132 },
1133 })?);
1134
1135 let models = runtime_models();
1136 test_drive_with_models(
1137 router.clone(),
1138 classify_request(),
1139 models.clone(),
1140 recorder.serve(),
1141 )
1142 .await?;
1143 test_drive_with_models(
1144 router,
1145 classify_follow_up_request(),
1146 models,
1147 recorder.serve(),
1148 )
1149 .await?;
1150
1151 assert_eq!(recorder.calls(), vec!["judge", "efficient", "efficient"]);
1152 Ok(())
1153 }
1154
1155 #[tokio::test]
1156 async fn one_classifier_uses_each_requests_runtime_models() -> Result<()> {
1157 let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1158 config: TaskClassifierConfig {
1159 classify_trigger: ClassifyTrigger::NewSession,
1160 ..test_config(TEST_THRESHOLD)
1161 },
1162 })?);
1163 let calls = Arc::new(Mutex::new(Vec::new()));
1164 let serve = |calls: Arc<Mutex<Vec<String>>>| {
1165 move |model: ModelId, _request: Request| {
1166 let calls = Arc::clone(&calls);
1167 async move {
1168 calls.lock().push(model.to_string());
1169 let text = if model.as_str().starts_with("judge-") {
1170 r#"{"crux":"bounded task","primary_rule":"SUP-1","capability_boundary":"supported","p_solve":0.9}"#.to_string()
1171 } else {
1172 model.to_string()
1173 };
1174 Ok(Response {
1175 llm_response: LlmResponse::Agg(text_response(None, text)),
1176 metadata: None,
1177 upstream_headers: Default::default(),
1178 })
1179 }
1180 }
1181 };
1182 let models = |suffix: &str| -> HashMap<Category, Vec<ModelId>> {
1183 [
1184 (
1185 Category::Judge,
1186 vec![ModelId::from(format!("judge-{suffix}"))],
1187 ),
1188 (
1189 Category::Efficient,
1190 vec![ModelId::from(format!("efficient-{suffix}"))],
1191 ),
1192 (
1193 Category::Capable,
1194 vec![ModelId::from(format!("capable-{suffix}"))],
1195 ),
1196 (
1197 Category::Any,
1198 vec![
1199 ModelId::from(format!("efficient-{suffix}")),
1200 ModelId::from(format!("capable-{suffix}")),
1201 ],
1202 ),
1203 ]
1204 .into()
1205 };
1206 let request = classify_session_request();
1207
1208 let (first, _) = test_drive_with_models(
1209 router.clone(),
1210 request.clone(),
1211 models("a"),
1212 serve(Arc::clone(&calls)),
1213 )
1214 .await?;
1215 let (second, _) =
1216 test_drive_with_models(router, request, models("b"), serve(Arc::clone(&calls))).await?;
1217
1218 assert_eq!(first, "efficient-a");
1219 assert_eq!(second, "efficient-b");
1220 assert_eq!(
1221 &*calls.lock(),
1222 &["judge-a", "efficient-a", "judge-b", "efficient-b"]
1223 );
1224 Ok(())
1225 }
1226
1227 #[test]
1228 fn the_threshold_boundary_is_inclusive() -> Result<()> {
1229 let policy = policy();
1230 let at_threshold = verdict(0.5, "supported", "SUP-1");
1231 let below_threshold = verdict(0.49, "supported", "SUP-1");
1232 assert_eq!(selected(&policy, Some(&at_threshold))?, "efficient");
1233 assert_eq!(selected(&policy, Some(&below_threshold))?, "capable");
1234 Ok(())
1235 }
1236
1237 #[test]
1238 fn the_threshold_moves_the_routing_boundary() -> Result<()> {
1239 let borderline = verdict(0.5, "supported", "SUP-1");
1240 let strict = TaskClassifierPolicy::new(&test_config(0.9));
1241 let lenient = TaskClassifierPolicy::new(&test_config(0.1));
1242 assert_eq!(selected(&strict, Some(&borderline))?, "capable");
1243 assert_eq!(selected(&lenient, Some(&borderline))?, "efficient");
1244 Ok(())
1245 }
1246
1247 #[test]
1248 fn classifier_config_rejects_unknown_fields() {
1249 let error = serde_json::from_value::<TaskClassifierConfig>(serde_json::json!({
1250 "base_threshold": 0.5,
1251 "classifier_magic": true,
1252 }))
1253 .expect_err("unknown classifier fields must be rejected");
1254
1255 assert!(
1256 error
1257 .to_string()
1258 .contains("unknown field `classifier_magic`"),
1259 "{error}"
1260 );
1261 }
1262
1263 #[test]
1264 fn invalid_classifier_config_is_rejected() -> Result<()> {
1265 for bad in [1.5, -0.1, f64::NAN, f64::INFINITY] {
1266 assert!(
1267 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1268 config: test_config(bad),
1269 })
1270 .is_err(),
1271 "base threshold {bad} should be rejected"
1272 );
1273 }
1274 for config in [
1275 TaskClassifierConfig {
1276 base_threshold: 0.5,
1277 threshold_step: -0.1,
1278 ..TaskClassifierConfig::default()
1279 },
1280 TaskClassifierConfig {
1281 base_threshold: 0.8,
1282 threshold_step: 0.11,
1283 ..TaskClassifierConfig::default()
1284 },
1285 TaskClassifierConfig {
1286 base_threshold: 0.5,
1287 message_hash_fallback: true,
1288 ..TaskClassifierConfig::default()
1289 },
1290 TaskClassifierConfig {
1291 base_threshold: 0.5,
1292 max_output_tokens: 0,
1293 ..TaskClassifierConfig::default()
1294 },
1295 ] {
1296 assert!(LlmTaskClassifier::new(LlmClassifierConfig::Capability { config }).is_err());
1297 }
1298 for base_threshold in [0.0, 1.0] {
1299 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1300 config: test_config(base_threshold),
1301 })?;
1302 }
1303 Ok(())
1304 }
1305
1306 #[test]
1307 fn message_hash_fallback_accepts_retaining_triggers() -> Result<()> {
1308 for trigger in [ClassifyTrigger::NewSession, ClassifyTrigger::UserTurn] {
1312 let config = TaskClassifierConfig {
1313 base_threshold: 0.5,
1314 classify_trigger: trigger,
1315 message_hash_fallback: true,
1316 ..TaskClassifierConfig::default()
1317 };
1318 LlmTaskClassifier::new(LlmClassifierConfig::Capability { config }).map_err(
1319 |error| LibsyError::AlgorithmError {
1320 message: format!("{trigger:?} with message_hash_fallback rejected: {error}"),
1321 },
1322 )?;
1323 }
1324 let every_request = TaskClassifierConfig {
1325 base_threshold: 0.5,
1326 classify_trigger: ClassifyTrigger::EveryRequest,
1327 message_hash_fallback: true,
1328 ..TaskClassifierConfig::default()
1329 };
1330 assert!(
1331 LlmTaskClassifier::new(LlmClassifierConfig::Capability {
1332 config: every_request
1333 })
1334 .is_err(),
1335 "every_request with message_hash_fallback should stay rejected"
1336 );
1337 Ok(())
1338 }
1339
1340 #[test]
1341 fn an_unusable_verdict_is_ambiguous() -> Result<()> {
1342 let policy = policy();
1343 let inconsistent_rule = TaskClassifierVerdict {
1344 capability_boundary: "uncertain".to_string(),
1345 ..verdict(1.0, "supported", "SUP-1")
1346 };
1347 let empty_crux = TaskClassifierVerdict {
1348 crux: " ".to_string(),
1349 ..verdict(1.0, "supported", "SUP-1")
1350 };
1351 let unusable = [
1352 Some(verdict(1.1, "supported", "SUP-1")),
1353 Some(inconsistent_rule),
1354 Some(empty_crux),
1355 None,
1356 ];
1357 for verdict in unusable {
1358 let classification = policy.to_classification(verdict.as_ref(), &policy_driver())?;
1359 assert!(matches!(classification, Classification::Ambiguous(_)));
1360 assert!(classification.argmax(false)?.is_none());
1361 assert!(classification.argmax(true)?.is_none());
1362 }
1363 Ok(())
1364 }
1365
1366 #[test]
1367 fn capability_boundaries_apply_monotonic_threshold_steps() -> Result<()> {
1368 let policy = TaskClassifierPolicy::new(&TaskClassifierConfig {
1369 threshold_step: 0.1,
1370 ..test_config(0.4)
1371 });
1372
1373 assert_eq!(
1374 selected(&policy, Some(&verdict(0.4, "supported", "SUP-2")))?,
1375 "efficient"
1376 );
1377 assert_eq!(
1378 selected(&policy, Some(&verdict(0.49, "uncertain", "UNC-1")))?,
1379 "capable"
1380 );
1381 assert_eq!(
1382 selected(&policy, Some(&verdict(0.5, "uncertain", "UNC-1")))?,
1383 "efficient"
1384 );
1385 assert_eq!(
1386 selected(&policy, Some(&verdict(0.5, "unmatched", "none")))?,
1387 "efficient"
1388 );
1389 assert_eq!(
1390 selected(&policy, Some(&verdict(0.59, "unsupported", "LIM-1")))?,
1391 "capable"
1392 );
1393 assert_eq!(
1394 selected(&policy, Some(&verdict(0.6, "unsupported", "LIM-1")))?,
1395 "efficient"
1396 );
1397 Ok(())
1398 }
1399
1400 fn capability_judge(recent_turn_window: Option<usize>) -> Result<CapabilityJudge> {
1403 Ok(StructuredJudge::new(
1404 TaskInput { recent_turn_window },
1405 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?,
1406 SerdeDecoder::new(),
1407 JudgeRuntimeConfig::new(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?,
1408 ))
1409 }
1410
1411 fn judged_contents(recent_turn_window: usize) -> Result<Vec<String>> {
1412 let judge = capability_judge(Some(recent_turn_window))?;
1413 let request = Request {
1414 llm_request: LlmRequest {
1415 messages: vec![
1416 Message::text(Role::System, "client instructions"),
1417 Message::text(Role::User, "initial task"),
1418 Message::text(Role::Assistant, "old response"),
1419 Message::text(Role::User, "old follow-up"),
1420 Message::text(Role::Assistant, "recent 1"),
1421 Message::text(Role::User, "recent 2"),
1422 ],
1423 ..LlmRequest::default()
1424 },
1425 raw_request: None,
1426 metadata: None,
1427 };
1428 Ok(judge
1429 .build_request(&State::default(), &request)
1430 .llm_request
1431 .messages
1432 .iter()
1433 .filter_map(|message| message.text_content("\n"))
1434 .collect())
1435 }
1436
1437 #[test]
1438 fn a_window_widens_the_judge_to_the_surrounding_conversation() -> Result<()> {
1439 let contents = judged_contents(2)?;
1441 assert!(contents.contains(&"client instructions".to_string()));
1442 assert!(contents.contains(&"initial task".to_string()));
1443 assert!(contents.contains(&"recent 1".to_string()));
1444 assert!(contents.contains(&"recent 2".to_string()));
1445 assert!(!contents.contains(&"old response".to_string()));
1446 Ok(())
1447 }
1448
1449 #[test]
1450 fn a_zero_window_keeps_only_the_instructions_and_the_task() -> Result<()> {
1451 let contents = judged_contents(0)?;
1452 assert!(contents.contains(&"client instructions".to_string()));
1453 assert!(contents.contains(&"initial task".to_string()));
1454 assert!(!contents.contains(&"recent 2".to_string()));
1455 Ok(())
1456 }
1457
1458 fn tool_call(id: &str) -> Message {
1459 Message {
1460 role: Role::Assistant,
1461 content: vec![ContentBlock::ToolCall(ToolCall {
1462 id: id.to_string(),
1463 name: "search".to_string(),
1464 arguments: Value::Null,
1465 })],
1466 }
1467 }
1468
1469 fn tool_result(id: &str) -> Message {
1470 Message {
1471 role: Role::Tool,
1472 content: vec![ContentBlock::ToolResult(ToolResult {
1473 tool_call_id: id.to_string(),
1474 content: vec![ContentBlock::Text {
1475 text: "tool output".to_string(),
1476 }],
1477 is_error: None,
1478 })],
1479 }
1480 }
1481
1482 #[test]
1483 fn default_task_input_keeps_user_content_around_tool_results() {
1484 let mut result = tool_result("call-1");
1485 result.role = Role::User;
1486 let mut mixed = result.clone();
1487 mixed.content.push(ContentBlock::Text {
1488 text: "latest follow-up".to_string(),
1489 });
1490 let input = TaskInput {
1491 recent_turn_window: None,
1492 };
1493 let mut request = Request {
1494 llm_request: LlmRequest {
1495 messages: vec![
1496 result.clone(),
1497 Message::text(Role::User, "initial task"),
1498 tool_call("call-1"),
1499 mixed,
1500 result.clone(),
1501 ],
1502 ..LlmRequest::default()
1503 },
1504 ..Request::default()
1505 };
1506 assert_eq!(
1507 input.build_messages(&State::default(), &request),
1508 vec![
1509 Message::text(Role::User, "initial task"),
1510 Message::text(Role::User, "latest follow-up"),
1511 ]
1512 );
1513 request.llm_request.messages = vec![result];
1514 assert!(input.build_messages(&State::default(), &request).is_empty());
1515 }
1516
1517 #[test]
1520 fn trimming_keeps_the_call_that_introduced_a_kept_tool_result() {
1521 let messages = vec![
1522 Message::text(Role::System, "client instructions"),
1523 Message::text(Role::User, "initial task"),
1524 Message::text(Role::Assistant, "old response"),
1525 tool_call("call-1"),
1526 tool_result("call-1"),
1527 Message::text(Role::Assistant, "recent 1"),
1528 Message::text(Role::User, "recent 2"),
1529 Message::text(Role::Assistant, "recent 3"),
1530 Message::text(Role::User, "recent 4"),
1531 ];
1532
1533 let kept = trim_messages(&messages, 5);
1535
1536 assert_eq!(
1537 kept,
1538 vec![
1539 Message::text(Role::System, "client instructions"),
1540 Message::text(Role::User, "initial task"),
1541 tool_call("call-1"),
1542 tool_result("call-1"),
1543 Message::text(Role::Assistant, "recent 1"),
1544 Message::text(Role::User, "recent 2"),
1545 Message::text(Role::Assistant, "recent 3"),
1546 Message::text(Role::User, "recent 4"),
1547 ]
1548 );
1549 }
1550
1551 #[test]
1554 fn trimming_pairs_a_repeated_id_with_the_call_that_precedes_it() {
1555 let messages = vec![
1556 Message::text(Role::System, "client instructions"),
1557 Message::text(Role::User, "initial task"),
1558 tool_call("x"),
1559 tool_result("x"),
1560 Message::text(Role::Assistant, "later"),
1561 tool_call("x"),
1562 tool_result("x"),
1563 ];
1564
1565 let kept = trim_messages(&messages, 4);
1567
1568 assert_eq!(
1569 kept,
1570 vec![
1571 Message::text(Role::System, "client instructions"),
1572 Message::text(Role::User, "initial task"),
1573 tool_call("x"),
1574 tool_result("x"),
1575 Message::text(Role::Assistant, "later"),
1576 tool_call("x"),
1577 tool_result("x"),
1578 ]
1579 );
1580 }
1581
1582 #[test]
1585 fn trimming_keeps_the_counted_window_when_a_result_cannot_be_paired() {
1586 let messages = vec![
1587 Message::text(Role::System, "client instructions"),
1588 tool_call("orphan"),
1589 Message::text(Role::User, "initial task"),
1590 Message::text(Role::Assistant, "old response"),
1591 tool_result("orphan"),
1592 Message::text(Role::Assistant, "recent 1"),
1593 Message::text(Role::User, "recent 2"),
1594 ];
1595
1596 let kept = trim_messages(&messages, 3);
1597
1598 assert_eq!(
1599 kept,
1600 vec![
1601 Message::text(Role::System, "client instructions"),
1602 Message::text(Role::User, "initial task"),
1603 tool_result("orphan"),
1604 Message::text(Role::Assistant, "recent 1"),
1605 Message::text(Role::User, "recent 2"),
1606 ]
1607 );
1608 }
1609
1610 #[test]
1611 fn a_window_restates_the_routing_instruction_last() -> Result<()> {
1612 let contents = judged_contents(2)?;
1613
1614 assert_eq!(
1617 contents.last().map(String::as_str),
1618 Some(TRAILING_ROUTING_INSTRUCTION)
1619 );
1620 Ok(())
1621 }
1622
1623 #[test]
1628 fn a_window_drops_reasoning_but_keeps_visible_text_and_tool_pairs() {
1629 let messages = vec![
1630 Message::text(Role::System, "client instructions"),
1631 Message::text(Role::User, "initial task"),
1632 Message {
1633 role: Role::Assistant,
1634 content: vec![
1635 ContentBlock::Reasoning {
1636 text: "private chain of thought".to_string(),
1637 signature: None,
1638 details: Vec::new(),
1639 },
1640 ContentBlock::Text {
1641 text: "visible answer".to_string(),
1642 },
1643 ],
1644 },
1645 tool_call("call-1"),
1646 tool_result("call-1"),
1647 Message {
1648 role: Role::Assistant,
1649 content: vec![ContentBlock::Reasoning {
1650 text: "reasoning-only turn".to_string(),
1651 signature: None,
1652 details: Vec::new(),
1653 }],
1654 },
1655 Message::text(Role::User, "follow-up"),
1656 ];
1657 let request = Request {
1658 llm_request: LlmRequest {
1659 messages,
1660 ..LlmRequest::default()
1661 },
1662 raw_request: None,
1663 metadata: None,
1664 };
1665
1666 let built = TaskInput {
1667 recent_turn_window: Some(10),
1668 }
1669 .build_messages(&State::default(), &request);
1670
1671 assert!(
1672 !built
1673 .iter()
1674 .flat_map(|message| &message.content)
1675 .any(|block| matches!(block, ContentBlock::Reasoning { .. })),
1676 "{built:?}"
1677 );
1678 assert!(
1679 built
1680 .iter()
1681 .any(|message| message.text_content("\n").as_deref() == Some("visible answer"))
1682 );
1683 assert!(built.contains(&tool_call("call-1")));
1684 assert!(built.contains(&tool_result("call-1")));
1685 assert_eq!(built.len(), 7);
1688 assert!(built.iter().all(|message| !message.content.is_empty()));
1689 }
1690
1691 #[test]
1692 fn the_default_path_is_left_unchanged() -> Result<()> {
1693 let judge = capability_judge(None)?;
1696 let request = Request {
1697 llm_request: LlmRequest {
1698 messages: vec![Message::text(Role::User, "the task")],
1699 ..LlmRequest::default()
1700 },
1701 raw_request: None,
1702 metadata: None,
1703 };
1704
1705 let built = judge.build_request(&State::default(), &request);
1706
1707 assert_eq!(built.llm_request.messages.len(), 1);
1710 assert!(!built.llm_request.instructions.is_empty());
1711 assert!(
1712 !built
1713 .llm_request
1714 .messages
1715 .iter()
1716 .filter_map(|message| message.text_content("\n"))
1717 .any(|text| text.contains(TRAILING_ROUTING_INSTRUCTION))
1718 );
1719 Ok(())
1720 }
1721
1722 #[test]
1723 fn capability_judge_builds_a_structured_request() -> Result<()> {
1724 let judge = capability_judge(None)?;
1725 let request = Request {
1726 llm_request: LlmRequest {
1727 model: Some("inbound".to_string()),
1728 messages: vec![
1729 Message::text(Role::System, "client instructions"),
1730 Message::text(Role::Developer, "client developer instructions"),
1731 Message::text(Role::User, "initial task"),
1732 Message::text(Role::Assistant, "old response"),
1733 Message::text(Role::User, "old follow-up"),
1734 Message::text(Role::Assistant, "recent 1"),
1735 Message::text(Role::User, "recent 2"),
1736 Message::text(Role::Assistant, "recent 3"),
1737 Message::text(Role::User, "recent 4"),
1738 Message::text(Role::Assistant, "recent 5"),
1739 ],
1740 ..LlmRequest::default()
1741 },
1742 raw_request: None,
1743 metadata: None,
1744 };
1745 let judge_request = judge.build_request(&State::default(), &request);
1746
1747 assert_eq!(judge_request.llm_request.model, request.llm_request.model);
1748 assert_eq!(judge_request.llm_request.instructions.len(), 1);
1749 assert_eq!(judge_request.llm_request.instructions[0].role, Role::System);
1750 assert_eq!(
1751 judge_request.llm_request.instructions[0].content,
1752 InstructionBlock {
1753 role: Role::System,
1754 content: Message::text(Role::System, judge.contract().system_prompt()).content,
1755 }
1756 .content,
1757 );
1758 assert_eq!(judge_request.llm_request.messages.len(), 2);
1759 let contents = judge_request
1760 .llm_request
1761 .messages
1762 .iter()
1763 .filter_map(|message| message.text_content("\n"))
1764 .collect::<Vec<_>>();
1765 assert!(contents.contains(&"recent 4".to_string()));
1766 assert!(contents.contains(&"initial task".to_string()));
1767 assert!(!contents.contains(&"recent 5".to_string()));
1768 assert!(!contents.contains(&"client instructions".to_string()));
1769 assert_eq!(
1770 judge_request.llm_request.output.response_format,
1771 Some(judge.contract().response_format().clone())
1772 );
1773 assert_eq!(
1774 judge_request.llm_request.output.max_output_tokens,
1775 Some(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)
1776 );
1777 Ok(())
1778 }
1779
1780 fn sample_value(spec: &Value) -> Value {
1781 if let Some(first) = spec
1782 .get("enum")
1783 .and_then(Value::as_array)
1784 .and_then(|values| values.first())
1785 {
1786 return first.clone();
1787 }
1788 match spec.get("type").and_then(Value::as_str) {
1789 Some("number") => serde_json::json!(0.5),
1790 Some("boolean") => serde_json::json!(false),
1791 _ => serde_json::json!("sample"),
1792 }
1793 }
1794
1795 fn schema_shaped_verdict(schema: &Value) -> Result<String> {
1796 let properties = schema
1797 .pointer("/json_schema/schema/properties")
1798 .and_then(Value::as_object)
1799 .ok_or_else(|| LibsyError::AlgorithmError {
1800 message: "packaged schema declares no properties".to_string(),
1801 })?;
1802 Ok(Value::Object(
1803 properties
1804 .iter()
1805 .map(|(name, spec)| (name.clone(), sample_value(spec)))
1806 .collect(),
1807 )
1808 .to_string())
1809 }
1810
1811 #[test]
1814 fn every_schema_property_round_trips_through_the_judge_parser() -> Result<()> {
1815 let contract =
1816 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?;
1817 let schema = contract.response_format();
1818 let reply = schema_shaped_verdict(schema)?;
1819 let judge: CapabilityJudge = StructuredJudge::new(
1820 TaskInput {
1821 recent_turn_window: None,
1822 },
1823 contract,
1824 SerdeDecoder::new(),
1825 JudgeRuntimeConfig::new(DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?,
1826 );
1827
1828 let verdict = judge.parse(&text_response(None, reply))?;
1829
1830 assert!(verdict.is_valid());
1831 assert!((0.0..=1.0).contains(&verdict.p_solve));
1832 Ok(())
1833 }
1834
1835 #[test]
1836 fn packaged_prompt_keeps_the_schema_in_the_structured_request() -> Result<()> {
1837 let contract =
1838 LlmTaskClassifier::load_capability_contract(&ClassifierContractConfig::default())?;
1839 let prompt = contract.system_prompt();
1840 let schema_name = contract
1841 .response_format()
1842 .pointer("/json_schema/name")
1843 .and_then(Value::as_str)
1844 .ok_or_else(|| LibsyError::AlgorithmError {
1845 message: "packaged response schema has no name".to_string(),
1846 })?;
1847 assert_eq!(schema_name, "CapabilityClassifierDecision");
1848 assert!(prompt.contains("SUP-1 [supported]"));
1849 assert!(prompt.contains("SUP-5 [supported]"));
1850 assert!(!prompt.contains("{{RESPONSE_SCHEMA}}"));
1851 assert!(!prompt.contains("\"type\": \"object\""));
1852 assert!(!prompt.contains("\"json_schema\""));
1853 assert!(!prompt.contains(schema_name));
1854 let rule_values = contract
1855 .response_format()
1856 .pointer("/json_schema/schema/properties/primary_rule/enum")
1857 .and_then(Value::as_array)
1858 .ok_or_else(|| LibsyError::AlgorithmError {
1859 message: "rendered response schema has no primary rule enum".to_string(),
1860 })?;
1861 assert!(
1862 rule_values
1863 .iter()
1864 .any(|value| value.as_str() == Some("SUP-1"))
1865 );
1866 assert!(
1867 rule_values
1868 .iter()
1869 .any(|value| value.as_str() == Some("none"))
1870 );
1871 Ok(())
1872 }
1873}