Skip to main content

switchyard_libsy/algorithms/
llm_class.rs

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