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    weight_count: Option<usize>,
28    rng: Mutex<StdRng>,
29}
30
31impl RandomClassifier {
32    /// Creates a classifier over ordered target names.
33    ///
34    /// Missing weights default to one per target. Explicit weights are relative,
35    /// follow target order, and need not sum to one. Zero disables a target.
36    /// Missing `seed` uses entropy-backed randomness.
37    ///
38    /// # Errors
39    ///
40    /// Returns an error if explicit weights are negative or non-finite, or contain no
41    /// positive value.
42    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        // All the available models
92        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            // The user gave us weights
107            distribution.sample(&mut *rng)
108        } else {
109            // No weights, assume equal probability
110            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
124/// Random router implemented as a stateless fall-through composition.
125pub struct Random {
126    inner: FallThrough<()>,
127}
128
129impl Random {
130    /// Creates a random router. The models themselves will be passed at runtime.
131    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        // Missing either target after 100 uniform draws has probability about 2^-99.
252        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}