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)> {
65 if raw.contains(',') {
66 raw.split(',')
67 .enumerate()
68 .filter_map(|(i, k)| {
69 let k = k.trim().to_string();
70 if k.is_empty() {
71 None
72 } else {
73 Some((format!("{}[{}]", env_var, i), k))
74 }
75 })
76 .collect()
77 } else {
78 vec![(env_var.to_string(), raw.trim().to_string())]
79 }
80}
81
82#[derive(Debug, Clone, Serialize, Deserialize)]
84pub struct KeyStats {
85 pub env_var: String,
87 pub total_requests: u64,
89 pub successes: u64,
91 pub failures: u64,
93 pub rate_limits: u64,
95 pub total_latency_ms: u64,
97 pub total_input_tokens: u64,
99 pub total_output_tokens: u64,
101 #[serde(skip)]
103 pub active_requests: u64,
104 #[serde(default)]
106 pub last_rate_limit_at: u64,
107 #[serde(default)]
109 pub last_success_at: u64,
110}
111
112impl KeyStats {
113 fn new(env_var: String) -> Self {
114 Self {
115 env_var,
116 total_requests: 0,
117 successes: 0,
118 failures: 0,
119 rate_limits: 0,
120 total_latency_ms: 0,
121 total_input_tokens: 0,
122 total_output_tokens: 0,
123 active_requests: 0,
124 last_rate_limit_at: 0,
125 last_success_at: 0,
126 }
127 }
128
129 pub fn avg_latency_ms(&self) -> f64 {
131 if self.successes == 0 {
132 return 0.0;
133 }
134 self.total_latency_ms as f64 / self.successes as f64
135 }
136
137 pub fn success_rate(&self) -> f64 {
139 let total = self.successes + self.failures;
140 if total == 0 {
141 return 0.5;
142 }
143 self.successes as f64 / total as f64
144 }
145
146 pub fn is_rate_limited(&self) -> bool {
148 if self.last_rate_limit_at == 0 {
149 return false;
150 }
151 let now = now_unix();
152 now.saturating_sub(self.last_rate_limit_at) < RATE_LIMIT_COOLDOWN_SECS
153 }
154
155 pub fn estimated_cost(&self, input_per_mtok: f64, output_per_mtok: f64) -> f64 {
157 let input_cost = (self.total_input_tokens as f64 / 1_000_000.0) * input_per_mtok;
158 let output_cost = (self.total_output_tokens as f64 / 1_000_000.0) * output_per_mtok;
159 input_cost + output_cost
160 }
161}
162
163#[derive(Debug, Clone)]
165pub struct KeyLease {
166 pub api_key: String,
168 pub env_var: String,
170}
171
172struct EndpointKeys {
174 keys: Vec<(String, String)>,
176 stats: HashMap<String, KeyStats>,
178 next_index: usize,
180}
181
182impl EndpointKeys {
183 fn new(env_vars: Vec<String>) -> Self {
184 let mut keys = Vec::new();
185 let mut stats = HashMap::new();
186
187 for env_var in env_vars {
188 if let Some(raw) = car_secrets::resolve_env_or_keychain(&env_var) {
189 for (sub_var, key) in expand_multi_key(&env_var, raw) {
190 stats.insert(sub_var.clone(), KeyStats::new(sub_var.clone()));
191 keys.push((sub_var, key));
192 }
193 }
194 }
195
196 Self {
197 keys,
198 stats,
199 next_index: 0,
200 }
201 }
202
203 fn lease(&mut self) -> Option<KeyLease> {
213 if self.keys.is_empty() {
214 return None;
215 }
216
217 let mut candidates: Vec<(usize, f64)> = Vec::new();
219 let mut all_cold = true;
220
221 for (idx, (ref env_var, _)) in self.keys.iter().enumerate() {
222 let stats = self.stats.get(env_var);
223
224 if let Some(s) = stats {
226 if s.is_rate_limited() {
227 continue;
228 }
229 if s.total_requests > 0 {
230 all_cold = false;
231 }
232 }
233
234 let score = match stats {
240 Some(s) if s.total_requests > 0 => {
241 let total_tokens = s.total_input_tokens + s.total_output_tokens;
242 let completed = s.successes + s.failures + s.rate_limits;
243 if completed > 0 {
244 let avg_tokens_per_req = total_tokens as f64 / completed as f64;
245 let inflight_estimate = s.active_requests as f64 * avg_tokens_per_req;
246 total_tokens as f64 + inflight_estimate
247 } else {
248 s.active_requests as f64 * 1000.0
250 }
251 }
252 _ => 0.0, };
254
255 candidates.push((idx, score));
256 }
257
258 if all_cold && !candidates.is_empty() {
260 let start = self.next_index % candidates.len();
261 let (idx, _) = candidates[start];
262 self.next_index = start + 1;
263 return self.issue_lease(idx);
264 }
265
266 if candidates.is_empty() {
267 let mut best_idx = 0;
269 let mut oldest_rl = u64::MAX;
270 for (idx, (ref env_var, _)) in self.keys.iter().enumerate() {
271 if let Some(stats) = self.stats.get(env_var) {
272 if stats.last_rate_limit_at < oldest_rl {
273 oldest_rl = stats.last_rate_limit_at;
274 best_idx = idx;
275 }
276 }
277 }
278 let env_var = &self.keys[best_idx].0;
279 warn!(env_var = %env_var, "all keys rate-limited, using oldest-cooldown key");
280 return self.issue_lease(best_idx);
281 }
282
283 candidates.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
285 let (best_idx, _score) = candidates[0];
286 self.issue_lease(best_idx)
287 }
288
289 fn issue_lease(&mut self, idx: usize) -> Option<KeyLease> {
291 let (ref env_var, ref key) = self.keys[idx];
292 let env_var = env_var.clone();
293 let api_key = key.clone();
294
295 if let Some(stats) = self.stats.get_mut(&env_var) {
296 stats.active_requests += 1;
297 stats.total_requests += 1;
298 }
299
300 Some(KeyLease { api_key, env_var })
301 }
302
303 fn report_success(
304 &mut self,
305 env_var: &str,
306 latency_ms: u64,
307 input_tokens: u64,
308 output_tokens: u64,
309 ) {
310 if let Some(stats) = self.stats.get_mut(env_var) {
311 stats.successes += 1;
312 stats.active_requests = stats.active_requests.saturating_sub(1);
313 stats.total_latency_ms += latency_ms;
314 stats.total_input_tokens += input_tokens;
315 stats.total_output_tokens += output_tokens;
316 stats.last_success_at = now_unix();
317 }
318 }
319
320 fn report_failure(&mut self, env_var: &str, is_rate_limit: bool) {
321 if let Some(stats) = self.stats.get_mut(env_var) {
322 stats.active_requests = stats.active_requests.saturating_sub(1);
323 if is_rate_limit {
324 stats.rate_limits += 1;
325 stats.last_rate_limit_at = now_unix();
326 } else {
327 stats.failures += 1;
328 }
329 }
330 }
331}
332
333pub struct KeyPool {
335 endpoints: RwLock<HashMap<String, EndpointKeys>>,
337}
338
339impl KeyPool {
340 pub fn new() -> Self {
341 Self {
342 endpoints: RwLock::new(HashMap::new()),
343 }
344 }
345
346 pub async fn register_endpoint(&self, endpoint: &str, env_vars: Vec<String>) {
349 let mut endpoints = self.endpoints.write().await;
350 let entry = endpoints
351 .entry(endpoint.to_string())
352 .or_insert_with(|| EndpointKeys::new(vec![]));
353
354 let existing_vars: std::collections::HashSet<String> =
355 entry.keys.iter().map(|(v, _)| v.clone()).collect();
356
357 let mut new_keys: Vec<(String, String)> = Vec::new();
358
359 for env_var in env_vars {
360 if existing_vars.contains(&env_var) {
364 continue;
365 }
366 if let Some(raw) = car_secrets::resolve_env_or_keychain(&env_var) {
367 for (sub_var, key) in expand_multi_key(&env_var, raw) {
368 if !existing_vars.contains(&sub_var) {
369 new_keys.push((sub_var, key));
370 }
371 }
372 }
373 }
374
375 for (var, key) in new_keys {
376 entry
377 .stats
378 .entry(var.clone())
379 .or_insert_with(|| KeyStats::new(var.clone()));
380 entry.keys.push((var, key));
381 }
382
383 debug!(
384 endpoint = %endpoint,
385 key_count = entry.keys.len(),
386 "registered endpoint keys"
387 );
388 }
389
390 pub async fn lease(&self, endpoint: &str) -> Option<KeyLease> {
392 let mut endpoints = self.endpoints.write().await;
393 endpoints.get_mut(endpoint)?.lease()
394 }
395
396 pub async fn lease_or_env(&self, endpoint: &str, fallback_env: &str) -> Option<KeyLease> {
401 if let Some(lease) = self.lease(endpoint).await {
403 return Some(lease);
404 }
405
406 if car_secrets::resolve_env_or_keychain(fallback_env).is_some() {
409 self.register_endpoint(endpoint, vec![fallback_env.to_string()])
410 .await;
411 return self.lease(endpoint).await;
412 }
413
414 None
415 }
416
417 pub async fn report_success(
419 &self,
420 endpoint: &str,
421 env_var: &str,
422 latency_ms: u64,
423 input_tokens: u64,
424 output_tokens: u64,
425 ) {
426 let mut endpoints = self.endpoints.write().await;
427 if let Some(ep) = endpoints.get_mut(endpoint) {
428 ep.report_success(env_var, latency_ms, input_tokens, output_tokens);
429 }
430 }
431
432 pub async fn report_failure(&self, endpoint: &str, env_var: &str, is_rate_limit: bool) {
434 let mut endpoints = self.endpoints.write().await;
435 if let Some(ep) = endpoints.get_mut(endpoint) {
436 ep.report_failure(env_var, is_rate_limit);
437 }
438 }
439
440 pub async fn endpoint_stats(&self, endpoint: &str) -> Vec<KeyStats> {
442 let endpoints = self.endpoints.read().await;
443 endpoints
444 .get(endpoint)
445 .map(|ep| ep.stats.values().cloned().collect())
446 .unwrap_or_default()
447 }
448
449 pub async fn all_stats(&self) -> HashMap<String, Vec<KeyStats>> {
451 let endpoints = self.endpoints.read().await;
452 endpoints
453 .iter()
454 .map(|(ep, keys)| (ep.clone(), keys.stats.values().cloned().collect()))
455 .collect()
456 }
457
458 pub async fn total_keys(&self) -> usize {
460 let endpoints = self.endpoints.read().await;
461 endpoints.values().map(|ep| ep.keys.len()).sum()
462 }
463
464 pub async fn available_keys(&self, endpoint: &str) -> usize {
466 let endpoints = self.endpoints.read().await;
467 endpoints
468 .get(endpoint)
469 .map(|ep| {
470 ep.keys
471 .iter()
472 .filter(|(env_var, _)| {
473 ep.stats
474 .get(env_var)
475 .map(|s| !s.is_rate_limited())
476 .unwrap_or(true)
477 })
478 .count()
479 })
480 .unwrap_or(0)
481 }
482
483 pub async fn save_stats(&self, path: &std::path::Path) -> Result<(), std::io::Error> {
485 let endpoints = self.endpoints.read().await;
486 let stats: HashMap<String, Vec<KeyStats>> = endpoints
487 .iter()
488 .map(|(ep, keys)| (ep.clone(), keys.stats.values().cloned().collect()))
489 .collect();
490 let json = serde_json::to_string_pretty(&stats).map_err(std::io::Error::other)?;
491 if let Some(parent) = path.parent() {
492 std::fs::create_dir_all(parent)?;
493 }
494 std::fs::write(path, json)
495 }
496
497 pub async fn load_stats(&self, path: &std::path::Path) -> Result<usize, std::io::Error> {
499 if !path.exists() {
500 return Ok(0);
501 }
502 let json = std::fs::read_to_string(path)?;
503 let saved: HashMap<String, Vec<KeyStats>> = serde_json::from_str(&json)
504 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
505
506 let mut endpoints = self.endpoints.write().await;
507 let mut count = 0;
508 for (endpoint, stats_list) in saved {
509 let ep = endpoints
510 .entry(endpoint)
511 .or_insert_with(|| EndpointKeys::new(vec![]));
512 for stats in stats_list {
513 ep.stats.insert(stats.env_var.clone(), stats);
514 count += 1;
515 }
516 }
517 Ok(count)
518 }
519}
520
521impl Default for KeyPool {
522 fn default() -> Self {
523 Self::new()
524 }
525}
526
527fn now_unix() -> u64 {
528 SystemTime::now()
529 .duration_since(UNIX_EPOCH)
530 .unwrap_or_default()
531 .as_secs()
532}
533
534#[cfg(test)]
535mod tests {
536 use super::*;
537
538 #[tokio::test]
539 async fn single_key_round_trip() {
540 std::env::set_var("TEST_KEY_POOL_1", "sk-test-111");
541
542 let pool = KeyPool::new();
543 pool.register_endpoint("https://api.test.com", vec!["TEST_KEY_POOL_1".into()])
544 .await;
545
546 let lease = pool.lease("https://api.test.com").await.unwrap();
547 assert_eq!(lease.api_key, "sk-test-111");
548 assert_eq!(lease.env_var, "TEST_KEY_POOL_1");
549
550 pool.report_success("https://api.test.com", &lease.env_var, 500, 100, 50)
551 .await;
552
553 let stats = pool.endpoint_stats("https://api.test.com").await;
554 assert_eq!(stats.len(), 1);
555 assert_eq!(stats[0].successes, 1);
556 assert_eq!(stats[0].total_latency_ms, 500);
557
558 std::env::remove_var("TEST_KEY_POOL_1");
559 }
560
561 #[tokio::test]
562 async fn multi_key_cold_start_round_robin() {
563 std::env::set_var("TEST_KEY_POOL_A", "sk-aaa");
564 std::env::set_var("TEST_KEY_POOL_B", "sk-bbb");
565
566 let pool = KeyPool::new();
567 pool.register_endpoint(
568 "https://api.test.com",
569 vec!["TEST_KEY_POOL_A".into(), "TEST_KEY_POOL_B".into()],
570 )
571 .await;
572
573 let l1 = pool.lease("https://api.test.com").await.unwrap();
575 pool.report_success("https://api.test.com", &l1.env_var, 100, 10, 5)
576 .await;
577
578 let l2 = pool.lease("https://api.test.com").await.unwrap();
579 pool.report_success("https://api.test.com", &l2.env_var, 100, 10, 5)
580 .await;
581
582 assert_ne!(l1.env_var, l2.env_var);
584
585 std::env::remove_var("TEST_KEY_POOL_A");
586 std::env::remove_var("TEST_KEY_POOL_B");
587 }
588
589 #[tokio::test]
590 async fn token_aware_prefers_least_used() {
591 std::env::set_var("TEST_KEY_POOL_TA1", "sk-ta1");
592 std::env::set_var("TEST_KEY_POOL_TA2", "sk-ta2");
593
594 let pool = KeyPool::new();
595 pool.register_endpoint(
596 "https://api.test.com",
597 vec!["TEST_KEY_POOL_TA1".into(), "TEST_KEY_POOL_TA2".into()],
598 )
599 .await;
600
601 let l1 = pool.lease("https://api.test.com").await.unwrap();
603 pool.report_success("https://api.test.com", &l1.env_var, 100, 1000, 500)
604 .await;
605
606 let l2 = pool.lease("https://api.test.com").await.unwrap();
607 pool.report_success("https://api.test.com", &l2.env_var, 100, 100, 50)
608 .await;
609
610 let l3 = pool.lease("https://api.test.com").await.unwrap();
613 assert_eq!(
614 l3.env_var, l2.env_var,
615 "should pick the key with fewer tokens"
616 );
617
618 pool.report_success("https://api.test.com", &l3.env_var, 100, 5000, 5000)
620 .await;
621
622 let l4 = pool.lease("https://api.test.com").await.unwrap();
624 assert_eq!(
625 l4.env_var, l1.env_var,
626 "should pick key with fewer tokens after rebalance"
627 );
628
629 std::env::remove_var("TEST_KEY_POOL_TA1");
630 std::env::remove_var("TEST_KEY_POOL_TA2");
631 }
632
633 #[tokio::test]
634 async fn comma_separated_keys() {
635 std::env::set_var("TEST_KEY_POOL_CSV", "sk-one, sk-two, sk-three");
636
637 let pool = KeyPool::new();
638 pool.register_endpoint("https://api.test.com", vec!["TEST_KEY_POOL_CSV".into()])
639 .await;
640
641 assert_eq!(pool.total_keys().await, 3);
642
643 let l1 = pool.lease("https://api.test.com").await.unwrap();
644 assert_eq!(l1.api_key, "sk-one");
645
646 let l2 = pool.lease("https://api.test.com").await.unwrap();
647 assert_eq!(l2.api_key, "sk-two");
648
649 let l3 = pool.lease("https://api.test.com").await.unwrap();
650 assert_eq!(l3.api_key, "sk-three");
651
652 std::env::remove_var("TEST_KEY_POOL_CSV");
653 }
654
655 #[tokio::test]
656 async fn rate_limited_key_skipped() {
657 std::env::set_var("TEST_KEY_POOL_RL1", "sk-rl1");
658 std::env::set_var("TEST_KEY_POOL_RL2", "sk-rl2");
659
660 let pool = KeyPool::new();
661 pool.register_endpoint(
662 "https://api.test.com",
663 vec!["TEST_KEY_POOL_RL1".into(), "TEST_KEY_POOL_RL2".into()],
664 )
665 .await;
666
667 let l1 = pool.lease("https://api.test.com").await.unwrap();
669 pool.report_failure("https://api.test.com", &l1.env_var, true)
670 .await;
671
672 let l2 = pool.lease("https://api.test.com").await.unwrap();
674 assert_ne!(l1.env_var, l2.env_var);
675
676 std::env::remove_var("TEST_KEY_POOL_RL1");
677 std::env::remove_var("TEST_KEY_POOL_RL2");
678 }
679
680 #[test]
690 fn a_single_key_is_trimmed_like_a_pooled_one() {
691 let pasted = expand_multi_key("PROVIDER_API_KEY", " sk-ant-abc123\n".to_string());
692 assert_eq!(
693 pasted,
694 vec![("PROVIDER_API_KEY".to_string(), "sk-ant-abc123".to_string())],
695 "a single key must be trimmed before it becomes an auth header"
696 );
697
698 let pooled = expand_multi_key("PROVIDER_API_KEY", " k1 , k2\n".to_string());
700 assert_eq!(
701 pooled,
702 vec![
703 ("PROVIDER_API_KEY[0]".to_string(), "k1".to_string()),
704 ("PROVIDER_API_KEY[1]".to_string(), "k2".to_string()),
705 ]
706 );
707 }
708
709 #[tokio::test]
710 async fn lease_or_env_fallback() {
711 std::env::set_var("TEST_KEY_POOL_FB", "sk-fallback");
712
713 let pool = KeyPool::new();
714
715 let lease = pool
717 .lease_or_env("https://api.new.com", "TEST_KEY_POOL_FB")
718 .await
719 .unwrap();
720 assert_eq!(lease.api_key, "sk-fallback");
721
722 assert_eq!(pool.total_keys().await, 1);
724
725 std::env::remove_var("TEST_KEY_POOL_FB");
726 }
727
728 #[tokio::test]
729 async fn missing_key_resolution_uses_the_test_file_backend_by_default() {
730 if !crate::run_in_isolated_test_process(
731 "key_pool::tests::missing_key_resolution_uses_the_test_file_backend_by_default",
732 "CAR_KEY_POOL_TEST_FILE_BACKEND_CHILD",
733 ) {
734 return;
735 }
736
737 const MISSING_KEY: &str = "CAR_KEY_POOL_MISSING_TEST_KEY";
738 unsafe {
739 std::env::remove_var(MISSING_KEY);
740 std::env::remove_var("CAR_SECRETS_FILE_DIR");
741 std::env::remove_var("CAR_TEST_NATIVE_KEYCHAIN");
742 std::env::remove_var("CAR_AGENT_ID");
743 std::env::remove_var("CAR_AGENT_TOKEN");
744 }
745 let before = car_secrets::secret_store_activity();
749 assert_eq!(car_secrets::resolve_env_or_keychain(MISSING_KEY), None);
750 let after = car_secrets::secret_store_activity();
751 assert_eq!(after.get_attempts, before.get_attempts + 1);
752 }
753
754 #[tokio::test]
755 async fn no_keys_returns_none() {
756 let pool = KeyPool::new();
757 assert!(pool.lease("https://nonexistent.com").await.is_none());
758 }
759}