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::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
24/// Stateless weighted classifier used by random fall-through routing.
25pub struct RandomClassifier {
26    targets: Vec<ModelId>,
27    distribution: WeightedIndex<f64>,
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 when targets are empty or duplicated, or when explicit
41    /// weights have the wrong length, are negative or non-finite, or contain no
42    /// positive value.
43    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
126/// Random router implemented as a stateless fall-through composition.
127pub struct Random {
128    inner: FallThrough<()>,
129}
130
131impl Random {
132    /// Creates a router over `targets`.
133    ///
134    /// # Errors
135    ///
136    /// Returns an error when targets or weights are invalid for [`RandomClassifier`].
137    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        // Missing either target after 100 uniform draws has probability about 2^-99.
261        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}