#[cfg(feature = "zstd")]
use std::io::{Read, Write};
#[cfg(feature = "zstd")]
use crate::CasError;
pub const DEFAULT_COMPRESSION_LEVEL: i32 = 3;
#[cfg(feature = "zstd")]
pub const MAX_DECOMPRESSED_SIZE: usize = 1024 * 1024 * 1024;
#[cfg(feature = "zstd")]
pub fn compress(data: &[u8], level: i32) -> Result<Vec<u8>, CasError> {
let mut encoder = zstd::Encoder::new(Vec::new(), level)
.map_err(|e| CasError::CompressionError(e.to_string()))?;
encoder
.write_all(data)
.map_err(|e| CasError::CompressionError(e.to_string()))?;
let compressed = encoder
.finish()
.map_err(|e| CasError::CompressionError(e.to_string()))?;
Ok(compressed)
}
#[cfg(feature = "zstd")]
pub fn compress_default(data: &[u8]) -> Result<Vec<u8>, CasError> {
compress(data, DEFAULT_COMPRESSION_LEVEL)
}
#[cfg(feature = "zstd")]
pub fn decompress(data: &[u8]) -> Result<Vec<u8>, CasError> {
let mut decoder =
zstd::Decoder::new(data).map_err(|e| CasError::DecompressionError(e.to_string()))?;
let mut output = Vec::with_capacity(data.len() * 2); let mut buffer = [0u8; 64 * 1024];
loop {
let n = decoder
.read(&mut buffer)
.map_err(|e| CasError::DecompressionError(e.to_string()))?;
if n == 0 {
break;
}
if output.len() + n > MAX_DECOMPRESSED_SIZE {
return Err(CasError::DecompressionTooLarge {
max: MAX_DECOMPRESSED_SIZE,
});
}
output.extend_from_slice(&buffer[..n]);
}
Ok(output)
}
#[cfg(all(test, feature = "zstd"))]
mod tests {
use super::*;
#[test]
fn test_compress_decompress_roundtrip() -> Result<(), CasError> {
let original = b"Hello! This is test data for compression roundtrip.";
let compressed = compress_default(original)?;
let decompressed = decompress(&compressed)?;
assert_eq!(original.as_slice(), decompressed.as_slice());
Ok(())
}
#[test]
fn test_compress_larger_data() -> Result<(), CasError> {
let original: Vec<u8> = (0..100_000).map(|i| (i % 256) as u8).collect();
let compressed = compress_default(&original)?;
let decompressed = decompress(&compressed)?;
assert_eq!(original, decompressed);
assert!(compressed.len() < original.len());
Ok(())
}
#[test]
fn test_compress_empty() -> Result<(), CasError> {
let original = b"";
let compressed = compress_default(original)?;
let decompressed = decompress(&compressed)?;
assert_eq!(original.as_slice(), decompressed.as_slice());
Ok(())
}
#[test]
fn test_is_zstd_compressed_magic() -> Result<(), CasError> {
let compressed = compress_default(b"hello")?;
assert!(compressed.len() >= 4);
assert_eq!(&compressed[..4], &[0x28, 0xB5, 0x2F, 0xFD]);
Ok::<_, CasError>(())
}
#[test]
fn test_decompress_invalid_data() {
let result = decompress(b"not zstd data at all!");
assert!(matches!(result, Err(CasError::DecompressionError(_))));
}
#[test]
fn test_compress_levels() -> Result<(), CasError> {
let data = "The quick brown fox jumps over the lazy dog. ".repeat(1000);
let bytes = data.as_bytes();
let c1 = compress(bytes, 1)?;
let c3 = compress(bytes, 3)?;
let c9 = compress(bytes, 9)?;
assert!(c9.len() <= c3.len());
assert!(c3.len() <= c1.len());
Ok(())
}
}