use std::collections::HashSet;
use anyhow::Result;
use crate::config::LlmCfg;
pub struct VaultEntry {
virtual_key: String,
provider_key: String,
allowed_models: HashSet<String>,
label: String,
}
impl VaultEntry {
pub fn provider_key(&self) -> &str {
&self.provider_key
}
pub fn label(&self) -> &str {
&self.label
}
pub fn model_allowed(&self, model: Option<&str>) -> bool {
self.allowed_models.is_empty() || model.is_some_and(|m| self.allowed_models.contains(m))
}
}
pub struct KeyVault {
entries: Vec<VaultEntry>,
}
impl KeyVault {
pub fn build(cfg: &LlmCfg) -> Result<Option<KeyVault>> {
if cfg.keys.is_empty() {
return Ok(None);
}
let mut entries: Vec<VaultEntry> = Vec::with_capacity(cfg.keys.len());
for (i, k) in cfg.keys.iter().enumerate() {
anyhow::ensure!(
!k.virtual_key.is_empty(),
"llm.keys[{i}].virtual_key must not be empty"
);
anyhow::ensure!(
k.virtual_key == k.virtual_key.trim(),
"llm.keys[{i}].virtual_key must not have leading/trailing whitespace"
);
anyhow::ensure!(
!k.provider_key.is_empty(),
"llm.keys[{i}].provider_key must not be empty"
);
anyhow::ensure!(
k.provider_key == k.provider_key.trim(),
"llm.keys[{i}].provider_key must not have leading/trailing whitespace"
);
anyhow::ensure!(
!entries.iter().any(|e| e.virtual_key == k.virtual_key),
"llm.keys[{i}]: duplicate virtual_key"
);
let label = if k.label.trim().is_empty() {
format!("key{i}")
} else {
k.label.clone()
};
entries.push(VaultEntry {
virtual_key: k.virtual_key.clone(),
provider_key: k.provider_key.clone(),
allowed_models: k.allowed_models.iter().cloned().collect(),
label,
});
}
Ok(Some(KeyVault { entries }))
}
pub fn lookup(&self, presented: &str) -> Option<&VaultEntry> {
let mut matched: Option<&VaultEntry> = None;
for entry in &self.entries {
if constant_time_eq(entry.virtual_key.as_bytes(), presented.as_bytes()) {
matched = Some(entry);
}
}
matched
}
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
let mut diff = a.len() ^ b.len();
for i in 0..a.len().max(b.len()) {
let x = a.get(i).copied().unwrap_or(0);
let y = b.get(i).copied().unwrap_or(0);
diff |= (x ^ y) as usize;
}
diff == 0
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::KeyEntryCfg;
fn cfg(keys: Vec<KeyEntryCfg>) -> LlmCfg {
LlmCfg {
enabled: true,
keys,
..Default::default()
}
}
fn entry(virtual_key: &str, provider_key: &str, models: &[&str]) -> KeyEntryCfg {
KeyEntryCfg {
virtual_key: virtual_key.into(),
provider_key: provider_key.into(),
allowed_models: models.iter().map(|s| s.to_string()).collect(),
label: String::new(),
}
}
#[test]
fn build_is_none_without_keys() {
assert!(KeyVault::build(&LlmCfg::default()).unwrap().is_none());
}
#[test]
fn lookup_maps_virtual_to_provider_key() {
let vault = KeyVault::build(&cfg(vec![entry("sk-virt-a", "sk-real-a", &[])]))
.unwrap()
.unwrap();
let e = vault.lookup("sk-virt-a").expect("known key resolves");
assert_eq!(e.provider_key(), "sk-real-a");
assert!(vault.lookup("sk-virt-unknown").is_none());
assert!(vault.lookup("sk-real-a").is_none());
}
#[test]
fn egress_allowlist_enforced_when_set() {
let vault = KeyVault::build(&cfg(vec![entry(
"sk-virt-a",
"sk-real-a",
&["gpt-4o", "gpt-4o-mini"],
)]))
.unwrap()
.unwrap();
let e = vault.lookup("sk-virt-a").unwrap();
assert!(e.model_allowed(Some("gpt-4o")));
assert!(!e.model_allowed(Some("o1-preview"))); assert!(!e.model_allowed(None)); }
#[test]
fn empty_allowlist_permits_any_model() {
let vault = KeyVault::build(&cfg(vec![entry("sk-virt-a", "sk-real-a", &[])]))
.unwrap()
.unwrap();
let e = vault.lookup("sk-virt-a").unwrap();
assert!(e.model_allowed(Some("anything-goes")));
assert!(e.model_allowed(None)); }
#[test]
fn rejects_empty_or_duplicate_keys() {
assert!(KeyVault::build(&cfg(vec![entry("sk-v", "", &[])])).is_err());
assert!(KeyVault::build(&cfg(vec![entry("", "sk-real", &[])])).is_err());
assert!(KeyVault::build(&cfg(vec![
entry("sk-dup", "sk-real-1", &[]),
entry("sk-dup", "sk-real-2", &[]),
]))
.is_err());
}
#[test]
fn rejects_whitespace_padded_keys() {
assert!(KeyVault::build(&cfg(vec![entry(" sk-v", "sk-real", &[])])).is_err());
assert!(KeyVault::build(&cfg(vec![entry("sk-v ", "sk-real", &[])])).is_err());
assert!(KeyVault::build(&cfg(vec![entry("sk-v", " sk-real", &[])])).is_err());
assert!(KeyVault::build(&cfg(vec![entry("sk-v", "sk-real ", &[])])).is_err());
}
#[test]
fn label_defaults_to_positional_id() {
let vault = KeyVault::build(&cfg(vec![entry("sk-v", "sk-r", &[])]))
.unwrap()
.unwrap();
assert_eq!(vault.lookup("sk-v").unwrap().label(), "key0");
}
}