Skip to main content

weft_core/
vkeys.rs

1use crate::config::VirtualKeyConfig;
2use anyhow::{bail, Result};
3use std::collections::HashMap;
4use std::sync::RwLock;
5
6/// Runtime state for a virtual key.
7#[derive(Debug, Clone)]
8pub struct VirtualKeyState {
9    pub key: String,
10    pub label: Option<String>,
11    pub provider: Option<String>,
12    pub model: Option<String>,
13    pub max_rpm: Option<u32>,
14    pub max_budget: Option<f64>,
15    pub request_count: u64,
16    pub minute_count: u32,
17    pub last_minute_reset: u64,
18}
19
20impl From<&VirtualKeyConfig> for VirtualKeyState {
21    fn from(cfg: &VirtualKeyConfig) -> Self {
22        Self {
23            key: cfg.key.clone(),
24            label: cfg.label.clone(),
25            provider: cfg.provider.clone(),
26            model: cfg.model.clone(),
27            max_rpm: cfg.max_rpm,
28            max_budget: cfg.max_budget,
29            request_count: 0,
30            minute_count: 0,
31            last_minute_reset: now_secs(),
32        }
33    }
34}
35
36pub struct VirtualKeyStore {
37    keys: RwLock<HashMap<String, VirtualKeyState>>,
38}
39
40impl Default for VirtualKeyStore {
41    fn default() -> Self {
42        Self::new()
43    }
44}
45
46impl VirtualKeyStore {
47    pub fn new() -> Self {
48        Self {
49            keys: RwLock::new(HashMap::new()),
50        }
51    }
52
53    /// Load virtual keys from config.
54    pub fn load_from_config(&self, configs: &[VirtualKeyConfig]) {
55        let mut keys = self.keys.write().unwrap();
56        for cfg in configs {
57            keys.insert(cfg.key.clone(), VirtualKeyState::from(cfg));
58        }
59    }
60
61    /// Validate a virtual key and return its state.
62    /// Also checks rate limits.
63    pub fn validate(&self, key: &str) -> Result<VirtualKeyState> {
64        let mut keys = self.keys.write().unwrap();
65        let state = keys
66            .get_mut(key)
67            .ok_or_else(|| anyhow::anyhow!("Invalid virtual key"))?;
68
69        // Reset minute counter if a new minute has started
70        let now = now_secs();
71        if now - state.last_minute_reset >= 60 {
72            state.minute_count = 0;
73            state.last_minute_reset = now;
74        }
75
76        // Check rate limit
77        if let Some(max_rpm) = state.max_rpm {
78            if state.minute_count >= max_rpm {
79                bail!("Rate limit exceeded for virtual key");
80            }
81        }
82
83        // Increment counters
84        state.minute_count += 1;
85        state.request_count += 1;
86
87        Ok(state.clone())
88    }
89
90    /// Create a new virtual key.
91    pub fn create(&self, cfg: VirtualKeyConfig) -> Result<()> {
92        let mut keys = self.keys.write().unwrap();
93        if keys.contains_key(&cfg.key) {
94            bail!("Virtual key '{}' already exists", cfg.key);
95        }
96        keys.insert(cfg.key.clone(), VirtualKeyState::from(&cfg));
97        Ok(())
98    }
99
100    /// Delete a virtual key.
101    pub fn delete(&self, key: &str) -> Result<()> {
102        let mut keys = self.keys.write().unwrap();
103        if keys.remove(key).is_none() {
104            bail!("Virtual key '{}' not found", key);
105        }
106        Ok(())
107    }
108
109    /// List all virtual keys (masked for display).
110    pub fn list(&self) -> Vec<VirtualKeySummary> {
111        let keys = self.keys.read().unwrap();
112        keys.values()
113            .map(|s| VirtualKeySummary {
114                key: s.key.clone(),
115                masked_key: mask_key(&s.key),
116                label: s.label.clone(),
117                provider: s.provider.clone(),
118                model: s.model.clone(),
119                max_rpm: s.max_rpm,
120                request_count: s.request_count,
121            })
122            .collect()
123    }
124
125    /// Get usage for a specific key.
126    pub fn usage(&self, key: &str) -> Option<VirtualKeyUsage> {
127        let keys = self.keys.read().unwrap();
128        keys.get(key).map(|s| VirtualKeyUsage {
129            key: mask_key(&s.key),
130            request_count: s.request_count,
131            minute_count: s.minute_count,
132        })
133    }
134}
135
136#[derive(Debug, Clone, serde::Serialize)]
137pub struct VirtualKeySummary {
138    pub key: String,
139    pub masked_key: String,
140    pub label: Option<String>,
141    pub provider: Option<String>,
142    pub model: Option<String>,
143    pub max_rpm: Option<u32>,
144    pub request_count: u64,
145}
146
147#[derive(Debug, Clone, serde::Serialize)]
148pub struct VirtualKeyUsage {
149    pub key: String,
150    pub request_count: u64,
151    pub minute_count: u32,
152}
153
154/// Mask a key for display: "vk-weft-abcdef" -> "vk-weft-abc***"
155fn mask_key(key: &str) -> String {
156    if key.len() <= 10 {
157        return format!("{}***", &key[..key.len().min(4)]);
158    }
159    format!("{}***", &key[..key.len() - 3])
160}
161
162fn now_secs() -> u64 {
163    std::time::SystemTime::now()
164        .duration_since(std::time::UNIX_EPOCH)
165        .unwrap_or_default()
166        .as_secs()
167}
168
169#[cfg(test)]
170mod tests {
171    use super::*;
172    use crate::config::VirtualKeyConfig;
173
174    fn make_vk(key: &str) -> VirtualKeyConfig {
175        VirtualKeyConfig {
176            key: key.into(),
177            label: Some("test".into()),
178            provider: Some("openrouter".into()),
179            model: Some("claude-sonnet-4".into()),
180            max_rpm: Some(10),
181            max_budget: None,
182        }
183    }
184
185    #[test]
186    fn test_validate_ok() {
187        let store = VirtualKeyStore::new();
188        store.create(make_vk("vk-test-001")).unwrap();
189        let state = store.validate("vk-test-001").unwrap();
190        assert_eq!(state.provider, Some("openrouter".into()));
191        assert_eq!(state.request_count, 1);
192    }
193
194    #[test]
195    fn test_validate_invalid_key() {
196        let store = VirtualKeyStore::new();
197        assert!(store.validate("vk-nonexistent").is_err());
198    }
199
200    #[test]
201    fn test_rate_limit() {
202        let store = VirtualKeyStore::new();
203        let mut cfg = make_vk("vk-rate-test");
204        cfg.max_rpm = Some(2);
205        store.create(cfg).unwrap();
206
207        store.validate("vk-rate-test").unwrap();
208        store.validate("vk-rate-test").unwrap();
209        // Third should fail
210        assert!(store.validate("vk-rate-test").is_err());
211    }
212
213    #[test]
214    fn test_create_duplicate() {
215        let store = VirtualKeyStore::new();
216        store.create(make_vk("vk-dup")).unwrap();
217        assert!(store.create(make_vk("vk-dup")).is_err());
218    }
219
220    #[test]
221    fn test_delete() {
222        let store = VirtualKeyStore::new();
223        store.create(make_vk("vk-del")).unwrap();
224        store.delete("vk-del").unwrap();
225        assert!(store.validate("vk-del").is_err());
226    }
227
228    #[test]
229    fn test_list_masks_keys() {
230        let store = VirtualKeyStore::new();
231        store.create(make_vk("vk-weft-abcdef")).unwrap();
232        let list = store.list();
233        assert_eq!(list.len(), 1);
234        assert_eq!(list[0].key, "vk-weft-abcdef");
235        assert!(list[0].masked_key.contains("***"));
236        assert!(!list[0].masked_key.contains("abcdef"));
237    }
238}