use flate2::Compression;
use flate2::read::{ZlibDecoder, ZlibEncoder};
use std::io::Read;
use crate::error::{Error, Result};
pub const COMPRESS_EXTFIELD: &str = "zlib";
pub fn compress(plaintext: &[u8]) -> Result<(Vec<u8>, bool)> {
let mut encoder = ZlibEncoder::new(plaintext, Compression::default());
let mut compressed = Vec::new();
encoder
.read_to_end(&mut compressed)
.map_err(|e| Error::Cipher(format!("compression failed: {e}")))?;
if compressed.len() < plaintext.len() {
Ok((compressed, true))
} else {
Ok((plaintext.to_vec(), false))
}
}
pub fn decompress(data: &[u8]) -> Result<Vec<u8>> {
let mut decoder = ZlibDecoder::new(data);
let mut decompressed = Vec::new();
decoder
.read_to_end(&mut decompressed)
.map_err(|e| Error::Cipher(format!("decompression failed: {e}")))?;
Ok(decompressed)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip_text() {
let input = b"Hello, World! This is a test of compression. ".repeat(20);
let (compressed, did_compress) = compress(&input).unwrap();
assert!(did_compress);
assert!(compressed.len() < input.len());
let recovered = decompress(&compressed).unwrap();
assert_eq!(recovered, input);
}
#[test]
fn round_trip_random_data() {
let input: Vec<u8> = (0..1024u32)
.flat_map(|i| i.wrapping_mul(2654435761).to_le_bytes())
.take(512)
.collect();
let (output, did_compress) = compress(&input).unwrap();
if did_compress {
let recovered = decompress(&output).unwrap();
assert_eq!(recovered, input);
} else {
assert_eq!(output, input);
}
}
#[test]
fn empty_input() {
let (output, did_compress) = compress(b"").unwrap();
assert!(!did_compress);
assert!(output.is_empty());
}
}