Skip to main content

weft_core/defaults/
key_selectors.rs

1use 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
8/// Round-robin key selection.
9pub 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        // State is tracked externally via ApiKeyState.failed
51    }
52
53    fn mark_success(&self, _provider: &str, _index: usize) {}
54}
55
56/// Failover: always pick the first non-failed key.
57pub 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
73/// Random key selection.
74pub 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}