Skip to main content

mongreldb_server/
vault_kms.rs

1//! HashiCorp Vault Transit implementation of the database KMS boundary.
2
3use std::io::Read;
4use std::time::Duration;
5
6use base64::engine::general_purpose::STANDARD;
7use base64::Engine as _;
8use mongreldb_core::{
9    KeyManagementError, KeyManagementHealth, KeyManagementProvider, KmsWrappedKey,
10};
11use reqwest::blocking::{Client, RequestBuilder};
12use reqwest::header::HeaderValue;
13use serde::{Deserialize, Serialize};
14use sha2::Digest as _;
15use zeroize::Zeroizing;
16
17const ALGORITHM: &str = "vault-transit";
18const MAX_VAULT_RESPONSE_BYTES: usize = 64 * 1024;
19
20pub struct VaultTransitConfig {
21    pub endpoint: String,
22    pub mount: String,
23    pub token: Zeroizing<String>,
24    pub namespace: Option<String>,
25    pub timeout: Duration,
26    pub ca_certificate_pem: Option<Vec<u8>>,
27}
28
29impl std::fmt::Debug for VaultTransitConfig {
30    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31        formatter
32            .debug_struct("VaultTransitConfig")
33            .field("endpoint", &self.endpoint)
34            .field("mount", &self.mount)
35            .field("token", &"<redacted>")
36            .field("namespace", &self.namespace)
37            .field(
38                "ca_certificate_pem",
39                &self.ca_certificate_pem.as_ref().map(|_| "<configured>"),
40            )
41            .field("timeout", &self.timeout)
42            .finish()
43    }
44}
45
46pub struct VaultTransitKeyManagementProvider {
47    client: Client,
48    endpoint: reqwest::Url,
49    mount: String,
50    token: Zeroizing<String>,
51    namespace: Option<String>,
52    provider_id: String,
53}
54
55impl std::fmt::Debug for VaultTransitKeyManagementProvider {
56    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
57        formatter
58            .debug_struct("VaultTransitKeyManagementProvider")
59            .field("endpoint", &self.endpoint)
60            .field("mount", &self.mount)
61            .field("token", &"<redacted>")
62            .field("namespace", &self.namespace)
63            .field("provider_id", &self.provider_id)
64            .finish()
65    }
66}
67
68impl VaultTransitKeyManagementProvider {
69    pub fn new(config: VaultTransitConfig) -> Result<Self, KeyManagementError> {
70        let mut endpoint = reqwest::Url::parse(&config.endpoint)
71            .map_err(|error| KeyManagementError::Failed(format!("invalid Vault URL: {error}")))?;
72        if endpoint.scheme() != "https" {
73            return Err(KeyManagementError::Failed(
74                "Vault endpoint must use HTTPS".into(),
75            ));
76        }
77        if !endpoint.username().is_empty()
78            || endpoint.password().is_some()
79            || endpoint.query().is_some()
80            || endpoint.fragment().is_some()
81        {
82            return Err(KeyManagementError::Failed(
83                "Vault endpoint must not contain credentials, query, or fragment".into(),
84            ));
85        }
86        if config.token.is_empty() {
87            return Err(KeyManagementError::Failed(
88                "Vault token must not be empty".into(),
89            ));
90        }
91        validate_path_segment("Vault mount", &config.mount)?;
92        let endpoint_path = endpoint.path().trim_end_matches('/').to_owned();
93        endpoint.set_path(&endpoint_path);
94
95        let mut client = Client::builder().timeout(config.timeout);
96        if let Some(pem) = config.ca_certificate_pem {
97            let certificate = reqwest::Certificate::from_pem(&pem).map_err(|error| {
98                KeyManagementError::Failed(format!("invalid Vault CA certificate: {error}"))
99            })?;
100            client = client.add_root_certificate(certificate);
101        }
102        let client = client
103            .build()
104            .map_err(|error| KeyManagementError::Failed(format!("Vault client: {error}")))?;
105        let provider_hash = sha2::Sha256::digest(format!(
106            "{endpoint}|{}|{}",
107            config.mount,
108            config.namespace.as_deref().unwrap_or_default()
109        ));
110        let provider_id = format!(
111            "vault-transit:{}",
112            provider_hash
113                .iter()
114                .map(|byte| format!("{byte:02x}"))
115                .collect::<String>()
116        );
117        Ok(Self {
118            client,
119            endpoint,
120            mount: config.mount,
121            token: config.token,
122            namespace: config.namespace,
123            provider_id,
124        })
125    }
126
127    fn url(&self, operation: &str, key_id: &str) -> Result<reqwest::Url, KeyManagementError> {
128        validate_path_segment("Vault operation", operation)?;
129        validate_path_segment("Vault key id", key_id)?;
130        let mut url = self.endpoint.clone();
131        {
132            let mut segments = url.path_segments_mut().map_err(|_| {
133                KeyManagementError::Failed("Vault endpoint cannot be a base URL".into())
134            })?;
135            segments.pop_if_empty();
136            segments.extend(["v1", &self.mount, operation, key_id]);
137        }
138        Ok(url)
139    }
140
141    fn authorize(&self, request: RequestBuilder) -> Result<RequestBuilder, KeyManagementError> {
142        let mut token = HeaderValue::from_str(self.token.as_str())
143            .map_err(|_| KeyManagementError::Failed("invalid Vault token header".into()))?;
144        token.set_sensitive(true);
145        let mut request = request.header("X-Vault-Token", token);
146        if let Some(namespace) = &self.namespace {
147            let mut namespace = HeaderValue::from_str(namespace)
148                .map_err(|_| KeyManagementError::Failed("invalid Vault namespace".into()))?;
149            namespace.set_sensitive(true);
150            request = request.header("X-Vault-Namespace", namespace);
151        }
152        Ok(request)
153    }
154
155    fn post<T: Serialize, R: for<'de> Deserialize<'de>>(
156        &self,
157        operation: &str,
158        key_id: &str,
159        body: &T,
160    ) -> Result<R, KeyManagementError> {
161        let request = self.client.post(self.url(operation, key_id)?).json(body);
162        let response = self
163            .authorize(request)?
164            .send()
165            .map_err(|error| KeyManagementError::Unavailable(format!("Vault request: {error}")))?;
166        let status = response.status();
167        if !status.is_success() {
168            return Err(KeyManagementError::Failed(format!(
169                "Vault {operation} returned HTTP {status}"
170            )));
171        }
172        decode_vault_response(response)
173    }
174}
175
176impl KeyManagementProvider for VaultTransitKeyManagementProvider {
177    fn provider_id(&self) -> &str {
178        &self.provider_id
179    }
180
181    fn wrap_key(
182        &self,
183        key_id: &str,
184        plaintext_key: &[u8],
185    ) -> Result<KmsWrappedKey, KeyManagementError> {
186        let response: EncryptResponse = self.post(
187            "encrypt",
188            key_id,
189            &EncryptRequest {
190                plaintext: STANDARD.encode(plaintext_key),
191            },
192        )?;
193        let key_version = vault_ciphertext_version(&response.data.ciphertext)?;
194        Ok(KmsWrappedKey {
195            kms_key_id: key_id.into(),
196            key_version,
197            wrapped_dek: response.data.ciphertext.into_bytes(),
198            algorithm: ALGORITHM.into(),
199        })
200    }
201
202    fn unwrap_key(
203        &self,
204        wrapped: &KmsWrappedKey,
205    ) -> Result<Zeroizing<Vec<u8>>, KeyManagementError> {
206        if wrapped.algorithm != ALGORITHM {
207            return Err(KeyManagementError::Failed(format!(
208                "unsupported wrapped-key algorithm {:?}",
209                wrapped.algorithm
210            )));
211        }
212        let ciphertext = std::str::from_utf8(&wrapped.wrapped_dek)
213            .map_err(|_| KeyManagementError::Failed("Vault ciphertext is not UTF-8".into()))?;
214        let response: DecryptResponse = self.post(
215            "decrypt",
216            &wrapped.kms_key_id,
217            &DecryptRequest { ciphertext },
218        )?;
219        let plaintext = STANDARD
220            .decode(response.data.plaintext)
221            .map_err(|_| KeyManagementError::Failed("Vault plaintext is not base64".into()))?;
222        Ok(Zeroizing::new(plaintext))
223    }
224
225    fn rewrap_key(
226        &self,
227        wrapped: &KmsWrappedKey,
228        new_key_id: &str,
229    ) -> Result<KmsWrappedKey, KeyManagementError> {
230        let plaintext = self.unwrap_key(wrapped)?;
231        self.wrap_key(new_key_id, plaintext.as_ref())
232    }
233
234    fn provider_health(&self) -> KeyManagementHealth {
235        let mut url = self.endpoint.clone();
236        let Ok(mut segments) = url.path_segments_mut() else {
237            return KeyManagementHealth::Unavailable;
238        };
239        segments.pop_if_empty();
240        segments.extend(["v1", "sys", "health"]);
241        drop(segments);
242        match self
243            .client
244            .get(url)
245            .send()
246            .map(|response| response.status())
247        {
248            Ok(status) if status.is_success() => KeyManagementHealth::Ready,
249            Ok(status) if matches!(status.as_u16(), 429 | 472 | 473) => {
250                KeyManagementHealth::Degraded
251            }
252            _ => KeyManagementHealth::Unavailable,
253        }
254    }
255}
256
257fn decode_vault_response<R, T>(reader: R) -> Result<T, KeyManagementError>
258where
259    R: Read,
260    T: for<'de> Deserialize<'de>,
261{
262    let mut bytes = Vec::new();
263    reader
264        .take((MAX_VAULT_RESPONSE_BYTES + 1) as u64)
265        .read_to_end(&mut bytes)
266        .map_err(|error| KeyManagementError::Failed(format!("Vault response: {error}")))?;
267    if bytes.len() > MAX_VAULT_RESPONSE_BYTES {
268        return Err(KeyManagementError::Failed(
269            "Vault response exceeds 64 KiB".into(),
270        ));
271    }
272    serde_json::from_slice(&bytes)
273        .map_err(|error| KeyManagementError::Failed(format!("Vault response: {error}")))
274}
275
276fn validate_path_segment(name: &str, value: &str) -> Result<(), KeyManagementError> {
277    if value.is_empty()
278        || value == "."
279        || value == ".."
280        || value.contains('/')
281        || value.contains('\\')
282    {
283        return Err(KeyManagementError::Failed(format!(
284            "{name} must be one non-empty path segment"
285        )));
286    }
287    Ok(())
288}
289
290fn vault_ciphertext_version(ciphertext: &str) -> Result<String, KeyManagementError> {
291    let mut parts = ciphertext.splitn(3, ':');
292    if parts.next() != Some("vault") {
293        return Err(KeyManagementError::Failed(
294            "Vault ciphertext has invalid prefix".into(),
295        ));
296    }
297    let version = parts
298        .next()
299        .filter(|value| {
300            value
301                .strip_prefix('v')
302                .is_some_and(|number| number.parse::<u64>().is_ok())
303        })
304        .ok_or_else(|| KeyManagementError::Failed("Vault ciphertext has invalid version".into()))?;
305    if parts.next().is_none() {
306        return Err(KeyManagementError::Failed(
307            "Vault ciphertext is truncated".into(),
308        ));
309    }
310    Ok(version.into())
311}
312
313#[derive(Serialize)]
314struct EncryptRequest {
315    plaintext: String,
316}
317
318#[derive(Deserialize)]
319struct EncryptResponse {
320    data: EncryptData,
321}
322
323#[derive(Deserialize)]
324struct EncryptData {
325    ciphertext: String,
326}
327
328#[derive(Serialize)]
329struct DecryptRequest<'a> {
330    ciphertext: &'a str,
331}
332
333#[derive(Deserialize)]
334struct DecryptResponse {
335    data: DecryptData,
336}
337
338#[derive(Deserialize)]
339struct DecryptData {
340    plaintext: String,
341}
342
343#[cfg(test)]
344mod tests {
345    use super::*;
346
347    #[test]
348    fn config_requires_https_and_safe_mount() {
349        let config = |endpoint: &str, mount: &str| VaultTransitConfig {
350            endpoint: endpoint.into(),
351            mount: mount.into(),
352            token: Zeroizing::new("token".into()),
353            namespace: None,
354            timeout: Duration::from_secs(1),
355            ca_certificate_pem: None,
356        };
357        assert!(
358            VaultTransitKeyManagementProvider::new(config("http://vault.example", "transit"))
359                .is_err()
360        );
361        assert!(VaultTransitKeyManagementProvider::new(config(
362            "https://vault.example",
363            "../transit"
364        ))
365        .is_err());
366        assert!(
367            VaultTransitKeyManagementProvider::new(config("https://vault.example", "transit"))
368                .is_ok()
369        );
370    }
371
372    #[test]
373    fn parses_vault_ciphertext_versions() {
374        assert_eq!(
375            vault_ciphertext_version("vault:v12:ciphertext").unwrap(),
376            "v12"
377        );
378        assert!(vault_ciphertext_version("invalid:v1:ciphertext").is_err());
379        assert!(vault_ciphertext_version("vault:latest:ciphertext").is_err());
380    }
381
382    #[test]
383    fn response_size_is_bounded_and_namespace_changes_identity() {
384        let oversized = vec![b' '; MAX_VAULT_RESPONSE_BYTES + 1];
385        assert!(decode_vault_response::<_, serde_json::Value>(oversized.as_slice()).is_err());
386
387        let provider = |namespace| {
388            VaultTransitKeyManagementProvider::new(VaultTransitConfig {
389                endpoint: "https://vault.example".into(),
390                mount: "transit".into(),
391                token: Zeroizing::new("token".into()),
392                namespace,
393                timeout: Duration::from_secs(1),
394                ca_certificate_pem: None,
395            })
396            .unwrap()
397        };
398        assert_ne!(
399            provider(Some("team-a".into())).provider_id(),
400            provider(Some("team-b".into())).provider_id()
401        );
402    }
403}