1use std::collections::HashSet;
23
24use anyhow::Result;
25
26use crate::config::LlmCfg;
27
28pub struct VaultEntry {
30 virtual_key: String,
32 provider_key: String,
33 allowed_models: HashSet<String>,
35 label: String,
37}
38
39impl VaultEntry {
40 pub fn provider_key(&self) -> &str {
42 &self.provider_key
43 }
44
45 pub fn label(&self) -> &str {
47 &self.label
48 }
49
50 pub fn model_allowed(&self, model: Option<&str>) -> bool {
54 self.allowed_models.is_empty() || model.is_some_and(|m| self.allowed_models.contains(m))
55 }
56}
57
58pub struct KeyVault {
61 entries: Vec<VaultEntry>,
62}
63
64impl KeyVault {
65 pub fn build(cfg: &LlmCfg) -> Result<Option<KeyVault>> {
69 if cfg.keys.is_empty() {
70 return Ok(None);
71 }
72 let mut entries: Vec<VaultEntry> = Vec::with_capacity(cfg.keys.len());
73 for (i, k) in cfg.keys.iter().enumerate() {
74 anyhow::ensure!(
75 !k.virtual_key.is_empty(),
76 "llm.keys[{i}].virtual_key must not be empty"
77 );
78 anyhow::ensure!(
79 k.virtual_key == k.virtual_key.trim(),
80 "llm.keys[{i}].virtual_key must not have leading/trailing whitespace"
81 );
82 anyhow::ensure!(
83 !k.provider_key.is_empty(),
84 "llm.keys[{i}].provider_key must not be empty"
85 );
86 anyhow::ensure!(
87 k.provider_key == k.provider_key.trim(),
88 "llm.keys[{i}].provider_key must not have leading/trailing whitespace"
89 );
90 anyhow::ensure!(
91 !entries.iter().any(|e| e.virtual_key == k.virtual_key),
92 "llm.keys[{i}]: duplicate virtual_key"
93 );
94 let label = if k.label.trim().is_empty() {
95 format!("key{i}")
96 } else {
97 k.label.clone()
98 };
99 entries.push(VaultEntry {
100 virtual_key: k.virtual_key.clone(),
101 provider_key: k.provider_key.clone(),
102 allowed_models: k.allowed_models.iter().cloned().collect(),
103 label,
104 });
105 }
106 Ok(Some(KeyVault { entries }))
107 }
108
109 pub fn lookup(&self, presented: &str) -> Option<&VaultEntry> {
113 let mut matched: Option<&VaultEntry> = None;
114 for entry in &self.entries {
115 if constant_time_eq(entry.virtual_key.as_bytes(), presented.as_bytes()) {
116 matched = Some(entry);
117 }
118 }
119 matched
120 }
121}
122
123fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
128 let mut diff = a.len() ^ b.len();
129 for i in 0..a.len().max(b.len()) {
130 let x = a.get(i).copied().unwrap_or(0);
131 let y = b.get(i).copied().unwrap_or(0);
132 diff |= (x ^ y) as usize;
133 }
134 diff == 0
135}
136
137#[cfg(test)]
138mod tests {
139 use super::*;
140 use crate::config::KeyEntryCfg;
141
142 fn cfg(keys: Vec<KeyEntryCfg>) -> LlmCfg {
143 LlmCfg {
144 enabled: true,
145 keys,
146 ..Default::default()
147 }
148 }
149
150 fn entry(virtual_key: &str, provider_key: &str, models: &[&str]) -> KeyEntryCfg {
151 KeyEntryCfg {
152 virtual_key: virtual_key.into(),
153 provider_key: provider_key.into(),
154 allowed_models: models.iter().map(|s| s.to_string()).collect(),
155 label: String::new(),
156 }
157 }
158
159 #[test]
160 fn build_is_none_without_keys() {
161 assert!(KeyVault::build(&LlmCfg::default()).unwrap().is_none());
162 }
163
164 #[test]
165 fn lookup_maps_virtual_to_provider_key() {
166 let vault = KeyVault::build(&cfg(vec![entry("sk-virt-a", "sk-real-a", &[])]))
167 .unwrap()
168 .unwrap();
169 let e = vault.lookup("sk-virt-a").expect("known key resolves");
170 assert_eq!(e.provider_key(), "sk-real-a");
171 assert!(vault.lookup("sk-virt-unknown").is_none());
173 assert!(vault.lookup("sk-real-a").is_none());
175 }
176
177 #[test]
178 fn egress_allowlist_enforced_when_set() {
179 let vault = KeyVault::build(&cfg(vec![entry(
180 "sk-virt-a",
181 "sk-real-a",
182 &["gpt-4o", "gpt-4o-mini"],
183 )]))
184 .unwrap()
185 .unwrap();
186 let e = vault.lookup("sk-virt-a").unwrap();
187 assert!(e.model_allowed(Some("gpt-4o")));
188 assert!(!e.model_allowed(Some("o1-preview"))); assert!(!e.model_allowed(None)); }
191
192 #[test]
193 fn empty_allowlist_permits_any_model() {
194 let vault = KeyVault::build(&cfg(vec![entry("sk-virt-a", "sk-real-a", &[])]))
195 .unwrap()
196 .unwrap();
197 let e = vault.lookup("sk-virt-a").unwrap();
198 assert!(e.model_allowed(Some("anything-goes")));
199 assert!(e.model_allowed(None)); }
201
202 #[test]
203 fn rejects_empty_or_duplicate_keys() {
204 assert!(KeyVault::build(&cfg(vec![entry("sk-v", "", &[])])).is_err());
206 assert!(KeyVault::build(&cfg(vec![entry("", "sk-real", &[])])).is_err());
208 assert!(KeyVault::build(&cfg(vec![
210 entry("sk-dup", "sk-real-1", &[]),
211 entry("sk-dup", "sk-real-2", &[]),
212 ]))
213 .is_err());
214 }
215
216 #[test]
217 fn rejects_whitespace_padded_keys() {
218 assert!(KeyVault::build(&cfg(vec![entry(" sk-v", "sk-real", &[])])).is_err());
222 assert!(KeyVault::build(&cfg(vec![entry("sk-v ", "sk-real", &[])])).is_err());
223 assert!(KeyVault::build(&cfg(vec![entry("sk-v", " sk-real", &[])])).is_err());
224 assert!(KeyVault::build(&cfg(vec![entry("sk-v", "sk-real ", &[])])).is_err());
225 }
226
227 #[test]
228 fn label_defaults_to_positional_id() {
229 let vault = KeyVault::build(&cfg(vec![entry("sk-v", "sk-r", &[])]))
230 .unwrap()
231 .unwrap();
232 assert_eq!(vault.lookup("sk-v").unwrap().label(), "key0");
233 }
234}