use super::{Compression, CompressionError, CompressionErrorKind, Result};
const DEFAULT_ZSTD_LEVEL: i32 = 3;
pub fn compress_with_level(payload: &[u8], level: Option<i32>) -> Result<Vec<u8>> {
let lvl = level.unwrap_or(DEFAULT_ZSTD_LEVEL);
::zstd::bulk::compress(payload, lvl).map_err(|e| encode_err(e.to_string()))
}
pub fn decompress(payload: &[u8]) -> Result<Vec<u8>> {
decompress_bounded(payload, super::MAX_DECOMPRESSED_LEN)
}
pub fn decompress_bounded(payload: &[u8], max_len: usize) -> Result<Vec<u8>> {
let decoder =
::zstd::stream::read::Decoder::new(payload).map_err(|e| decode_err(e.to_string()))?;
super::read_to_end_bounded(decoder, max_len, Compression::Zstd)
}
const fn encode_err(message: String) -> CompressionError {
CompressionError::new(
Compression::Zstd,
CompressionErrorKind::EncodeFailed { message },
)
}
const fn decode_err(message: String) -> CompressionError {
CompressionError::new(
Compression::Zstd,
CompressionErrorKind::DecodeFailed { message },
)
}
#[cfg(test)]
mod tests {
use super::{super::CompressionErrorKind, compress_with_level, decompress, decompress_bounded};
#[test]
fn decompress_bounded_rejects_a_decompression_bomb() {
let payload = vec![0u8; 4096];
let compressed = compress_with_level(&payload, None).unwrap();
let err = decompress_bounded(&compressed, 64).unwrap_err();
assert!(
matches!(
err.kind,
CompressionErrorKind::DecompressedTooLarge { limit: 64 }
),
"expected DecompressedTooLarge, got {:?}",
err.kind
);
}
#[test]
fn decompress_bounded_allows_output_at_exactly_the_limit() {
let payload = vec![0u8; 4096];
let compressed = compress_with_level(&payload, None).unwrap();
assert_eq!(decompress_bounded(&compressed, 4096).unwrap(), payload);
assert_eq!(decompress(&compressed).unwrap(), payload);
}
}