use crate::error::PersistenceError;
use crate::error::Result;
use std::io::Read;
use std::io::Write;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CompressionLevel {
Fast,
Balanced,
Maximum,
Custom(i32),
}
impl CompressionLevel {
pub fn to_level(self) -> i32 {
match self {
Self::Fast => 1,
Self::Balanced => 3,
Self::Maximum => 9,
Self::Custom(level) => level.clamp(1, 22),
}
}
}
impl Default for CompressionLevel {
fn default() -> Self {
Self::Balanced
}
}
pub struct Compressor {
level: CompressionLevel,
}
impl Compressor {
pub const fn new(level: CompressionLevel) -> Self {
Self { level }
}
pub fn compress(&self, data: &[u8]) -> Result<Vec<u8>> {
let mut encoder = zstd::Encoder::new(Vec::new(), self.level.to_level())
.map_err(|e| PersistenceError::Compression(e.to_string()))?;
encoder
.write_all(data)
.map_err(|e| PersistenceError::Compression(e.to_string()))?;
encoder
.finish()
.map_err(|e| PersistenceError::Compression(e.to_string()))
}
pub fn decompress(&self, compressed: &[u8]) -> Result<Vec<u8>> {
let mut decoder = zstd::Decoder::new(compressed)
.map_err(|e| PersistenceError::Compression(e.to_string()))?;
let mut decompressed = Vec::new();
decoder
.read_to_end(&mut decompressed)
.map_err(|e| PersistenceError::Compression(e.to_string()))?;
Ok(decompressed)
}
pub fn compress_with_dict(&self, data: &[u8], _dict: &[u8]) -> Result<Vec<u8>> {
zstd::encode_all(std::io::Cursor::new(data), self.level.to_level())
.map_err(|e| PersistenceError::Compression(e.to_string()))
}
pub fn compression_ratio(original_size: usize, compressed_size: usize) -> f32 {
if compressed_size == 0 {
return 0.0;
}
1.0 - (compressed_size as f32 / original_size as f32)
}
}
pub struct StreamCompressor {
level: CompressionLevel,
}
impl StreamCompressor {
pub const fn new(level: CompressionLevel) -> Self {
Self { level }
}
pub fn compress_stream<R: Read, W: Write>(&self, mut reader: R, writer: W) -> Result<u64> {
let mut encoder = zstd::Encoder::new(writer, self.level.to_level())
.map_err(|e| PersistenceError::Compression(e.to_string()))?;
let bytes_written = std::io::copy(&mut reader, &mut encoder)?;
encoder
.finish()
.map_err(|e| PersistenceError::Compression(e.to_string()))?;
Ok(bytes_written)
}
pub fn decompress_stream<R: Read, W: Write>(&self, reader: R, mut writer: W) -> Result<u64> {
let mut decoder =
zstd::Decoder::new(reader).map_err(|e| PersistenceError::Compression(e.to_string()))?;
let bytes_written = std::io::copy(&mut decoder, &mut writer)?;
Ok(bytes_written)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_compression_roundtrip() {
let compressor = Compressor::new(CompressionLevel::Balanced);
let data = "Hello, AGCodex! This is a test of the compression system. ".repeat(100);
let data = data.as_bytes();
let compressed = compressor.compress(data).unwrap();
assert!(compressed.len() < data.len());
let decompressed = compressor.decompress(&compressed).unwrap();
assert_eq!(decompressed, data);
}
#[test]
fn test_compression_levels() {
let data = vec![b'A'; 10000];
let fast = Compressor::new(CompressionLevel::Fast);
let balanced = Compressor::new(CompressionLevel::Balanced);
let maximum = Compressor::new(CompressionLevel::Maximum);
let fast_compressed = fast.compress(&data).unwrap();
let balanced_compressed = balanced.compress(&data).unwrap();
let maximum_compressed = maximum.compress(&data).unwrap();
assert!(maximum_compressed.len() <= fast_compressed.len());
assert_eq!(fast.decompress(&fast_compressed).unwrap(), data);
assert_eq!(balanced.decompress(&balanced_compressed).unwrap(), data);
assert_eq!(maximum.decompress(&maximum_compressed).unwrap(), data);
}
#[test]
fn test_compression_ratio() {
let ratio = Compressor::compression_ratio(1000, 100);
assert!((ratio - 0.9).abs() < 0.001);
let ratio = Compressor::compression_ratio(1000, 500);
assert!((ratio - 0.5).abs() < 0.001);
}
}