use crate::error::RiegeliError;
use crate::varint::{decode_u64, encode_u64};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum CompressionType {
None = 0,
Brotli = b'b',
Zstd = b'z',
Snappy = b's',
}
impl TryFrom<u8> for CompressionType {
type Error = RiegeliError;
fn try_from(b: u8) -> Result<Self, Self::Error> {
match b {
0 => Ok(CompressionType::None),
b'b' => Ok(CompressionType::Brotli),
b'z' => Ok(CompressionType::Zstd),
b's' => Ok(CompressionType::Snappy),
_ => Err(RiegeliError::UnknownCompressionType(b)),
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct CompressOptions {
pub level: Option<i32>,
pub window_log: Option<u32>,
}
#[cfg(feature = "brotli")]
pub(crate) fn compress_brotli(
input: &[u8],
opts: CompressOptions,
) -> Result<Vec<u8>, RiegeliError> {
use std::io::Write as _;
let quality = opts.level.unwrap_or(6).clamp(0, 11) as u32;
let lgwin = opts.window_log.unwrap_or(22).clamp(10, 30);
let mut output = Vec::new();
{
let mut writer = brotli::CompressorWriter::new(&mut output, 4096, quality, lgwin);
writer.write_all(input).map_err(|e| {
RiegeliError::MalformedData(format!("brotli compress error: {e}").into())
})?;
}
Ok(output)
}
#[cfg(feature = "brotli")]
pub(crate) fn decompress_brotli(input: &[u8], max_len: u64) -> Result<Vec<u8>, RiegeliError> {
use brotli::{BrotliDecompressStream, BrotliResult, BrotliState, HeapAlloc};
const MAX_DECOMPRESS_PREALLOC: u64 = 1 << 24; let hard_cap = usize::try_from(max_len.saturating_add(1)).unwrap_or(usize::MAX);
let mut state = BrotliState::new(
HeapAlloc::<u8>::new(0),
HeapAlloc::<u32>::new(0),
HeapAlloc::new(Default::default()),
);
let mut available_in = input.len();
let mut input_offset = 0usize;
let mut output = vec![0u8; hard_cap.min(MAX_DECOMPRESS_PREALLOC as usize).max(64)];
let mut output_offset = 0usize;
let mut total_out = 0usize;
loop {
let mut available_out = output.len() - output_offset;
match BrotliDecompressStream(
&mut available_in,
&mut input_offset,
input,
&mut available_out,
&mut output_offset,
&mut output,
&mut total_out,
&mut state,
) {
BrotliResult::ResultSuccess => {
if available_in != 0 {
return Err(RiegeliError::MalformedData(
"trailing data after Brotli-compressed stream".into(),
));
}
output.truncate(output_offset);
if output.len() as u64 > max_len {
return Err(RiegeliError::MalformedData(
format!("decompressed data exceeds its declared size ({max_len} bytes)")
.into(),
));
}
return Ok(output);
}
BrotliResult::NeedsMoreOutput => {
if output.len() >= hard_cap {
return Err(RiegeliError::MalformedData(
format!("decompressed data exceeds its declared size ({max_len} bytes)")
.into(),
));
}
let new_len = output.len().saturating_mul(2).max(4096).min(hard_cap);
output.resize(new_len, 0);
}
BrotliResult::NeedsMoreInput => {
return Err(RiegeliError::MalformedData(
"brotli decompress error: truncated stream".into(),
));
}
BrotliResult::ResultFailure => {
return Err(RiegeliError::MalformedData(
"brotli decompress error: invalid stream".into(),
));
}
}
}
}
#[cfg(feature = "zstd")]
pub(crate) fn compress_zstd(input: &[u8], opts: CompressOptions) -> Result<Vec<u8>, RiegeliError> {
let level = opts.level.unwrap_or(3).clamp(-131072, 22);
let compressed = if let Some(wlog) = opts.window_log {
use std::io::Write as _;
let mut output = Vec::new();
{
let mut encoder = zstd::Encoder::new(&mut output, level).map_err(|e| {
RiegeliError::MalformedData(format!("zstd encoder init: {e}").into())
})?;
encoder
.window_log(wlog)
.map_err(|e| RiegeliError::MalformedData(format!("zstd window_log: {e}").into()))?;
encoder.write_all(input).map_err(|e| {
RiegeliError::MalformedData(format!("zstd compress error: {e}").into())
})?;
encoder.finish().map_err(|e| {
RiegeliError::MalformedData(format!("zstd finish error: {e}").into())
})?;
}
output
} else {
zstd::encode_all(input, level)
.map_err(|e| RiegeliError::MalformedData(format!("zstd compress error: {e}").into()))?
};
Ok(compressed)
}
#[cfg(feature = "zstd")]
pub(crate) fn decompress_zstd(input: &[u8], max_len: u64) -> Result<Vec<u8>, RiegeliError> {
use std::io::Read as _;
const MAX_DECOMPRESS_PREALLOC: u64 = 1 << 24; let cursor = std::io::Cursor::new(input);
let mut decoder = zstd::stream::read::Decoder::with_buffer(cursor)
.map_err(|e| RiegeliError::MalformedData(format!("zstd decompress error: {e}").into()))?
.single_frame();
let window_log_max: u32 = if cfg!(target_pointer_width = "32") {
30
} else {
31
};
decoder
.window_log_max(window_log_max)
.map_err(|e| RiegeliError::MalformedData(format!("zstd window_log_max: {e}").into()))?;
let mut output = Vec::with_capacity(max_len.min(MAX_DECOMPRESS_PREALLOC) as usize);
(&mut decoder)
.take(max_len.saturating_add(1))
.read_to_end(&mut output)
.map_err(|e| RiegeliError::MalformedData(format!("zstd decompress error: {e}").into()))?;
if output.len() as u64 > max_len {
return Err(RiegeliError::MalformedData(
format!("decompressed data exceeds its declared size ({max_len} bytes)").into(),
));
}
let consumed = decoder.finish().position() as usize;
if consumed < input.len() {
return Err(RiegeliError::MalformedData(
"trailing data after Zstd-compressed stream".into(),
));
}
Ok(output)
}
#[cfg(feature = "snappy")]
pub(crate) fn compress_snappy(input: &[u8]) -> Result<Vec<u8>, RiegeliError> {
let mut encoder = snap::raw::Encoder::new();
encoder
.compress_vec(input)
.map_err(|e| RiegeliError::MalformedData(format!("snappy compress error: {e}").into()))
}
#[cfg(feature = "snappy")]
pub(crate) fn decompress_snappy(input: &[u8]) -> Result<Vec<u8>, RiegeliError> {
let mut decoder = snap::raw::Decoder::new();
decoder
.decompress_vec(input)
.map_err(|e| RiegeliError::MalformedData(format!("snappy decompress error: {e}").into()))
}
pub(crate) fn compress_data(
data: &[u8],
compression: CompressionType,
opts: CompressOptions,
) -> Result<Vec<u8>, RiegeliError> {
match compression {
CompressionType::None => Ok(data.to_vec()),
CompressionType::Brotli => {
#[cfg(feature = "brotli")]
{
compress_brotli(data, opts)
}
#[cfg(not(feature = "brotli"))]
{
Err(RiegeliError::UnsupportedCompression(
CompressionType::Brotli as u8,
))
}
}
CompressionType::Zstd => {
#[cfg(feature = "zstd")]
{
compress_zstd(data, opts)
}
#[cfg(not(feature = "zstd"))]
{
Err(RiegeliError::UnsupportedCompression(
CompressionType::Zstd as u8,
))
}
}
CompressionType::Snappy => {
#[cfg(feature = "snappy")]
{
compress_snappy(data)
}
#[cfg(not(feature = "snappy"))]
{
Err(RiegeliError::UnsupportedCompression(
CompressionType::Snappy as u8,
))
}
}
}
}
pub(crate) fn compress_with_prefix(
data: &[u8],
compression: CompressionType,
opts: CompressOptions,
) -> Result<Vec<u8>, RiegeliError> {
let compressed = compress_data(data, compression, opts)?;
if compression == CompressionType::None {
return Ok(compressed);
}
let mut result = Vec::new();
result.extend_from_slice(&encode_u64(data.len() as u64));
result.extend_from_slice(&compressed);
Ok(result)
}
pub(crate) fn compress_length_prefixed(
data: &[u8],
compression: CompressionType,
opts: CompressOptions,
) -> Result<Vec<u8>, RiegeliError> {
use crate::varint::length_varint_u64;
let compressed = compress_data(data, compression, opts)?;
let mut blob_len = compressed.len() as u64;
if compression != CompressionType::None {
blob_len += length_varint_u64(data.len() as u64) as u64;
}
let mut result = Vec::new();
result.extend_from_slice(&encode_u64(blob_len));
if compression != CompressionType::None {
result.extend_from_slice(&encode_u64(data.len() as u64));
}
result.extend_from_slice(&compressed);
Ok(result)
}
pub(crate) fn decompress_with_prefix(
data: &[u8],
compression: CompressionType,
) -> Result<Vec<u8>, RiegeliError> {
if compression == CompressionType::None {
return Ok(data.to_vec());
}
let (uncompressed_size, consumed) = decode_u64(data).map_err(|e| {
RiegeliError::MalformedData(format!("reading uncompressed_size prefix: {e}").into())
})?;
let out = decompress_data_capped(&data[consumed..], compression, uncompressed_size)?;
if out.len() as u64 != uncompressed_size {
return Err(RiegeliError::MalformedData(
format!(
"decompressed size {} != declared {uncompressed_size}",
out.len()
)
.into(),
));
}
Ok(out)
}
pub(crate) fn decompress_data_capped(
data: &[u8],
compression: CompressionType,
max_len: u64,
) -> Result<Vec<u8>, RiegeliError> {
let check = |out: Vec<u8>| {
if out.len() as u64 > max_len {
Err(RiegeliError::MalformedData(
format!("decompressed data exceeds its declared size ({max_len} bytes)").into(),
))
} else {
Ok(out)
}
};
match compression {
CompressionType::None => check(data.to_vec()),
CompressionType::Brotli => {
#[cfg(feature = "brotli")]
{
decompress_brotli(data, max_len)
}
#[cfg(not(feature = "brotli"))]
{
Err(RiegeliError::UnsupportedCompression(
CompressionType::Brotli as u8,
))
}
}
CompressionType::Zstd => {
#[cfg(feature = "zstd")]
{
decompress_zstd(data, max_len)
}
#[cfg(not(feature = "zstd"))]
{
Err(RiegeliError::UnsupportedCompression(
CompressionType::Zstd as u8,
))
}
}
CompressionType::Snappy => {
#[cfg(feature = "snappy")]
{
let declared = snap::raw::decompress_len(data).map_err(|e| {
RiegeliError::MalformedData(format!("snappy length error: {e}").into())
})?;
if declared as u64 > max_len {
return Err(RiegeliError::MalformedData(
format!("decompressed data exceeds its declared size ({max_len} bytes)")
.into(),
));
}
decompress_snappy(data)
}
#[cfg(not(feature = "snappy"))]
{
Err(RiegeliError::UnsupportedCompression(
CompressionType::Snappy as u8,
))
}
}
}
}
#[cfg(test)]
mod tests {
#![allow(unused_imports)]
use super::*;
const INPUT: &[u8] = b"hello world hello world hello world";
#[test]
#[cfg(feature = "brotli")]
fn brotli_round_trip() {
let compressed = compress_brotli(INPUT, CompressOptions::default()).unwrap();
let out = decompress_data_capped(&compressed, CompressionType::Brotli, 1 << 20).unwrap();
assert_eq!(out, INPUT);
}
#[test]
#[cfg(feature = "brotli")]
fn brotli_trailing_garbage_rejected() {
let mut compressed = compress_brotli(INPUT, CompressOptions::default()).unwrap();
compressed.extend_from_slice(&[0xde, 0xad, 0xbe, 0xef]);
let result = decompress_data_capped(&compressed, CompressionType::Brotli, 1 << 20);
assert!(
result.is_err(),
"trailing bytes after the Brotli stream must be rejected, got {result:?}"
);
}
#[test]
#[cfg(feature = "brotli")]
fn brotli_truncated_stream_rejected() {
let compressed = compress_brotli(INPUT, CompressOptions::default()).unwrap();
let truncated = &compressed[..compressed.len() - 1];
let result = decompress_data_capped(truncated, CompressionType::Brotli, 1 << 20);
assert!(
result.is_err(),
"truncated Brotli stream must be rejected, got {result:?}"
);
}
#[test]
#[cfg(feature = "zstd")]
fn zstd_round_trip() {
let compressed = compress_zstd(INPUT, CompressOptions::default()).unwrap();
let out = decompress_data_capped(&compressed, CompressionType::Zstd, 1 << 20).unwrap();
assert_eq!(out, INPUT);
}
#[test]
#[cfg(feature = "zstd")]
fn zstd_trailing_garbage_rejected() {
let mut compressed = compress_zstd(INPUT, CompressOptions::default()).unwrap();
compressed.extend_from_slice(&[0xde, 0xad, 0xbe, 0xef]);
let result = decompress_data_capped(&compressed, CompressionType::Zstd, 1 << 20);
assert!(
result.is_err(),
"trailing bytes after the Zstd frame must be rejected, got {result:?}"
);
}
#[test]
#[cfg(feature = "zstd")]
fn zstd_second_frame_rejected() {
let mut compressed = compress_zstd(INPUT, CompressOptions::default()).unwrap();
let empty_frame = compress_zstd(b"", CompressOptions::default()).unwrap();
compressed.extend_from_slice(&empty_frame);
let result = decompress_data_capped(&compressed, CompressionType::Zstd, 1 << 20);
assert!(
result.is_err(),
"a second Zstd frame after the first must be rejected, got {result:?}"
);
}
#[test]
#[cfg(feature = "snappy")]
fn snappy_round_trip() {
let compressed = compress_snappy(INPUT).unwrap();
let out = decompress_data_capped(&compressed, CompressionType::Snappy, 1 << 20).unwrap();
assert_eq!(out, INPUT);
}
#[test]
#[cfg(feature = "snappy")]
fn snappy_trailing_garbage_rejected() {
let mut compressed = compress_snappy(INPUT).unwrap();
compressed.extend_from_slice(&[0xde, 0xad, 0xbe, 0xef]);
let result = decompress_data_capped(&compressed, CompressionType::Snappy, 1 << 20);
assert!(
result.is_err(),
"trailing bytes after the Snappy data must be rejected, got {result:?}"
);
}
#[cfg(feature = "zstd")]
#[test]
fn zstd_large_window_log_roundtrips() {
let input = vec![7u8; 4096];
for window_log in [28u32, 30] {
let compressed = compress_zstd(
&input,
CompressOptions {
level: None,
window_log: Some(window_log),
},
)
.expect("compress");
let out =
decompress_data_capped(&compressed, CompressionType::Zstd, input.len() as u64)
.unwrap_or_else(|e| {
panic!("zstd window_log {window_log} stream failed to decompress: {e}")
});
assert_eq!(out, input);
}
}
#[cfg(feature = "brotli")]
#[test]
fn brotli_large_window_stream_decompresses() {
let input = vec![3u8; 4096];
let mut params = brotli::enc::BrotliEncoderParams::default();
params.lgwin = 28;
params.large_window = true;
let mut compressed = Vec::new();
brotli::BrotliCompress(&mut &input[..], &mut compressed, ¶ms).expect("compress");
let out = decompress_data_capped(&compressed, CompressionType::Brotli, input.len() as u64)
.expect("large-window brotli stream failed to decompress");
assert_eq!(out, input);
}
}