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