use argon2::{Algorithm, Argon2, Params, Version};
use chacha20poly1305::aead::{Aead, KeyInit};
use chacha20poly1305::{XChaCha20Poly1305, XNonce};
use color_eyre::Result;
use color_eyre::eyre::eyre;
use hmac::{Hmac, Mac};
use serde::{Deserialize, Serialize};
use sha2::Sha256;
use zeroize::Zeroize;
pub type HmacSha256 = Hmac<Sha256>;
#[derive(Clone, Debug)]
pub struct KeyMaterial(pub [u8; 32]);
impl KeyMaterial {
#[allow(clippy::expect_used)]
#[must_use]
pub fn random() -> Self {
let mut k = [0u8; 32];
getrandom::fill(&mut k).expect("Failed to get random bytes");
Self(k)
}
}
impl Drop for KeyMaterial {
fn drop(&mut self) {
self.0.zeroize();
}
}
#[derive(Clone, Serialize, Deserialize)]
pub struct KdfParams {
pub salt: Vec<u8>,
pub m_cost_kib: u32,
pub t_cost: u32,
pub p_cost: u32,
}
impl KdfParams {
#[allow(clippy::expect_used)]
#[must_use]
pub fn default_secure() -> Self {
let mut salt = vec![0u8; 16];
getrandom::fill(&mut salt).expect("Failed to get random bytes");
Self {
salt,
m_cost_kib: 19456,
t_cost: 3,
p_cost: 1,
} }
}
pub fn derive_key(master: &str, kdf: &KdfParams) -> Result<KeyMaterial> {
let argon2 = Argon2::new(
Algorithm::Argon2id,
Version::V0x13,
Params::new(kdf.m_cost_kib, kdf.t_cost, kdf.p_cost, Some(32)).map_err(|e| eyre!("{e}"))?,
);
let mut out = [0u8; 32];
argon2
.hash_password_into(master.as_bytes(), &kdf.salt, &mut out)
.map_err(|e| eyre!("{e}"))?;
Ok(KeyMaterial(out))
}
#[derive(Serialize, Deserialize)]
pub struct WrappedVaultKey {
pub nonce: Vec<u8>,
pub ciphertext: Vec<u8>,
}
pub fn wrap_vault_key(master_derived: &KeyMaterial, vault_key: &KeyMaterial) -> Result<(WrappedVaultKey, Vec<u8>)> {
let aead = XChaCha20Poly1305::new((&master_derived.0).into());
let mut nonce = [0u8; 24];
getrandom::fill(&mut nonce)?;
let ct = aead
.encrypt(XNonce::from_slice(&nonce), vault_key.0.as_ref())
.map_err(|_| eyre!("AEAD encrypt failed"))?;
let wrapped = WrappedVaultKey {
nonce: nonce.to_vec(),
ciphertext: ct,
};
let mut mac = <HmacSha256 as Mac>::new_from_slice(&master_derived.0)?;
mac.update(b"chamber-verifier");
let tag = mac.finalize().into_bytes().to_vec();
Ok((wrapped, tag))
}
pub fn unwrap_vault_key(
master_derived: &KeyMaterial,
wrapped: &WrappedVaultKey,
verifier: Option<&[u8]>,
) -> Result<KeyMaterial> {
if let Some(v) = verifier {
let mut mac = <HmacSha256 as Mac>::new_from_slice(&master_derived.0)?;
mac.update(b"chamber-verifier");
mac.verify_slice(v).map_err(|_| eyre!("Verifier mismatch"))?;
}
let aead = XChaCha20Poly1305::new((&master_derived.0).into());
let nonce = XNonce::from_slice(&wrapped.nonce);
let pt = aead
.decrypt(nonce, wrapped.ciphertext.as_ref())
.map_err(|_| eyre!("AEAD decrypt failed"))?;
let mut key = [0u8; 32];
key.copy_from_slice(&pt);
Ok(KeyMaterial(key))
}
pub fn aead_encrypt(vault_key: &KeyMaterial, plaintext: &[u8], ad: &[u8]) -> Result<(Vec<u8>, Vec<u8>)> {
let aead = XChaCha20Poly1305::new((&vault_key.0).into());
let mut nonce = [0u8; 24];
getrandom::fill(&mut nonce)?;
let ct = aead
.encrypt(
XNonce::from_slice(&nonce),
chacha20poly1305::aead::Payload {
msg: plaintext,
aad: ad,
},
)
.map_err(|_| eyre!("encrypt failed"))?;
Ok((nonce.to_vec(), ct))
}
pub fn aead_decrypt(vault_key: &KeyMaterial, nonce: &[u8], ciphertext: &[u8], ad: &[u8]) -> Result<Vec<u8>> {
let aead = XChaCha20Poly1305::new((&vault_key.0).into());
let pt = aead
.decrypt(
XNonce::from_slice(nonce),
chacha20poly1305::aead::Payload {
msg: ciphertext,
aad: ad,
},
)
.map_err(|_| eyre!("decrypt failed"))?;
Ok(pt)
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
use super::*;
use hex::encode as hex_encode;
fn small_kdf(salt: &[u8]) -> KdfParams {
let mut s = salt.to_vec();
if s.len() < 8 {
s.resize(8, 0); }
KdfParams {
salt: s,
m_cost_kib: 8, t_cost: 1,
p_cost: 1,
}
}
#[test]
fn test_keymaterial_random_and_length() {
let k1 = KeyMaterial::random();
let k2 = KeyMaterial::random();
assert_eq!(k1.0.len(), 32);
assert_eq!(k2.0.len(), 32);
assert_ne!(hex_encode(k1.0), hex_encode(k2.0));
}
#[test]
fn test_derive_key_deterministic_and_salt_sensitive() {
let kdf1 = small_kdf(b"salt-1");
let kdf2 = small_kdf(b"salt-2");
let master = "correct horse battery staple";
let a = derive_key(master, &kdf1).unwrap();
let b = derive_key(master, &kdf1).unwrap();
let c = derive_key(master, &kdf2).unwrap();
assert_eq!(hex_encode(a.0), hex_encode(b.0));
assert_ne!(hex_encode(a.0), hex_encode(c.0));
}
#[test]
fn test_aead_encrypt_decrypt_roundtrip_with_ad() {
let key = KeyMaterial::random();
let msg = b"secret message";
let ad = b"associated-data";
let (nonce, ct) = aead_encrypt(&key, msg, ad).unwrap();
let pt = aead_decrypt(&key, &nonce, &ct, ad).unwrap();
assert_eq!(pt, msg);
}
#[test]
fn test_aead_decrypt_wrong_ad_fails() {
let key = KeyMaterial::random();
let msg = b"message";
let ad_ok = b"ad-ok";
let ad_bad = b"ad-bad";
let (nonce, ct) = aead_encrypt(&key, msg, ad_ok).unwrap();
let err = aead_decrypt(&key, &nonce, &ct, ad_bad).unwrap_err();
assert!(err.to_string().to_lowercase().contains("decrypt"));
}
#[test]
fn test_aead_decrypt_wrong_key_fails() {
let key1 = KeyMaterial::random();
let key2 = KeyMaterial::random();
let (nonce, ct) = aead_encrypt(&key1, b"data", b"ad").unwrap();
let err = aead_decrypt(&key2, &nonce, &ct, b"ad").unwrap_err();
assert!(err.to_string().to_lowercase().contains("decrypt"));
}
#[test]
fn test_aead_tamper_detection() {
let key = KeyMaterial::random();
let (nonce, mut ct) = aead_encrypt(&key, b"payload", b"ad").unwrap();
if let Some(byte) = ct.get_mut(0) {
*byte ^= 0x01;
}
let err = aead_decrypt(&key, &nonce, &ct, b"ad").unwrap_err();
assert!(err.to_string().to_lowercase().contains("decrypt"));
}
#[test]
fn test_wrap_unwrap_vault_key_roundtrip_and_verifier() {
let master = "test-master";
let kdf = small_kdf(b"wrapsalt");
let master_derived = derive_key(master, &kdf).unwrap();
let vk = KeyMaterial::random();
let (wrapped, verifier) = wrap_vault_key(&master_derived, &vk).unwrap();
let unwrapped = unwrap_vault_key(&master_derived, &wrapped, Some(&verifier)).unwrap();
assert_eq!(hex_encode(vk.0), hex_encode(unwrapped.0));
let unwrapped2 = unwrap_vault_key(&master_derived, &wrapped, None).unwrap();
assert_eq!(hex_encode(vk.0), hex_encode(unwrapped2.0));
}
#[test]
fn test_unwrap_verifier_mismatch_fails() {
let master_ok = "master-ok";
let master_bad = "master-bad";
let kdf = small_kdf(b"v-salt");
let md_ok = derive_key(master_ok, &kdf).unwrap();
let md_bad = derive_key(master_bad, &kdf).unwrap();
let vk = KeyMaterial::random();
let (wrapped, verifier) = wrap_vault_key(&md_ok, &vk).unwrap();
let err = unwrap_vault_key(&md_bad, &wrapped, Some(&verifier)).unwrap_err();
assert!(err.to_string().to_lowercase().contains("verifier"));
}
#[test]
fn test_unwrap_with_tampered_ciphertext_fails() {
let master = "master";
let kdf = small_kdf(b"salt-x");
let md = derive_key(master, &kdf).unwrap();
let vk = KeyMaterial::random();
let (mut wrapped, verifier) = wrap_vault_key(&md, &vk).unwrap();
if let Some(byte) = wrapped.ciphertext.get_mut(0) {
*byte ^= 0x80;
}
let err = unwrap_vault_key(&md, &wrapped, Some(&verifier)).unwrap_err();
assert!(err.to_string().to_lowercase().contains("aead"));
}
#[test]
fn test_kdfparams_default_secure_has_expected_shape() {
let kdf = KdfParams::default_secure();
assert_eq!(kdf.salt.len(), 16);
assert!(kdf.m_cost_kib >= 1024);
assert!(kdf.t_cost >= 1);
assert!(kdf.p_cost >= 1);
let km = derive_key("pw", &kdf).unwrap();
assert_eq!(km.0.len(), 32);
}
#[test]
fn test_hmac_verifier_stable_for_same_key() {
let master = "verifier-master";
let kdf = small_kdf(b"vsalt");
let md = derive_key(master, &kdf).unwrap();
let vk = KeyMaterial::random();
let (_, tag1) = wrap_vault_key(&md, &vk).unwrap();
let (_, tag2) = wrap_vault_key(&md, &vk).unwrap();
assert_eq!(hex_encode(tag1), hex_encode(tag2));
}
#[test]
fn test_hmac_verifier_differs_for_different_master_keys() {
let kdf = small_kdf(b"vsalt");
let md1 = derive_key("m1", &kdf).unwrap();
let md2 = derive_key("m2", &kdf).unwrap();
let vk = KeyMaterial::random();
let (_, tag1) = wrap_vault_key(&md1, &vk).unwrap();
let (_, tag2) = wrap_vault_key(&md2, &vk).unwrap();
assert_ne!(hex_encode(tag1), hex_encode(tag2));
}
mod hex {
#[allow(clippy::format_collect)]
pub fn encode<T: AsRef<[u8]>>(data: T) -> String {
data.as_ref().iter().map(|b| format!("{b:02x}")).collect()
}
}
}