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 = "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 = "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> {
use std::io::Read as _;
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")]
{
const MAX_DECOMPRESS_PREALLOC: u64 = 1 << 24; let mut out = Vec::with_capacity(max_len.min(MAX_DECOMPRESS_PREALLOC) as usize);
let reader = brotli::Decompressor::new(data, 4096);
reader
.take(max_len.saturating_add(1))
.read_to_end(&mut out)
.map_err(|e| {
RiegeliError::MalformedData(format!("brotli decompress error: {e}").into())
})?;
check(out)
}
#[cfg(not(feature = "brotli"))]
{
Err(RiegeliError::UnsupportedCompression(
CompressionType::Brotli as u8,
))
}
}
CompressionType::Zstd => {
#[cfg(feature = "zstd")]
{
const MAX_DECOMPRESS_PREALLOC: u64 = 1 << 24; let mut out = Vec::with_capacity(max_len.min(MAX_DECOMPRESS_PREALLOC) as usize);
let reader = zstd::stream::read::Decoder::new(data).map_err(|e| {
RiegeliError::MalformedData(format!("zstd decoder init: {e}").into())
})?;
reader
.take(max_len.saturating_add(1))
.read_to_end(&mut out)
.map_err(|e| {
RiegeliError::MalformedData(format!("zstd decompress error: {e}").into())
})?;
check(out)
}
#[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,
))
}
}
}
}