1use std::collections::HashMap;
37use std::time::{SystemTime, UNIX_EPOCH};
38
39use serde::{Deserialize, Serialize};
40use tokio::sync::RwLock;
41use tracing::{debug, warn};
42
43const RATE_LIMIT_COOLDOWN_SECS: u64 = 60;
45
46fn expand_multi_key(env_var: &str, raw: String) -> Vec<(String, String)> {
49 if raw.contains(',') {
50 raw.split(',')
51 .enumerate()
52 .filter_map(|(i, k)| {
53 let k = k.trim().to_string();
54 if k.is_empty() {
55 None
56 } else {
57 Some((format!("{}[{}]", env_var, i), k))
58 }
59 })
60 .collect()
61 } else {
62 vec![(env_var.to_string(), raw)]
63 }
64}
65
66#[derive(Debug, Clone, Serialize, Deserialize)]
68pub struct KeyStats {
69 pub env_var: String,
71 pub total_requests: u64,
73 pub successes: u64,
75 pub failures: u64,
77 pub rate_limits: u64,
79 pub total_latency_ms: u64,
81 pub total_input_tokens: u64,
83 pub total_output_tokens: u64,
85 #[serde(skip)]
87 pub active_requests: u64,
88 #[serde(default)]
90 pub last_rate_limit_at: u64,
91 #[serde(default)]
93 pub last_success_at: u64,
94}
95
96impl KeyStats {
97 fn new(env_var: String) -> Self {
98 Self {
99 env_var,
100 total_requests: 0,
101 successes: 0,
102 failures: 0,
103 rate_limits: 0,
104 total_latency_ms: 0,
105 total_input_tokens: 0,
106 total_output_tokens: 0,
107 active_requests: 0,
108 last_rate_limit_at: 0,
109 last_success_at: 0,
110 }
111 }
112
113 pub fn avg_latency_ms(&self) -> f64 {
115 if self.successes == 0 {
116 return 0.0;
117 }
118 self.total_latency_ms as f64 / self.successes as f64
119 }
120
121 pub fn success_rate(&self) -> f64 {
123 let total = self.successes + self.failures;
124 if total == 0 {
125 return 0.5;
126 }
127 self.successes as f64 / total as f64
128 }
129
130 pub fn is_rate_limited(&self) -> bool {
132 if self.last_rate_limit_at == 0 {
133 return false;
134 }
135 let now = now_unix();
136 now.saturating_sub(self.last_rate_limit_at) < RATE_LIMIT_COOLDOWN_SECS
137 }
138
139 pub fn estimated_cost(&self, input_per_mtok: f64, output_per_mtok: f64) -> f64 {
141 let input_cost = (self.total_input_tokens as f64 / 1_000_000.0) * input_per_mtok;
142 let output_cost = (self.total_output_tokens as f64 / 1_000_000.0) * output_per_mtok;
143 input_cost + output_cost
144 }
145}
146
147#[derive(Debug, Clone)]
149pub struct KeyLease {
150 pub api_key: String,
152 pub env_var: String,
154}
155
156struct EndpointKeys {
158 keys: Vec<(String, String)>,
160 stats: HashMap<String, KeyStats>,
162 next_index: usize,
164}
165
166impl EndpointKeys {
167 fn new(env_vars: Vec<String>) -> Self {
168 let mut keys = Vec::new();
169 let mut stats = HashMap::new();
170
171 for env_var in env_vars {
172 if let Some(raw) = car_secrets::resolve_env_or_keychain(&env_var) {
173 for (sub_var, key) in expand_multi_key(&env_var, raw) {
174 stats.insert(sub_var.clone(), KeyStats::new(sub_var.clone()));
175 keys.push((sub_var, key));
176 }
177 }
178 }
179
180 Self {
181 keys,
182 stats,
183 next_index: 0,
184 }
185 }
186
187 fn lease(&mut self) -> Option<KeyLease> {
197 if self.keys.is_empty() {
198 return None;
199 }
200
201 let mut candidates: Vec<(usize, f64)> = Vec::new();
203 let mut all_cold = true;
204
205 for (idx, (ref env_var, _)) in self.keys.iter().enumerate() {
206 let stats = self.stats.get(env_var);
207
208 if let Some(s) = stats {
210 if s.is_rate_limited() {
211 continue;
212 }
213 if s.total_requests > 0 {
214 all_cold = false;
215 }
216 }
217
218 let score = match stats {
224 Some(s) if s.total_requests > 0 => {
225 let total_tokens = s.total_input_tokens + s.total_output_tokens;
226 let completed = s.successes + s.failures + s.rate_limits;
227 if completed > 0 {
228 let avg_tokens_per_req = total_tokens as f64 / completed as f64;
229 let inflight_estimate = s.active_requests as f64 * avg_tokens_per_req;
230 total_tokens as f64 + inflight_estimate
231 } else {
232 s.active_requests as f64 * 1000.0
234 }
235 }
236 _ => 0.0, };
238
239 candidates.push((idx, score));
240 }
241
242 if all_cold && !candidates.is_empty() {
244 let start = self.next_index % candidates.len();
245 let (idx, _) = candidates[start];
246 self.next_index = start + 1;
247 return self.issue_lease(idx);
248 }
249
250 if candidates.is_empty() {
251 let mut best_idx = 0;
253 let mut oldest_rl = u64::MAX;
254 for (idx, (ref env_var, _)) in self.keys.iter().enumerate() {
255 if let Some(stats) = self.stats.get(env_var) {
256 if stats.last_rate_limit_at < oldest_rl {
257 oldest_rl = stats.last_rate_limit_at;
258 best_idx = idx;
259 }
260 }
261 }
262 let env_var = &self.keys[best_idx].0;
263 warn!(env_var = %env_var, "all keys rate-limited, using oldest-cooldown key");
264 return self.issue_lease(best_idx);
265 }
266
267 candidates.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
269 let (best_idx, _score) = candidates[0];
270 self.issue_lease(best_idx)
271 }
272
273 fn issue_lease(&mut self, idx: usize) -> Option<KeyLease> {
275 let (ref env_var, ref key) = self.keys[idx];
276 let env_var = env_var.clone();
277 let api_key = key.clone();
278
279 if let Some(stats) = self.stats.get_mut(&env_var) {
280 stats.active_requests += 1;
281 stats.total_requests += 1;
282 }
283
284 Some(KeyLease { api_key, env_var })
285 }
286
287 fn report_success(
288 &mut self,
289 env_var: &str,
290 latency_ms: u64,
291 input_tokens: u64,
292 output_tokens: u64,
293 ) {
294 if let Some(stats) = self.stats.get_mut(env_var) {
295 stats.successes += 1;
296 stats.active_requests = stats.active_requests.saturating_sub(1);
297 stats.total_latency_ms += latency_ms;
298 stats.total_input_tokens += input_tokens;
299 stats.total_output_tokens += output_tokens;
300 stats.last_success_at = now_unix();
301 }
302 }
303
304 fn report_failure(&mut self, env_var: &str, is_rate_limit: bool) {
305 if let Some(stats) = self.stats.get_mut(env_var) {
306 stats.active_requests = stats.active_requests.saturating_sub(1);
307 if is_rate_limit {
308 stats.rate_limits += 1;
309 stats.last_rate_limit_at = now_unix();
310 } else {
311 stats.failures += 1;
312 }
313 }
314 }
315}
316
317pub struct KeyPool {
319 endpoints: RwLock<HashMap<String, EndpointKeys>>,
321}
322
323impl KeyPool {
324 pub fn new() -> Self {
325 Self {
326 endpoints: RwLock::new(HashMap::new()),
327 }
328 }
329
330 pub async fn register_endpoint(&self, endpoint: &str, env_vars: Vec<String>) {
333 let mut endpoints = self.endpoints.write().await;
334 let entry = endpoints
335 .entry(endpoint.to_string())
336 .or_insert_with(|| EndpointKeys::new(vec![]));
337
338 let existing_vars: std::collections::HashSet<String> =
339 entry.keys.iter().map(|(v, _)| v.clone()).collect();
340
341 let mut new_keys: Vec<(String, String)> = Vec::new();
342
343 for env_var in env_vars {
344 if existing_vars.contains(&env_var) {
348 continue;
349 }
350 if let Some(raw) = car_secrets::resolve_env_or_keychain(&env_var) {
351 for (sub_var, key) in expand_multi_key(&env_var, raw) {
352 if !existing_vars.contains(&sub_var) {
353 new_keys.push((sub_var, key));
354 }
355 }
356 }
357 }
358
359 for (var, key) in new_keys {
360 entry
361 .stats
362 .entry(var.clone())
363 .or_insert_with(|| KeyStats::new(var.clone()));
364 entry.keys.push((var, key));
365 }
366
367 debug!(
368 endpoint = %endpoint,
369 key_count = entry.keys.len(),
370 "registered endpoint keys"
371 );
372 }
373
374 pub async fn lease(&self, endpoint: &str) -> Option<KeyLease> {
376 let mut endpoints = self.endpoints.write().await;
377 endpoints.get_mut(endpoint)?.lease()
378 }
379
380 pub async fn lease_or_env(&self, endpoint: &str, fallback_env: &str) -> Option<KeyLease> {
385 if let Some(lease) = self.lease(endpoint).await {
387 return Some(lease);
388 }
389
390 if car_secrets::resolve_env_or_keychain(fallback_env).is_some() {
393 self.register_endpoint(endpoint, vec![fallback_env.to_string()])
394 .await;
395 return self.lease(endpoint).await;
396 }
397
398 None
399 }
400
401 pub async fn report_success(
403 &self,
404 endpoint: &str,
405 env_var: &str,
406 latency_ms: u64,
407 input_tokens: u64,
408 output_tokens: u64,
409 ) {
410 let mut endpoints = self.endpoints.write().await;
411 if let Some(ep) = endpoints.get_mut(endpoint) {
412 ep.report_success(env_var, latency_ms, input_tokens, output_tokens);
413 }
414 }
415
416 pub async fn report_failure(&self, endpoint: &str, env_var: &str, is_rate_limit: bool) {
418 let mut endpoints = self.endpoints.write().await;
419 if let Some(ep) = endpoints.get_mut(endpoint) {
420 ep.report_failure(env_var, is_rate_limit);
421 }
422 }
423
424 pub async fn endpoint_stats(&self, endpoint: &str) -> Vec<KeyStats> {
426 let endpoints = self.endpoints.read().await;
427 endpoints
428 .get(endpoint)
429 .map(|ep| ep.stats.values().cloned().collect())
430 .unwrap_or_default()
431 }
432
433 pub async fn all_stats(&self) -> HashMap<String, Vec<KeyStats>> {
435 let endpoints = self.endpoints.read().await;
436 endpoints
437 .iter()
438 .map(|(ep, keys)| (ep.clone(), keys.stats.values().cloned().collect()))
439 .collect()
440 }
441
442 pub async fn total_keys(&self) -> usize {
444 let endpoints = self.endpoints.read().await;
445 endpoints.values().map(|ep| ep.keys.len()).sum()
446 }
447
448 pub async fn available_keys(&self, endpoint: &str) -> usize {
450 let endpoints = self.endpoints.read().await;
451 endpoints
452 .get(endpoint)
453 .map(|ep| {
454 ep.keys
455 .iter()
456 .filter(|(env_var, _)| {
457 ep.stats
458 .get(env_var)
459 .map(|s| !s.is_rate_limited())
460 .unwrap_or(true)
461 })
462 .count()
463 })
464 .unwrap_or(0)
465 }
466
467 pub async fn save_stats(&self, path: &std::path::Path) -> Result<(), std::io::Error> {
469 let endpoints = self.endpoints.read().await;
470 let stats: HashMap<String, Vec<KeyStats>> = endpoints
471 .iter()
472 .map(|(ep, keys)| (ep.clone(), keys.stats.values().cloned().collect()))
473 .collect();
474 let json = serde_json::to_string_pretty(&stats).map_err(std::io::Error::other)?;
475 if let Some(parent) = path.parent() {
476 std::fs::create_dir_all(parent)?;
477 }
478 std::fs::write(path, json)
479 }
480
481 pub async fn load_stats(&self, path: &std::path::Path) -> Result<usize, std::io::Error> {
483 if !path.exists() {
484 return Ok(0);
485 }
486 let json = std::fs::read_to_string(path)?;
487 let saved: HashMap<String, Vec<KeyStats>> = serde_json::from_str(&json)
488 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
489
490 let mut endpoints = self.endpoints.write().await;
491 let mut count = 0;
492 for (endpoint, stats_list) in saved {
493 let ep = endpoints
494 .entry(endpoint)
495 .or_insert_with(|| EndpointKeys::new(vec![]));
496 for stats in stats_list {
497 ep.stats.insert(stats.env_var.clone(), stats);
498 count += 1;
499 }
500 }
501 Ok(count)
502 }
503}
504
505impl Default for KeyPool {
506 fn default() -> Self {
507 Self::new()
508 }
509}
510
511fn now_unix() -> u64 {
512 SystemTime::now()
513 .duration_since(UNIX_EPOCH)
514 .unwrap_or_default()
515 .as_secs()
516}
517
518#[cfg(test)]
519mod tests {
520 use super::*;
521
522 #[tokio::test]
523 async fn single_key_round_trip() {
524 std::env::set_var("TEST_KEY_POOL_1", "sk-test-111");
525
526 let pool = KeyPool::new();
527 pool.register_endpoint("https://api.test.com", vec!["TEST_KEY_POOL_1".into()])
528 .await;
529
530 let lease = pool.lease("https://api.test.com").await.unwrap();
531 assert_eq!(lease.api_key, "sk-test-111");
532 assert_eq!(lease.env_var, "TEST_KEY_POOL_1");
533
534 pool.report_success("https://api.test.com", &lease.env_var, 500, 100, 50)
535 .await;
536
537 let stats = pool.endpoint_stats("https://api.test.com").await;
538 assert_eq!(stats.len(), 1);
539 assert_eq!(stats[0].successes, 1);
540 assert_eq!(stats[0].total_latency_ms, 500);
541
542 std::env::remove_var("TEST_KEY_POOL_1");
543 }
544
545 #[tokio::test]
546 async fn multi_key_cold_start_round_robin() {
547 std::env::set_var("TEST_KEY_POOL_A", "sk-aaa");
548 std::env::set_var("TEST_KEY_POOL_B", "sk-bbb");
549
550 let pool = KeyPool::new();
551 pool.register_endpoint(
552 "https://api.test.com",
553 vec!["TEST_KEY_POOL_A".into(), "TEST_KEY_POOL_B".into()],
554 )
555 .await;
556
557 let l1 = pool.lease("https://api.test.com").await.unwrap();
559 pool.report_success("https://api.test.com", &l1.env_var, 100, 10, 5)
560 .await;
561
562 let l2 = pool.lease("https://api.test.com").await.unwrap();
563 pool.report_success("https://api.test.com", &l2.env_var, 100, 10, 5)
564 .await;
565
566 assert_ne!(l1.env_var, l2.env_var);
568
569 std::env::remove_var("TEST_KEY_POOL_A");
570 std::env::remove_var("TEST_KEY_POOL_B");
571 }
572
573 #[tokio::test]
574 async fn token_aware_prefers_least_used() {
575 std::env::set_var("TEST_KEY_POOL_TA1", "sk-ta1");
576 std::env::set_var("TEST_KEY_POOL_TA2", "sk-ta2");
577
578 let pool = KeyPool::new();
579 pool.register_endpoint(
580 "https://api.test.com",
581 vec!["TEST_KEY_POOL_TA1".into(), "TEST_KEY_POOL_TA2".into()],
582 )
583 .await;
584
585 let l1 = pool.lease("https://api.test.com").await.unwrap();
587 pool.report_success("https://api.test.com", &l1.env_var, 100, 1000, 500)
588 .await;
589
590 let l2 = pool.lease("https://api.test.com").await.unwrap();
591 pool.report_success("https://api.test.com", &l2.env_var, 100, 100, 50)
592 .await;
593
594 let l3 = pool.lease("https://api.test.com").await.unwrap();
597 assert_eq!(
598 l3.env_var, l2.env_var,
599 "should pick the key with fewer tokens"
600 );
601
602 pool.report_success("https://api.test.com", &l3.env_var, 100, 5000, 5000)
604 .await;
605
606 let l4 = pool.lease("https://api.test.com").await.unwrap();
608 assert_eq!(
609 l4.env_var, l1.env_var,
610 "should pick key with fewer tokens after rebalance"
611 );
612
613 std::env::remove_var("TEST_KEY_POOL_TA1");
614 std::env::remove_var("TEST_KEY_POOL_TA2");
615 }
616
617 #[tokio::test]
618 async fn comma_separated_keys() {
619 std::env::set_var("TEST_KEY_POOL_CSV", "sk-one, sk-two, sk-three");
620
621 let pool = KeyPool::new();
622 pool.register_endpoint("https://api.test.com", vec!["TEST_KEY_POOL_CSV".into()])
623 .await;
624
625 assert_eq!(pool.total_keys().await, 3);
626
627 let l1 = pool.lease("https://api.test.com").await.unwrap();
628 assert_eq!(l1.api_key, "sk-one");
629
630 let l2 = pool.lease("https://api.test.com").await.unwrap();
631 assert_eq!(l2.api_key, "sk-two");
632
633 let l3 = pool.lease("https://api.test.com").await.unwrap();
634 assert_eq!(l3.api_key, "sk-three");
635
636 std::env::remove_var("TEST_KEY_POOL_CSV");
637 }
638
639 #[tokio::test]
640 async fn rate_limited_key_skipped() {
641 std::env::set_var("TEST_KEY_POOL_RL1", "sk-rl1");
642 std::env::set_var("TEST_KEY_POOL_RL2", "sk-rl2");
643
644 let pool = KeyPool::new();
645 pool.register_endpoint(
646 "https://api.test.com",
647 vec!["TEST_KEY_POOL_RL1".into(), "TEST_KEY_POOL_RL2".into()],
648 )
649 .await;
650
651 let l1 = pool.lease("https://api.test.com").await.unwrap();
653 pool.report_failure("https://api.test.com", &l1.env_var, true)
654 .await;
655
656 let l2 = pool.lease("https://api.test.com").await.unwrap();
658 assert_ne!(l1.env_var, l2.env_var);
659
660 std::env::remove_var("TEST_KEY_POOL_RL1");
661 std::env::remove_var("TEST_KEY_POOL_RL2");
662 }
663
664 #[tokio::test]
665 async fn lease_or_env_fallback() {
666 std::env::set_var("TEST_KEY_POOL_FB", "sk-fallback");
667
668 let pool = KeyPool::new();
669
670 let lease = pool
672 .lease_or_env("https://api.new.com", "TEST_KEY_POOL_FB")
673 .await
674 .unwrap();
675 assert_eq!(lease.api_key, "sk-fallback");
676
677 assert_eq!(pool.total_keys().await, 1);
679
680 std::env::remove_var("TEST_KEY_POOL_FB");
681 }
682
683 #[tokio::test]
684 async fn no_keys_returns_none() {
685 let pool = KeyPool::new();
686 assert!(pool.lease("https://nonexistent.com").await.is_none());
687 }
688}