weft_core/defaults/
key_selectors.rs1use crate::layers::key_selector::{ApiKeyState, KeySelectorLayer};
2use anyhow::{bail, Result};
3use async_trait::async_trait;
4use std::collections::HashMap;
5use std::sync::atomic::{AtomicUsize, Ordering};
6use std::sync::Mutex;
7
8pub struct RoundRobinSelector {
10 counters: Mutex<HashMap<String, AtomicUsize>>,
11}
12
13impl Default for RoundRobinSelector {
14 fn default() -> Self {
15 Self::new()
16 }
17}
18
19impl RoundRobinSelector {
20 pub fn new() -> Self {
21 Self {
22 counters: Mutex::new(HashMap::new()),
23 }
24 }
25}
26
27#[async_trait]
28impl KeySelectorLayer for RoundRobinSelector {
29 async fn select(&self, provider: &str, keys: &[ApiKeyState]) -> Result<usize> {
30 let available: Vec<usize> = keys
31 .iter()
32 .enumerate()
33 .filter(|(_, k)| k.enabled && !k.failed)
34 .map(|(i, _)| i)
35 .collect();
36
37 if available.is_empty() {
38 bail!("No available keys for provider '{}'", provider);
39 }
40
41 let mut counters = self.counters.lock().unwrap();
42 let counter = counters
43 .entry(provider.to_string())
44 .or_insert_with(|| AtomicUsize::new(0));
45 let idx = counter.fetch_add(1, Ordering::Relaxed) % available.len();
46 Ok(available[idx])
47 }
48
49 fn mark_failed(&self, _provider: &str, _index: usize) {
50 }
52
53 fn mark_success(&self, _provider: &str, _index: usize) {}
54}
55
56pub struct FailoverSelector;
58
59#[async_trait]
60impl KeySelectorLayer for FailoverSelector {
61 async fn select(&self, provider: &str, keys: &[ApiKeyState]) -> Result<usize> {
62 keys.iter()
63 .enumerate()
64 .find(|(_, k)| k.enabled && !k.failed)
65 .map(|(i, _)| i)
66 .ok_or_else(|| anyhow::anyhow!("No available keys for provider '{}'", provider))
67 }
68
69 fn mark_failed(&self, _provider: &str, _index: usize) {}
70 fn mark_success(&self, _provider: &str, _index: usize) {}
71}
72
73pub struct RandomSelector;
75
76#[async_trait]
77impl KeySelectorLayer for RandomSelector {
78 async fn select(&self, provider: &str, keys: &[ApiKeyState]) -> Result<usize> {
79 let available: Vec<usize> = keys
80 .iter()
81 .enumerate()
82 .filter(|(_, k)| k.enabled && !k.failed)
83 .map(|(i, _)| i)
84 .collect();
85
86 if available.is_empty() {
87 bail!("No available keys for provider '{}'", provider);
88 }
89
90 use rand::Rng;
91 let idx = rand::thread_rng().gen_range(0..available.len());
92 Ok(available[idx])
93 }
94
95 fn mark_failed(&self, _provider: &str, _index: usize) {}
96 fn mark_success(&self, _provider: &str, _index: usize) {}
97}
98
99#[cfg(test)]
100mod tests {
101 use super::*;
102
103 fn make_keys() -> Vec<ApiKeyState> {
104 vec![
105 ApiKeyState {
106 index: 0,
107 value: "sk-aaa".into(),
108 label: None,
109 enabled: true,
110 failed: false,
111 usage_count: 0,
112 },
113 ApiKeyState {
114 index: 1,
115 value: "sk-bbb".into(),
116 label: None,
117 enabled: true,
118 failed: false,
119 usage_count: 0,
120 },
121 ]
122 }
123
124 #[tokio::test]
125 async fn test_round_robin() {
126 let sel = RoundRobinSelector::new();
127 let keys = make_keys();
128 let a = sel.select("p", &keys).await.unwrap();
129 let b = sel.select("p", &keys).await.unwrap();
130 assert_ne!(a, b);
131 }
132
133 #[tokio::test]
134 async fn test_failover_picks_first() {
135 let sel = FailoverSelector;
136 let keys = make_keys();
137 let idx = sel.select("p", &keys).await.unwrap();
138 assert_eq!(idx, 0);
139 }
140
141 #[tokio::test]
142 async fn test_failover_skips_failed() {
143 let sel = FailoverSelector;
144 let mut keys = make_keys();
145 keys[0].failed = true;
146 let idx = sel.select("p", &keys).await.unwrap();
147 assert_eq!(idx, 1);
148 }
149
150 #[tokio::test]
151 async fn test_all_failed_errors() {
152 let sel = FailoverSelector;
153 let mut keys = make_keys();
154 keys[0].failed = true;
155 keys[1].failed = true;
156 assert!(sel.select("p", &keys).await.is_err());
157 }
158}