Skip to main content

edgeguard/
keyvault.rs

1//! BYO-key vault + egress governance (gateway L2).
2//!
3//! The security wedge: clients authenticate with a **virtual key**; the real **provider key** (the
4//! upstream's `sk-…` secret) lives only on the edge and is injected into the upstream request on the
5//! way out. A client — or a compromised one — never sees a provider key, and a leaked virtual key is
6//! revoked by deleting one vault entry without rotating the provider credential.
7//!
8//! Two controls per key:
9//!   * **key swap** — the presented virtual key is replaced by its mapped provider key in the
10//!     upstream `Authorization`, so the provider secret never appears in the client surface or logs;
11//!   * **egress allowlist** — an optional set of model names the key may reach; a request for any
12//!     other model is denied `403` (fail-closed once a list is set).
13//!
14//! Virtual keys are matched by a **constant-time** comparison that scans every entry (mirroring the
15//! API-key gate in [`crate::auth`]), so the match time doesn't reveal which key — if any — was hit.
16//! Provider keys are held in memory as configured (encryption-at-rest is a property of wherever the
17//! config/secret is stored, e.g. the control-plane secret store that pushes them).
18//!
19//! When any `[[llm.keys]]` is configured the vault is **enabled for all proxied traffic**: a request
20//! without a known virtual key is rejected `401` before it reaches the upstream.
21
22use std::collections::HashSet;
23
24use anyhow::Result;
25
26use crate::config::LlmCfg;
27
28/// One resolved vault entry: the provider secret to inject and the model egress policy.
29pub struct VaultEntry {
30    /// The client-facing secret this entry matches (compared constant-time on lookup).
31    virtual_key: String,
32    provider_key: String,
33    /// Egress allowlist of model names; empty means unrestricted.
34    allowed_models: HashSet<String>,
35    /// Non-secret label for logs / metrics / audit.
36    label: String,
37}
38
39impl VaultEntry {
40    /// The real provider secret to inject upstream. Not exposed beyond the proxy boundary.
41    pub fn provider_key(&self) -> &str {
42        &self.provider_key
43    }
44
45    /// The non-secret label (e.g. a team / tenant name) for logging.
46    pub fn label(&self) -> &str {
47        &self.label
48    }
49
50    /// Whether `model` is permitted for this key. An empty allowlist permits any model; a non-empty
51    /// one permits only listed models (fail-closed for everything else, including `None` — a request
52    /// whose model cannot be parsed is denied when an allowlist is set).
53    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
58/// The configured vault entries. Built once per config (re)load and carried on the proxy
59/// [`Runtime`](crate::proxy::Runtime).
60pub struct KeyVault {
61    entries: Vec<VaultEntry>,
62}
63
64impl KeyVault {
65    /// Build the vault from `[llm].keys`. Returns `Ok(None)` when no keys are configured (the proxy
66    /// then skips vault enforcement entirely). Rejects an empty virtual/provider key or a duplicate
67    /// virtual key, so a broken config fails at startup/reload rather than per-request.
68    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    /// Resolve a presented virtual key to its entry, or `None` if unknown. Scans **every** entry
110    /// with a constant-time comparison so the lookup time doesn't reveal which key matched (the same
111    /// discipline as [`crate::auth`]'s API-key check).
112    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
123/// Constant-time byte comparison: folds every byte *and the length mismatch* into one `usize`
124/// accumulator, so the execution path is identical regardless of where (or whether) the bytes
125/// differ. Returning early on a length mismatch would let an attacker distinguish "wrong length"
126/// from "right length, wrong bytes" by timing.
127fn 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        // An unknown virtual key resolves to nothing (caller rejects 401).
172        assert!(vault.lookup("sk-virt-unknown").is_none());
173        // The provider key is never itself a valid virtual key.
174        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"))); // not on the allowlist → denied
189        assert!(!e.model_allowed(None)); // unparseable model → fail-closed when allowlist set
190    }
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)); // empty allowlist → unrestricted even without a model
200    }
201
202    #[test]
203    fn rejects_empty_or_duplicate_keys() {
204        // Empty provider key.
205        assert!(KeyVault::build(&cfg(vec![entry("sk-v", "", &[])])).is_err());
206        // Empty virtual key.
207        assert!(KeyVault::build(&cfg(vec![entry("", "sk-real", &[])])).is_err());
208        // Duplicate virtual key across two entries.
209        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        // A virtual key with a leading space would never match a client's presented key (which
219        // wouldn't have the space), so catch it at build time rather than silently creating a
220        // dead entry.
221        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}