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