Skip to main content

switchyard_libsy/algorithms/llm_class/
decision.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Compare two candidates by the chance that the capable candidate alone succeeds.
5
6use std::collections::{BTreeMap, HashSet};
7
8use async_trait::async_trait;
9use serde_json::{Value, json};
10use switchyard_protocol::{
11    Category, ChoiceOption, DecisionKind, DecisionQuestion, DecisionRequest, DecisionValue,
12    ModelId, Request, Response,
13};
14
15use crate::algorithms::llm_class::TaskInput;
16use crate::algorithms::util::llm_judge::{libsy_error_reason, report_fail_open};
17use crate::algorithms::util::robustness::safe_error_summary;
18use crate::{Classification, Classifier, Driver, LibsyError, Result, Score, State};
19
20/// Evidence and policy for a relative-advantage decision judge.
21///
22/// Candidate labels keep model names out of the generated context. Evidence must use
23/// the same labels, with unknown outcomes left unknown. Only the first runtime
24/// capable and efficient targets are compared; extra candidates do not add routes.
25#[derive(Clone, Debug)]
26pub struct DecisionJudgeConfig {
27    /// Route capable only when its advantage score is strictly above this cutoff.
28    /// This score is not a calibrated solve probability. Choose a cutoff from evaluations.
29    pub cutoff: f64,
30    /// Replaces the packaged structured instructions; must keep the meaning of
31    /// `advantage` (capable succeeds and efficient fails) and `no_advantage`.
32    pub instructions: Option<Value>,
33    /// Maps anonymous labels used in evidence (e.g. `"a"`) to runtime model IDs.
34    /// Every compared target needs a unique label; extra candidates provide context.
35    pub candidates: BTreeMap<String, ModelId>,
36    /// JSON passed unchanged to the judge, such as candidate descriptions, reference
37    /// cases, outcome/cost summaries, and selection notes. The judge interprets it
38    /// using the instructions; the router neither reads its fields nor derives statistics.
39    pub evidence: Value,
40}
41
42impl DecisionJudgeConfig {
43    pub(super) fn validate(&self) -> Result<()> {
44        if !(0.0..=1.0).contains(&self.cutoff) {
45            return Err(LibsyError::AlgorithmError {
46                message: "decision cutoff must be between 0 and 1".into(),
47            });
48        }
49        if self.candidates.values().collect::<HashSet<_>>().len() != self.candidates.len() {
50            return Err(LibsyError::AlgorithmError {
51                message: "decision candidates must map to distinct targets".into(),
52            });
53        }
54        Ok(())
55    }
56
57    fn candidate(&self, model: &ModelId) -> Result<&str> {
58        self.candidates
59            .iter()
60            .find(|(_, target)| *target == model)
61            .map(|(id, _)| id.as_str())
62            .ok_or_else(|| LibsyError::AlgorithmError {
63                message: format!("decision candidate is missing for target {model}"),
64            })
65    }
66}
67
68pub(super) struct DecisionClassifier {
69    config: DecisionJudgeConfig,
70    input: TaskInput,
71    fail_open: bool,
72    question: DecisionQuestion,
73}
74
75impl DecisionClassifier {
76    pub(super) fn new(
77        mut config: DecisionJudgeConfig,
78        input: TaskInput,
79        fail_open: bool,
80    ) -> Result<Self> {
81        let instructions = match config.instructions.take() {
82            Some(instructions) => instructions,
83            None => serde_json::from_str(include_str!(
84                "../../prompts/capability-classifier/relative_advantage.json"
85            ))
86            .map_err(|error| LibsyError::external("loading decision judge instructions", error))?,
87        };
88        let question = DecisionQuestion {
89            instructions,
90            kind: DecisionKind::Choice {
91                options: [
92                    ("advantage", "The capable candidate succeeds AND the efficient candidate fails."),
93                    ("no_advantage", "The efficient candidate succeeds OR the capable candidate fails, including shared failure."),
94                ]
95                .into_iter()
96                .map(|(id, description)| ChoiceOption {
97                    id: id.into(),
98                    description: Some(json!(description)),
99                })
100                .collect(),
101            },
102        };
103        Ok(Self {
104            config,
105            input,
106            fail_open,
107            question,
108        })
109    }
110}
111
112#[async_trait]
113impl Classifier<State> for DecisionClassifier {
114    async fn score(
115        &self,
116        _state: &mut State,
117        request: &mut Request,
118        driver: &Driver,
119    ) -> Result<(Classification, Option<Response>)> {
120        let judge = driver.first_model_for(&Category::Judge)?;
121        let capable = driver.first_model_for(&Category::Capable)?;
122        let efficient = driver.first_model_for(&Category::Efficient)?;
123        let decision = DecisionRequest {
124            model: None,
125            context: json!({
126                "task": self.input.messages(request),
127                "candidates": self.config.candidates.keys().collect::<Vec<_>>(),
128                "comparison": {
129                    "capable": self.config.candidate(capable)?,
130                    "efficient": self.config.candidate(efficient)?,
131                },
132                "evidence": self.config.evidence,
133            }),
134            questions: BTreeMap::from([("route".into(), self.question.clone())]),
135        };
136        let response = match driver.call_decision(decision, judge.clone()).await {
137            Ok(response) => response,
138            Err(error) if self.fail_open => {
139                return Ok(unavailable(
140                    driver,
141                    judge,
142                    safe_error_summary(&error),
143                    libsy_error_reason(&error),
144                ));
145            }
146            Err(error) => return Err(error),
147        };
148        // The provider's selected option can differ from the application's cutoff.
149        // Only the requested event's score is used; confidence is a separate signal.
150        let advantage = response
151            .answers
152            .get("route")
153            .and_then(|answer| match &answer.value {
154                DecisionValue::Choice {
155                    probabilities: Some(probabilities),
156                    ..
157                } => probabilities.get("advantage").map(|p| p.0),
158                _ => None,
159            })
160            .filter(|p| (0.0..=1.0).contains(p));
161        let Some(advantage) = advantage else {
162            return Ok(unavailable(
163                driver,
164                judge,
165                "missing or invalid advantage score".into(),
166                "invalid_verdict",
167            ));
168        };
169        let (target, category) = if advantage > self.config.cutoff {
170            (capable, Category::Capable)
171        } else {
172            (efficient, Category::Efficient)
173        };
174        driver.set_evidence(json!({
175            "source": "decision_classifier",
176            "verdict": "relative_advantage",
177            "score": advantage,
178            "threshold": self.config.cutoff,
179        }));
180        Ok((
181            Classification::Scores(vec![Score {
182                target: target.clone(),
183                confidence: 1.0,
184                category: Some(category),
185            }]),
186            None,
187        ))
188    }
189}
190
191fn unavailable(
192    driver: &Driver,
193    judge: &ModelId,
194    error: String,
195    reason: &'static str,
196) -> (Classification, Option<Response>) {
197    report_fail_open(judge.as_str(), error, reason);
198    driver.set_evidence_if_empty(json!({"source": "fail_open", "reason_code": reason}));
199    (Classification::Ambiguous(vec![]), None)
200}