tokio_prompt_orchestrator/
load_balancer.rs1use std::collections::HashMap;
21use std::sync::{Arc, Mutex};
22use std::time::Duration;
23
24#[derive(Debug, Clone)]
28pub struct ModelEndpoint {
29 pub id: String,
31 pub url: String,
33 pub weight: u32,
36 pub max_rps: f64,
38 pub healthy: bool,
40 pub latency_p99_ms: f64,
42}
43
44#[derive(Debug, Clone)]
46pub struct BalancerConfig {
47 pub endpoints: Vec<ModelEndpoint>,
49 pub health_check_interval: Duration,
51 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#[derive(Debug, Clone, Default)]
68pub struct EndpointStats {
69 pub requests: u64,
71 pub failures: u64,
73 pub consecutive_failures: u32,
75 pub avg_latency_ms: f64,
77 pub is_healthy: bool,
79}
80
81#[derive(Debug, Clone, Default)]
83pub struct LoadBalancerStats {
84 pub total_requests: u64,
86 pub by_endpoint: HashMap<String, EndpointStats>,
88 pub unhealthy_count: usize,
90}
91
92const FAILURE_THRESHOLD: u32 = 3;
96
97const 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 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 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 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 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 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#[derive(Clone)]
266pub struct LoadBalancer {
267 inner: Arc<Mutex<BalancerInner>>,
268}
269
270impl LoadBalancer {
271 pub fn new(config: BalancerConfig) -> Self {
273 Self {
274 inner: Arc::new(Mutex::new(BalancerInner::new(config))),
275 }
276 }
277
278 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 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 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 pub fn stats(&self) -> LoadBalancerStats {
309 self.inner
310 .lock()
311 .map(|g| g.stats())
312 .unwrap_or_default()
313 }
314
315 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#[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 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()); b.mark_failure("a");
393 assert!(b.select().is_none()); }
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 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 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 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 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 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 let b = lb(vec![ep("zero", 0), ep("ok", 1)], true);
574 for _ in 0..10 {
576 let id = b.select().unwrap().id;
577 assert_eq!(id, "ok");
578 }
579 }
580}