1use crate::core::algorithm::Driver;
5use crate::{LibsyError, Result};
6use async_trait::async_trait;
7use switchyard_protocol::{Category, ModelId, Request, Response};
8
9#[derive(Debug, Clone, PartialEq)]
11pub struct Score {
12 pub confidence: f64,
14 pub target: ModelId,
16 pub category: Option<Category>,
20}
21
22pub enum Classification {
25 Scores(Vec<Score>),
27 Ambiguous(Vec<Score>),
30}
31
32impl Classification {
33 pub fn argmax(&self, ignore_ambiguous: bool) -> Result<Option<Score>> {
39 match self {
40 Classification::Scores(scores) => argmax(scores),
41 Classification::Ambiguous(scores) => {
42 if ignore_ambiguous {
43 argmax(scores)
44 } else {
45 Ok(None)
46 }
47 }
48 }
49 }
50}
51
52fn argmax(scores: &[Score]) -> Result<Option<Score>> {
56 let mut best: Option<&Score> = None;
57 for score in scores.iter() {
58 if score.confidence.is_nan() {
59 return Err(LibsyError::AlgorithmError {
60 message: format!(
61 "classifier returned NaN confidence for target {:?}",
62 score.target
63 ),
64 });
65 }
66 match best {
67 Some(cur_best) if score.confidence > cur_best.confidence => best = Some(score),
68 None => best = Some(score),
69 _ => {}
70 }
71 }
72 Ok(best.cloned())
73}
74
75#[async_trait]
77pub trait Classifier<S = ()>: Send + Sync {
78 async fn score(
88 &self,
89 state: &mut S,
90 request: &mut Request,
91 driver: &Driver,
92 ) -> Result<(Classification, Option<Response>)>;
93}
94
95#[cfg(test)]
96mod tests {
97 use super::*;
98 use crate::core::testing::empty_driver;
99 use switchyard_protocol::text_request;
100
101 fn score(target: &str, confidence: f64) -> Score {
103 Score {
104 target: ModelId::from(target),
105 confidence,
106 category: None,
107 }
108 }
109
110 #[test]
111 fn argmax_picks_the_highest_confidence_score() -> Result<()> {
112 let scores = vec![score("weak", 0.2), score("strong", 0.9), score("mid", 0.5)];
113 let best = Classification::Scores(scores).argmax(false)?;
114 assert_eq!(best, Some(score("strong", 0.9)));
115 Ok(())
116 }
117
118 #[test]
119 fn argmax_breaks_ties_by_cascade_order() -> Result<()> {
120 let scores = vec![score("first", 0.7), score("second", 0.7)];
122 let best = Classification::Scores(scores).argmax(false)?;
123 assert_eq!(best.map(|s| s.target), Some(ModelId::from("first")));
124 Ok(())
125 }
126
127 #[test]
128 fn argmax_on_an_empty_set_abstains() -> Result<()> {
129 assert_eq!(Classification::Scores(vec![]).argmax(false)?, None);
131 assert_eq!(Classification::Ambiguous(vec![]).argmax(true)?, None);
132 Ok(())
133 }
134
135 #[test]
136 fn argmax_errors_on_nan_confidence() {
137 let scores = vec![score("weak", 0.3), score("strong", f64::NAN)];
139 assert!(matches!(
140 Classification::Scores(scores).argmax(false),
141 Err(LibsyError::AlgorithmError { message })
142 if message == "classifier returned NaN confidence for target \"strong\""
143 ));
144 assert!(matches!(
146 Classification::Scores(vec![score("only", f64::NAN)]).argmax(false),
147 Err(LibsyError::AlgorithmError { message })
148 if message == "classifier returned NaN confidence for target \"only\""
149 ));
150 }
151
152 #[test]
153 fn ambiguous_without_ignore_makes_no_choice() -> Result<()> {
154 let scores = vec![score("strong", 0.9)];
156 assert_eq!(Classification::Ambiguous(scores).argmax(false)?, None);
157 Ok(())
158 }
159
160 #[test]
161 fn ambiguous_with_ignore_falls_back_to_argmax() -> Result<()> {
162 let scores = vec![score("weak", 0.3), score("strong", 0.8)];
163 let best = Classification::Ambiguous(scores).argmax(true)?;
164 assert_eq!(best, Some(score("strong", 0.8)));
165 Ok(())
166 }
167
168 #[test]
169 fn scores_variant_ignores_the_ambiguous_flag() -> Result<()> {
170 let scores = vec![score("a", 0.4), score("b", 0.6)];
172 let with_ignore = Classification::Scores(scores.clone()).argmax(true)?;
173 let without_ignore = Classification::Scores(scores).argmax(false)?;
174 assert_eq!(with_ignore, without_ignore);
175 assert_eq!(with_ignore, Some(score("b", 0.6)));
176 Ok(())
177 }
178
179 struct RecordingClassifier;
181
182 #[async_trait]
183 impl Classifier<bool> for RecordingClassifier {
184 async fn score(
185 &self,
186 state: &mut bool,
187 request: &mut Request,
188 _driver: &Driver,
189 ) -> Result<(Classification, Option<Response>)> {
190 *state = true;
191 let target = request.model_id().unwrap_or(ModelId::from("auto"));
192 Ok((
193 Classification::Scores(vec![Score {
194 target,
195 confidence: 1.0,
196 category: None,
197 }]),
198 None,
199 ))
200 }
201 }
202
203 #[tokio::test]
204 async fn classifier_reads_request_and_mutates_state() -> Result<()> {
205 let mut state = false;
206 let mut request = Request {
207 llm_request: text_request(Some("strong".to_string()), "hi"),
208 raw_request: None,
209 metadata: None,
210 };
211 let (classification, _) = RecordingClassifier
212 .score(&mut state, &mut request, &empty_driver())
213 .await?;
214 assert_eq!(
215 classification.argmax(false)?.map(|s| s.target),
216 Some(ModelId::from("strong"))
217 );
218 assert!(state);
219 Ok(())
220 }
221
222 struct RewritingClassifier;
224
225 #[async_trait]
226 impl Classifier for RewritingClassifier {
227 async fn score(
228 &self,
229 _state: &mut (),
230 request: &mut Request,
231 _driver: &Driver,
232 ) -> Result<(Classification, Option<Response>)> {
233 request.llm_request.model = Some("rewritten".to_string());
234 Ok((
235 Classification::Scores(vec![Score {
236 target: ModelId::from("rewritten"),
237 confidence: 1.0,
238 category: None,
239 }]),
240 None,
241 ))
242 }
243 }
244
245 #[tokio::test]
246 async fn classifier_rewrites_the_request_in_place() -> Result<()> {
247 let mut state = ();
248 let mut request = Request {
249 llm_request: text_request(Some("auto".to_string()), "hi"),
250 raw_request: None,
251 metadata: None,
252 };
253
254 RewritingClassifier
255 .score(&mut state, &mut request, &empty_driver())
256 .await?;
257
258 assert_eq!(request.model_id().as_deref(), Some("rewritten"));
261 Ok(())
262 }
263}