Skip to main content

tokio_prompt_orchestrator/
provider_health.rs

1//! # Provider Health Monitor
2//!
3//! Tracks latency, error rates, and availability for each configured
4//! LLM provider. Exposes a health score (0.0–1.0) per provider.
5//!
6//! The monitor maintains a rolling window of latency samples (in milliseconds)
7//! per provider. After every [`record`] call the p50/p95/p99 latencies and
8//! the rolling error rate are recomputed from the current window, and a fresh
9//! health score is derived via the formula:
10//!
11//! ```text
12//! score = (1.0 - error_rate) * latency_weight
13//! where latency_weight = 1.0 / (1.0 + p95_latency_ms / 1000.0)
14//! ```
15//!
16//! [`record`]: ProviderHealthMonitor::record
17
18use std::{
19    collections::HashMap,
20    sync::Arc,
21    time::Instant,
22};
23use tokio::sync::RwLock;
24
25// ---------------------------------------------------------------------------
26// ProviderHealth snapshot
27// ---------------------------------------------------------------------------
28
29/// A point-in-time health snapshot for a single provider.
30#[derive(Debug, Clone)]
31pub struct ProviderHealth {
32    /// Stable identifier matching the provider key used when calling [`record`].
33    ///
34    /// [`record`]: ProviderHealthMonitor::record
35    pub provider_id: String,
36    /// `true` while `consecutive_failures` is below the hard-failure threshold
37    /// (currently 5) **and** `health_score > 0.0`.
38    pub is_reachable: bool,
39    /// 50th-percentile latency over the current sample window (milliseconds).
40    pub p50_latency_ms: f64,
41    /// 95th-percentile latency over the current sample window (milliseconds).
42    pub p95_latency_ms: f64,
43    /// 99th-percentile latency over the current sample window (milliseconds).
44    pub p99_latency_ms: f64,
45    /// Rolling error rate in `[0.0, 1.0]` computed from the sample window.
46    pub error_rate: f32,
47    /// Composite health score in `[0.0, 1.0]`; higher is better.
48    pub health_score: f32,
49    /// Monotonic timestamp of the most recent [`record`] call.
50    ///
51    /// [`record`]: ProviderHealthMonitor::record
52    pub last_check: Instant,
53    /// Number of consecutive failures since the last success.
54    pub consecutive_failures: u32,
55    /// Total successful + failed requests recorded since monitor creation.
56    pub total_requests: u64,
57    /// Total failed requests recorded since monitor creation.
58    pub total_errors: u64,
59}
60
61/// Hard-failure threshold: providers with this many consecutive failures in a
62/// row are marked `is_reachable = false` regardless of their score.
63const CONSECUTIVE_FAILURE_LIMIT: u32 = 5;
64
65impl ProviderHealth {
66    /// Returns `true` when `health_score > 0.5`.
67    ///
68    /// # Panics
69    ///
70    /// Never panics.
71    pub fn is_healthy(&self) -> bool {
72        self.health_score > 0.5
73    }
74
75    /// Returns `true` when `0.2 < health_score <= 0.5`.
76    ///
77    /// # Panics
78    ///
79    /// Never panics.
80    pub fn is_degraded(&self) -> bool {
81        self.health_score > 0.2 && self.health_score <= 0.5
82    }
83
84    /// Returns `true` when `health_score <= 0.2`.
85    ///
86    /// # Panics
87    ///
88    /// Never panics.
89    pub fn is_critical(&self) -> bool {
90        self.health_score <= 0.2
91    }
92}
93
94// ---------------------------------------------------------------------------
95// Internal per-provider mutable state
96// ---------------------------------------------------------------------------
97
98#[derive(Debug)]
99struct ProviderState {
100    /// Rolling window of round-trip latency samples (milliseconds).
101    /// Only successful requests contribute a latency sample; failures
102    /// still increment `error_count` and `request_count`.
103    latency_window: Vec<u64>,
104    /// Parallel window tracking whether each slot was an error (`true`) or
105    /// success (`false`).  Kept in sync with `latency_window` length by
106    /// storing `0` latency for errors; however we track errors separately
107    /// so percentile computation is not skewed by zero-latency error slots.
108    error_window: Vec<bool>,
109    consecutive_failures: u32,
110    total_requests: u64,
111    total_errors: u64,
112    last_check: Instant,
113}
114
115impl ProviderState {
116    fn new() -> Self {
117        Self {
118            latency_window: Vec::new(),
119            error_window: Vec::new(),
120            consecutive_failures: 0,
121            total_requests: 0,
122            total_errors: 0,
123            last_check: Instant::now(),
124        }
125    }
126}
127
128// ---------------------------------------------------------------------------
129// ProviderHealthMonitor
130// ---------------------------------------------------------------------------
131
132/// A thread-safe, cloneable monitor that tracks per-provider health.
133///
134/// All public methods are `async` and acquire an internal `RwLock` for the
135/// minimum duration necessary.
136///
137/// # Clone semantics
138///
139/// Cloning shares the same underlying `Arc`-wrapped state, so all clones
140/// observe the same data.
141#[derive(Clone)]
142pub struct ProviderHealthMonitor {
143    /// Keyed by provider_id.
144    state: Arc<RwLock<HashMap<String, ProviderState>>>,
145    /// Maximum number of samples retained per provider.
146    window_size: usize,
147}
148
149impl ProviderHealthMonitor {
150    /// Create a new monitor with the given rolling-window size.
151    ///
152    /// `window_size` is the maximum number of samples (requests) retained per
153    /// provider when computing percentiles and error rate.  Older samples are
154    /// evicted when the window is full (FIFO).  A minimum of `1` is enforced.
155    ///
156    /// # Panics
157    ///
158    /// Never panics.
159    pub fn new(window_size: usize) -> Self {
160        Self {
161            state: Arc::new(RwLock::new(HashMap::new())),
162            window_size: window_size.max(1),
163        }
164    }
165
166    /// Record the outcome of a single completed request for `provider_id`.
167    ///
168    /// * `latency_ms` — round-trip time in milliseconds.  Ignored (not added
169    ///   to the latency window) when `success` is `false`.
170    /// * `success` — `true` for a 2xx response, `false` for any error.
171    ///
172    /// # Panics
173    ///
174    /// Never panics.
175    pub async fn record(&self, provider_id: &str, latency_ms: u64, success: bool) {
176        let mut map = self.state.write().await;
177        let entry = map
178            .entry(provider_id.to_string())
179            .or_insert_with(ProviderState::new);
180
181        entry.total_requests += 1;
182        entry.last_check = Instant::now();
183
184        if success {
185            entry.consecutive_failures = 0;
186            // Add latency sample, evict oldest if window is full.
187            if entry.latency_window.len() >= self.window_size {
188                entry.latency_window.remove(0);
189            }
190            entry.latency_window.push(latency_ms);
191        } else {
192            entry.total_errors += 1;
193            entry.consecutive_failures += 1;
194        }
195
196        // Error window tracks success/failure for every request.
197        if entry.error_window.len() >= self.window_size {
198            entry.error_window.remove(0);
199        }
200        entry.error_window.push(!success);
201    }
202
203    /// Return a [`ProviderHealth`] snapshot for `provider_id`, or `None` if no
204    /// data has been recorded for that provider yet.
205    ///
206    /// # Panics
207    ///
208    /// Never panics.
209    pub async fn get_health(&self, provider_id: &str) -> Option<ProviderHealth> {
210        let map = self.state.read().await;
211        let entry = map.get(provider_id)?;
212        Some(Self::build_health(provider_id, entry))
213    }
214
215    /// Return all known providers sorted by health score, best first.
216    ///
217    /// # Panics
218    ///
219    /// Never panics.
220    pub async fn ranked_providers(&self) -> Vec<ProviderHealth> {
221        let map = self.state.read().await;
222        let mut snapshots: Vec<ProviderHealth> = map
223            .iter()
224            .map(|(id, state)| Self::build_health(id, state))
225            .collect();
226        // Sort descending by health_score; ties broken by p95 latency ascending.
227        snapshots.sort_by(|a, b| {
228            b.health_score
229                .partial_cmp(&a.health_score)
230                .unwrap_or(std::cmp::Ordering::Equal)
231                .then(
232                    a.p95_latency_ms
233                        .partial_cmp(&b.p95_latency_ms)
234                        .unwrap_or(std::cmp::Ordering::Equal),
235                )
236        });
237        snapshots
238    }
239
240    /// Return the healthiest provider from the given `candidates` list.
241    ///
242    /// Returns `None` if `candidates` is empty or none of the candidates have
243    /// recorded data.
244    ///
245    /// # Panics
246    ///
247    /// Never panics.
248    pub async fn best_provider<'a>(&self, candidates: &'a [String]) -> Option<&'a String> {
249        if candidates.is_empty() {
250            return None;
251        }
252        let map = self.state.read().await;
253        let mut best_idx: Option<usize> = None;
254        let mut best_score = -1.0_f32;
255        for (i, id) in candidates.iter().enumerate() {
256            let score = map
257                .get(id.as_str())
258                .map(|s| Self::build_health(id, s).health_score)
259                .unwrap_or(0.0);
260            if score > best_score {
261                best_score = score;
262                best_idx = Some(i);
263            }
264        }
265        best_idx.map(|i| &candidates[i])
266    }
267
268    /// Returns `true` if `provider_id` has a health score above the "healthy"
269    /// threshold (`> 0.5`) and is not in a hard-failure state.
270    ///
271    /// Providers with no recorded data are considered **usable** (optimistic
272    /// default) so new providers are tried immediately.
273    ///
274    /// # Panics
275    ///
276    /// Never panics.
277    pub async fn is_usable(&self, provider_id: &str) -> bool {
278        let map = self.state.read().await;
279        match map.get(provider_id) {
280            None => true, // no data yet — assume healthy
281            Some(entry) => {
282                let health = Self::build_health(provider_id, entry);
283                health.is_reachable && health.is_healthy()
284            }
285        }
286    }
287
288    // -----------------------------------------------------------------------
289    // Private helpers
290    // -----------------------------------------------------------------------
291
292    /// Build a [`ProviderHealth`] snapshot from a [`ProviderState`] reference.
293    fn build_health(provider_id: &str, entry: &ProviderState) -> ProviderHealth {
294        let p50 = Self::compute_percentile(&entry.latency_window, 50.0);
295        let p95 = Self::compute_percentile(&entry.latency_window, 95.0);
296        let p99 = Self::compute_percentile(&entry.latency_window, 99.0);
297
298        let error_rate = if entry.error_window.is_empty() {
299            0.0_f32
300        } else {
301            let errors = entry.error_window.iter().filter(|&&e| e).count();
302            errors as f32 / entry.error_window.len() as f32
303        };
304
305        let health_score = Self::compute_health_score(error_rate, p95);
306
307        let is_reachable = entry.consecutive_failures < CONSECUTIVE_FAILURE_LIMIT
308            && health_score > 0.0;
309
310        ProviderHealth {
311            provider_id: provider_id.to_string(),
312            is_reachable,
313            p50_latency_ms: p50,
314            p95_latency_ms: p95,
315            p99_latency_ms: p99,
316            error_rate,
317            health_score,
318            last_check: entry.last_check,
319            consecutive_failures: entry.consecutive_failures,
320            total_requests: entry.total_requests,
321            total_errors: entry.total_errors,
322        }
323    }
324
325    /// Compute the `pct`-th percentile of `samples` (e.g. `95.0` for p95).
326    ///
327    /// Returns `0.0` when `samples` is empty.
328    ///
329    /// Uses the nearest-rank method: sort the slice then index at
330    /// `ceil(pct/100 * n) - 1` (clamped to `[0, n-1]`).
331    ///
332    /// # Panics
333    ///
334    /// Never panics.
335    fn compute_percentile(samples: &[u64], pct: f64) -> f64 {
336        if samples.is_empty() {
337            return 0.0;
338        }
339        let mut sorted = samples.to_vec();
340        sorted.sort_unstable();
341        let n = sorted.len();
342        // nearest-rank formula
343        let rank = ((pct / 100.0) * n as f64).ceil() as usize;
344        let idx = rank.saturating_sub(1).min(n - 1);
345        sorted[idx] as f64
346    }
347
348    /// Compute the composite health score from an error rate and p95 latency.
349    ///
350    /// ```text
351    /// score = (1.0 - error_rate) * (1.0 / (1.0 + p95_latency_ms / 1000.0))
352    /// ```
353    ///
354    /// Result is clamped to `[0.0, 1.0]`.
355    ///
356    /// # Panics
357    ///
358    /// Never panics.
359    fn compute_health_score(error_rate: f32, p95_latency_ms: f64) -> f32 {
360        let availability = (1.0 - error_rate).clamp(0.0, 1.0);
361        let latency_weight = (1.0 / (1.0 + p95_latency_ms / 1000.0)) as f32;
362        (availability * latency_weight).clamp(0.0, 1.0)
363    }
364}
365
366// ---------------------------------------------------------------------------
367// Unit tests
368// ---------------------------------------------------------------------------
369
370#[cfg(test)]
371mod tests {
372    use super::*;
373
374    // Helper: build a monitor and record N successes with the given latency.
375    async fn monitor_with_successes(latencies: &[u64]) -> ProviderHealthMonitor {
376        let m = ProviderHealthMonitor::new(100);
377        for &ms in latencies {
378            m.record("p1", ms, true).await;
379        }
380        m
381    }
382
383    #[tokio::test]
384    async fn test_no_data_returns_none() {
385        let m = ProviderHealthMonitor::new(10);
386        assert!(m.get_health("unknown").await.is_none());
387    }
388
389    #[tokio::test]
390    async fn test_is_usable_with_no_data() {
391        let m = ProviderHealthMonitor::new(10);
392        // Providers with no data are optimistically considered usable.
393        assert!(m.is_usable("new_provider").await);
394    }
395
396    #[tokio::test]
397    async fn test_single_success_is_healthy() {
398        let m = monitor_with_successes(&[200]).await;
399        let h = m.get_health("p1").await.unwrap();
400        assert!(h.is_healthy(), "single success should be healthy");
401        assert_eq!(h.error_rate, 0.0);
402        assert_eq!(h.consecutive_failures, 0);
403        assert_eq!(h.total_requests, 1);
404        assert_eq!(h.total_errors, 0);
405    }
406
407    #[tokio::test]
408    async fn test_percentile_computation() {
409        // 10 samples: 100..1000 ms in steps of 100.
410        let latencies: Vec<u64> = (1..=10).map(|i| i * 100).collect();
411        let m = monitor_with_successes(&latencies).await;
412        let h = m.get_health("p1").await.unwrap();
413        // p50 = 500 ms (nearest-rank of 10 samples at 50%)
414        assert_eq!(h.p50_latency_ms, 500.0);
415        // p95 = 1000 ms (ceil(0.95*10)=10, idx=9 → 1000)
416        assert_eq!(h.p95_latency_ms, 1000.0);
417        // p99 = 1000 ms
418        assert_eq!(h.p99_latency_ms, 1000.0);
419    }
420
421    #[tokio::test]
422    async fn test_error_rate_all_failures() {
423        let m = ProviderHealthMonitor::new(10);
424        for _ in 0..5 {
425            m.record("p1", 0, false).await;
426        }
427        let h = m.get_health("p1").await.unwrap();
428        assert_eq!(h.error_rate, 1.0);
429        assert!(h.is_critical(), "all failures → score should be critical");
430    }
431
432    #[tokio::test]
433    async fn test_mixed_error_rate() {
434        let m = ProviderHealthMonitor::new(10);
435        // 8 successes, 2 failures → error_rate = 0.2
436        for _ in 0..8 {
437            m.record("p1", 100, true).await;
438        }
439        for _ in 0..2 {
440            m.record("p1", 0, false).await;
441        }
442        let h = m.get_health("p1").await.unwrap();
443        assert!((h.error_rate - 0.2).abs() < 1e-5, "error_rate should be 0.2");
444    }
445
446    #[tokio::test]
447    async fn test_consecutive_failures_marks_unreachable() {
448        let m = ProviderHealthMonitor::new(20);
449        for _ in 0..CONSECUTIVE_FAILURE_LIMIT {
450            m.record("p1", 0, false).await;
451        }
452        let h = m.get_health("p1").await.unwrap();
453        assert!(!h.is_reachable, "should be unreachable after too many consecutive failures");
454    }
455
456    #[tokio::test]
457    async fn test_ranked_providers_order() {
458        let m = ProviderHealthMonitor::new(20);
459        // good provider: fast, no errors
460        for _ in 0..10 {
461            m.record("good", 50, true).await;
462        }
463        // bad provider: high error rate
464        for _ in 0..8 {
465            m.record("bad", 50, false).await;
466        }
467        for _ in 0..2 {
468            m.record("bad", 50, true).await;
469        }
470        let ranked = m.ranked_providers().await;
471        assert_eq!(ranked[0].provider_id, "good");
472        assert_eq!(ranked[1].provider_id, "bad");
473    }
474
475    #[tokio::test]
476    async fn test_best_provider_from_candidates() {
477        let m = ProviderHealthMonitor::new(20);
478        for _ in 0..10 {
479            m.record("fast", 50, true).await;
480        }
481        for _ in 0..5 {
482            m.record("slow", 0, false).await;
483        }
484        for _ in 0..5 {
485            m.record("slow", 5000, true).await;
486        }
487        let candidates = vec!["slow".to_string(), "fast".to_string()];
488        let best = m.best_provider(&candidates).await;
489        assert_eq!(best.map(|s| s.as_str()), Some("fast"));
490    }
491
492    #[tokio::test]
493    async fn test_window_eviction() {
494        // Window size of 3: only the last 3 samples should affect percentiles.
495        let m = ProviderHealthMonitor::new(3);
496        // Push 10 large latencies, then 3 small ones.
497        for _ in 0..10 {
498            m.record("p1", 9000, true).await;
499        }
500        for _ in 0..3 {
501            m.record("p1", 10, true).await;
502        }
503        let h = m.get_health("p1").await.unwrap();
504        // All window samples should now be 10 ms.
505        assert_eq!(h.p50_latency_ms, 10.0);
506        assert_eq!(h.p95_latency_ms, 10.0);
507    }
508
509    #[tokio::test]
510    async fn test_health_score_formula() {
511        // Zero error rate, 0 ms latency → perfect score of 1.0
512        let score = ProviderHealthMonitor::compute_health_score(0.0, 0.0);
513        assert!((score - 1.0).abs() < 1e-5, "perfect conditions should score 1.0");
514
515        // 100% error rate → score of 0.0
516        let score = ProviderHealthMonitor::compute_health_score(1.0, 0.0);
517        assert!((score - 0.0).abs() < 1e-5, "all errors should score 0.0");
518
519        // 0% error, 1000 ms p95 → latency_weight = 1/(1+1) = 0.5
520        let score = ProviderHealthMonitor::compute_health_score(0.0, 1000.0);
521        assert!((score - 0.5).abs() < 1e-4, "1000 ms p95 should give score 0.5");
522    }
523}