use crate::{AvroResult, Error, error::Details, types::Value};
use strum::{EnumIter, EnumString, IntoStaticStr};
#[derive(Clone, Copy, Eq, PartialEq, Debug)]
pub struct DeflateSettings {
compression_level: miniz_oxide::deflate::CompressionLevel,
}
impl DeflateSettings {
pub fn new(compression_level: miniz_oxide::deflate::CompressionLevel) -> Self {
DeflateSettings { compression_level }
}
pub fn compression_level(&self) -> u8 {
self.compression_level as u8
}
}
impl Default for DeflateSettings {
fn default() -> Self {
Self::new(miniz_oxide::deflate::CompressionLevel::DefaultCompression)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, EnumIter, EnumString, IntoStaticStr)]
#[strum(serialize_all = "kebab_case")]
pub enum Codec {
Null,
Deflate(DeflateSettings),
#[cfg(feature = "snappy")]
Snappy,
#[cfg(feature = "zstandard")]
Zstandard(zstandard::ZstandardSettings),
#[cfg(feature = "bzip")]
Bzip2(bzip::Bzip2Settings),
#[cfg(feature = "xz")]
Xz(xz::XzSettings),
}
impl From<Codec> for Value {
fn from(value: Codec) -> Self {
Self::Bytes(<&str>::from(value).as_bytes().to_vec())
}
}
impl Codec {
pub fn compress(self, stream: &mut Vec<u8>) -> AvroResult<()> {
match self {
Codec::Null => (),
Codec::Deflate(settings) => {
let compressed =
miniz_oxide::deflate::compress_to_vec(stream, settings.compression_level());
*stream = compressed;
}
#[cfg(feature = "snappy")]
Codec::Snappy => {
let mut encoded: Vec<u8> = vec![0; snap::raw::max_compress_len(stream.len())];
let compressed_size = snap::raw::Encoder::new()
.compress(&stream[..], &mut encoded[..])
.map_err(Details::SnappyCompress)?;
let mut hasher = crc32fast::Hasher::new();
hasher.update(&stream[..]);
let checksum = hasher.finalize();
let checksum_as_bytes = checksum.to_be_bytes();
let checksum_len = checksum_as_bytes.len();
encoded.truncate(compressed_size + checksum_len);
encoded[compressed_size..].copy_from_slice(&checksum_as_bytes);
*stream = encoded;
}
#[cfg(feature = "zstandard")]
Codec::Zstandard(settings) => {
use std::io::Write;
let mut encoder = zstd::Encoder::new(Vec::new(), settings.compression_level as i32)
.map_err(Details::ZstdCompress)?;
encoder.write_all(stream).map_err(Details::ZstdCompress)?;
*stream = encoder.finish().map_err(Details::ZstdCompress)?;
}
#[cfg(feature = "bzip")]
Codec::Bzip2(settings) => {
use bzip2::read::BzEncoder;
use std::io::Read;
let mut encoder = BzEncoder::new(&stream[..], settings.compression());
let mut buffer = Vec::new();
encoder
.read_to_end(&mut buffer)
.unwrap_or_else(|_| unreachable!("No I/O errors possible with Vec<u8>"));
*stream = buffer;
}
#[cfg(feature = "xz")]
Codec::Xz(settings) => {
use liblzma::read::XzEncoder;
use std::io::Read;
let mut encoder = XzEncoder::new(&stream[..], settings.compression_level as u32);
let mut buffer = Vec::new();
encoder
.read_to_end(&mut buffer)
.unwrap_or_else(|_| unreachable!("No I/O errors possible with Vec<u8>"));
*stream = buffer;
}
};
Ok(())
}
pub fn decompress(self, stream: &mut Vec<u8>) -> AvroResult<()> {
let max_bytes =
crate::util::max_allocation_bytes(crate::util::DEFAULT_MAX_ALLOCATION_BYTES);
*stream = match self {
Codec::Null => return Ok(()),
Codec::Deflate(_settings) => miniz_oxide::inflate::decompress_to_vec_with_limit(stream, max_bytes).map_err(|e| {
use std::io::ErrorKind;
use miniz_oxide::inflate::TINFLStatus;
let details = match e.status {
TINFLStatus::FailedCannotMakeProgress | TINFLStatus::NeedsMoreInput => Details::DeflateDecompress(ErrorKind::UnexpectedEof.into()),
TINFLStatus::Adler32Mismatch | TINFLStatus::Failed | TINFLStatus::BadParam => Details::DeflateDecompress(ErrorKind::InvalidData.into()),
TINFLStatus::Done => Details::DeflateDecompress(std::io::Error::other("Unexpected error: miniz_oxide reported an error with a success status. Please report this to avro-rs developers.")),
TINFLStatus::HasMoreOutput => Details::MemoryAllocation {
desired: None,
maximum: max_bytes,
},
other => Details::DeflateDecompress(std::io::Error::other(format!("Unexpected error: {other:?}")))
};
Error::new(details)
})?,
#[cfg(feature = "snappy")]
Codec::Snappy => {
let data_end = stream
.len()
.checked_sub(4)
.ok_or(Details::BadSnappyLength(stream.len()))?;
let decompressed_size = snap::raw::decompress_len(&stream[..data_end])
.map_err(Details::GetSnappyDecompressLen)?;
let decompressed_size = crate::util::safe_len(decompressed_size)?;
let mut decoded = vec![0; decompressed_size];
snap::raw::Decoder::new()
.decompress(&stream[..data_end], &mut decoded[..])
.map_err(Details::SnappyDecompress)?;
let mut last_four: [u8; 4] = [0; 4];
last_four.copy_from_slice(&stream[data_end..]);
let expected: u32 = u32::from_be_bytes(last_four);
let mut hasher = crc32fast::Hasher::new();
hasher.update(&decoded);
let actual = hasher.finalize();
if expected != actual {
return Err(Details::SnappyCrc32{expected, actual}.into());
}
decoded
}
#[cfg(feature = "zstandard")]
Codec::Zstandard(_settings) => {
use std::io::{BufReader, Read};
use zstd::zstd_safe;
let mut decoded = Vec::new();
let buffer_size = zstd_safe::DCtx::in_size();
let buffer = BufReader::with_capacity(buffer_size, &stream[..]);
let decoder = zstd::Decoder::new(buffer).map_err(Details::ZstdDecompress)?;
decoder
.take((max_bytes as u64).saturating_add(1))
.read_to_end(&mut decoded)
.map_err(Details::ZstdDecompress)?;
if decoded.len() > max_bytes {
return Err(Details::MemoryAllocation { desired: None, maximum: max_bytes }.into());
}
decoded
}
#[cfg(feature = "bzip")]
Codec::Bzip2(_) => {
use bzip2::read::BzDecoder;
use std::io::Read;
let mut decoded = Vec::new();
BzDecoder::new(&stream[..])
.take((max_bytes as u64).saturating_add(1))
.read_to_end(&mut decoded)
.map_err(Details::Bzip2Decompress)?;
if decoded.len() > max_bytes {
return Err(Details::MemoryAllocation { desired: None, maximum: max_bytes }.into());
}
decoded
}
#[cfg(feature = "xz")]
Codec::Xz(_) => {
use liblzma::read::XzDecoder;
use std::io::Read;
let mut decoded: Vec<u8> = Vec::new();
XzDecoder::new(&stream[..])
.take((max_bytes as u64).saturating_add(1))
.read_to_end(&mut decoded)
.map_err(Details::XzDecompress)?;
if decoded.len() > max_bytes {
return Err(Details::MemoryAllocation { desired: None, maximum: max_bytes }.into());
}
decoded
}
};
Ok(())
}
}
#[cfg(feature = "bzip")]
pub mod bzip {
use bzip2::Compression;
#[derive(Clone, Copy, Eq, PartialEq, Debug)]
pub struct Bzip2Settings {
pub compression_level: u8,
}
impl Bzip2Settings {
pub fn new(compression_level: u8) -> Self {
Self { compression_level }
}
pub(crate) fn compression(&self) -> Compression {
Compression::new(self.compression_level as u32)
}
}
impl Default for Bzip2Settings {
fn default() -> Self {
Bzip2Settings::new(Compression::best().level() as u8)
}
}
}
#[cfg(feature = "zstandard")]
pub mod zstandard {
#[derive(Clone, Copy, Eq, PartialEq, Debug)]
pub struct ZstandardSettings {
pub compression_level: u8,
}
impl ZstandardSettings {
pub fn new(compression_level: u8) -> Self {
Self { compression_level }
}
}
impl Default for ZstandardSettings {
fn default() -> Self {
Self::new(0)
}
}
}
#[cfg(feature = "xz")]
pub mod xz {
#[derive(Clone, Copy, Eq, PartialEq, Debug)]
pub struct XzSettings {
pub compression_level: u8,
}
impl XzSettings {
pub fn new(compression_level: u8) -> Self {
Self { compression_level }
}
}
impl Default for XzSettings {
fn default() -> Self {
XzSettings::new(9)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use apache_avro_test_helper::TestResult;
use miniz_oxide::deflate::CompressionLevel;
use pretty_assertions::{assert_eq, assert_ne};
const INPUT: &[u8] = b"theanswertolifetheuniverseandeverythingis42theanswertolifetheuniverseandeverythingis4theanswertolifetheuniverseandeverythingis2";
#[test]
fn null_compress_and_decompress() -> TestResult {
let codec = Codec::Null;
let mut stream = INPUT.to_vec();
codec.compress(&mut stream)?;
assert_eq!(INPUT, stream.as_slice());
codec.decompress(&mut stream)?;
assert_eq!(INPUT, stream.as_slice());
Ok(())
}
#[test]
fn deflate_compress_and_decompress() -> TestResult {
compress_and_decompress(Codec::Deflate(DeflateSettings::new(
CompressionLevel::BestCompression,
)))
}
#[cfg(feature = "snappy")]
#[test]
fn snappy_compress_and_decompress() -> TestResult {
compress_and_decompress(Codec::Snappy)
}
#[cfg(feature = "snappy")]
#[test]
fn snappy_decompress_short_block_errors_without_panicking() {
for len in 0..4usize {
let mut stream = vec![0u8; len];
let result = Codec::Snappy.decompress(&mut stream);
assert!(result.is_err(), "len={len} should error, got {result:?}");
}
}
#[cfg(feature = "zstandard")]
#[test]
fn zstd_compress_and_decompress() -> TestResult {
compress_and_decompress(Codec::Zstandard(zstandard::ZstandardSettings::default()))
}
#[cfg(feature = "bzip")]
#[test]
fn bzip_compress_and_decompress() -> TestResult {
compress_and_decompress(Codec::Bzip2(bzip::Bzip2Settings::default()))
}
#[cfg(feature = "xz")]
#[test]
fn xz_compress_and_decompress() -> TestResult {
compress_and_decompress(Codec::Xz(xz::XzSettings::default()))
}
fn compress_and_decompress(codec: Codec) -> TestResult {
let mut stream = INPUT.to_vec();
codec.compress(&mut stream)?;
assert_ne!(INPUT, stream.as_slice());
assert!(INPUT.len() > stream.len());
codec.decompress(&mut stream)?;
assert_eq!(INPUT, stream.as_slice());
Ok(())
}
#[test]
fn codec_to_str() {
assert_eq!(<&str>::from(Codec::Null), "null");
assert_eq!(
<&str>::from(Codec::Deflate(DeflateSettings::default())),
"deflate"
);
#[cfg(feature = "snappy")]
assert_eq!(<&str>::from(Codec::Snappy), "snappy");
#[cfg(feature = "zstandard")]
assert_eq!(
<&str>::from(Codec::Zstandard(zstandard::ZstandardSettings::default())),
"zstandard"
);
#[cfg(feature = "bzip")]
assert_eq!(
<&str>::from(Codec::Bzip2(bzip::Bzip2Settings::default())),
"bzip2"
);
#[cfg(feature = "xz")]
assert_eq!(<&str>::from(Codec::Xz(xz::XzSettings::default())), "xz");
}
#[test]
fn codec_from_str() {
use std::str::FromStr;
assert_eq!(Codec::from_str("null").unwrap(), Codec::Null);
assert_eq!(
Codec::from_str("deflate").unwrap(),
Codec::Deflate(DeflateSettings::default())
);
#[cfg(feature = "snappy")]
assert_eq!(Codec::from_str("snappy").unwrap(), Codec::Snappy);
#[cfg(feature = "zstandard")]
assert_eq!(
Codec::from_str("zstandard").unwrap(),
Codec::Zstandard(zstandard::ZstandardSettings::default())
);
#[cfg(feature = "bzip")]
assert_eq!(
Codec::from_str("bzip2").unwrap(),
Codec::Bzip2(bzip::Bzip2Settings::default())
);
#[cfg(feature = "xz")]
assert_eq!(
Codec::from_str("xz").unwrap(),
Codec::Xz(xz::XzSettings::default())
);
assert!(Codec::from_str("not a codec").is_err());
}
}