use aead::{Aead, KeyInit};
use aes_gcm_siv::{Aes256GcmSiv, Nonce};
use base64::Engine as _;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use serde_json::Value as JsonValue;
use std::fmt;
use uuid::Uuid;
use crate::runtime::config::EncryptionSettings;
#[derive(Clone)]
pub(super) struct EncryptionRuntime {
keys: Vec<EncryptionKey>,
active_version: u8,
}
#[derive(Clone)]
struct EncryptionKey {
version: u8,
key: [u8; 32],
}
impl fmt::Debug for EncryptionRuntime {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EncryptionRuntime")
.field("active_version", &self.active_version)
.field("keys", &RedactedKeyVersions(&self.keys))
.finish()
}
}
impl fmt::Debug for EncryptionKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EncryptionKey")
.field("version", &self.version)
.field("key", &"[redacted]")
.finish()
}
}
struct RedactedKeyVersions<'a>(&'a [EncryptionKey]);
impl fmt::Debug for RedactedKeyVersions<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let versions = self.0.iter().map(|key| key.version).collect::<Vec<_>>();
f.debug_struct("RedactedKeyVersions")
.field("versions", &versions)
.field("material", &"[redacted]")
.finish()
}
}
impl EncryptionRuntime {
pub(super) async fn from_settings(
settings: &EncryptionSettings,
) -> Result<Option<Self>, String> {
if let Some(runtime) = Self::from_key_materials(settings)? {
return Ok(Some(runtime));
}
Self::from_vault_export(settings).await
}
fn from_key_materials(settings: &EncryptionSettings) -> Result<Option<Self>, String> {
let mut keys = settings
.keys
.iter()
.map(|(version, value)| {
Ok(EncryptionKey {
version: *version,
key: decode_encryption_key(value).map_err(|err| {
format!("{} failed to decode: {err}", key_label(*version))
})?,
})
})
.collect::<Result<Vec<_>, String>>()?;
if keys.is_empty() {
return Ok(None);
}
keys.sort_by_key(|key| key.version);
let active_version = settings
.active_version
.unwrap_or_else(|| keys.last().map(|key| key.version).unwrap_or(1));
if !keys.iter().any(|key| key.version == active_version) {
return Err(format!(
"active encryption version {active_version} has no matching key"
));
}
Ok(Some(Self {
keys,
active_version,
}))
}
async fn from_vault_export(settings: &EncryptionSettings) -> Result<Option<Self>, String> {
let (Some(addr), Some(token)) = (&settings.vault_addr, &settings.vault_token) else {
return Ok(None);
};
#[cfg(not(feature = "http-client"))]
{
let _ = (addr, token);
return Err(
"Vault key export requires the `http-client` Cargo feature; \
rebuild with `--features http-client` or supply static keys \
via UDB_ENCRYPTION_KEYS / UDB_ENCRYPTION_ACTIVE_VERSION."
.to_string(),
);
}
#[cfg(feature = "http-client")]
if addr.starts_with("http://") && !settings.dev_mode {
return Err(
"UDB_VAULT_ADDR uses plain HTTP — X-Vault-Token would be sent unencrypted. \
Set https:// or set UDB_DEV_MODE=true to override."
.to_string(),
);
}
#[cfg(feature = "http-client")]
let url = format!(
"{}/v1/{}/export/encryption-key/{}",
addr.trim_end_matches('/'),
settings.vault_transit_mount.trim_matches('/'),
settings.vault_transit_key_name.trim_matches('/')
);
#[cfg(feature = "http-client")]
let vault_timeout = std::time::Duration::from_secs(settings.vault_timeout_secs.max(1));
#[cfg(feature = "http-client")]
{
let client = reqwest::Client::builder()
.timeout(vault_timeout)
.build()
.map_err(|e| format!("Vault HTTP client build failed: {e}"))?;
let payload: JsonValue = client
.get(url)
.header("X-Vault-Token", token)
.send()
.await
.map_err(|err| format!("Vault Transit key export request failed: {err}"))?
.error_for_status()
.map_err(|err| format!("Vault Transit key export failed: {err}"))?
.json()
.await
.map_err(|err| format!("Vault Transit key export JSON decode failed: {err}"))?;
let key_map = payload
.pointer("/data/keys")
.and_then(JsonValue::as_object)
.ok_or_else(|| "Vault Transit export response missing data.keys".to_string())?;
let mut keys = Vec::new();
for (version, value) in key_map {
let version = version
.parse::<u8>()
.map_err(|err| format!("invalid Vault key version {version}: {err}"))?;
let encoded = value
.as_str()
.ok_or_else(|| format!("Vault key version {version} is not a string"))?;
keys.push(EncryptionKey {
version,
key: decode_encryption_key(encoded).map_err(|err| {
format!("Vault key version {version} failed to decode: {err}")
})?,
});
}
if keys.is_empty() {
return Ok(None);
}
keys.sort_by_key(|key| key.version);
let active_version = settings
.active_version
.unwrap_or_else(|| keys.last().map(|k| k.version).unwrap_or(1));
if !keys.iter().any(|k| k.version == active_version) {
return Err(format!(
"UDB_ENCRYPTION_ACTIVE_VERSION={active_version} has no matching key in Vault export"
));
}
Ok(Some(Self {
keys,
active_version,
}))
}
}
pub(super) fn encrypt_json_value(&self, value: &JsonValue) -> Result<String, String> {
let key = self
.key(self.active_version)
.ok_or_else(|| "active encryption key is missing".to_string())?;
let mut nonce_bytes = [0u8; 12];
nonce_bytes.copy_from_slice(&Uuid::new_v4().as_bytes()[..12]);
let cipher = Aes256GcmSiv::new_from_slice(&key.key)
.map_err(|err| format!("invalid AEAD key: {err}"))?;
let plaintext = serde_json::to_vec(value)
.map_err(|err| format!("JSON plaintext serialization failed: {err}"))?;
let ciphertext = cipher
.encrypt(Nonce::from_slice(&nonce_bytes), plaintext.as_ref())
.map_err(|err| format!("AEAD encryption failed: {err}"))?;
let mut envelope = Vec::with_capacity(nonce_bytes.len() + ciphertext.len());
envelope.extend_from_slice(&nonce_bytes);
envelope.extend_from_slice(&ciphertext);
Ok(format!(
"udb-aead:v{}:{}",
key.version,
BASE64_STANDARD.encode(envelope)
))
}
pub(super) fn decrypt_json_value(&self, value: &str) -> Result<JsonValue, String> {
let Some((version, encoded)) = parse_ciphertext(value) else {
return Ok(JsonValue::String(value.to_string()));
};
let envelope = BASE64_STANDARD
.decode(encoded)
.map_err(|err| format!("ciphertext base64 decode failed: {err}"))?;
if envelope.len() <= 12 {
return Err("ciphertext envelope is too short".to_string());
}
let (nonce_bytes, ciphertext) = envelope.split_at(12);
let mut last_error = None;
for key in self.keys_for_decrypt(version) {
let cipher = Aes256GcmSiv::new_from_slice(&key.key)
.map_err(|err| format!("invalid AEAD key: {err}"))?;
match cipher.decrypt(Nonce::from_slice(nonce_bytes), ciphertext) {
Ok(plaintext) => {
return serde_json::from_slice(&plaintext)
.map_err(|err| format!("JSON plaintext decode failed: {err}"));
}
Err(err) => last_error = Some(format!("AEAD decryption failed: {err}")),
}
}
Err(last_error.unwrap_or_else(|| format!("no key for ciphertext version {version}")))
}
fn key(&self, version: u8) -> Option<&EncryptionKey> {
self.keys.iter().find(|key| key.version == version)
}
fn keys_for_decrypt(&self, version: u8) -> Vec<&EncryptionKey> {
let mut keys = self
.keys
.iter()
.filter(|key| key.version == version)
.collect::<Vec<_>>();
keys.extend(self.keys.iter().filter(|key| key.version != version));
keys
}
}
fn parse_ciphertext(value: &str) -> Option<(u8, &str)> {
let rest = value.strip_prefix("udb-aead:v")?;
let (version, encoded) = rest.split_once(':')?;
Some((version.parse().ok()?, encoded))
}
fn key_label(version: u8) -> String {
if version == 1 {
"UDB_ENCRYPTION_KEY/UDB_ENCRYPTION_KEY_V1".to_string()
} else {
format!("UDB_ENCRYPTION_KEY_V{version}")
}
}
fn decode_encryption_key(value: &str) -> Result<[u8; 32], String> {
let trimmed = value.trim();
let base64_len = BASE64_STANDARD.decode(trimmed).ok().map(|bytes| {
let len = bytes.len();
if len == 32 {
return Ok(bytes.try_into().expect("length checked"));
}
Err(len)
});
if let Some(Ok(key)) = base64_len {
return Ok(key);
}
let hex_len = decode_hex(trimmed).ok().map(|bytes| {
let len = bytes.len();
if len == 32 {
return Ok(bytes.try_into().expect("length checked"));
}
Err(len)
});
if let Some(Ok(key)) = hex_len {
return Ok(key);
}
let raw = trimmed.as_bytes();
if raw.len() == 32 {
let mut key = [0u8; 32];
key.copy_from_slice(raw);
return Ok(key);
}
let mut attempts = Vec::new();
if let Some(Err(len)) = base64_len {
attempts.push(format!("base64 decoded to {len} bytes"));
}
if let Some(Err(len)) = hex_len {
attempts.push(format!("hex decoded to {len} bytes"));
}
attempts.push(format!("raw value is {} bytes", raw.len()));
Err(format!(
"encryption key must be 32 bytes ({})",
attempts.join("; ")
))
}
fn decode_hex(value: &str) -> Result<Vec<u8>, base64::DecodeError> {
let value = value.strip_prefix("0x").unwrap_or(value);
if !value.len().is_multiple_of(2) {
return Err(base64::DecodeError::InvalidLength(value.len()));
}
let mut bytes = Vec::with_capacity(value.len() / 2);
for chunk in value.as_bytes().chunks(2) {
let pair = std::str::from_utf8(chunk)
.map_err(|_| base64::DecodeError::InvalidByte(0, chunk[0]))?;
let byte = u8::from_str_radix(pair, 16)
.map_err(|_| base64::DecodeError::InvalidByte(0, chunk[0]))?;
bytes.push(byte);
}
Ok(bytes)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn encryption_runtime_round_trips_json_with_active_key_version() {
let runtime = EncryptionRuntime {
active_version: 2,
keys: vec![
EncryptionKey {
version: 1,
key: [1; 32],
},
EncryptionKey {
version: 2,
key: [2; 32],
},
],
};
let value = json!({"nid": "1234567890", "score": 7});
let ciphertext = runtime.encrypt_json_value(&value).unwrap();
assert!(ciphertext.starts_with("udb-aead:v2:"));
assert_eq!(runtime.decrypt_json_value(&ciphertext).unwrap(), value);
}
#[test]
fn encryption_runtime_decrypts_old_key_versions() {
let old_runtime = EncryptionRuntime {
active_version: 1,
keys: vec![EncryptionKey {
version: 1,
key: [1; 32],
}],
};
let rotated_runtime = EncryptionRuntime {
active_version: 2,
keys: vec![
EncryptionKey {
version: 1,
key: [1; 32],
},
EncryptionKey {
version: 2,
key: [2; 32],
},
],
};
let value = JsonValue::String("legacy secret".to_string());
let ciphertext = old_runtime.encrypt_json_value(&value).unwrap();
assert_eq!(
rotated_runtime.decrypt_json_value(&ciphertext).unwrap(),
value
);
}
#[test]
fn decode_encryption_key_accepts_base64_hex_and_raw_32_byte_keys() {
let bytes = [0xabu8; 32];
let base64 = BASE64_STANDARD.encode(bytes);
assert_eq!(decode_encryption_key(&base64).unwrap(), bytes);
let hex = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef";
let expected_hex = [
0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0x01, 0x23, 0x45, 0x67, 0x89, 0xab,
0xcd, 0xef, 0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0x01, 0x23, 0x45, 0x67,
0x89, 0xab, 0xcd, 0xef,
];
assert_eq!(decode_encryption_key(hex).unwrap(), expected_hex);
assert_eq!(
decode_encryption_key(&format!("0x{hex}")).unwrap(),
expected_hex
);
let raw = "0123456789abcdef0123456789abcdef";
assert_eq!(decode_encryption_key(raw).unwrap(), *raw.as_bytes());
}
#[test]
fn decode_encryption_key_reports_attempted_lengths() {
let err = decode_encryption_key(
"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef0000",
)
.unwrap_err();
assert!(err.contains("base64 decoded to"));
assert!(err.contains("hex decoded to 34 bytes"));
assert!(err.contains("raw value is 68 bytes"));
}
#[tokio::test]
async fn inline_key_decode_error_names_configured_key_source() {
let mut settings = EncryptionSettings::default();
settings.keys.insert(1, "short".to_string());
let err = EncryptionRuntime::from_settings(&settings)
.await
.unwrap_err();
assert!(err.contains("UDB_ENCRYPTION_KEY/UDB_ENCRYPTION_KEY_V1 failed to decode"));
assert!(err.contains("encryption key must be 32 bytes"));
}
#[test]
fn encryption_debug_redacts_key_material() {
let runtime = EncryptionRuntime {
active_version: 7,
keys: vec![EncryptionKey {
version: 7,
key: [7; 32],
}],
};
let runtime_debug = format!("{runtime:?}");
let key_debug = format!("{:?}", runtime.keys[0]);
assert!(runtime_debug.contains("active_version"));
assert!(runtime_debug.contains("7"));
assert!(runtime_debug.contains("[redacted]"));
assert!(key_debug.contains("[redacted]"));
assert!(!runtime_debug.contains("[7, 7"));
assert!(!key_debug.contains("[7, 7"));
}
}