use std::collections::HashMap;
use std::sync::Arc;
use chacha20poly1305::{
aead::{Aead, AeadCore, KeyInit, OsRng},
ChaCha20Poly1305, Key, Nonce,
};
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use thiserror::Error;
use zeroize::ZeroizeOnDrop;
use super::ports::{AuditAction, AuditEvent, AuditSink, KeyManagementService};
#[derive(Debug, Error)]
pub enum VaultError {
#[error("Encryption failed")]
EncryptionFailed,
#[error("Decryption failed: bad key or tampered ciphertext")]
DecryptionFailed,
#[error("Secret not found: {0}")]
NotFound(String),
#[error("Secret has expired: {0}")]
Expired(String),
#[error("Vault key must be exactly 32 bytes")]
InvalidKeyLength,
}
#[derive(Clone, ZeroizeOnDrop)]
pub struct VaultKey {
raw: [u8; 32],
}
impl VaultKey {
pub fn generate() -> Self {
let key = ChaCha20Poly1305::generate_key(&mut OsRng);
let mut raw = [0u8; 32];
raw.copy_from_slice(&key);
Self { raw }
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, VaultError> {
if bytes.len() != 32 {
return Err(VaultError::InvalidKeyLength);
}
let mut raw = [0u8; 32];
raw.copy_from_slice(bytes);
Ok(Self { raw })
}
fn cipher(&self) -> ChaCha20Poly1305 {
ChaCha20Poly1305::new(Key::from_slice(&self.raw))
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EncryptedBlob(Vec<u8>);
impl EncryptedBlob {
pub fn seal(key: &VaultKey, plaintext: &[u8]) -> Result<Self, VaultError> {
let cipher = key.cipher();
let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng);
let mut out = nonce.to_vec(); let ciphertext =
cipher.encrypt(&nonce, plaintext).map_err(|_| VaultError::EncryptionFailed)?;
out.extend_from_slice(&ciphertext);
Ok(Self(out))
}
pub fn open(&self, key: &VaultKey) -> Result<Vec<u8>, VaultError> {
if self.0.len() < 12 {
return Err(VaultError::DecryptionFailed);
}
let (nonce_bytes, ciphertext) = self.0.split_at(12);
let nonce = Nonce::from_slice(nonce_bytes);
let cipher = key.cipher();
cipher.decrypt(nonce, ciphertext).map_err(|_| VaultError::DecryptionFailed)
}
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn nonce_bytes(&self) -> &[u8] {
&self.0[..12]
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VaultEntry {
pub name: String,
pub blob: EncryptedBlob,
pub created_at: DateTime<Utc>,
pub expires_at: Option<DateTime<Utc>>,
pub version: u32,
}
impl VaultEntry {
fn is_expired(&self) -> bool {
self.expires_at.map(|exp| Utc::now() > exp).unwrap_or(false)
}
}
pub struct SecretVault {
key: VaultKey,
entries: HashMap<String, VaultEntry>,
audit_sink: Option<Arc<dyn AuditSink>>,
}
impl SecretVault {
pub fn new(key: VaultKey) -> Self {
Self { key, entries: HashMap::new(), audit_sink: None }
}
pub fn with_audit_sink(mut self, sink: Arc<dyn AuditSink>) -> Self {
self.audit_sink = Some(sink);
self
}
fn audit(&self, event: AuditEvent) {
if let Some(sink) = &self.audit_sink {
sink.record(event);
}
}
pub fn put(
&mut self,
name: impl Into<String>,
plaintext: &[u8],
ttl_seconds: Option<i64>,
) -> Result<(), VaultError> {
let name = name.into();
let blob = EncryptedBlob::seal(&self.key, plaintext);
match blob {
Err(e) => {
self.audit(AuditEvent::failure(
None,
name.clone(),
AuditAction::VaultWrite,
e.to_string(),
));
Err(e)
}
Ok(blob) => {
let expires_at =
ttl_seconds.map(|secs| Utc::now() + chrono::Duration::seconds(secs));
let version = self.entries.get(&name).map(|e| e.version + 1).unwrap_or(1);
self.entries.insert(
name.clone(),
VaultEntry {
name: name.clone(),
blob,
created_at: Utc::now(),
expires_at,
version,
},
);
self.audit(AuditEvent::success(None, name, AuditAction::VaultWrite));
Ok(())
}
}
}
pub fn get(&self, name: &str) -> Result<Vec<u8>, VaultError> {
let entry = self.entries.get(name).ok_or_else(|| VaultError::NotFound(name.to_owned()))?;
if entry.is_expired() {
self.audit(AuditEvent::failure(None, name, AuditAction::VaultRead, "secret expired"));
return Err(VaultError::Expired(name.to_owned()));
}
let result = entry.blob.open(&self.key);
match &result {
Ok(_) => self.audit(AuditEvent::success(None, name, AuditAction::VaultRead)),
Err(e) => {
self.audit(AuditEvent::failure(None, name, AuditAction::VaultRead, e.to_string()))
}
}
result
}
pub fn rotate(&mut self, name: &str) -> Result<(), VaultError> {
let plaintext = self.get(name)?; self.put(name, &plaintext, None)
}
pub fn remove(&mut self, name: &str) -> bool {
let removed = self.entries.remove(name).is_some();
if removed {
self.audit(AuditEvent::success(None, name, AuditAction::VaultWrite));
}
removed
}
pub fn list(&self) -> Vec<&str> {
self.entries.values().filter(|e| !e.is_expired()).map(|e| e.name.as_str()).collect()
}
pub fn entry(&self, name: &str) -> Option<&VaultEntry> {
self.entries.get(name)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct KmsVaultEntry {
pub name: String,
pub blob: EncryptedBlob,
pub wrapped_dek: Vec<u8>,
pub created_at: DateTime<Utc>,
pub expires_at: Option<DateTime<Utc>>,
pub version: u32,
}
impl KmsVaultEntry {
fn is_expired(&self) -> bool {
self.expires_at.map(|exp| Utc::now() > exp).unwrap_or(false)
}
}
pub struct KmsSecretVault {
kms: Arc<dyn KeyManagementService>,
#[doc(hidden)]
pub entries: HashMap<String, KmsVaultEntry>,
}
impl KmsSecretVault {
pub fn new(kms: Arc<dyn KeyManagementService>) -> Self {
Self { kms, entries: HashMap::new() }
}
pub fn put(
&mut self,
name: impl Into<String>,
plaintext: &[u8],
ttl_seconds: Option<i64>,
) -> Result<(), VaultError> {
let name = name.into();
let dk = self.kms.generate_data_key().map_err(|_| VaultError::EncryptionFailed)?;
let dek_key = VaultKey::from_bytes(&dk.plaintext)?;
let blob = EncryptedBlob::seal(&dek_key, plaintext)?;
let wrapped_dek = dk.wrapped.clone();
let expires_at = ttl_seconds.map(|secs| Utc::now() + chrono::Duration::seconds(secs));
let version = self.entries.get(&name).map(|e| e.version + 1).unwrap_or(1);
self.entries.insert(
name.clone(),
KmsVaultEntry { name, blob, wrapped_dek, created_at: Utc::now(), expires_at, version },
);
Ok(())
}
pub fn get(&self, name: &str) -> Result<Vec<u8>, VaultError> {
let entry = self.entries.get(name).ok_or_else(|| VaultError::NotFound(name.to_owned()))?;
if entry.is_expired() {
return Err(VaultError::Expired(name.to_owned()));
}
let plaintext_dek = self
.kms
.decrypt_data_key(&entry.wrapped_dek)
.map_err(|_| VaultError::DecryptionFailed)?;
let dek_key = VaultKey::from_bytes(&plaintext_dek)?;
entry.blob.open(&dek_key)
}
pub fn rotate(&mut self, name: &str) -> Result<(), VaultError> {
let plaintext = self.get(name)?;
self.put(name, &plaintext, None)
}
pub fn remove(&mut self, name: &str) -> bool {
self.entries.remove(name).is_some()
}
pub fn list(&self) -> Vec<&str> {
self.entries.values().filter(|e| !e.is_expired()).map(|e| e.name.as_str()).collect()
}
pub fn entry(&self, name: &str) -> Option<&KmsVaultEntry> {
self.entries.get(name)
}
}
#[cfg(test)]
mod tests {
use std::collections::HashSet;
use super::*;
fn make_vault() -> SecretVault {
SecretVault::new(VaultKey::generate())
}
#[test]
fn round_trip_bytes() {
let mut v = make_vault();
v.put("db_password", b"s3cr3t!", None).unwrap();
let plain = v.get("db_password").unwrap();
assert_eq!(plain, b"s3cr3t!");
}
#[test]
fn round_trip_unicode() {
let mut v = make_vault();
let secret = "日本語テスト🔑";
v.put("unicode_key", secret.as_bytes(), None).unwrap();
let plain = v.get("unicode_key").unwrap();
assert_eq!(String::from_utf8(plain).unwrap(), secret);
}
#[test]
fn round_trip_empty_value() {
let mut v = make_vault();
v.put("empty", b"", None).unwrap();
let plain = v.get("empty").unwrap();
assert_eq!(plain, b"");
}
#[test]
fn wrong_key_fails_to_decrypt() {
let key_a = VaultKey::generate();
let key_b = VaultKey::generate();
let blob = EncryptedBlob::seal(&key_a, b"my secret").unwrap();
let result = blob.open(&key_b);
assert!(
matches!(result, Err(VaultError::DecryptionFailed)),
"expected DecryptionFailed, got {result:?}"
);
}
#[test]
fn tampered_ciphertext_rejected() {
let key = VaultKey::generate();
let mut blob = EncryptedBlob::seal(&key, b"sensitive data").unwrap();
blob.0[12] ^= 0xFF;
let result = blob.open(&key);
assert!(
matches!(result, Err(VaultError::DecryptionFailed)),
"expected DecryptionFailed, got {result:?}"
);
}
#[test]
fn tampered_tag_rejected() {
let key = VaultKey::generate();
let mut blob = EncryptedBlob::seal(&key, b"sensitive data").unwrap();
let last = blob.0.len() - 1;
blob.0[last] ^= 0xFF;
let result = blob.open(&key);
assert!(
matches!(result, Err(VaultError::DecryptionFailed)),
"expected DecryptionFailed, got {result:?}"
);
}
#[test]
fn tampered_nonce_rejected() {
let key = VaultKey::generate();
let mut blob = EncryptedBlob::seal(&key, b"sensitive data").unwrap();
blob.0[0] ^= 0xFF;
let result = blob.open(&key);
assert!(
matches!(result, Err(VaultError::DecryptionFailed)),
"expected DecryptionFailed, got {result:?}"
);
}
#[test]
fn nonces_are_unique_across_encryptions() {
let key = VaultKey::generate();
let n = 1000;
let mut nonces: HashSet<Vec<u8>> = HashSet::new();
for _ in 0..n {
let blob = EncryptedBlob::seal(&key, b"same plaintext").unwrap();
nonces.insert(blob.nonce_bytes().to_vec());
}
assert_eq!(nonces.len(), n, "nonce collision detected");
}
#[test]
fn same_plaintext_yields_different_ciphertexts() {
let key = VaultKey::generate();
let blob1 = EncryptedBlob::seal(&key, b"plaintext").unwrap();
let blob2 = EncryptedBlob::seal(&key, b"plaintext").unwrap();
assert_ne!(
blob1.as_bytes(),
blob2.as_bytes(),
"ciphertexts must differ due to unique nonces"
);
}
#[test]
fn ttl_secret_accessible_before_expiry() {
let mut v = make_vault();
v.put("api_key", b"live-key", Some(60)).unwrap();
let plain = v.get("api_key").unwrap();
assert_eq!(plain, b"live-key");
}
#[test]
fn expired_secret_returns_error() {
let mut v = make_vault();
v.put("stale_key", b"old-value", Some(-1)).unwrap();
let result = v.get("stale_key");
assert!(matches!(result, Err(VaultError::Expired(_))), "expected Expired, got {result:?}");
}
#[test]
fn missing_secret_returns_not_found() {
let v = make_vault();
let result = v.get("nonexistent");
assert!(
matches!(result, Err(VaultError::NotFound(_))),
"expected NotFound, got {result:?}"
);
}
#[test]
fn version_increments_on_put() {
let mut v = make_vault();
v.put("key", b"v1", None).unwrap();
assert_eq!(v.entry("key").unwrap().version, 1);
v.put("key", b"v2", None).unwrap();
assert_eq!(v.entry("key").unwrap().version, 2);
}
#[test]
fn rotate_re_encrypts_with_new_nonce() {
let mut v = make_vault();
v.put("rot_key", b"my-secret", None).unwrap();
let nonce_before = v.entry("rot_key").unwrap().blob.nonce_bytes().to_vec();
v.rotate("rot_key").unwrap();
let nonce_after = v.entry("rot_key").unwrap().blob.nonce_bytes().to_vec();
assert_ne!(nonce_before, nonce_after, "rotation must use a fresh nonce");
assert_eq!(v.get("rot_key").unwrap(), b"my-secret");
}
#[test]
fn rotate_increments_version() {
let mut v = make_vault();
v.put("ver_key", b"data", None).unwrap();
v.rotate("ver_key").unwrap();
assert_eq!(v.entry("ver_key").unwrap().version, 2);
}
#[test]
fn remove_deletes_entry() {
let mut v = make_vault();
v.put("tmp", b"temporary", None).unwrap();
assert!(v.remove("tmp"));
assert!(matches!(v.get("tmp"), Err(VaultError::NotFound(_))));
}
#[test]
fn list_excludes_expired_entries() {
let mut v = make_vault();
v.put("live", b"a", Some(60)).unwrap();
v.put("dead", b"b", Some(-1)).unwrap();
let names = v.list();
assert!(names.contains(&"live"));
assert!(!names.contains(&"dead"));
}
#[test]
fn vault_key_from_bytes_wrong_length_fails() {
let result = VaultKey::from_bytes(&[0u8; 16]);
assert!(matches!(result, Err(VaultError::InvalidKeyLength)));
}
#[test]
fn vault_key_from_bytes_round_trip() {
let key = VaultKey::generate();
let raw = key.raw;
let restored = VaultKey::from_bytes(&raw).unwrap();
let blob = EncryptedBlob::seal(&key, b"test").unwrap();
assert_eq!(blob.open(&restored).unwrap(), b"test");
}
use std::sync::Arc;
use crate::adapters::audit::InMemoryAuditSink;
use crate::domain::ports::{AuditAction, AuditOutcome, AuditSink};
fn make_audited_vault() -> (SecretVault, Arc<InMemoryAuditSink>) {
let sink = Arc::new(InMemoryAuditSink::new());
let vault = SecretVault::new(VaultKey::generate())
.with_audit_sink(Arc::clone(&sink) as Arc<dyn AuditSink>);
(vault, sink)
}
#[test]
fn vault_put_emits_vault_write_success() {
let (mut vault, sink) = make_audited_vault();
vault.put("api_key", b"value", None).unwrap();
let events = sink.events();
assert_eq!(events.len(), 1);
assert_eq!(events[0].action, AuditAction::VaultWrite);
assert_eq!(events[0].outcome, AuditOutcome::Success);
assert_eq!(events[0].subject, "api_key");
}
#[test]
fn vault_get_emits_vault_read_success() {
let (mut vault, sink) = make_audited_vault();
vault.put("k", b"v", None).unwrap();
sink.drain();
vault.get("k").unwrap();
let events = sink.events();
assert_eq!(events.len(), 1);
assert_eq!(events[0].action, AuditAction::VaultRead);
assert_eq!(events[0].outcome, AuditOutcome::Success);
}
#[test]
fn vault_get_expired_emits_vault_read_failure_with_reason() {
let (mut vault, sink) = make_audited_vault();
vault.put("exp_key", b"v", Some(-1)).unwrap();
sink.drain();
let _ = vault.get("exp_key");
let events = sink.events();
assert_eq!(events.len(), 1);
assert_eq!(events[0].action, AuditAction::VaultRead);
assert_eq!(events[0].outcome, AuditOutcome::Failure);
assert_eq!(events[0].reason.as_deref(), Some("secret expired"));
}
#[test]
fn vault_rotate_emits_read_then_write() {
let (mut vault, sink) = make_audited_vault();
vault.put("r", b"data", None).unwrap();
sink.drain();
vault.rotate("r").unwrap();
let events = sink.events();
assert!(events.len() >= 2, "rotate must emit VaultRead + VaultWrite");
assert!(events.iter().any(|e| e.action == AuditAction::VaultRead));
assert!(events.iter().any(|e| e.action == AuditAction::VaultWrite));
}
#[test]
fn vault_audit_event_contains_no_plaintext_secret() {
let (mut vault, sink) = make_audited_vault();
let plaintext = b"super-secret-password-12345";
vault.put("pw", plaintext, None).unwrap();
vault.get("pw").unwrap();
for event in sink.events() {
let ptext = std::str::from_utf8(plaintext).unwrap();
assert_ne!(event.subject, ptext, "plaintext in subject");
if let Some(reason) = &event.reason {
assert_ne!(reason.as_str(), ptext, "plaintext in reason");
}
}
}
use chacha20poly1305::{aead::OsRng as ChaChaOsRng, KeyInit};
use crate::adapters::kms::LocalKmsAdapter;
fn make_kms_vault() -> KmsSecretVault {
let kek_raw = {
let k = chacha20poly1305::ChaCha20Poly1305::generate_key(&mut ChaChaOsRng);
let mut raw = [0u8; 32];
raw.copy_from_slice(&k);
raw
};
let kms = Arc::new(LocalKmsAdapter::new(kek_raw));
KmsSecretVault::new(kms)
}
#[test]
fn kms_vault_round_trip() {
let mut vault = make_kms_vault();
vault.put("api_key", b"super-secret", None).unwrap();
let plain = vault.get("api_key").unwrap();
assert_eq!(plain, b"super-secret");
}
#[test]
fn kms_vault_round_trip_unicode() {
let mut vault = make_kms_vault();
let secret = "日本語テスト🔑";
vault.put("uni", secret.as_bytes(), None).unwrap();
assert_eq!(vault.get("uni").unwrap(), secret.as_bytes());
}
#[test]
fn kms_vault_round_trip_empty() {
let mut vault = make_kms_vault();
vault.put("empty", b"", None).unwrap();
assert_eq!(vault.get("empty").unwrap(), b"");
}
#[test]
fn kms_vault_wrong_kek_fails_to_decrypt() {
let kek_a = {
let k = chacha20poly1305::ChaCha20Poly1305::generate_key(&mut ChaChaOsRng);
let mut r = [0u8; 32];
r.copy_from_slice(&k);
r
};
let kek_b = {
let k = chacha20poly1305::ChaCha20Poly1305::generate_key(&mut ChaChaOsRng);
let mut r = [0u8; 32];
r.copy_from_slice(&k);
r
};
let kms_a = Arc::new(LocalKmsAdapter::new(kek_a));
let kms_b = Arc::new(LocalKmsAdapter::new(kek_b));
let mut vault_a = KmsSecretVault::new(Arc::clone(&kms_a) as Arc<dyn KeyManagementService>);
vault_a.put("secret", b"value", None).unwrap();
let entry_a = vault_a.entry("secret").unwrap().clone();
let mut vault_b = KmsSecretVault::new(Arc::clone(&kms_b) as Arc<dyn KeyManagementService>);
vault_b.entries.insert("secret".to_string(), entry_a);
let result = vault_b.get("secret");
assert!(
matches!(result, Err(VaultError::DecryptionFailed)),
"wrong KEK must not decrypt: {result:?}"
);
}
#[test]
fn kms_vault_each_secret_has_distinct_dek() {
let mut vault = make_kms_vault();
vault.put("s1", b"value1", None).unwrap();
vault.put("s2", b"value2", None).unwrap();
let wrapped1 = vault.entry("s1").unwrap().wrapped_dek.clone();
let wrapped2 = vault.entry("s2").unwrap().wrapped_dek.clone();
assert_ne!(wrapped1, wrapped2, "each secret must use a distinct DEK");
}
#[test]
fn kms_vault_tampered_wrapped_dek_rejected() {
let mut vault = make_kms_vault();
vault.put("tok", b"payload", None).unwrap();
if let Some(entry) = vault.entries.get_mut("tok") {
entry.wrapped_dek[12] ^= 0xFF;
}
let result = vault.get("tok");
assert!(
matches!(result, Err(VaultError::DecryptionFailed)),
"tampered wrapped DEK must be rejected: {result:?}"
);
}
#[test]
fn kms_vault_rotate_uses_new_dek() {
let mut vault = make_kms_vault();
vault.put("key", b"data", None).unwrap();
let wrapped_before = vault.entry("key").unwrap().wrapped_dek.clone();
vault.rotate("key").unwrap();
let wrapped_after = vault.entry("key").unwrap().wrapped_dek.clone();
assert_ne!(wrapped_before, wrapped_after, "rotation must generate a new DEK");
assert_eq!(vault.get("key").unwrap(), b"data");
}
}