Skip to main content

switchyard_libsy/algorithms/
rand.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Random routing as a stateless [`FallThrough`] composition.
5//!
6//! [`RandomClassifier`] selects one target; [`FallThrough`] owns the common
7//! processor/classifier/target-call orchestration.
8
9use 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
24/// Stateless weighted classifier used by random fall-through routing.
25pub struct RandomClassifier {
26    distribution: Option<WeightedIndex<f64>>,
27    rng: Mutex<StdRng>,
28}
29
30impl RandomClassifier {
31    /// Creates a classifier over ordered target names.
32    ///
33    /// Missing weights default to one per target. Explicit weights are relative,
34    /// follow target order, and need not sum to one. Zero disables a target.
35    /// Missing `seed` uses entropy-backed randomness.
36    ///
37    /// # Errors
38    ///
39    /// Returns an error explicit weights are negative or non-finite, or contain no
40    /// positive value.
41    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            // Temp until can remove Option from driver
90            return Err(LibsyError::NoTargets);
91        };
92        // All the available models
93        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            // The user gave us weights
100            distribution.sample(&mut *rng)
101        } else {
102            // No weights, assume equal probability
103            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
116/// Random router implemented as a stateless fall-through composition.
117pub struct Random {
118    inner: FallThrough<()>,
119}
120
121impl Random {
122    /// Creates a random router. The models themselves will be passed at runtime.
123    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    /*
177    fn target_set(names: &[&str]) -> Vec<ModelId> {
178        names.iter().map(|name| ModelId::from(*name)).collect()
179    }
180    */
181
182    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        // Missing either target after 100 uniform draws has probability about 2^-99.
254        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}