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