switchyard_libsy/algorithms/llm_class/
decision.rs1use 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#[derive(Clone, Debug)]
26pub struct DecisionJudgeConfig {
27 pub cutoff: f64,
30 pub instructions: Option<Value>,
33 pub candidates: BTreeMap<String, ModelId>,
36 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 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}