use argon2::{Algorithm, Argon2, Params, Version};
use chacha20poly1305::{
XChaCha20Poly1305,
aead::{Aead, KeyInit},
};
use secrecy::{ExposeSecret, SecretBox};
use zeroize::ZeroizeOnDrop;
use crate::error::{Error, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct KdfParams {
pub m_kib: u32,
pub t_cost: u32,
pub p_cost: u32,
}
impl Default for KdfParams {
fn default() -> Self {
Self {
m_kib: 65_536, t_cost: 3,
p_cost: 4,
}
}
}
#[derive(ZeroizeOnDrop)]
pub struct MasterKey(pub(crate) SecretBox<[u8; 32]>);
impl MasterKey {
pub fn as_bytes(&self) -> &[u8; 32] {
self.0.expose_secret()
}
#[cfg_attr(not(unix), allow(dead_code))]
pub(crate) fn from_bytes(arr: [u8; 32]) -> Self {
MasterKey(SecretBox::new(Box::new(arr)))
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Encrypted {
pub nonce: [u8; 24],
pub ciphertext: Vec<u8>,
}
pub fn derive_key(passphrase: &[u8], salt: &[u8; 16], params: &KdfParams) -> Result<MasterKey> {
let argon_params = Params::new(params.m_kib, params.t_cost, params.p_cost, Some(32))
.map_err(|e| Error::Crypto(format!("invalid argon2 params: {e}")))?;
let argon = Argon2::new(Algorithm::Argon2id, Version::V0x13, argon_params);
let mut out = [0u8; 32];
argon
.hash_password_into(passphrase, salt, &mut out)
.map_err(|e| Error::Crypto(format!("argon2 derivation failed: {e}")))?;
Ok(MasterKey(SecretBox::new(Box::new(out))))
}
pub fn encrypt(key: &MasterKey, plaintext: &[u8]) -> Result<Encrypted> {
let cipher = XChaCha20Poly1305::new(key.as_bytes().into());
let nonce = gen_nonce();
let nonce_obj = chacha20poly1305::XNonce::from_slice(&nonce);
let ciphertext = cipher
.encrypt(nonce_obj, plaintext)
.map_err(|e| Error::Crypto(format!("encryption failed: {e}")))?;
Ok(Encrypted { nonce, ciphertext })
}
pub fn decrypt(key: &MasterKey, enc: &Encrypted) -> Result<Vec<u8>> {
let cipher = XChaCha20Poly1305::new(key.as_bytes().into());
let nonce_obj = chacha20poly1305::XNonce::from_slice(&enc.nonce);
cipher
.decrypt(nonce_obj, enc.ciphertext.as_ref())
.map_err(|_| Error::BadPassphrase)
}
pub fn gen_salt() -> [u8; 16] {
let mut buf = [0u8; 16];
getrandom::fill(&mut buf).expect("getrandom::fill failed for salt generation");
buf
}
pub fn gen_nonce() -> [u8; 24] {
let mut buf = [0u8; 24];
getrandom::fill(&mut buf).expect("getrandom::fill failed for nonce generation");
buf
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn kdf_params_default_matches_prd() {
let p = KdfParams::default();
assert_eq!(p.m_kib, 65_536); assert_eq!(p.t_cost, 3);
assert_eq!(p.p_cost, 4);
}
#[test]
fn derive_key_is_deterministic_for_same_inputs() {
let params = KdfParams {
m_kib: 4_096, t_cost: 1,
p_cost: 1,
};
let salt = gen_salt();
let k1 = derive_key(b"correct horse battery staple", &salt, ¶ms).unwrap();
let k2 = derive_key(b"correct horse battery staple", &salt, ¶ms).unwrap();
assert_eq!(k1.as_bytes(), k2.as_bytes());
let k3 = derive_key(b"different passphrase", &salt, ¶ms).unwrap();
assert_ne!(k1.as_bytes(), k3.as_bytes());
}
#[test]
fn encrypt_decrypt_roundtrip() {
let params = KdfParams {
m_kib: 4_096,
t_cost: 1,
p_cost: 1,
};
let salt = gen_salt();
let key = derive_key(b"round-trip-passphrase", &salt, ¶ms).unwrap();
let plaintext = b"hello zkv \x00 binary \xff data";
let enc = encrypt(&key, plaintext).unwrap();
let dec = decrypt(&key, &enc).unwrap();
assert_eq!(dec, plaintext);
assert_eq!(enc.ciphertext.len(), plaintext.len() + 16);
}
#[test]
fn decrypt_with_wrong_key_fails() {
let params = KdfParams {
m_kib: 4_096,
t_cost: 1,
p_cost: 1,
};
let salt = gen_salt();
let good = derive_key(b"the right passphrase", &salt, ¶ms).unwrap();
let bad = derive_key(b"the wrong passphrase", &salt, ¶ms).unwrap();
let enc = encrypt(&good, b"secret payload").unwrap();
let res = decrypt(&bad, &enc);
assert!(matches!(res, Err(Error::BadPassphrase)));
}
#[test]
fn decrypt_tampered_ciphertext_fails() {
let params = KdfParams {
m_kib: 4_096,
t_cost: 1,
p_cost: 1,
};
let salt = gen_salt();
let key = derive_key(b"tamper-test-passphrase", &salt, ¶ms).unwrap();
let mut enc = encrypt(&key, b"tamper me").unwrap();
enc.ciphertext[0] ^= 0xff;
let res = decrypt(&key, &enc);
assert!(matches!(res, Err(Error::BadPassphrase)));
}
#[test]
fn decrypt_tampered_nonce_fails() {
let params = KdfParams {
m_kib: 4_096,
t_cost: 1,
p_cost: 1,
};
let salt = gen_salt();
let key = derive_key(b"nonce-tamper", &salt, ¶ms).unwrap();
let mut enc = encrypt(&key, b"payload").unwrap();
enc.nonce[0] ^= 0x01;
let res = decrypt(&key, &enc);
assert!(matches!(res, Err(Error::BadPassphrase)));
}
#[test]
fn gen_salt_and_nonce_are_random() {
let s1 = gen_salt();
let s2 = gen_salt();
assert_ne!(s1, s2);
let n1 = gen_nonce();
let n2 = gen_nonce();
assert_ne!(n1, n2);
}
}