use crate::core::error::{Error, Result};
use crate::storage::unified_memory::CompressionType;
pub trait CompressionEngine: Send + Sync {
fn compress(&self, data: &[u8]) -> Result<Vec<u8>>;
fn decompress(&self, data: &[u8]) -> Result<Vec<u8>>;
fn name(&self) -> &'static str;
fn compression_ratio(&self, original_size: usize, compressed_size: usize) -> f64;
}
const LZ4_MAX_EXPANSION: usize = 256;
const LZ4_EXPANSION_SLACK: usize = 64 * 1024;
const MAX_DECOMPRESSED_BLOCK: usize = 1 << 32;
const LZ4_LEN_PREFIX: usize = 8;
pub struct Lz4CompressionEngine;
impl CompressionEngine for Lz4CompressionEngine {
fn compress(&self, data: &[u8]) -> Result<Vec<u8>> {
if data.is_empty() {
return Ok(Vec::new());
}
let body = oxiarc_lz4::compress_bytes(data)
.map_err(|e| Error::InvalidValue(format!("LZ4 compression failed: {}", e)))?;
let mut out = Vec::with_capacity(LZ4_LEN_PREFIX + body.len());
out.extend_from_slice(&(data.len() as u64).to_le_bytes());
out.extend_from_slice(&body);
Ok(out)
}
fn decompress(&self, data: &[u8]) -> Result<Vec<u8>> {
if data.is_empty() {
return Ok(Vec::new());
}
if data.len() < LZ4_LEN_PREFIX {
return Err(Error::InvalidValue(format!(
"Truncated LZ4 block: {} bytes, need at least {}",
data.len(),
LZ4_LEN_PREFIX
)));
}
let mut len_bytes = [0u8; LZ4_LEN_PREFIX];
len_bytes.copy_from_slice(&data[..LZ4_LEN_PREFIX]);
let original_len = u64::from_le_bytes(len_bytes);
if original_len == 0 {
return Ok(Vec::new());
}
let body = &data[LZ4_LEN_PREFIX..];
let bound = body
.len()
.saturating_mul(LZ4_MAX_EXPANSION)
.saturating_add(LZ4_EXPANSION_SLACK)
.min(MAX_DECOMPRESSED_BLOCK);
if original_len > bound as u64 {
return Err(Error::InvalidValue(format!(
"Refusing LZ4 decompression bomb: header claims {} bytes from a {} byte body (bound {})",
original_len,
body.len(),
bound
)));
}
let decoded = oxiarc_lz4::decompress_bytes(body, original_len as usize)
.map_err(|e| Error::InvalidValue(format!("LZ4 decompression failed: {}", e)))?;
if decoded.len() != original_len as usize {
return Err(Error::InvalidValue(format!(
"LZ4 decompression produced {} bytes, header claimed {}",
decoded.len(),
original_len
)));
}
Ok(decoded)
}
fn name(&self) -> &'static str {
"LZ4"
}
fn compression_ratio(&self, original_size: usize, compressed_size: usize) -> f64 {
if compressed_size == 0 {
0.0
} else {
original_size as f64 / compressed_size as f64
}
}
}
pub struct ZstdCompressionEngine {
compression_level: i32,
}
impl ZstdCompressionEngine {
pub fn new(level: i32) -> Self {
Self {
compression_level: level,
}
}
pub fn level(&self) -> i32 {
self.compression_level
}
}
impl CompressionEngine for ZstdCompressionEngine {
fn compress(&self, data: &[u8]) -> Result<Vec<u8>> {
if data.is_empty() {
return Ok(Vec::new());
}
oxiarc_zstd::compress_with_level(data, self.compression_level)
.map_err(|e| Error::InvalidValue(format!("ZSTD compression failed: {}", e)))
}
fn decompress(&self, data: &[u8]) -> Result<Vec<u8>> {
if data.is_empty() {
return Ok(Vec::new());
}
oxiarc_zstd::decompress(data)
.map_err(|e| Error::InvalidValue(format!("ZSTD decompression failed: {}", e)))
}
fn name(&self) -> &'static str {
"ZSTD"
}
fn compression_ratio(&self, original_size: usize, compressed_size: usize) -> f64 {
if compressed_size == 0 {
0.0
} else {
original_size as f64 / compressed_size as f64
}
}
}
pub struct NoCompressionEngine;
impl CompressionEngine for NoCompressionEngine {
fn compress(&self, data: &[u8]) -> Result<Vec<u8>> {
Ok(data.to_vec())
}
fn decompress(&self, data: &[u8]) -> Result<Vec<u8>> {
Ok(data.to_vec())
}
fn name(&self) -> &'static str {
"None"
}
fn compression_ratio(&self, _original_size: usize, _compressed_size: usize) -> f64 {
1.0
}
}
pub fn build_engines(
zstd_level: i32,
) -> std::collections::HashMap<CompressionType, Box<dyn CompressionEngine>> {
let mut engines: std::collections::HashMap<CompressionType, Box<dyn CompressionEngine>> =
std::collections::HashMap::new();
engines.insert(CompressionType::None, Box::new(NoCompressionEngine));
engines.insert(CompressionType::Lz4, Box::new(Lz4CompressionEngine));
engines.insert(
CompressionType::Zstd,
Box::new(ZstdCompressionEngine::new(zstd_level)),
);
engines.insert(
CompressionType::Snappy,
Box::new(Lz4CompressionEngine), );
engines.insert(
CompressionType::Gzip,
Box::new(ZstdCompressionEngine::new(zstd_level)),
);
engines
}
pub fn resolve_compression(requested: CompressionType) -> CompressionType {
match requested {
CompressionType::Auto => CompressionType::Zstd,
other => other,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn lz4_roundtrip() {
let engine = Lz4CompressionEngine;
let data = b"Hello, World! This is a test string for compression.";
let compressed = engine.compress(data).expect("compress");
assert_eq!(engine.decompress(&compressed).expect("decompress"), data);
}
#[test]
fn lz4_rejects_decompression_bomb() {
let engine = Lz4CompressionEngine;
let mut hostile = u64::MAX.to_le_bytes().to_vec();
hostile.extend_from_slice(&[0, 0, 0, 0]);
let err = engine.decompress(&hostile).expect_err("must reject");
assert!(
format!("{}", err).contains("bomb"),
"unexpected error: {}",
err
);
}
#[test]
fn lz4_rejects_truncated_header() {
let engine = Lz4CompressionEngine;
assert!(engine.decompress(&[1, 2, 3]).is_err());
}
#[test]
fn codecs_actually_shrink_compressible_data() {
let original: Vec<u8> = b"PandRS compression test payload. "
.iter()
.copied()
.cycle()
.take(16 * 1024)
.collect();
for engine in [
Box::new(Lz4CompressionEngine) as Box<dyn CompressionEngine>,
Box::new(ZstdCompressionEngine::new(3)) as Box<dyn CompressionEngine>,
] {
let compressed = engine.compress(&original).expect("compress");
assert!(
compressed.len() < original.len(),
"{} did not shrink data: {} -> {}",
engine.name(),
original.len(),
compressed.len()
);
assert_eq!(
engine.decompress(&compressed).expect("decompress"),
original
);
}
}
#[test]
fn every_compression_type_has_an_engine() {
let engines = build_engines(3);
for ty in [
CompressionType::None,
CompressionType::Lz4,
CompressionType::Zstd,
CompressionType::Snappy,
CompressionType::Gzip,
] {
assert!(engines.contains_key(&ty), "missing engine for {:?}", ty);
}
assert_eq!(
resolve_compression(CompressionType::Auto),
CompressionType::Zstd
);
assert!(engines.contains_key(&resolve_compression(CompressionType::Auto)));
}
}