1use crate::config::VirtualKeyConfig;
2use anyhow::{bail, Result};
3use std::collections::HashMap;
4use std::sync::RwLock;
5
6#[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 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 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 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 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 state.minute_count += 1;
85 state.request_count += 1;
86
87 Ok(state.clone())
88 }
89
90 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 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 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 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
154fn 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 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}