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