1use std::collections::HashMap;
4
5#[derive(Debug, Clone)]
7pub struct Provider {
8 pub id: String,
10 pub name: String,
12 pub api_endpoint: String,
14 pub max_rpm: u32,
16 pub max_tpm: u64,
18 pub priority: u8,
20 pub enabled: bool,
22}
23
24#[derive(Debug, Clone)]
26pub struct ProviderHealth {
27 pub provider_id: String,
29 pub is_healthy: bool,
31 pub error_rate: f64,
33 pub avg_latency_ms: u64,
35 pub last_checked: u64,
37}
38
39#[derive(Debug, Clone)]
41pub enum ProviderSelection {
42 Primary,
44 Fallback {
46 reason: String,
48 },
49 LoadBalanced {
51 weights: Vec<(String, f64)>,
53 },
54}
55
56#[derive(Debug, Clone, Default)]
58pub struct ProviderStats {
59 pub requests: u64,
61 pub errors: u64,
63 pub total_tokens: u64,
65 pub avg_latency_ms: u64,
67}
68
69#[derive(Debug, Default)]
71struct RateLimitState {
72 requests_this_window: u32,
74 tokens_this_window: u64,
76 window_start_ms: u64,
78}
79
80impl RateLimitState {
81 #[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#[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 cost_per_token: HashMap<String, f64>,
115}
116
117impl ProviderManager {
118 pub fn new() -> Self {
120 Self::default()
121 }
122
123 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 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 pub fn update_health(&mut self, health: ProviderHealth) {
139 self.health.insert(health.provider_id.clone(), health);
140 }
141
142 pub fn select_provider(
147 &mut self,
148 selection: &ProviderSelection,
149 now: u64,
150 ) -> Option<&Provider> {
151 match selection {
152 ProviderSelection::Primary => {
153 self.failover_ids(now).into_iter().next().map(|id| &self.providers[&id])
155 }
156 ProviderSelection::Fallback { .. } => {
157 let chain = self.failover_ids(now);
159 chain.into_iter().nth(1).map(|id| &self.providers[&id])
160 }
161 ProviderSelection::LoadBalanced { weights } => {
162 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 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 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 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 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 pub fn provider_stats(&self, provider_id: &str) -> Option<ProviderStats> {
211 self.stats.get(provider_id).cloned()
212 }
213
214 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 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 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 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 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 rl.record(100, 0);
451 assert!(rl.would_exceed(1, 10_000, 100, 0));
453 assert!(!rl.would_exceed(1, 10_000, 100, 60_001));
455 }
456}