1use 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}