use crate::{Error, Result, MAGIC_COMPRESSED};
use alloc::vec::Vec;
use serde::{Deserialize, Serialize};
pub const MAX_DECOMPRESSED_SIZE: usize = 512 * 1024 * 1024;
pub const MAX_COMPRESSION_RATIO: usize = 256;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[repr(u8)]
pub enum CompressionAlgorithm {
None = 0x00,
Gzip = 0x01,
Brotli = 0x02,
}
impl CompressionAlgorithm {
pub fn from_byte(byte: u8) -> Result<Self> {
match byte {
0x00 => Ok(CompressionAlgorithm::None),
0x01 => Ok(CompressionAlgorithm::Gzip),
0x02 => Ok(CompressionAlgorithm::Brotli),
_ => Err(Error::UnsupportedAlgorithm(byte)),
}
}
pub fn name(&self) -> &'static str {
match self {
CompressionAlgorithm::None => "none",
CompressionAlgorithm::Gzip => "gzip",
CompressionAlgorithm::Brotli => "brotli",
}
}
pub fn from_name(name: &str) -> Result<Self> {
match name.to_lowercase().as_str() {
"none" => Ok(CompressionAlgorithm::None),
"gzip" => Ok(CompressionAlgorithm::Gzip),
"brotli" => Ok(CompressionAlgorithm::Brotli),
_ => Err(Error::UnsupportedAlgorithm(0)),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressionResult {
pub compressed: Vec<u8>,
pub algorithm: CompressionAlgorithm,
pub original_size: usize,
pub compressed_size: usize,
pub compression_ratio: f64,
}
#[derive(Debug, Clone)]
pub struct CompressionOptions {
pub algorithm: CompressionAlgorithm,
pub min_size_threshold: usize,
pub level: u32,
}
impl Default for CompressionOptions {
fn default() -> Self {
Self {
algorithm: CompressionAlgorithm::Brotli,
min_size_threshold: 100,
level: 6,
}
}
}
pub fn compress(data: &[u8], options: Option<CompressionOptions>) -> Result<CompressionResult> {
let opts = options.unwrap_or_default();
validate_compression_level(opts.algorithm, opts.level)?;
let original_size = data.len();
if original_size < opts.min_size_threshold {
return Ok(CompressionResult {
compressed: data.to_vec(),
algorithm: CompressionAlgorithm::None,
original_size,
compressed_size: original_size,
compression_ratio: 1.0,
});
}
if opts.algorithm == CompressionAlgorithm::None {
return Ok(CompressionResult {
compressed: data.to_vec(),
algorithm: CompressionAlgorithm::None,
original_size,
compressed_size: original_size,
compression_ratio: 1.0,
});
}
let (compressed, algorithm) = match opts.algorithm {
CompressionAlgorithm::Brotli => compress_brotli(data, opts.level)?,
CompressionAlgorithm::Gzip => compress_gzip(data, opts.level)?,
CompressionAlgorithm::None => (data.to_vec(), CompressionAlgorithm::None),
};
let compressed_size = compressed.len();
let compression_ratio = compressed_size as f64 / original_size as f64;
let within_expansion_policy = compressed_size > 0
&& original_size <= compressed_size.saturating_mul(MAX_COMPRESSION_RATIO);
if compression_ratio < 0.9 && within_expansion_policy {
Ok(CompressionResult {
compressed,
algorithm,
original_size,
compressed_size,
compression_ratio,
})
} else {
Ok(CompressionResult {
compressed: data.to_vec(),
algorithm: CompressionAlgorithm::None,
original_size,
compressed_size: original_size,
compression_ratio: 1.0,
})
}
}
fn validate_compression_level(algorithm: CompressionAlgorithm, level: u32) -> Result<()> {
let valid = match algorithm {
CompressionAlgorithm::None => level == 0 || level == CompressionOptions::default().level,
CompressionAlgorithm::Gzip => level <= 9,
CompressionAlgorithm::Brotli => level <= 11,
};
if valid {
Ok(())
} else {
Err(Error::InvalidConfiguration(format!(
"invalid compression level {level} for {}",
algorithm.name()
)))
}
}
pub fn decompress(data: &[u8], algorithm: CompressionAlgorithm) -> Result<Vec<u8>> {
decompress_with_limits(
data,
algorithm,
MAX_DECOMPRESSED_SIZE,
MAX_COMPRESSION_RATIO,
)
}
pub fn decompress_with_limits(
data: &[u8],
algorithm: CompressionAlgorithm,
max_output_size: usize,
max_ratio: usize,
) -> Result<Vec<u8>> {
if max_ratio == 0 {
return Err(Error::InvalidConfiguration(
"decompression ratio limit must be greater than zero".to_string(),
));
}
let ratio_limit = data.len().saturating_mul(max_ratio);
let effective_limit = max_output_size.min(ratio_limit);
match algorithm {
CompressionAlgorithm::None => {
if data.len() > max_output_size {
return Err(Error::PayloadTooLarge {
size: data.len(),
limit: max_output_size,
});
}
Ok(data.to_vec())
}
CompressionAlgorithm::Gzip => decompress_gzip(data, effective_limit),
CompressionAlgorithm::Brotli => decompress_brotli(data, effective_limit),
}
}
pub fn decompress_exact(
data: &[u8],
algorithm: CompressionAlgorithm,
expected_size: usize,
) -> Result<Vec<u8>> {
if expected_size > MAX_DECOMPRESSED_SIZE {
return Err(Error::PayloadTooLarge {
size: expected_size,
limit: MAX_DECOMPRESSED_SIZE,
});
}
let output = decompress_with_limits(data, algorithm, expected_size, MAX_COMPRESSION_RATIO)?;
if output.len() != expected_size {
return Err(Error::SizeMismatch {
expected: expected_size,
actual: output.len(),
});
}
Ok(output)
}
fn read_decompressed_with_limit<R: std::io::Read>(
reader: R,
max_output_size: usize,
) -> Result<Vec<u8>> {
use std::io::Read;
let read_limit = max_output_size.saturating_add(1) as u64;
let mut limited = reader.take(read_limit);
let mut output = Vec::with_capacity(max_output_size.min(64 * 1024));
limited
.read_to_end(&mut output)
.map_err(|e| Error::DecompressionFailed(e.to_string()))?;
if output.len() > max_output_size {
return Err(Error::PayloadTooLarge {
size: output.len(),
limit: max_output_size,
});
}
Ok(output)
}
fn compress_brotli(data: &[u8], level: u32) -> Result<(Vec<u8>, CompressionAlgorithm)> {
use brotli::enc::BrotliEncoderParams;
let mut output = Vec::new();
let mut params = BrotliEncoderParams::default();
params.quality = level as i32;
brotli::BrotliCompress(&mut std::io::Cursor::new(data), &mut output, ¶ms)
.map_err(|e| Error::CompressionFailed(e.to_string()))?;
Ok((output, CompressionAlgorithm::Brotli))
}
fn decompress_brotli(data: &[u8], max_output_size: usize) -> Result<Vec<u8>> {
let decoder = brotli::Decompressor::new(std::io::Cursor::new(data), 4096);
read_decompressed_with_limit(decoder, max_output_size)
}
fn compress_gzip(data: &[u8], level: u32) -> Result<(Vec<u8>, CompressionAlgorithm)> {
use flate2::write::GzEncoder;
use flate2::Compression;
use std::io::Write;
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(level));
encoder
.write_all(data)
.map_err(|e| Error::CompressionFailed(e.to_string()))?;
let output = encoder
.finish()
.map_err(|e| Error::CompressionFailed(e.to_string()))?;
Ok((output, CompressionAlgorithm::Gzip))
}
fn decompress_gzip(data: &[u8], max_output_size: usize) -> Result<Vec<u8>> {
use flate2::read::GzDecoder;
read_decompressed_with_limit(GzDecoder::new(data), max_output_size)
}
pub fn serialize_with_header(result: &CompressionResult) -> Result<Vec<u8>> {
let original_size =
u32::try_from(result.original_size).map_err(|_| Error::PayloadTooLarge {
size: result.original_size,
limit: u32::MAX as usize,
})?;
let mut output = Vec::with_capacity(7 + result.compressed.len());
output.extend_from_slice(MAGIC_COMPRESSED);
output.push(result.algorithm as u8);
output.extend_from_slice(&original_size.to_be_bytes());
output.extend_from_slice(&result.compressed);
Ok(output)
}
pub fn deserialize_with_header(data: &[u8]) -> Result<(Vec<u8>, CompressionAlgorithm, usize)> {
if data.len() < 7 {
return Err(Error::TruncatedPayload {
expected: 7,
actual: data.len(),
});
}
if &data[0..2] != MAGIC_COMPRESSED {
return Err(Error::InvalidFormat);
}
let algorithm = CompressionAlgorithm::from_byte(data[2])?;
let original_size = u32::from_be_bytes([data[3], data[4], data[5], data[6]]) as usize;
let compressed = data[7..].to_vec();
Ok((compressed, algorithm, original_size))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_gzip_roundtrip() {
let data = b"Hello, World! This is a test message that should be compressed.";
let result = compress(
data,
Some(CompressionOptions {
algorithm: CompressionAlgorithm::Gzip,
min_size_threshold: 10,
level: 6,
}),
)
.unwrap();
let decompressed = decompress(&result.compressed, result.algorithm).unwrap();
assert_eq!(data, &decompressed[..]);
}
#[test]
fn test_brotli_roundtrip() {
let data = b"Hello, World! This is a test message that should be compressed with Brotli.";
let result = compress(
data,
Some(CompressionOptions {
algorithm: CompressionAlgorithm::Brotli,
min_size_threshold: 10,
level: 6,
}),
)
.unwrap();
let decompressed = decompress(&result.compressed, result.algorithm).unwrap();
assert_eq!(data, &decompressed[..]);
}
#[test]
fn test_skip_small_data() {
let data = b"tiny";
let result = compress(
data,
Some(CompressionOptions {
algorithm: CompressionAlgorithm::Brotli,
min_size_threshold: 100, level: 6,
}),
)
.unwrap();
assert_eq!(result.algorithm, CompressionAlgorithm::None);
assert_eq!(result.compressed, data);
}
#[test]
fn test_header_serialization() {
let data = b"Test data for header serialization test with enough content.";
let result = compress(
data,
Some(CompressionOptions {
algorithm: CompressionAlgorithm::Gzip,
min_size_threshold: 10,
level: 6,
}),
)
.unwrap();
let serialized = serialize_with_header(&result).unwrap();
let (compressed, algorithm, original_size) = deserialize_with_header(&serialized).unwrap();
assert_eq!(algorithm, result.algorithm);
assert_eq!(original_size, result.original_size);
assert_eq!(compressed, result.compressed);
}
#[test]
fn test_decompression_rejects_bomb_before_full_expansion() {
let plaintext = vec![0u8; 2 * 1024 * 1024];
let (compressed, _) = compress_gzip(&plaintext, 6).unwrap();
assert!(compressed.len().saturating_mul(MAX_COMPRESSION_RATIO) < plaintext.len());
assert!(matches!(
decompress(&compressed, CompressionAlgorithm::Gzip),
Err(Error::PayloadTooLarge { .. })
));
}
#[test]
fn test_decompression_exact_enforces_expected_size() {
let plaintext = b"bounded decompression".repeat(64);
let (compressed, _) = compress_gzip(&plaintext, 6).unwrap();
assert_eq!(
decompress_exact(&compressed, CompressionAlgorithm::Gzip, plaintext.len()).unwrap(),
plaintext
);
assert!(matches!(
decompress_exact(&compressed, CompressionAlgorithm::Gzip, plaintext.len() - 1),
Err(Error::PayloadTooLarge { .. }) | Err(Error::SizeMismatch { .. })
));
}
#[test]
fn test_none_decompression_obeys_output_bound() {
assert!(matches!(
decompress_with_limits(&[0u8; 9], CompressionAlgorithm::None, 8, 256),
Err(Error::PayloadTooLarge { size: 9, limit: 8 })
));
}
#[test]
fn test_compression_rejects_invalid_levels_before_processing() {
assert!(matches!(
compress(
b"small",
Some(CompressionOptions {
algorithm: CompressionAlgorithm::Gzip,
min_size_threshold: usize::MAX,
level: 10,
})
),
Err(Error::InvalidConfiguration(_))
));
assert!(matches!(
compress(
b"small",
Some(CompressionOptions {
algorithm: CompressionAlgorithm::Brotli,
min_size_threshold: usize::MAX,
level: 12,
})
),
Err(Error::InvalidConfiguration(_))
));
}
}