use crate::constants::{MAX_DECOMPRESSED_BYTES, MIN_COMPRESS_BYTES, ZSTD_LEVEL};
use crate::error::{MessageError, Result};
pub const COMPRESSION_NONE: u8 = 0;
pub const COMPRESSION_ZSTD: u8 = 1;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CompressedPayload {
pub compression: u8,
pub bytes: Vec<u8>,
pub uncompressed_len: u32,
}
pub fn compress_payload(payload: &[u8]) -> Result<CompressedPayload> {
let uncompressed_len =
u32::try_from(payload.len()).map_err(|_| MessageError::PayloadTooLarge(payload.len()))?;
if payload.len() < MIN_COMPRESS_BYTES {
return Ok(raw(payload, uncompressed_len));
}
let compressed = zstd::bulk::compress(payload, ZSTD_LEVEL)
.map_err(|e| MessageError::Codec(e.to_string()))?;
if compressed.len() >= payload.len() {
return Ok(raw(payload, uncompressed_len));
}
Ok(CompressedPayload {
compression: COMPRESSION_ZSTD,
bytes: compressed,
uncompressed_len,
})
}
pub fn decompress_payload(compression: u8, data: &[u8], uncompressed_len: u32) -> Result<Vec<u8>> {
let declared = uncompressed_len as usize;
if declared > MAX_DECOMPRESSED_BYTES {
return Err(MessageError::DecompressionBomb {
declared,
max: MAX_DECOMPRESSED_BYTES,
});
}
match compression {
COMPRESSION_NONE => {
if data.len() != declared {
return Err(MessageError::DecompressedLengthMismatch {
expected: declared,
actual: data.len(),
});
}
Ok(data.to_vec())
}
COMPRESSION_ZSTD => {
let decoded = zstd::bulk::decompress(data, declared)
.map_err(|e| MessageError::Codec(e.to_string()))?;
if decoded.len() != declared {
return Err(MessageError::DecompressedLengthMismatch {
expected: declared,
actual: decoded.len(),
});
}
Ok(decoded)
}
other => Err(MessageError::UnsupportedCompression(other)),
}
}
fn raw(payload: &[u8], uncompressed_len: u32) -> CompressedPayload {
CompressedPayload {
compression: COMPRESSION_NONE,
bytes: payload.to_vec(),
uncompressed_len,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn compressible(len: usize) -> Vec<u8> {
(0..len).map(|i| (i % 7) as u8).collect()
}
#[test]
fn small_payload_stays_raw() {
let payload = compressible(MIN_COMPRESS_BYTES - 1);
let out = compress_payload(&payload).unwrap();
assert_eq!(out.compression, COMPRESSION_NONE);
assert_eq!(out.bytes, payload);
assert_eq!(out.uncompressed_len as usize, payload.len());
}
#[test]
fn large_compressible_payload_uses_zstd_and_round_trips() {
let payload = compressible(4096);
let out = compress_payload(&payload).unwrap();
assert_eq!(out.compression, COMPRESSION_ZSTD);
assert!(
out.bytes.len() < payload.len(),
"zstd should shrink a low-entropy payload"
);
let restored =
decompress_payload(out.compression, &out.bytes, out.uncompressed_len).unwrap();
assert_eq!(restored, payload);
}
#[test]
fn raw_round_trips() {
let payload = compressible(16);
let out = compress_payload(&payload).unwrap();
let restored =
decompress_payload(out.compression, &out.bytes, out.uncompressed_len).unwrap();
assert_eq!(restored, payload);
}
#[test]
fn incompressible_payload_falls_back_to_raw() {
use sha2::{Digest, Sha256};
let mut payload = Vec::new();
let mut ctr = 0u64;
while payload.len() < 1024 {
let mut h = Sha256::new();
h.update(b"incompressible");
h.update(ctr.to_le_bytes());
payload.extend_from_slice(&h.finalize());
ctr += 1;
}
let out = compress_payload(&payload).unwrap();
assert_eq!(out.compression, COMPRESSION_NONE);
assert_eq!(out.bytes, payload);
}
#[test]
fn compression_is_deterministic() {
let payload = compressible(2048);
assert_eq!(
compress_payload(&payload).unwrap(),
compress_payload(&payload).unwrap()
);
}
#[test]
fn unknown_compression_id_is_rejected() {
let err = decompress_payload(2, &[1, 2, 3], 3).unwrap_err();
assert_eq!(err, MessageError::UnsupportedCompression(2));
}
#[test]
fn declared_length_over_cap_is_a_bomb() {
let declared = (MAX_DECOMPRESSED_BYTES + 1) as u32;
let err = decompress_payload(COMPRESSION_ZSTD, &[], declared).unwrap_err();
assert_eq!(
err,
MessageError::DecompressionBomb {
declared: declared as usize,
max: MAX_DECOMPRESSED_BYTES
}
);
}
#[test]
fn zstd_bomb_that_underdeclares_output_is_rejected() {
let payload = compressible(8192);
let out = compress_payload(&payload).unwrap();
assert_eq!(out.compression, COMPRESSION_ZSTD);
let err = decompress_payload(out.compression, &out.bytes, 16).unwrap_err();
assert!(matches!(
err,
MessageError::Codec(_) | MessageError::DecompressedLengthMismatch { .. }
));
}
#[test]
fn raw_length_mismatch_is_rejected() {
let err = decompress_payload(COMPRESSION_NONE, &[1, 2, 3], 4).unwrap_err();
assert_eq!(
err,
MessageError::DecompressedLengthMismatch {
expected: 4,
actual: 3
}
);
}
}