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