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