use crate::constants;
use crate::error::{Error, Result};
use crate::primitives::{aead, random};
use std::collections::HashMap;
use std::io::Read;
use subtle::ConstantTimeEq;
use zeroize::{Zeroize, ZeroizeOnDrop, Zeroizing};
const FLAG_COMPRESSED: u8 = 0x01;
const COMPRESSION_LEVEL: ruzstd::encoding::CompressionLevel =
ruzstd::encoding::CompressionLevel::Fastest;
const MIN_BLOB_LEN: usize =
1 + 1 + crate::constants::AEAD_NONCE_SIZE + crate::constants::AEAD_TAG_SIZE;
#[cfg(target_arch = "wasm32")]
const MAX_DECOMPRESSED_SIZE: u64 = 16 * 1024 * 1024;
#[cfg(not(target_arch = "wasm32"))]
const MAX_DECOMPRESSED_SIZE: u64 = 256 * 1024 * 1024;
#[cfg(target_arch = "wasm32")]
const MAX_DECOMPRESSED_SIZE_USIZE: usize = 16 * 1024 * 1024;
#[cfg(not(target_arch = "wasm32"))]
const MAX_DECOMPRESSED_SIZE_USIZE: usize = 256 * 1024 * 1024;
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct StorageKey {
#[zeroize(skip)]
version: u8,
key: [u8; 32],
}
impl StorageKey {
pub fn new(version: u8, mut key: [u8; 32]) -> Result<Self> {
if version == 0 {
key.zeroize();
return Err(Error::UnsupportedVersion);
}
if bool::from(key.ct_eq(&[0u8; 32])) {
key.zeroize();
return Err(Error::InvalidData);
}
let result = Self { version, key };
key.zeroize();
Ok(result)
}
pub fn version(&self) -> u8 {
self.version
}
pub fn key(&self) -> &[u8; 32] {
&self.key
}
}
pub struct StorageKeyRing {
active_version: u8,
keys: HashMap<u8, StorageKey>,
}
impl StorageKeyRing {
pub fn new(key: StorageKey) -> Result<Self> {
let version = key.version();
Ok(Self {
active_version: version,
keys: HashMap::from([(version, key)]),
})
}
pub fn add_key(&mut self, key: StorageKey, make_active: bool) -> Result<bool> {
let version = key.version();
if version == self.active_version && !make_active {
return Err(Error::InvalidData);
}
if make_active {
self.active_version = version;
}
let replaced = self.keys.insert(version, key).is_some();
Ok(replaced)
}
pub fn active_key(&self) -> Option<&StorageKey> {
self.keys.get(&self.active_version)
}
pub fn get_key(&self, version: u8) -> Option<&StorageKey> {
self.keys.get(&version)
}
pub fn remove_key(&mut self, version: u8) -> Result<bool> {
if version == 0 {
return Err(Error::UnsupportedVersion);
}
if version == self.active_version {
return Err(Error::InvalidData);
}
Ok(self.keys.remove(&version).is_some())
}
}
impl Drop for StorageKeyRing {
fn drop(&mut self) {
for key in self.keys.values_mut() {
key.key.zeroize();
}
}
}
#[must_use = "encrypted blob must be stored; discarding it loses the data"]
pub fn encrypt_blob(
key: &StorageKey,
plaintext: &[u8],
channel_id: &str,
segment_id: &str,
compress: bool,
) -> Result<Vec<u8>> {
if plaintext.len() as u64 > MAX_DECOMPRESSED_SIZE {
return Err(Error::InvalidData);
}
let (compressed, flags) = if compress {
let c = ruzstd::encoding::compress_to_vec(plaintext, COMPRESSION_LEVEL);
(Some(Zeroizing::new(c)), FLAG_COMPRESSED)
} else {
(None, 0u8)
};
let data: &[u8] = match &compressed {
Some(c) => c,
None => plaintext,
};
let mut nonce = [0u8; 24];
random::random_bytes(&mut nonce);
let aad = build_storage_aad(key.version(), flags, channel_id, segment_id)?;
let ciphertext = aead::aead_encrypt(key.key(), &nonce, data, &aad)?;
let mut blob = Vec::with_capacity(1 + 1 + 24 + ciphertext.len());
blob.push(key.version());
blob.push(flags);
blob.extend_from_slice(&nonce);
blob.extend_from_slice(&ciphertext);
Ok(blob)
}
#[must_use = "decrypted data contains sensitive plaintext that must be consumed or zeroized"]
pub fn decrypt_blob(
keyring: &StorageKeyRing,
blob: &[u8],
channel_id: &str,
segment_id: &str,
) -> Result<Zeroizing<Vec<u8>>> {
if blob.len() < MIN_BLOB_LEN {
return Err(Error::AeadFailed);
}
let version = blob[0];
let flags = blob[1];
let known_flags = FLAG_COMPRESSED;
if flags & !known_flags != 0 {
return Err(Error::AeadFailed);
}
let nonce: &[u8; 24] = blob[2..26].try_into().map_err(|_| Error::AeadFailed)?;
let ciphertext = &blob[26..];
let key = keyring.get_key(version).ok_or(Error::AeadFailed)?;
let aad =
build_storage_aad(version, flags, channel_id, segment_id).map_err(|_| Error::AeadFailed)?;
let data = aead::aead_decrypt(key.key(), nonce, ciphertext, &aad)?;
if flags & FLAG_COMPRESSED != 0 {
if data.is_empty() {
return Ok(Zeroizing::new(Vec::new()));
}
let decoder = ruzstd::decoding::StreamingDecoder::new(data.as_slice())
.map_err(|_| Error::AeadFailed)?;
let mut limited = decoder.take(MAX_DECOMPRESSED_SIZE + 1);
let hint = data
.len()
.saturating_mul(4)
.min(MAX_DECOMPRESSED_SIZE_USIZE);
let mut decompressed = Zeroizing::new(Vec::with_capacity(hint));
limited
.read_to_end(&mut decompressed)
.map_err(|_| Error::AeadFailed)?;
if decompressed.len() as u64 > MAX_DECOMPRESSED_SIZE {
return Err(Error::AeadFailed);
}
Ok(decompressed)
} else {
Ok(data)
}
}
fn build_storage_aad(
version: u8,
flags: u8,
channel_id: &str,
segment_id: &str,
) -> Result<Vec<u8>> {
let ch = channel_id.as_bytes();
let seg = segment_id.as_bytes();
if ch.len() > u16::MAX as usize || seg.len() > u16::MAX as usize {
return Err(Error::InvalidData);
}
let mut aad =
Vec::with_capacity(constants::STORAGE_AAD.len() + 1 + 1 + 2 + ch.len() + 2 + seg.len());
aad.extend_from_slice(constants::STORAGE_AAD);
aad.push(version);
aad.push(flags);
let ch_len = u16::try_from(ch.len()).expect("ch.len() ≤ u16::MAX validated above");
let seg_len = u16::try_from(seg.len()).expect("seg.len() ≤ u16::MAX validated above");
aad.extend_from_slice(&ch_len.to_be_bytes());
aad.extend_from_slice(ch);
aad.extend_from_slice(&seg_len.to_be_bytes());
aad.extend_from_slice(seg);
Ok(aad)
}
fn build_dm_queue_aad(
version: u8,
flags: u8,
recipient_fp: &[u8; 32],
batch_id: &str,
) -> Result<Vec<u8>> {
let bid = batch_id.as_bytes();
if bid.len() > u16::MAX as usize {
return Err(Error::InvalidData);
}
let mut aad =
Vec::with_capacity(constants::DM_QUEUE_AAD.len() + 1 + 1 + 2 + 32 + 2 + bid.len());
aad.extend_from_slice(constants::DM_QUEUE_AAD);
aad.push(version);
aad.push(flags);
aad.extend_from_slice(&32u16.to_be_bytes());
aad.extend_from_slice(recipient_fp);
let bid_len = u16::try_from(bid.len()).expect("bid.len() ≤ u16::MAX validated above");
aad.extend_from_slice(&bid_len.to_be_bytes());
aad.extend_from_slice(bid);
Ok(aad)
}
#[must_use = "encrypted blob must be stored; discarding it loses the data"]
pub fn encrypt_dm_queue_blob(
key: &StorageKey,
plaintext: &[u8],
recipient_fp: &[u8; 32],
batch_id: &str,
compress: bool,
) -> Result<Vec<u8>> {
if plaintext.len() as u64 > MAX_DECOMPRESSED_SIZE {
return Err(Error::InvalidData);
}
let (compressed, flags) = if compress {
let c = ruzstd::encoding::compress_to_vec(plaintext, COMPRESSION_LEVEL);
(Some(Zeroizing::new(c)), FLAG_COMPRESSED)
} else {
(None, 0u8)
};
let data: &[u8] = match &compressed {
Some(c) => c,
None => plaintext,
};
let mut nonce = [0u8; 24];
random::random_bytes(&mut nonce);
let aad = build_dm_queue_aad(key.version(), flags, recipient_fp, batch_id)?;
let ciphertext = aead::aead_encrypt(key.key(), &nonce, data, &aad)?;
let mut blob = Vec::with_capacity(1 + 1 + 24 + ciphertext.len());
blob.push(key.version());
blob.push(flags);
blob.extend_from_slice(&nonce);
blob.extend_from_slice(&ciphertext);
Ok(blob)
}
#[must_use = "decrypted data contains sensitive plaintext that must be consumed or zeroized"]
pub fn decrypt_dm_queue_blob(
keyring: &StorageKeyRing,
blob: &[u8],
recipient_fp: &[u8; 32],
batch_id: &str,
) -> Result<Zeroizing<Vec<u8>>> {
if blob.len() < MIN_BLOB_LEN {
return Err(Error::AeadFailed);
}
let version = blob[0];
let flags = blob[1];
let known_flags = FLAG_COMPRESSED;
if flags & !known_flags != 0 {
return Err(Error::AeadFailed);
}
let nonce: &[u8; 24] = blob[2..26].try_into().map_err(|_| Error::AeadFailed)?;
let ciphertext = &blob[26..];
let key = keyring.get_key(version).ok_or(Error::AeadFailed)?;
let aad = build_dm_queue_aad(version, flags, recipient_fp, batch_id)
.map_err(|_| Error::AeadFailed)?;
let data = aead::aead_decrypt(key.key(), nonce, ciphertext, &aad)?;
if flags & FLAG_COMPRESSED != 0 {
if data.is_empty() {
return Ok(Zeroizing::new(Vec::new()));
}
let decoder = ruzstd::decoding::StreamingDecoder::new(data.as_slice())
.map_err(|_| Error::AeadFailed)?;
let mut limited = decoder.take(MAX_DECOMPRESSED_SIZE + 1);
let hint = data
.len()
.saturating_mul(4)
.min(MAX_DECOMPRESSED_SIZE_USIZE);
let mut decompressed = Zeroizing::new(Vec::with_capacity(hint));
limited
.read_to_end(&mut decompressed)
.map_err(|_| Error::AeadFailed)?;
if decompressed.len() as u64 > MAX_DECOMPRESSED_SIZE {
return Err(Error::AeadFailed);
}
Ok(decompressed)
} else {
Ok(data)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
fn test_key(version: u8) -> StorageKey {
StorageKey::new(version, crate::primitives::random::random_array()).unwrap()
}
fn test_keyring(version: u8) -> (StorageKey, StorageKeyRing) {
let key = test_key(version);
let key_copy = key.clone();
let ring = StorageKeyRing::new(key).unwrap();
(key_copy, ring)
}
#[test]
fn encrypt_decrypt_uncompressed() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_blob(key, b"hello world", "chan", "seg", false).unwrap();
let pt = decrypt_blob(&ring, &blob, "chan", "seg").unwrap();
assert_eq!(&*pt, b"hello world");
}
#[test]
fn encrypt_decrypt_compressed() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_blob(key, b"hello world", "chan", "seg", true).unwrap();
let pt = decrypt_blob(&ring, &blob, "chan", "seg").unwrap();
assert_eq!(&*pt, b"hello world");
}
#[test]
fn wrong_channel_id() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_blob(key, b"data", "chan_a", "seg", false).unwrap();
assert!(matches!(
decrypt_blob(&ring, &blob, "chan_b", "seg"),
Err(Error::AeadFailed)
));
}
#[test]
fn wrong_segment_id() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_blob(key, b"data", "chan", "seg_1", false).unwrap();
assert!(matches!(
decrypt_blob(&ring, &blob, "chan", "seg_2"),
Err(Error::AeadFailed)
));
}
#[test]
fn tampered_blob() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let mut blob = encrypt_blob(key, b"data", "chan", "seg", false).unwrap();
blob[14] ^= 0xFF;
assert!(matches!(
decrypt_blob(&ring, &blob, "chan", "seg"),
Err(Error::AeadFailed)
));
}
#[test]
fn truncated_blob() {
let (_, ring) = test_keyring(1);
let short = vec![0u8; MIN_BLOB_LEN - 1];
assert!(matches!(
decrypt_blob(&ring, &short, "chan", "seg"),
Err(Error::AeadFailed)
));
}
#[test]
fn unknown_flags() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let mut blob = encrypt_blob(key, b"data", "chan", "seg", false).unwrap();
blob[1] |= 0x02;
assert!(matches!(
decrypt_blob(&ring, &blob, "chan", "seg"),
Err(Error::AeadFailed)
));
blob[1] |= 0x01;
assert!(matches!(
decrypt_blob(&ring, &blob, "chan", "seg"),
Err(Error::AeadFailed)
));
}
#[test]
fn storage_key_version_0_rejected() {
assert!(matches!(
StorageKey::new(0, [0u8; 32]),
Err(Error::UnsupportedVersion)
));
}
#[test]
fn storage_key_zero_key_rejected() {
assert!(matches!(
StorageKey::new(1, [0u8; 32]),
Err(Error::InvalidData)
));
}
#[test]
fn flag_tampering_detected_by_aead() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let mut blob = encrypt_blob(key, b"not zstd data", "chan", "seg", false).unwrap();
blob[1] |= FLAG_COMPRESSED;
assert!(matches!(
decrypt_blob(&ring, &blob, "chan", "seg"),
Err(Error::AeadFailed)
));
}
#[test]
fn key_rotation() {
let (_, mut ring) = test_keyring(1);
let key1 = ring.active_key().unwrap();
let blob_v1 = encrypt_blob(key1, b"v1 data", "chan", "seg", false).unwrap();
ring.add_key(test_key(2), true).unwrap();
let pt = decrypt_blob(&ring, &blob_v1, "chan", "seg").unwrap();
assert_eq!(&*pt, b"v1 data");
let key2 = ring.active_key().unwrap();
assert_eq!(key2.version(), 2);
let blob_v2 = encrypt_blob(key2, b"v2 data", "chan", "seg", false).unwrap();
let pt2 = decrypt_blob(&ring, &blob_v2, "chan", "seg").unwrap();
assert_eq!(&*pt2, b"v2 data");
}
#[test]
fn key_not_found() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let mut blob = encrypt_blob(key, b"data", "chan", "seg", false).unwrap();
blob[0] = 99;
assert!(matches!(
decrypt_blob(&ring, &blob, "chan", "seg"),
Err(Error::AeadFailed)
));
}
#[test]
fn keyring_add_replace_active_without_flag_rejected() {
let (_, mut ring) = test_keyring(1);
assert!(matches!(
ring.add_key(test_key(1), false),
Err(Error::InvalidData)
));
}
#[test]
fn keyring_add_replace_active_with_flag_succeeds() {
let (_, mut ring) = test_keyring(1);
let replaced = ring.add_key(test_key(1), true).unwrap();
assert!(replaced);
}
#[test]
fn keyring_remove_active_rejected() {
let (_, mut ring) = test_keyring(1);
assert!(matches!(ring.remove_key(1), Err(Error::InvalidData)));
}
#[test]
fn keyring_remove_nonexistent() {
let (_, mut ring) = test_keyring(1);
let removed = ring.remove_key(99).unwrap();
assert!(!removed);
}
#[test]
fn keyring_remove_version_0_rejected() {
let (_, mut ring) = test_keyring(1);
assert!(matches!(ring.remove_key(0), Err(Error::UnsupportedVersion)));
}
#[test]
fn keyring_active_key_returns_correct_version() {
let (_, ring) = test_keyring(5);
let active = ring.active_key().unwrap();
assert_eq!(active.version(), 5);
}
#[test]
fn keyring_get_key_returns_none_for_absent() {
let (_, ring) = test_keyring(1);
assert!(ring.get_key(99).is_none());
}
#[test]
fn empty_plaintext_both_modes() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
for compress in [false, true] {
let blob = encrypt_blob(key, b"", "chan", "seg", compress).unwrap();
let pt = decrypt_blob(&ring, &blob, "chan", "seg").unwrap();
assert!(pt.is_empty());
}
}
fn make_blob_with_plaintext(
key: &StorageKey,
flags: u8,
plaintext: &[u8],
channel_id: &str,
segment_id: &str,
) -> Vec<u8> {
let nonce = [0x42u8; 24];
let aad = build_storage_aad(key.version(), flags, channel_id, segment_id).unwrap();
let ct = crate::primitives::aead::aead_encrypt(key.key(), &nonce, plaintext, &aad).unwrap();
let mut blob = Vec::with_capacity(1 + 1 + 24 + ct.len());
blob.push(key.version());
blob.push(flags);
blob.extend_from_slice(&nonce);
blob.extend_from_slice(&ct);
blob
}
#[test]
fn invalid_compressed_data() {
let (key, ring) = test_keyring(1);
let garbage = b"this is definitely not valid zstd";
let blob = make_blob_with_plaintext(&key, FLAG_COMPRESSED, garbage, "chan", "seg");
assert!(matches!(
decrypt_blob(&ring, &blob, "chan", "seg"),
Err(Error::AeadFailed)
));
}
#[test]
#[ignore = "allocates 256 MiB — run explicitly to verify zip-bomb limit"]
fn decompression_bomb_rejected() {
let plaintext = vec![0u8; MAX_DECOMPRESSED_SIZE_USIZE + 1];
let compressed = ruzstd::encoding::compress_to_vec(plaintext.as_slice(), COMPRESSION_LEVEL);
let (key, ring) = test_keyring(1);
let blob = make_blob_with_plaintext(&key, FLAG_COMPRESSED, &compressed, "bomb", "0");
assert!(matches!(
decrypt_blob(&ring, &blob, "bomb", "0"),
Err(Error::AeadFailed)
));
}
#[test]
fn storage_aad_structure() {
let aad = build_storage_aad(1, FLAG_COMPRESSED, "my-channel", "seg-42").unwrap();
let mut expected = Vec::new();
expected.extend_from_slice(b"lo-storage-v1");
expected.push(1); expected.push(FLAG_COMPRESSED); expected.extend_from_slice(&10u16.to_be_bytes()); expected.extend_from_slice(b"my-channel");
expected.extend_from_slice(&6u16.to_be_bytes()); expected.extend_from_slice(b"seg-42");
assert_eq!(aad, expected);
}
#[test]
fn storage_aad_empty_channel_and_segment_are_distinct() {
let aad1 = build_storage_aad(1, 0, "", "x").unwrap();
let aad2 = build_storage_aad(1, 0, "x", "").unwrap();
assert_ne!(aad1, aad2, "empty channel vs empty segment must differ");
let aad = build_storage_aad(1, 0, "", "").unwrap();
let prefix_len = b"lo-storage-v1".len() + 1 + 1; assert_eq!(aad.len(), prefix_len + 2 + 2);
}
#[test]
fn build_storage_aad_rejects_oversized_ids() {
let long = "x".repeat(u16::MAX as usize + 1);
assert!(matches!(
build_storage_aad(1, 0, &long, "seg"),
Err(Error::InvalidData)
));
assert!(matches!(
build_storage_aad(1, 0, "ch", &long),
Err(Error::InvalidData)
));
}
#[test]
fn storage_encrypt_decrypt_empty_strings() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_blob(key, b"hello", "", "", false).unwrap();
let decrypted = decrypt_blob(&ring, &blob, "", "").unwrap();
assert_eq!(&*decrypted, b"hello");
}
const TEST_FP: [u8; 32] = [0xABu8; 32];
#[test]
fn dm_queue_encrypt_decrypt_uncompressed() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_dm_queue_blob(key, b"dm payload", &TEST_FP, "batch-0", false).unwrap();
let pt = decrypt_dm_queue_blob(&ring, &blob, &TEST_FP, "batch-0").unwrap();
assert_eq!(&*pt, b"dm payload");
}
#[test]
fn dm_queue_encrypt_decrypt_compressed() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_dm_queue_blob(key, b"dm payload", &TEST_FP, "batch-0", true).unwrap();
let pt = decrypt_dm_queue_blob(&ring, &blob, &TEST_FP, "batch-0").unwrap();
assert_eq!(&*pt, b"dm payload");
}
#[test]
fn dm_queue_wrong_recipient_fp() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_dm_queue_blob(key, b"data", &TEST_FP, "batch-0", false).unwrap();
let wrong_fp = [0xCDu8; 32];
assert!(matches!(
decrypt_dm_queue_blob(&ring, &blob, &wrong_fp, "batch-0"),
Err(Error::AeadFailed)
));
}
#[test]
fn dm_queue_wrong_batch_id() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob = encrypt_dm_queue_blob(key, b"data", &TEST_FP, "batch-0", false).unwrap();
assert!(matches!(
decrypt_dm_queue_blob(&ring, &blob, &TEST_FP, "batch-1"),
Err(Error::AeadFailed)
));
}
#[test]
fn dm_queue_tampered_blob() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let mut blob = encrypt_dm_queue_blob(key, b"data", &TEST_FP, "batch-0", false).unwrap();
blob[14] ^= 0xFF;
assert!(matches!(
decrypt_dm_queue_blob(&ring, &blob, &TEST_FP, "batch-0"),
Err(Error::AeadFailed)
));
}
#[test]
fn dm_queue_empty_plaintext() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
for compress in [false, true] {
let blob = encrypt_dm_queue_blob(key, b"", &TEST_FP, "batch-0", compress).unwrap();
let pt = decrypt_dm_queue_blob(&ring, &blob, &TEST_FP, "batch-0").unwrap();
assert!(pt.is_empty());
}
}
#[test]
fn dm_queue_key_rotation() {
let (_, mut ring) = test_keyring(1);
let key1 = ring.active_key().unwrap();
let blob_v1 = encrypt_dm_queue_blob(key1, b"v1 dm", &TEST_FP, "b0", false).unwrap();
ring.add_key(test_key(2), true).unwrap();
let pt = decrypt_dm_queue_blob(&ring, &blob_v1, &TEST_FP, "b0").unwrap();
assert_eq!(&*pt, b"v1 dm");
let key2 = ring.active_key().unwrap();
let blob_v2 = encrypt_dm_queue_blob(key2, b"v2 dm", &TEST_FP, "b0", false).unwrap();
let pt2 = decrypt_dm_queue_blob(&ring, &blob_v2, &TEST_FP, "b0").unwrap();
assert_eq!(&*pt2, b"v2 dm");
}
#[test]
fn dm_queue_aad_structure() {
let fp = [0x42u8; 32];
let aad = build_dm_queue_aad(1, FLAG_COMPRESSED, &fp, "batch-7").unwrap();
let mut expected = Vec::new();
expected.extend_from_slice(b"lo-dm-queue-v1");
expected.push(1); expected.push(FLAG_COMPRESSED); expected.extend_from_slice(&32u16.to_be_bytes()); expected.extend_from_slice(&fp);
expected.extend_from_slice(&7u16.to_be_bytes()); expected.extend_from_slice(b"batch-7");
assert_eq!(aad, expected);
}
#[test]
fn dm_queue_aad_rejects_oversized_batch_id() {
let long = "x".repeat(u16::MAX as usize + 1);
let fp = [0x00u8; 32];
assert!(matches!(
build_dm_queue_aad(1, 0, &fp, &long),
Err(Error::InvalidData)
));
}
#[test]
fn community_and_dm_queue_blobs_not_interchangeable() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let community_blob = encrypt_blob(key, b"data", "chan", "seg", false).unwrap();
assert!(matches!(
decrypt_dm_queue_blob(&ring, &community_blob, &TEST_FP, "seg"),
Err(Error::AeadFailed)
));
let dm_blob = encrypt_dm_queue_blob(key, b"data", &TEST_FP, "batch-0", false).unwrap();
assert!(matches!(
decrypt_blob(&ring, &dm_blob, "chan", "batch-0"),
Err(Error::AeadFailed)
));
}
#[test]
fn encrypt_blob_nonce_freshness() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let blob1 = encrypt_blob(key, b"same input", "chan", "seg", false).unwrap();
let blob2 = encrypt_blob(key, b"same input", "chan", "seg", false).unwrap();
assert_ne!(
blob1, blob2,
"encrypt_blob must produce different blobs due to random nonce"
);
}
#[test]
fn encrypt_dm_queue_blob_nonce_freshness() {
let (_, ring) = test_keyring(1);
let key = ring.active_key().unwrap();
let fp = [0xABu8; 32];
let blob1 = encrypt_dm_queue_blob(key, b"same input", &fp, "batch", false).unwrap();
let blob2 = encrypt_dm_queue_blob(key, b"same input", &fp, "batch", false).unwrap();
assert_ne!(
blob1, blob2,
"encrypt_dm_queue_blob must produce different blobs due to random nonce"
);
}
mod proptests {
use super::*;
use proptest::prelude::*;
proptest! {
#[test]
fn encrypt_decrypt_roundtrip(
plaintext in proptest::collection::vec(any::<u8>(), 0..4096),
compress in any::<bool>(),
channel in "[a-z]{1,16}",
segment in "[a-z]{1,16}",
version in 1u8..=255u8,
) {
let (_, ring) = test_keyring(version);
let key = ring.active_key().unwrap();
let blob = encrypt_blob(key, &plaintext, &channel, &segment, compress).unwrap();
let decrypted = decrypt_blob(&ring, &blob, &channel, &segment).unwrap();
prop_assert_eq!(&*decrypted, &plaintext);
}
#[test]
fn dm_queue_encrypt_decrypt_roundtrip(
plaintext in proptest::collection::vec(any::<u8>(), 0..4096),
compress in any::<bool>(),
batch_id in "[a-z0-9]{1,16}",
version in 1u8..=255u8,
) {
let fp = [0xABu8; 32];
let (_, ring) = test_keyring(version);
let key = ring.active_key().unwrap();
let blob = encrypt_dm_queue_blob(key, &plaintext, &fp, &batch_id, compress).unwrap();
let decrypted = decrypt_dm_queue_blob(&ring, &blob, &fp, &batch_id).unwrap();
prop_assert_eq!(&*decrypted, &plaintext);
}
}
}
}