use crate::error::{Error, Result};
use zeroize::Zeroizing;
use chacha20poly1305::aead::AeadInPlace;
use chacha20poly1305::{KeyInit, Tag, XChaCha20Poly1305, XNonce};
const TAG_LEN: usize = 16;
#[must_use = "ciphertext must be transmitted or stored"]
#[allow(clippy::explicit_auto_deref)] pub fn aead_encrypt(
key: &[u8; 32],
nonce: &[u8; 24],
plaintext: &[u8],
aad: &[u8],
) -> Result<Vec<u8>> {
let cipher = XChaCha20Poly1305::new(key.into());
let nonce = XNonce::from_slice(nonce);
let mut buffer = Zeroizing::new({
let cap = plaintext
.len()
.checked_add(TAG_LEN)
.ok_or(Error::AeadFailed)?;
let mut v = Vec::with_capacity(cap);
v.extend_from_slice(plaintext);
v
});
let tag = cipher
.encrypt_in_place_detached(nonce, aad, &mut *buffer)
.map_err(|_| Error::AeadFailed)?;
debug_assert!(
buffer.capacity() >= buffer.len() + TAG_LEN,
"aead_encrypt: unexpected capacity — reallocation would bypass Zeroizing"
);
buffer.extend_from_slice(&tag);
Ok(std::mem::take(&mut *buffer))
}
#[must_use = "decrypted plaintext must be consumed or zeroized"]
#[allow(clippy::explicit_auto_deref)] pub fn aead_decrypt(
key: &[u8; 32],
nonce: &[u8; 24],
ciphertext: &[u8],
aad: &[u8],
) -> Result<Zeroizing<Vec<u8>>> {
if ciphertext.len() < TAG_LEN {
return Err(Error::AeadFailed);
}
let cipher = XChaCha20Poly1305::new(key.into());
let nonce = XNonce::from_slice(nonce);
let ct_len = ciphertext.len() - TAG_LEN;
let mut buffer = Zeroizing::new(ciphertext[..ct_len].to_vec());
let tag = Tag::from_slice(&ciphertext[ct_len..]);
cipher
.decrypt_in_place_detached(nonce, aad, &mut *buffer, tag)
.map_err(|_| Error::AeadFailed)?;
Ok(buffer)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
use hex_literal::hex;
#[test]
fn deterministic_consistency() {
let key = hex!("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20");
let nonce = hex!("000000000000000000000000000000000000000000000001");
let pt = b"soliton-aead-kat";
let aad = b"lo-test-v1";
let ct1 = aead_encrypt(&key, &nonce, pt, aad).unwrap();
let ct2 = aead_encrypt(&key, &nonce, pt, aad).unwrap();
assert_eq!(ct1, ct2);
assert_eq!(ct1.len(), pt.len() + 16);
let decrypted = aead_decrypt(&key, &nonce, &ct1, aad).unwrap();
assert_eq!(&*decrypted, pt);
}
#[test]
fn round_trip() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let plaintext = b"round trip test";
let ct = aead_encrypt(&key, &nonce, plaintext, b"").unwrap();
let pt = aead_decrypt(&key, &nonce, &ct, b"").unwrap();
assert_eq!(&*pt, plaintext);
}
#[test]
fn round_trip_with_aad() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let aad = b"additional data";
let ct = aead_encrypt(&key, &nonce, b"secret", aad).unwrap();
let pt = aead_decrypt(&key, &nonce, &ct, aad).unwrap();
assert_eq!(&*pt, b"secret");
}
#[test]
fn round_trip_empty_plaintext() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let ct = aead_encrypt(&key, &nonce, b"", b"").unwrap();
assert_eq!(ct.len(), 16); let pt = aead_decrypt(&key, &nonce, &ct, b"").unwrap();
assert!(pt.is_empty());
}
#[test]
fn round_trip_large_plaintext() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let plaintext = vec![0xABu8; 65536];
let ct = aead_encrypt(&key, &nonce, &plaintext, b"").unwrap();
let pt = aead_decrypt(&key, &nonce, &ct, b"").unwrap();
assert_eq!(&*pt, &plaintext);
}
#[test]
fn tampered_ciphertext() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let mut ct = aead_encrypt(&key, &nonce, b"plaintext", b"").unwrap();
ct[0] ^= 0xFF; assert!(matches!(
aead_decrypt(&key, &nonce, &ct, b""),
Err(Error::AeadFailed)
));
}
#[test]
fn tampered_tag() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let mut ct = aead_encrypt(&key, &nonce, b"plaintext", b"").unwrap();
let last = ct.len() - 1;
ct[last] ^= 0xFF; assert!(matches!(
aead_decrypt(&key, &nonce, &ct, b""),
Err(Error::AeadFailed)
));
}
#[test]
fn wrong_key() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let ct = aead_encrypt(&key, &nonce, b"plaintext", b"").unwrap();
let wrong_key: [u8; 32] = crate::primitives::random::random_array();
assert!(matches!(
aead_decrypt(&wrong_key, &nonce, &ct, b""),
Err(Error::AeadFailed)
));
}
#[test]
fn wrong_nonce() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let ct = aead_encrypt(&key, &nonce, b"plaintext", b"").unwrap();
let wrong_nonce: [u8; 24] = crate::primitives::random::random_array();
assert!(matches!(
aead_decrypt(&key, &wrong_nonce, &ct, b""),
Err(Error::AeadFailed)
));
}
#[test]
fn wrong_aad() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let ct = aead_encrypt(&key, &nonce, b"plaintext", b"correct").unwrap();
assert!(matches!(
aead_decrypt(&key, &nonce, &ct, b"wrong"),
Err(Error::AeadFailed)
));
}
#[test]
fn too_short_ciphertext() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let short = vec![0u8; 15]; assert!(matches!(
aead_decrypt(&key, &nonce, &short, b""),
Err(Error::AeadFailed)
));
}
#[test]
fn decrypt_returns_zeroizing() {
let key: [u8; 32] = crate::primitives::random::random_array();
let nonce: [u8; 24] = crate::primitives::random::random_array();
let ct = aead_encrypt(&key, &nonce, b"test", b"").unwrap();
let result: Zeroizing<Vec<u8>> = aead_decrypt(&key, &nonce, &ct, b"").unwrap();
assert_eq!(&*result, b"test");
}
}