switchyard_libsy/algorithms/
rand.rs1use std::collections::BTreeSet;
10use std::sync::Arc;
11
12use async_trait::async_trait;
13use parking_lot::Mutex;
14use rand::SeedableRng;
15use rand::distr::{Distribution, weighted::WeightedIndex};
16use rand::rngs::StdRng;
17
18use crate::algorithms::fall_through::FallThrough;
19use crate::core::algorithm::{Algorithm, Driver};
20use crate::core::classifier::{Classification, Classifier, Score};
21use crate::{LibsyError, Result};
22use switchyard_protocol::{ModelId, Request, Response};
23
24pub struct RandomClassifier {
26 targets: Vec<ModelId>,
27 distribution: WeightedIndex<f64>,
28 rng: Mutex<StdRng>,
29}
30
31impl RandomClassifier {
32 pub fn new(
44 targets: Vec<ModelId>,
45 weights: Option<Vec<f64>>,
46 seed: Option<u64>,
47 ) -> Result<Self> {
48 let target_count = targets.len();
49 if target_count == 0 {
50 return Err(LibsyError::NoTargets);
51 }
52 let unique_targets = targets.iter().map(ModelId::as_str).collect::<BTreeSet<_>>();
53 if unique_targets.len() != target_count {
54 return Err(LibsyError::AlgorithmError {
55 message: "random targets must be unique".to_string(),
56 });
57 }
58
59 let weights = weights.unwrap_or_else(|| vec![1.0; target_count]);
60 if weights.len() != target_count {
61 return Err(invalid_weights(format!(
62 "expected {target_count} weights, got {}",
63 weights.len()
64 )));
65 }
66 if weights
67 .iter()
68 .any(|weight| !weight.is_finite() || *weight < 0.0)
69 {
70 return Err(invalid_weights(
71 "weights must be finite and nonnegative".to_string(),
72 ));
73 }
74 if !weights.iter().any(|weight| *weight > 0.0) {
75 return Err(invalid_weights(
76 "at least one weight must be positive".to_string(),
77 ));
78 }
79 let distribution =
80 WeightedIndex::new(weights).map_err(|error| invalid_weights(error.to_string()))?;
81 let rng = match seed {
82 Some(seed) => StdRng::seed_from_u64(seed),
83 None => rand::make_rng(),
84 };
85 Ok(Self {
86 targets,
87 distribution,
88 rng: Mutex::new(rng),
89 })
90 }
91
92 fn select_target(&self) -> ModelId {
93 let mut rng = self.rng.lock();
94 let index = self.distribution.sample(&mut *rng);
95 self.targets[index].clone()
96 }
97}
98
99fn invalid_weights(message: String) -> LibsyError {
100 LibsyError::AlgorithmError {
101 message: format!("invalid random weights: {message}"),
102 }
103}
104
105#[async_trait]
106impl<S> Classifier<S> for RandomClassifier
107where
108 S: Send + 'static,
109{
110 async fn score(
111 &self,
112 _state: &mut S,
113 _request: &mut Request,
114 _driver: Option<&Driver>,
115 ) -> Result<(Classification, Option<Response>)> {
116 Ok((
117 Classification::Scores(vec![Score {
118 confidence: 1.0,
119 target: self.select_target(),
120 }]),
121 None,
122 ))
123 }
124}
125
126pub struct Random {
128 inner: FallThrough<()>,
129}
130
131impl Random {
132 pub fn new(
138 targets: Vec<ModelId>,
139 weights: Option<Vec<f64>>,
140 seed: Option<u64>,
141 ) -> Result<Self> {
142 let classifier = Arc::new(RandomClassifier::new(targets.clone(), weights, seed)?);
143 let inner = FallThrough::<()>::new(targets)
144 .with_name("random")
145 .with_decision_reason(random_decision_reason)
146 .with_classifier(classifier);
147 Ok(Self { inner })
148 }
149}
150
151fn random_decision_reason(_name: &str, winner: &Score) -> String {
152 format!("random routing selected target '{}'", winner.target)
153}
154
155#[async_trait]
156impl Algorithm for Random {
157 fn name(&self) -> &str {
158 "random"
159 }
160
161 async fn route(
162 self: Arc<Self>,
163 driver: Driver,
164 request: Request,
165 ) -> Result<crate::RoutingOutcome> {
166 self.inner.execute(driver, request).await
167 }
168}
169
170#[cfg(test)]
171mod tests {
172 use super::*;
173 use std::collections::HashSet;
174
175 use switchyard_protocol::{Metadata, completion_text, text_request};
176
177 use crate::algorithms::util::affinity::AffinityRouter;
178 use crate::core::testing::{echo, test_drive};
179 use switchyard_protocol::Request;
180
181 fn request() -> Request {
182 Request {
183 llm_request: text_request(Some("auto".to_string()), "hi"),
184 raw_request: None,
185 metadata: None,
186 }
187 }
188
189 fn request_for_session(session_id: &str) -> Request {
190 Request {
191 metadata: Some(Metadata {
192 session_id: Some(session_id.to_string()),
193 ..Metadata::default()
194 }),
195 ..request()
196 }
197 }
198
199 fn target_set(names: &[&str]) -> Vec<ModelId> {
200 names.iter().map(|name| ModelId::from(*name)).collect()
201 }
202
203 fn algorithm(names: &[&str], weights: Option<Vec<f64>>, seed: Option<u64>) -> Result<Random> {
204 Random::new(target_set(names), weights, seed)
205 }
206
207 fn shared_algorithm(names: &[&str]) -> Result<Arc<dyn Algorithm>> {
208 Ok(Arc::new(algorithm(names, None, None)?))
209 }
210
211 async fn selected_models(algorithm: Arc<dyn Algorithm>, count: usize) -> Result<Vec<String>> {
212 let mut selected = Vec::with_capacity(count);
213 for _ in 0..count {
214 let (_, response) = test_drive(algorithm.clone(), request(), echo()).await?;
215 selected.push(
216 response
217 .llm_response
218 .as_agg()
219 .map(completion_text)
220 .unwrap_or_default(),
221 );
222 }
223 Ok(selected)
224 }
225
226 #[tokio::test]
227 async fn single_target_is_always_selected_and_called() -> Result<()> {
228 let algorithm = shared_algorithm(&["only/model"])?;
229 let (selected_model, response) = test_drive(algorithm, request(), echo()).await?;
230
231 assert_eq!(
232 response
233 .llm_response
234 .as_agg()
235 .map(completion_text)
236 .unwrap_or_default(),
237 "only/model"
238 );
239 assert_eq!(selected_model, "only/model");
240 Ok(())
241 }
242
243 #[tokio::test]
244 async fn selection_covers_all_targets_over_many_runs() -> Result<()> {
245 let algorithm = shared_algorithm(&["a/model", "b/model"])?;
246 let mut seen = HashSet::new();
247
248 for _ in 0..100 {
249 let (selected_model, response) =
250 test_drive(algorithm.clone(), request(), echo()).await?;
251 let served_model = response
252 .llm_response
253 .as_agg()
254 .map(completion_text)
255 .unwrap_or_default();
256 assert_eq!(selected_model, served_model.as_str());
257 seen.insert(served_model);
258 }
259
260 assert_eq!(
262 seen.len(),
263 2,
264 "expected both targets to be selected, saw {seen:?}"
265 );
266 Ok(())
267 }
268
269 #[tokio::test]
270 async fn weighted_seeded_selection_is_reproducible() -> Result<()> {
271 let first: Arc<dyn Algorithm> = Arc::new(algorithm(
272 &["a/model", "b/model"],
273 Some(vec![1.0, 3.0]),
274 Some(42),
275 )?);
276 let second: Arc<dyn Algorithm> = Arc::new(algorithm(
277 &["a/model", "b/model"],
278 Some(vec![1.0, 3.0]),
279 Some(42),
280 )?);
281
282 let first_selections = selected_models(first, 1_000).await?;
283 let second_selections = selected_models(second, 1_000).await?;
284 assert_eq!(first_selections, second_selections);
285
286 let second_count = first_selections
287 .iter()
288 .filter(|model| model.as_str() == "b/model")
289 .count();
290 assert!(
291 (700..=800).contains(&second_count),
292 "expected a roughly 25/75 split, selected b/model {second_count} times"
293 );
294 Ok(())
295 }
296
297 #[tokio::test]
298 async fn affinity_reuses_the_initial_random_selection() -> Result<()> {
299 let names = ["a/model", "b/model"];
300 let affinity = Arc::new(AffinityRouter::new());
301 let random = Arc::new(RandomClassifier::new(
302 names.iter().map(|name| ModelId::from(*name)).collect(),
303 None,
304 Some(42),
305 )?);
306 let algorithm: Arc<dyn Algorithm> = Arc::new(
307 FallThrough::<()>::new(target_set(&names))
308 .with_name("affinity_random")
309 .with_processor(affinity.clone())
310 .with_classifier(affinity.clone())
311 .with_classifier(random),
312 );
313
314 let (_, first) =
315 test_drive(algorithm.clone(), request_for_session("session-1"), echo()).await?;
316 let selected = first
317 .llm_response
318 .as_agg()
319 .map(completion_text)
320 .unwrap_or_default();
321
322 let mut state = ();
323 let mut request = request_for_session("session-1");
324 let retained = affinity
325 .score(&mut state, &mut request, None)
326 .await?
327 .0
328 .argmax(false)?;
329 assert_eq!(
330 retained.map(|score| score.target),
331 Some(ModelId::from(selected.clone()))
332 );
333
334 let (_, second) = test_drive(algorithm, request_for_session("session-1"), echo()).await?;
335 assert_eq!(
336 second
337 .llm_response
338 .as_agg()
339 .map(completion_text)
340 .unwrap_or_default(),
341 selected
342 );
343 Ok(())
344 }
345
346 #[test]
347 fn rejects_invalid_weights() {
348 let cases = [
349 (vec![1.0], "expected 2 weights"),
350 (vec![1.0, -1.0], "finite and nonnegative"),
351 (vec![0.0, 0.0], "at least one weight must be positive"),
352 (vec![1.0, f64::INFINITY], "finite and nonnegative"),
353 ];
354
355 for (weights, expected) in cases {
356 let error = algorithm(&["a/model", "b/model"], Some(weights), None)
357 .err()
358 .map(|error| error.to_string())
359 .unwrap_or_default();
360 assert!(error.contains(expected), "unexpected error: {error}");
361 }
362 }
363
364 #[test]
365 fn rejects_invalid_targets() {
366 let error = algorithm(&[], None, None).err();
367 assert!(matches!(error, Some(LibsyError::NoTargets)));
368
369 let error = algorithm(&["same/model", "same/model"], None, None)
370 .err()
371 .map(|error| error.to_string())
372 .unwrap_or_default();
373 assert!(error.contains("random targets must be unique"));
374 }
375
376 #[tokio::test]
377 async fn decision_is_inspectable() -> Result<()> {
378 let algorithm = shared_algorithm(&["only/model"])?;
379 let (selected_model, _) = test_drive(algorithm, request(), echo()).await?;
380 assert_eq!(selected_model, "only/model");
381 Ok(())
382 }
383}