use std::fmt;
use std::io::Read;
use chacha20poly1305::aead::{Aead, KeyInit, Payload};
use chacha20poly1305::XChaCha20Poly1305;
use zeroize::Zeroizing;
pub(crate) const STENOXIDE_AAD: &[u8] = b"STENOXIDE-v1";
#[cfg(feature = "pqc")]
pub(crate) const STENOXIDE_IDENTITY_AAD: &[u8] = b"STENOXIDE-identity-v1";
const ZSTD_LEVEL: i32 = 19;
const MAX_DECOMPRESSED_BYTES: u64 = 512 * 1024 * 1024;
#[derive(Debug)]
pub enum AEADError {
AuthenticationFailed,
CipherError(String),
}
impl fmt::Display for AEADError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
AEADError::AuthenticationFailed => {
write!(f, "authentication failed: wrong password or corrupted data")
}
AEADError::CipherError(message) => write!(f, "cipher error: {message}"),
}
}
}
impl std::error::Error for AEADError {}
#[derive(Debug)]
pub enum CryptoError {
CompressionError(String),
DecompressionError(String),
AEADError(AEADError),
}
impl fmt::Display for CryptoError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
CryptoError::CompressionError(message) => {
write!(f, "failed to compress the payload: {message}")
}
CryptoError::DecompressionError(message) => {
write!(f, "failed to decompress the payload: {message}")
}
CryptoError::AEADError(err) => write!(f, "{err}"),
}
}
}
impl std::error::Error for CryptoError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
CryptoError::AEADError(err) => Some(err),
_ => None,
}
}
}
impl From<AEADError> for CryptoError {
fn from(err: AEADError) -> Self {
CryptoError::AEADError(err)
}
}
pub trait AEADCipher: Send + Sync {
fn encrypt(
&self,
key: &[u8; 32],
nonce: &[u8; 24],
plaintext: &[u8],
aad: &[u8],
) -> Result<Zeroizing<Vec<u8>>, AEADError>;
fn decrypt(
&self,
key: &[u8; 32],
nonce: &[u8; 24],
ciphertext: &[u8],
aad: &[u8],
) -> Result<Zeroizing<Vec<u8>>, AEADError>;
}
#[derive(Debug, Default, Clone, Copy)]
pub struct XChaCha20Poly1305Cipher;
impl XChaCha20Poly1305Cipher {
pub fn new() -> Self {
Self
}
}
impl AEADCipher for XChaCha20Poly1305Cipher {
fn encrypt(
&self,
key: &[u8; 32],
nonce: &[u8; 24],
plaintext: &[u8],
aad: &[u8],
) -> Result<Zeroizing<Vec<u8>>, AEADError> {
let cipher = XChaCha20Poly1305::new(key.into());
let ciphertext = cipher
.encrypt(
nonce.into(),
Payload {
msg: plaintext,
aad,
},
)
.map_err(|err| AEADError::CipherError(err.to_string()))?;
Ok(Zeroizing::new(ciphertext))
}
fn decrypt(
&self,
key: &[u8; 32],
nonce: &[u8; 24],
ciphertext: &[u8],
aad: &[u8],
) -> Result<Zeroizing<Vec<u8>>, AEADError> {
let cipher = XChaCha20Poly1305::new(key.into());
let plaintext = cipher
.decrypt(
nonce.into(),
Payload {
msg: ciphertext,
aad,
},
)
.map_err(|_| AEADError::AuthenticationFailed)?;
Ok(Zeroizing::new(plaintext))
}
}
pub(crate) fn compress(plaintext: &[u8]) -> Result<Zeroizing<Vec<u8>>, CryptoError> {
zstd::encode_all(plaintext, ZSTD_LEVEL)
.map(Zeroizing::new)
.map_err(|err| CryptoError::CompressionError(err.to_string()))
}
pub(crate) fn decompress(compressed: &[u8]) -> Result<Zeroizing<Vec<u8>>, CryptoError> {
decompress_within(compressed, MAX_DECOMPRESSED_BYTES)
}
fn decompress_within(compressed: &[u8], ceiling: u64) -> Result<Zeroizing<Vec<u8>>, CryptoError> {
let decoder = zstd::stream::read::Decoder::new(compressed)
.map_err(|err| CryptoError::DecompressionError(err.to_string()))?;
let mut plaintext = Zeroizing::new(Vec::new());
decoder
.take(ceiling + 1)
.read_to_end(&mut plaintext)
.map_err(|err| CryptoError::DecompressionError(err.to_string()))?;
if plaintext.len() as u64 > ceiling {
return Err(CryptoError::DecompressionError(format!(
"the payload expands past the {ceiling}-byte ceiling"
)));
}
Ok(plaintext)
}
pub fn compress_and_encrypt(
plaintext: &[u8],
enc_key: &[u8; 32],
nonce: &[u8; 24],
cipher: &dyn AEADCipher,
) -> Result<Zeroizing<Vec<u8>>, CryptoError> {
let compressed = compress(plaintext)?;
let ciphertext = cipher.encrypt(enc_key, nonce, &compressed, STENOXIDE_AAD)?;
drop(compressed);
Ok(ciphertext)
}
pub fn decrypt_and_decompress(
ciphertext: &[u8],
enc_key: &[u8; 32],
nonce: &[u8; 24],
cipher: &dyn AEADCipher,
) -> Result<Zeroizing<Vec<u8>>, CryptoError> {
let compressed = cipher.decrypt(enc_key, nonce, ciphertext, STENOXIDE_AAD)?;
let plaintext = decompress(compressed.as_slice())?;
drop(compressed);
Ok(plaintext)
}
#[cfg(test)]
mod tests {
#![allow(clippy::expect_used)]
#![allow(clippy::panic)]
use super::*;
const KEY: [u8; 32] = [0x2Bu8; 32];
const NONCE: [u8; 24] = [0x7Fu8; 24];
fn plaintext() -> Vec<u8> {
b"the same sentence, over and over. ".repeat(32)
}
#[test]
fn the_cipher_round_trips_its_own_output() {
let cipher = XChaCha20Poly1305Cipher::new();
let message = b"a message";
let sealed = cipher
.encrypt(&KEY, &NONCE, message, b"aad")
.expect("encryption must succeed");
assert_eq!(sealed.len(), message.len() + 16);
let opened = cipher
.decrypt(&KEY, &NONCE, &sealed, b"aad")
.expect("decryption must succeed");
assert_eq!(opened.as_slice(), message.as_slice());
}
#[test]
fn every_way_of_being_wrong_looks_the_same() {
let cipher = XChaCha20Poly1305Cipher::new();
let sealed = cipher
.encrypt(&KEY, &NONCE, b"a message", STENOXIDE_AAD)
.expect("encryption must succeed");
let mut damaged = sealed.to_vec();
damaged[0] ^= 0x40;
let attempts = [
cipher.decrypt(&[0u8; 32], &NONCE, &sealed, STENOXIDE_AAD),
cipher.decrypt(&KEY, &[0u8; 24], &sealed, STENOXIDE_AAD),
cipher.decrypt(&KEY, &NONCE, &sealed, b"other-construction"),
cipher.decrypt(&KEY, &NONCE, &damaged, STENOXIDE_AAD),
cipher.decrypt(&KEY, &NONCE, &sealed[..4], STENOXIDE_AAD),
];
for attempt in attempts {
match attempt.map(|_| ()) {
Err(AEADError::AuthenticationFailed) => {}
Err(other) => panic!("expected an authentication failure, got: {other:?}"),
Ok(()) => panic!("a wrong input must not authenticate"),
}
}
}
#[test]
fn the_payload_is_compressed_before_it_is_encrypted() {
let cipher = XChaCha20Poly1305Cipher::new();
let plaintext = plaintext();
let ciphertext = compress_and_encrypt(&plaintext, &KEY, &NONCE, &cipher)
.expect("compression and encryption must succeed");
assert!(
ciphertext.len() < plaintext.len(),
"a repetitive payload must shrink: {} against {}",
ciphertext.len(),
plaintext.len()
);
let recovered = decrypt_and_decompress(&ciphertext, &KEY, &NONCE, &cipher)
.expect("decryption and decompression must succeed");
assert_eq!(recovered.as_slice(), plaintext.as_slice());
}
#[test]
fn authentication_runs_before_decompression() {
let cipher = XChaCha20Poly1305Cipher::new();
let ciphertext = compress_and_encrypt(&plaintext(), &KEY, &NONCE, &cipher)
.expect("compression and encryption must succeed");
let error = decrypt_and_decompress(&ciphertext, &[9u8; 32], &NONCE, &cipher)
.map(|_| ())
.expect_err("a wrong key must not authenticate");
assert!(
matches!(
error,
CryptoError::AEADError(AEADError::AuthenticationFailed)
),
"got: {error:?}"
);
}
#[test]
fn a_verified_payload_that_will_not_decompress_is_a_decompression_failure() {
let cipher = XChaCha20Poly1305Cipher::new();
let sealed = cipher
.encrypt(&KEY, &NONCE, b"not a zstandard frame", STENOXIDE_AAD)
.expect("encryption must succeed");
let error = decrypt_and_decompress(&sealed, &KEY, &NONCE, &cipher)
.map(|_| ())
.expect_err("authenticated nonsense must not decompress");
assert!(
matches!(error, CryptoError::DecompressionError(_)),
"got: {error:?}"
);
}
fn bomb(plaintext_bytes: u64) -> Vec<u8> {
use std::io::Write;
const CHUNK: usize = 64 * 1024;
let zeros = [0u8; CHUNK];
let mut encoder =
zstd::stream::write::Encoder::new(Vec::new(), 1).expect("the encoder must start");
let mut written = 0u64;
while written < plaintext_bytes {
let step = CHUNK.min((plaintext_bytes - written) as usize);
encoder.write_all(&zeros[..step]).expect("the sink is memory");
written += step as u64;
}
encoder.finish().expect("the frame must close")
}
#[test]
fn a_payload_that_ends_at_the_ceiling_is_returned() {
const CEILING: u64 = 64 * 1024;
let frame = bomb(CEILING);
let plaintext = decompress_within(&frame, CEILING)
.expect("a payload that ends at the ceiling must decompress");
assert_eq!(plaintext.len() as u64, CEILING);
}
#[test]
fn a_payload_that_expands_past_the_ceiling_is_refused() {
const CEILING: u64 = 64 * 1024;
let frame = bomb(CEILING * 1_000);
assert!(
(frame.len() as u64) < CEILING / 10,
"the bomb must be far smaller than the ceiling: {} bytes",
frame.len()
);
let error = decompress_within(&frame, CEILING)
.map(|_| ())
.expect_err("a frame past the ceiling must be refused");
match error {
CryptoError::DecompressionError(message) => {
assert!(message.contains("ceiling"), "got: {message}");
}
other => panic!("expected a decompression failure, got: {other:?}"),
}
}
#[test]
fn a_bomb_and_a_damaged_payload_are_the_same_kind_of_failure() {
const CEILING: u64 = 64 * 1024;
let from_bomb = decompress_within(&bomb(CEILING * 1_000), CEILING).map(|_| ());
let from_garbage = decompress_within(b"not a zstandard frame", CEILING).map(|_| ());
for outcome in [from_bomb, from_garbage] {
assert!(
matches!(outcome, Err(CryptoError::DecompressionError(_))),
"got: {outcome:?}"
);
}
}
#[test]
fn the_ceiling_clears_the_largest_container_by_an_order_of_magnitude() {
let largest_ciphertext = crate::image_io::validate::MAX_PIXELS * 3 / 8;
assert!(
MAX_DECOMPRESSED_BYTES > largest_ciphertext * 10,
"{MAX_DECOMPRESSED_BYTES} against {largest_ciphertext}"
);
}
#[test]
fn an_ordinary_payload_is_untouched_by_the_ceiling() {
let cipher = XChaCha20Poly1305Cipher::new();
let plaintext = plaintext();
let ciphertext = compress_and_encrypt(&plaintext, &KEY, &NONCE, &cipher)
.expect("compression and encryption must succeed");
let recovered = decrypt_and_decompress(&ciphertext, &KEY, &NONCE, &cipher)
.expect("an ordinary payload must survive the ceiling");
assert_eq!(recovered.as_slice(), plaintext.as_slice());
}
#[test]
fn every_failure_explains_itself() {
assert!(AEADError::AuthenticationFailed
.to_string()
.contains("corrupted"));
assert!(AEADError::CipherError("no key".to_owned())
.to_string()
.contains("no key"));
assert!(CryptoError::CompressionError("level".to_owned())
.to_string()
.contains("level"));
assert!(CryptoError::DecompressionError("truncated".to_owned())
.to_string()
.contains("truncated"));
let wrapped = CryptoError::from(AEADError::AuthenticationFailed);
assert_eq!(
wrapped.to_string(),
AEADError::AuthenticationFailed.to_string()
);
assert!(std::error::Error::source(&wrapped).is_some());
assert!(
std::error::Error::source(&CryptoError::CompressionError("x".to_owned())).is_none()
);
}
}