Skip to main content

tokio_prompt_orchestrator/
load_balancer.rs

1//! # Multi-Model Load Balancer
2//!
3//! Weighted round-robin selection of model endpoints with health tracking
4//! and automatic failover.
5//!
6//! ## Selection Algorithm
7//!
8//! Uses the smooth weighted round-robin algorithm (Nginx-style):
9//! each endpoint's `current_weight` is incremented by its `weight` on every
10//! call to [`LoadBalancer::select`]; the endpoint with the highest
11//! `current_weight` is chosen and then has `total_weight` subtracted from
12//! its `current_weight`.
13//!
14//! ## Health Tracking
15//!
16//! - After **3 consecutive failures** an endpoint is marked unhealthy.
17//! - **1 success** recovers an endpoint.
18//! - When `failover = true` unhealthy endpoints are skipped during selection.
19
20use std::collections::HashMap;
21use std::sync::{Arc, Mutex};
22use std::time::Duration;
23
24// ── Types ─────────────────────────────────────────────────────────────────────
25
26/// A single model inference endpoint.
27#[derive(Debug, Clone)]
28pub struct ModelEndpoint {
29    /// Unique identifier for this endpoint (e.g. `"gpt-4o-primary"`).
30    pub id: String,
31    /// Base URL of the endpoint (e.g. `"https://api.openai.com/v1"`).
32    pub url: String,
33    /// Relative weight used for weighted round-robin selection.
34    /// Higher weight = more traffic. Must be > 0.
35    pub weight: u32,
36    /// Maximum requests per second this endpoint can sustain.
37    pub max_rps: f64,
38    /// Whether this endpoint is currently considered healthy.
39    pub healthy: bool,
40    /// 99th-percentile observed latency in milliseconds.
41    pub latency_p99_ms: f64,
42}
43
44/// Configuration for the [`LoadBalancer`].
45#[derive(Debug, Clone)]
46pub struct BalancerConfig {
47    /// The set of model endpoints to balance across.
48    pub endpoints: Vec<ModelEndpoint>,
49    /// How often to poll endpoints for health status.
50    pub health_check_interval: Duration,
51    /// When `true`, unhealthy endpoints are skipped during selection.
52    /// When `false`, all endpoints participate regardless of health.
53    pub failover: bool,
54}
55
56impl Default for BalancerConfig {
57    fn default() -> Self {
58        Self {
59            endpoints: Vec::new(),
60            health_check_interval: Duration::from_secs(30),
61            failover: true,
62        }
63    }
64}
65
66/// Per-endpoint statistics tracked by the [`LoadBalancer`].
67#[derive(Debug, Clone, Default)]
68pub struct EndpointStats {
69    /// Total number of requests dispatched to this endpoint.
70    pub requests: u64,
71    /// Total number of failures recorded for this endpoint.
72    pub failures: u64,
73    /// Current consecutive failure run (reset on success).
74    pub consecutive_failures: u32,
75    /// Exponential-moving-average latency in milliseconds.
76    pub avg_latency_ms: f64,
77    /// Whether the endpoint is currently healthy.
78    pub is_healthy: bool,
79}
80
81/// Aggregate statistics for the entire balancer.
82#[derive(Debug, Clone, Default)]
83pub struct LoadBalancerStats {
84    /// Total requests dispatched across all endpoints.
85    pub total_requests: u64,
86    /// Per-endpoint breakdown.
87    pub by_endpoint: HashMap<String, EndpointStats>,
88    /// Number of endpoints currently marked unhealthy.
89    pub unhealthy_count: usize,
90}
91
92// ── Internal state ────────────────────────────────────────────────────────────
93
94/// Number of consecutive failures before an endpoint is marked unhealthy.
95const FAILURE_THRESHOLD: u32 = 3;
96
97/// Exponential moving average α for latency tracking.
98const LATENCY_EMA_ALPHA: f64 = 0.2;
99
100struct EndpointState {
101    endpoint: ModelEndpoint,
102    current_weight: i64,
103    stats: EndpointStats,
104}
105
106struct BalancerInner {
107    endpoints: Vec<EndpointState>,
108    config: BalancerConfig,
109    total_requests: u64,
110}
111
112impl BalancerInner {
113    fn new(config: BalancerConfig) -> Self {
114        let endpoints = config
115            .endpoints
116            .iter()
117            .map(|ep| EndpointState {
118                endpoint: ep.clone(),
119                current_weight: 0,
120                stats: EndpointStats {
121                    is_healthy: ep.healthy,
122                    ..Default::default()
123                },
124            })
125            .collect();
126        Self {
127            endpoints,
128            config,
129            total_requests: 0,
130        }
131    }
132
133    /// Weighted round-robin selection. Returns the index of the selected endpoint.
134    fn select_index(&mut self) -> Option<usize> {
135        if self.endpoints.is_empty() {
136            return None;
137        }
138
139        let total_weight: i64 = self
140            .endpoints
141            .iter()
142            .filter(|s| !self.config.failover || s.endpoint.healthy)
143            .map(|s| s.endpoint.weight as i64)
144            .sum();
145
146        if total_weight == 0 {
147            return None;
148        }
149
150        // Increment current_weight by own weight for all eligible endpoints.
151        for state in self.endpoints.iter_mut() {
152            if !self.config.failover || state.endpoint.healthy {
153                state.current_weight += state.endpoint.weight as i64;
154            }
155        }
156
157        // Pick the endpoint with the highest current_weight.
158        let idx = self
159            .endpoints
160            .iter()
161            .enumerate()
162            .filter(|(_, s)| !self.config.failover || s.endpoint.healthy)
163            .max_by_key(|(_, s)| s.current_weight)
164            .map(|(i, _)| i)?;
165
166        // Subtract total_weight from the winner.
167        self.endpoints[idx].current_weight -= total_weight;
168        self.endpoints[idx].stats.requests += 1;
169        self.total_requests += 1;
170
171        Some(idx)
172    }
173
174    fn mark_success(&mut self, id: &str, latency_ms: f64) {
175        if let Some(state) = self.endpoints.iter_mut().find(|s| s.endpoint.id == id) {
176            state.stats.consecutive_failures = 0;
177            state.endpoint.healthy = true;
178            state.stats.is_healthy = true;
179            // Update EMA latency.
180            if state.stats.avg_latency_ms == 0.0 {
181                state.stats.avg_latency_ms = latency_ms;
182            } else {
183                state.stats.avg_latency_ms = LATENCY_EMA_ALPHA * latency_ms
184                    + (1.0 - LATENCY_EMA_ALPHA) * state.stats.avg_latency_ms;
185            }
186            state.endpoint.latency_p99_ms = state.stats.avg_latency_ms;
187        }
188    }
189
190    fn mark_failure(&mut self, id: &str) {
191        if let Some(state) = self.endpoints.iter_mut().find(|s| s.endpoint.id == id) {
192            state.stats.failures += 1;
193            state.stats.consecutive_failures += 1;
194            if state.stats.consecutive_failures >= FAILURE_THRESHOLD {
195                state.endpoint.healthy = false;
196                state.stats.is_healthy = false;
197            }
198        }
199    }
200
201    fn stats(&self) -> LoadBalancerStats {
202        let by_endpoint: HashMap<String, EndpointStats> = self
203            .endpoints
204            .iter()
205            .map(|s| (s.endpoint.id.clone(), s.stats.clone()))
206            .collect();
207
208        let unhealthy_count = self
209            .endpoints
210            .iter()
211            .filter(|s| !s.endpoint.healthy)
212            .count();
213
214        LoadBalancerStats {
215            total_requests: self.total_requests,
216            by_endpoint,
217            unhealthy_count,
218        }
219    }
220
221    fn endpoint_at(&self, idx: usize) -> Option<ModelEndpoint> {
222        self.endpoints.get(idx).map(|s| s.endpoint.clone())
223    }
224}
225
226// ── Public handle ─────────────────────────────────────────────────────────────
227
228/// Thread-safe weighted round-robin load balancer for model endpoints.
229///
230/// # Example
231///
232/// ```rust
233/// use tokio_prompt_orchestrator::load_balancer::{
234///     BalancerConfig, LoadBalancer, ModelEndpoint,
235/// };
236/// use std::time::Duration;
237///
238/// let config = BalancerConfig {
239///     endpoints: vec![
240///         ModelEndpoint {
241///             id: "ep-a".to_string(),
242///             url: "http://a".to_string(),
243///             weight: 2,
244///             max_rps: 100.0,
245///             healthy: true,
246///             latency_p99_ms: 0.0,
247///         },
248///         ModelEndpoint {
249///             id: "ep-b".to_string(),
250///             url: "http://b".to_string(),
251///             weight: 1,
252///             max_rps: 50.0,
253///             healthy: true,
254///             latency_p99_ms: 0.0,
255///         },
256///     ],
257///     health_check_interval: Duration::from_secs(30),
258///     failover: true,
259/// };
260///
261/// let lb = LoadBalancer::new(config);
262/// let ep = lb.select();
263/// assert!(ep.is_some());
264/// ```
265#[derive(Clone)]
266pub struct LoadBalancer {
267    inner: Arc<Mutex<BalancerInner>>,
268}
269
270impl LoadBalancer {
271    /// Create a new load balancer from the given configuration.
272    pub fn new(config: BalancerConfig) -> Self {
273        Self {
274            inner: Arc::new(Mutex::new(BalancerInner::new(config))),
275        }
276    }
277
278    /// Select the next endpoint using weighted round-robin.
279    ///
280    /// Returns `None` when:
281    /// - No endpoints are configured.
282    /// - All endpoints are unhealthy and `failover = true`.
283    pub fn select(&self) -> Option<ModelEndpoint> {
284        let mut guard = self.inner.lock().ok()?;
285        let idx = guard.select_index()?;
286        guard.endpoint_at(idx)
287    }
288
289    /// Record a successful response from endpoint `id` with the given latency.
290    ///
291    /// Resets the consecutive-failure counter and marks the endpoint healthy.
292    pub fn mark_success(&self, id: &str, latency_ms: f64) {
293        if let Ok(mut g) = self.inner.lock() {
294            g.mark_success(id, latency_ms);
295        }
296    }
297
298    /// Record a failure from endpoint `id`.
299    ///
300    /// After 3 consecutive failures the endpoint is marked unhealthy.
301    pub fn mark_failure(&self, id: &str) {
302        if let Ok(mut g) = self.inner.lock() {
303            g.mark_failure(id);
304        }
305    }
306
307    /// Return a snapshot of current balancer statistics.
308    pub fn stats(&self) -> LoadBalancerStats {
309        self.inner
310            .lock()
311            .map(|g| g.stats())
312            .unwrap_or_default()
313    }
314
315    /// Return the configured failover setting.
316    pub fn failover_enabled(&self) -> bool {
317        self.inner
318            .lock()
319            .map(|g| g.config.failover)
320            .unwrap_or(true)
321    }
322}
323
324// ── Tests ─────────────────────────────────────────────────────────────────────
325
326#[cfg(test)]
327#[allow(clippy::unwrap_used, clippy::expect_used)]
328mod tests {
329    use super::*;
330    use std::time::Duration;
331
332    fn ep(id: &str, weight: u32) -> ModelEndpoint {
333        ModelEndpoint {
334            id: id.to_string(),
335            url: format!("http://{id}"),
336            weight,
337            max_rps: 100.0,
338            healthy: true,
339            latency_p99_ms: 0.0,
340        }
341    }
342
343    fn lb(endpoints: Vec<ModelEndpoint>, failover: bool) -> LoadBalancer {
344        LoadBalancer::new(BalancerConfig {
345            endpoints,
346            health_check_interval: Duration::from_secs(30),
347            failover,
348        })
349    }
350
351    #[test]
352    fn test_empty_returns_none() {
353        let b = lb(vec![], true);
354        assert!(b.select().is_none());
355    }
356
357    #[test]
358    fn test_single_endpoint_always_selected() {
359        let b = lb(vec![ep("a", 1)], true);
360        for _ in 0..10 {
361            assert_eq!(b.select().unwrap().id, "a");
362        }
363    }
364
365    #[test]
366    fn test_equal_weights_round_robin() {
367        let b = lb(vec![ep("a", 1), ep("b", 1)], true);
368        let ids: Vec<String> = (0..4).map(|_| b.select().unwrap().id).collect();
369        // Each endpoint should appear exactly twice in 4 selections.
370        let a_count = ids.iter().filter(|s| s.as_str() == "a").count();
371        let b_count = ids.iter().filter(|s| s.as_str() == "b").count();
372        assert_eq!(a_count, 2);
373        assert_eq!(b_count, 2);
374    }
375
376    #[test]
377    fn test_weighted_distribution() {
378        let b = lb(vec![ep("heavy", 2), ep("light", 1)], true);
379        let selections: Vec<String> = (0..9).map(|_| b.select().unwrap().id).collect();
380        let heavy = selections.iter().filter(|s| s.as_str() == "heavy").count();
381        let light = selections.iter().filter(|s| s.as_str() == "light").count();
382        assert_eq!(heavy, 6);
383        assert_eq!(light, 3);
384    }
385
386    #[test]
387    fn test_mark_failure_three_times_marks_unhealthy() {
388        let b = lb(vec![ep("a", 1)], true);
389        b.mark_failure("a");
390        b.mark_failure("a");
391        assert!(b.select().is_some()); // still healthy after 2
392        b.mark_failure("a");
393        assert!(b.select().is_none()); // unhealthy after 3
394    }
395
396    #[test]
397    fn test_mark_success_recovers_unhealthy() {
398        let b = lb(vec![ep("a", 1)], true);
399        b.mark_failure("a");
400        b.mark_failure("a");
401        b.mark_failure("a");
402        assert!(b.select().is_none());
403        b.mark_success("a", 10.0);
404        assert!(b.select().is_some());
405    }
406
407    #[test]
408    fn test_failover_skips_unhealthy() {
409        let b = lb(vec![ep("a", 1), ep("b", 1)], true);
410        b.mark_failure("a");
411        b.mark_failure("a");
412        b.mark_failure("a");
413        // Only "b" should be selected now.
414        for _ in 0..5 {
415            assert_eq!(b.select().unwrap().id, "b");
416        }
417    }
418
419    #[test]
420    fn test_no_failover_selects_unhealthy() {
421        let b = lb(vec![ep("a", 1), ep("b", 1)], false);
422        b.mark_failure("a");
423        b.mark_failure("a");
424        b.mark_failure("a");
425        // Both should still participate.
426        let ids: Vec<String> = (0..4).map(|_| b.select().unwrap().id).collect();
427        assert!(ids.iter().any(|s| s == "a"));
428        assert!(ids.iter().any(|s| s == "b"));
429    }
430
431    #[test]
432    fn test_total_request_counter() {
433        let b = lb(vec![ep("a", 1)], true);
434        b.select();
435        b.select();
436        b.select();
437        assert_eq!(b.stats().total_requests, 3);
438    }
439
440    #[test]
441    fn test_per_endpoint_request_counter() {
442        let b = lb(vec![ep("a", 1)], true);
443        b.select();
444        b.select();
445        let stats = b.stats();
446        assert_eq!(stats.by_endpoint["a"].requests, 2);
447    }
448
449    #[test]
450    fn test_failure_counter_increments() {
451        let b = lb(vec![ep("a", 1)], true);
452        b.mark_failure("a");
453        b.mark_failure("a");
454        let stats = b.stats();
455        assert_eq!(stats.by_endpoint["a"].failures, 2);
456    }
457
458    #[test]
459    fn test_consecutive_failures_reset_on_success() {
460        let b = lb(vec![ep("a", 1)], true);
461        b.mark_failure("a");
462        b.mark_failure("a");
463        b.mark_success("a", 5.0);
464        let stats = b.stats();
465        assert_eq!(stats.by_endpoint["a"].consecutive_failures, 0);
466    }
467
468    #[test]
469    fn test_unhealthy_count_in_stats() {
470        let b = lb(vec![ep("a", 1), ep("b", 1)], true);
471        b.mark_failure("a");
472        b.mark_failure("a");
473        b.mark_failure("a");
474        let stats = b.stats();
475        assert_eq!(stats.unhealthy_count, 1);
476    }
477
478    #[test]
479    fn test_latency_tracking() {
480        let b = lb(vec![ep("a", 1)], true);
481        b.mark_success("a", 100.0);
482        let stats = b.stats();
483        assert!(stats.by_endpoint["a"].avg_latency_ms > 0.0);
484    }
485
486    #[test]
487    fn test_latency_ema_updates() {
488        let b = lb(vec![ep("a", 1)], true);
489        b.mark_success("a", 100.0);
490        b.mark_success("a", 200.0);
491        let stats = b.stats();
492        // EMA should be between 100 and 200.
493        let avg = stats.by_endpoint["a"].avg_latency_ms;
494        assert!(avg > 100.0 && avg < 200.0);
495    }
496
497    #[test]
498    fn test_all_unhealthy_returns_none_with_failover() {
499        let b = lb(vec![ep("a", 1), ep("b", 1)], true);
500        for _ in 0..3 {
501            b.mark_failure("a");
502            b.mark_failure("b");
503        }
504        assert!(b.select().is_none());
505    }
506
507    #[test]
508    fn test_recovery_after_all_unhealthy() {
509        let b = lb(vec![ep("a", 1)], true);
510        for _ in 0..3 {
511            b.mark_failure("a");
512        }
513        assert!(b.select().is_none());
514        b.mark_success("a", 10.0);
515        assert!(b.select().is_some());
516    }
517
518    #[test]
519    fn test_stats_healthy_status_reflects_mark_failure() {
520        let b = lb(vec![ep("a", 1)], true);
521        assert!(b.stats().by_endpoint["a"].is_healthy);
522        for _ in 0..3 {
523            b.mark_failure("a");
524        }
525        assert!(!b.stats().by_endpoint["a"].is_healthy);
526    }
527
528    #[test]
529    fn test_failover_enabled_flag() {
530        let b = lb(vec![ep("a", 1)], true);
531        assert!(b.failover_enabled());
532        let b2 = lb(vec![ep("a", 1)], false);
533        assert!(!b2.failover_enabled());
534    }
535
536    #[test]
537    fn test_three_endpoints_weighted() {
538        let b = lb(
539            vec![ep("a", 3), ep("b", 2), ep("c", 1)],
540            true,
541        );
542        let selections: Vec<String> = (0..6).map(|_| b.select().unwrap().id).collect();
543        let a = selections.iter().filter(|s| s.as_str() == "a").count();
544        let b_cnt = selections.iter().filter(|s| s.as_str() == "b").count();
545        let c = selections.iter().filter(|s| s.as_str() == "c").count();
546        assert_eq!(a, 3);
547        assert_eq!(b_cnt, 2);
548        assert_eq!(c, 1);
549    }
550
551    #[test]
552    fn test_clone_shares_state() {
553        let b = lb(vec![ep("a", 1)], true);
554        let b2 = b.clone();
555        b.mark_failure("a");
556        b.mark_failure("a");
557        b.mark_failure("a");
558        // b2 should see the same state.
559        assert!(b2.select().is_none());
560    }
561
562    #[test]
563    fn test_unknown_endpoint_mark_does_not_panic() {
564        let b = lb(vec![ep("a", 1)], true);
565        // Should not panic.
566        b.mark_failure("does-not-exist");
567        b.mark_success("does-not-exist", 0.0);
568    }
569
570    #[test]
571    fn test_zero_weight_endpoint_ignored() {
572        // An endpoint with weight 0 contributes 0 to total_weight and is never selected.
573        let b = lb(vec![ep("zero", 0), ep("ok", 1)], true);
574        // Should never return "zero".
575        for _ in 0..10 {
576            let id = b.select().unwrap().id;
577            assert_eq!(id, "ok");
578        }
579    }
580}