Skip to main content

tokio_prompt_orchestrator/
provider_manager.rs

1//! Multi-provider LLM manager with failover, load-balancing, and rate-limit tracking.
2
3use std::collections::HashMap;
4
5/// Configuration for a single LLM provider.
6#[derive(Debug, Clone)]
7pub struct Provider {
8    /// Unique stable identifier (e.g. `"openai"`, `"anthropic"`).
9    pub id: String,
10    /// Human-readable display name.
11    pub name: String,
12    /// Base URL of the provider's inference API.
13    pub api_endpoint: String,
14    /// Maximum requests per minute allowed by the provider.
15    pub max_rpm: u32,
16    /// Maximum tokens per minute allowed by the provider.
17    pub max_tpm: u64,
18    /// Lower number = higher priority (0 is highest).
19    pub priority: u8,
20    /// Whether this provider is currently enabled.
21    pub enabled: bool,
22}
23
24/// Point-in-time health snapshot for a provider.
25#[derive(Debug, Clone)]
26pub struct ProviderHealth {
27    /// ID of the provider this snapshot belongs to.
28    pub provider_id: String,
29    /// Whether the provider is currently reachable and responding.
30    pub is_healthy: bool,
31    /// Fraction of recent requests that failed (0.0 – 1.0).
32    pub error_rate: f64,
33    /// Rolling average round-trip latency in milliseconds.
34    pub avg_latency_ms: u64,
35    /// Wall-clock timestamp when the health check ran (ms since epoch).
36    pub last_checked: u64,
37}
38
39/// Strategy used to choose which provider handles a request.
40#[derive(Debug, Clone)]
41pub enum ProviderSelection {
42    /// Always use the highest-priority healthy provider.
43    Primary,
44    /// Use the next available provider; `reason` explains why the primary was skipped.
45    Fallback {
46        /// Human-readable reason for the fallback (e.g. `"primary rate-limited"`).
47        reason: String,
48    },
49    /// Distribute load proportionally according to the supplied weights.
50    LoadBalanced {
51        /// `(provider_id, weight)` pairs; weights need not sum to 1.
52        weights: Vec<(String, f64)>,
53    },
54}
55
56/// Cumulative statistics for a single provider.
57#[derive(Debug, Clone, Default)]
58pub struct ProviderStats {
59    /// Total requests sent to this provider.
60    pub requests: u64,
61    /// Total failed requests.
62    pub errors: u64,
63    /// Total tokens consumed.
64    pub total_tokens: u64,
65    /// Exponentially-weighted moving average of latency in ms.
66    pub avg_latency_ms: u64,
67}
68
69/// Per-provider rate-limit window state (sliding 60-second window, simplified).
70#[derive(Debug, Default)]
71struct RateLimitState {
72    /// Requests issued in the current minute window.
73    requests_this_window: u32,
74    /// Tokens issued in the current minute window.
75    tokens_this_window: u64,
76    /// Start of the current 60-second window (ms since epoch).
77    window_start_ms: u64,
78}
79
80impl RateLimitState {
81    /// Advance the window if more than 60 seconds have passed, then return
82    /// whether a request consuming `tokens` would stay within the limits.
83    #[cfg(test)]
84    fn would_exceed(&mut self, max_rpm: u32, max_tpm: u64, tokens: u64, now: u64) -> bool {
85        const WINDOW_MS: u64 = 60_000;
86        if now.saturating_sub(self.window_start_ms) >= WINDOW_MS {
87            self.requests_this_window = 0;
88            self.tokens_this_window = 0;
89            self.window_start_ms = now;
90        }
91        self.requests_this_window >= max_rpm || self.tokens_this_window + tokens > max_tpm
92    }
93
94    fn record(&mut self, tokens: u64, now: u64) {
95        const WINDOW_MS: u64 = 60_000;
96        if now.saturating_sub(self.window_start_ms) >= WINDOW_MS {
97            self.requests_this_window = 0;
98            self.tokens_this_window = 0;
99            self.window_start_ms = now;
100        }
101        self.requests_this_window += 1;
102        self.tokens_this_window += tokens;
103    }
104}
105
106/// Central registry of LLM providers with health, stats, and rate-limit tracking.
107#[derive(Debug, Default)]
108pub struct ProviderManager {
109    providers: HashMap<String, Provider>,
110    health: HashMap<String, ProviderHealth>,
111    stats: HashMap<String, ProviderStats>,
112    rate_limits: HashMap<String, RateLimitState>,
113    /// Estimated cost-per-token for each provider (USD).  Set externally.
114    cost_per_token: HashMap<String, f64>,
115}
116
117impl ProviderManager {
118    /// Create an empty manager.
119    pub fn new() -> Self {
120        Self::default()
121    }
122
123    /// Register (or replace) a provider.
124    pub fn register(&mut self, provider: Provider) {
125        let id = provider.id.clone();
126        self.providers.insert(id.clone(), provider);
127        self.rate_limits.entry(id).or_default();
128    }
129
130    /// Set the estimated cost per token (USD) for a provider.
131    ///
132    /// Used by `best_provider_for_budget`.
133    pub fn set_cost_per_token(&mut self, provider_id: &str, cost: f64) {
134        self.cost_per_token.insert(provider_id.to_string(), cost);
135    }
136
137    /// Replace the stored health snapshot for a provider.
138    pub fn update_health(&mut self, health: ProviderHealth) {
139        self.health.insert(health.provider_id.clone(), health);
140    }
141
142    /// Choose a provider according to `selection`.
143    ///
144    /// Rate-limited or disabled providers are skipped.  Returns `None` if no
145    /// eligible provider exists.
146    pub fn select_provider(
147        &mut self,
148        selection: &ProviderSelection,
149        now: u64,
150    ) -> Option<&Provider> {
151        match selection {
152            ProviderSelection::Primary => {
153                // Highest-priority (lowest priority number) healthy enabled provider.
154                self.failover_ids(now).into_iter().next().map(|id| &self.providers[&id])
155            }
156            ProviderSelection::Fallback { .. } => {
157                // Same as failover but skip the first (primary).
158                let chain = self.failover_ids(now);
159                chain.into_iter().nth(1).map(|id| &self.providers[&id])
160            }
161            ProviderSelection::LoadBalanced { weights } => {
162                // Pick the highest-weight enabled+healthy provider from the list.
163                let best = weights
164                    .iter()
165                    .filter(|(pid, _)| self.is_eligible(pid, 0, now))
166                    .max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
167                    .map(|(pid, _)| pid.clone());
168                best.and_then(|id| self.providers.get(&id))
169            }
170        }
171    }
172
173    /// Return healthy, enabled providers sorted by ascending priority (lowest = best).
174    pub fn failover_chain(&self) -> Vec<&Provider> {
175        let mut eligible: Vec<&Provider> = self
176            .providers
177            .values()
178            .filter(|p| p.enabled && self.is_healthy(&p.id))
179            .collect();
180        eligible.sort_by_key(|p| p.priority);
181        eligible
182    }
183
184    /// Record the outcome of a request to a provider.
185    pub fn record_request(
186        &mut self,
187        provider_id: &str,
188        tokens: u64,
189        latency_ms: u64,
190        success: bool,
191    ) {
192        let rl = self.rate_limits.entry(provider_id.to_string()).or_default();
193        // Use a dummy `now` of 0 for recording — the window was already checked at selection time.
194        rl.record(tokens, 0);
195
196        let stats = self.stats.entry(provider_id.to_string()).or_default();
197        stats.requests += 1;
198        if !success {
199            stats.errors += 1;
200        }
201        stats.total_tokens += tokens;
202        // Exponential moving average: α = 0.1
203        const ALPHA: f64 = 0.1;
204        stats.avg_latency_ms = (ALPHA * latency_ms as f64
205            + (1.0 - ALPHA) * stats.avg_latency_ms as f64)
206            .round() as u64;
207    }
208
209    /// Return the cumulative stats for a provider, if any requests have been recorded.
210    pub fn provider_stats(&self, provider_id: &str) -> Option<ProviderStats> {
211        self.stats.get(provider_id).cloned()
212    }
213
214    /// Return the cheapest healthy provider whose cost-per-token is ≤ `cost_per_token_limit`.
215    pub fn best_provider_for_budget(&self, cost_per_token_limit: f64) -> Option<&Provider> {
216        self.providers
217            .values()
218            .filter(|p| p.enabled && self.is_healthy(&p.id))
219            .filter(|p| {
220                self.cost_per_token
221                    .get(&p.id)
222                    .is_some_and(|&c| c <= cost_per_token_limit)
223            })
224            .min_by(|a, b| {
225                let ca = self.cost_per_token.get(&a.id).copied().unwrap_or(f64::MAX);
226                let cb = self.cost_per_token.get(&b.id).copied().unwrap_or(f64::MAX);
227                ca.partial_cmp(&cb).unwrap_or(std::cmp::Ordering::Equal)
228            })
229    }
230
231    // ── internal helpers ──────────────────────────────────────────────────────
232
233    fn is_healthy(&self, provider_id: &str) -> bool {
234        self.health
235            .get(provider_id)
236            .is_none_or(|h| h.is_healthy)
237    }
238
239    fn is_eligible(&self, provider_id: &str, tokens: u64, now: u64) -> bool {
240        let provider = match self.providers.get(provider_id) {
241            Some(p) => p,
242            None => return false,
243        };
244        if !provider.enabled {
245            return false;
246        }
247        if !self.is_healthy(provider_id) {
248            return false;
249        }
250        // We need mutable access for the rate-limit check; we approximate by
251        // checking the current window counts without advancing the window.
252        if let Some(rl) = self.rate_limits.get(provider_id) {
253            const WINDOW_MS: u64 = 60_000;
254            let in_window = now.saturating_sub(rl.window_start_ms) < WINDOW_MS;
255            if in_window {
256                if rl.requests_this_window >= provider.max_rpm {
257                    return false;
258                }
259                if rl.tokens_this_window + tokens > provider.max_tpm {
260                    return false;
261                }
262            }
263        }
264        true
265    }
266
267    /// Build the failover chain as a list of IDs (cheaply cloned).
268    fn failover_ids(&mut self, now: u64) -> Vec<String> {
269        let mut eligible: Vec<(u8, String)> = self
270            .providers
271            .values()
272            .filter(|p| p.enabled && self.is_healthy(&p.id))
273            .map(|p| (p.priority, p.id.clone()))
274            .collect();
275        eligible.sort_by_key(|(pri, _)| *pri);
276
277        // Filter by rate limits (read-only check; we don't consume quota here).
278        eligible
279            .into_iter()
280            .filter(|(_, id)| self.is_eligible(id, 0, now))
281            .map(|(_, id)| id)
282            .collect()
283    }
284}
285
286#[cfg(test)]
287mod tests {
288    use super::*;
289
290    fn make_provider(id: &str, priority: u8, enabled: bool) -> Provider {
291        Provider {
292            id: id.to_string(),
293            name: id.to_string(),
294            api_endpoint: format!("https://{}.example.com/v1", id),
295            max_rpm: 60,
296            max_tpm: 100_000,
297            priority,
298            enabled,
299        }
300    }
301
302    fn healthy(id: &str) -> ProviderHealth {
303        ProviderHealth {
304            provider_id: id.to_string(),
305            is_healthy: true,
306            error_rate: 0.0,
307            avg_latency_ms: 50,
308            last_checked: 0,
309        }
310    }
311
312    fn unhealthy(id: &str) -> ProviderHealth {
313        ProviderHealth {
314            provider_id: id.to_string(),
315            is_healthy: false,
316            error_rate: 1.0,
317            avg_latency_ms: 5000,
318            last_checked: 0,
319        }
320    }
321
322    #[test]
323    fn register_and_select_primary() {
324        let mut mgr = ProviderManager::new();
325        mgr.register(make_provider("openai", 0, true));
326        mgr.update_health(healthy("openai"));
327        let p = mgr.select_provider(&ProviderSelection::Primary, 0).unwrap();
328        assert_eq!(p.id, "openai");
329    }
330
331    #[test]
332    fn select_primary_skips_disabled() {
333        let mut mgr = ProviderManager::new();
334        mgr.register(make_provider("openai", 0, false));
335        mgr.register(make_provider("anthropic", 1, true));
336        mgr.update_health(healthy("openai"));
337        mgr.update_health(healthy("anthropic"));
338        let p = mgr.select_provider(&ProviderSelection::Primary, 0).unwrap();
339        assert_eq!(p.id, "anthropic");
340    }
341
342    #[test]
343    fn select_primary_skips_unhealthy() {
344        let mut mgr = ProviderManager::new();
345        mgr.register(make_provider("openai", 0, true));
346        mgr.register(make_provider("anthropic", 1, true));
347        mgr.update_health(unhealthy("openai"));
348        mgr.update_health(healthy("anthropic"));
349        let p = mgr.select_provider(&ProviderSelection::Primary, 0).unwrap();
350        assert_eq!(p.id, "anthropic");
351    }
352
353    #[test]
354    fn select_primary_none_when_all_unhealthy() {
355        let mut mgr = ProviderManager::new();
356        mgr.register(make_provider("openai", 0, true));
357        mgr.update_health(unhealthy("openai"));
358        assert!(mgr.select_provider(&ProviderSelection::Primary, 0).is_none());
359    }
360
361    #[test]
362    fn failover_chain_sorted_by_priority() {
363        let mut mgr = ProviderManager::new();
364        mgr.register(make_provider("b", 2, true));
365        mgr.register(make_provider("a", 0, true));
366        mgr.register(make_provider("c", 1, true));
367        mgr.update_health(healthy("a"));
368        mgr.update_health(healthy("b"));
369        mgr.update_health(healthy("c"));
370        let chain: Vec<&str> = mgr.failover_chain().iter().map(|p| p.id.as_str()).collect();
371        assert_eq!(chain, vec!["a", "c", "b"]);
372    }
373
374    #[test]
375    fn fallback_selection_skips_primary() {
376        let mut mgr = ProviderManager::new();
377        mgr.register(make_provider("primary", 0, true));
378        mgr.register(make_provider("backup", 1, true));
379        mgr.update_health(healthy("primary"));
380        mgr.update_health(healthy("backup"));
381        let p = mgr
382            .select_provider(
383                &ProviderSelection::Fallback {
384                    reason: "test".to_string(),
385                },
386                0,
387            )
388            .unwrap();
389        assert_eq!(p.id, "backup");
390    }
391
392    #[test]
393    fn load_balanced_picks_highest_weight() {
394        let mut mgr = ProviderManager::new();
395        mgr.register(make_provider("a", 0, true));
396        mgr.register(make_provider("b", 1, true));
397        mgr.update_health(healthy("a"));
398        mgr.update_health(healthy("b"));
399        let weights = vec![("a".to_string(), 0.3), ("b".to_string(), 0.7)];
400        let p = mgr
401            .select_provider(&ProviderSelection::LoadBalanced { weights }, 0)
402            .unwrap();
403        assert_eq!(p.id, "b");
404    }
405
406    #[test]
407    fn record_request_updates_stats() {
408        let mut mgr = ProviderManager::new();
409        mgr.register(make_provider("openai", 0, true));
410        mgr.record_request("openai", 500, 100, true);
411        mgr.record_request("openai", 200, 200, false);
412        let stats = mgr.provider_stats("openai").unwrap();
413        assert_eq!(stats.requests, 2);
414        assert_eq!(stats.errors, 1);
415        assert_eq!(stats.total_tokens, 700);
416    }
417
418    #[test]
419    fn provider_stats_none_for_unknown() {
420        let mgr = ProviderManager::new();
421        assert!(mgr.provider_stats("ghost").is_none());
422    }
423
424    #[test]
425    fn best_provider_for_budget_finds_cheapest() {
426        let mut mgr = ProviderManager::new();
427        mgr.register(make_provider("cheap", 1, true));
428        mgr.register(make_provider("expensive", 0, true));
429        mgr.update_health(healthy("cheap"));
430        mgr.update_health(healthy("expensive"));
431        mgr.set_cost_per_token("cheap", 0.000001);
432        mgr.set_cost_per_token("expensive", 0.00001);
433        let p = mgr.best_provider_for_budget(0.000005).unwrap();
434        assert_eq!(p.id, "cheap");
435    }
436
437    #[test]
438    fn best_provider_for_budget_none_when_all_exceed() {
439        let mut mgr = ProviderManager::new();
440        mgr.register(make_provider("pricey", 0, true));
441        mgr.update_health(healthy("pricey"));
442        mgr.set_cost_per_token("pricey", 0.01);
443        assert!(mgr.best_provider_for_budget(0.000001).is_none());
444    }
445
446    #[test]
447    fn rate_limit_state_advances_window() {
448        let mut rl = RateLimitState::default();
449        // Within the window, one request should not exceed rpm=1.
450        rl.record(100, 0);
451        // Now at rpm=1; another request should exceed.
452        assert!(rl.would_exceed(1, 10_000, 100, 0));
453        // After 60s, window resets and request is allowed again.
454        assert!(!rl.would_exceed(1, 10_000, 100, 60_001));
455    }
456}