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