use std::io::{Cursor, Read, Write};
use flate2::{Compression, read::MultiGzDecoder, write::GzEncoder};
use lz4_flex::frame::{BlockMode, BlockSize, FrameDecoder, FrameEncoder, FrameInfo};
use snap::raw::{Decoder as SnappyDecoder, Encoder as SnappyEncoder, decompress_len};
use zstd::bulk::Compressor as ZstdCompressor;
use zstd::zstd_safe::CParameter;
use crate::native::{KafkaClientError, KafkaClientResult};
pub(crate) const MAX_DECOMPRESSED_RECORD_BATCH_BYTES: usize = 64 * 1024 * 1024;
const GZIP_ATTRIBUTE: i16 = 1;
const SNAPPY_ATTRIBUTE: i16 = 2;
const LZ4_ATTRIBUTE: i16 = 3;
const ZSTD_ATTRIBUTE: i16 = 4;
const ZSTD_WINDOW_LOG_MAX: u32 = 26;
const XERIAL_SNAPPY_MAGIC: &[u8; 8] = b"\x82SNAPPY\0";
const XERIAL_SNAPPY_HEADER_LEN: usize = 16;
const XERIAL_SNAPPY_VERSION: u32 = 1;
const XERIAL_SNAPPY_COMPATIBLE_VERSION: u32 = 1;
const XERIAL_SNAPPY_BLOCK_BYTES: usize = 32 * 1024;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub(crate) enum CompressionCodec {
#[default]
None,
Gzip,
Snappy,
Lz4,
Zstd,
}
impl CompressionCodec {
pub(crate) fn from_attribute(attribute: i16) -> KafkaClientResult<Self> {
match attribute {
0 => Ok(Self::None),
GZIP_ATTRIBUTE => Ok(Self::Gzip),
SNAPPY_ATTRIBUTE => Ok(Self::Snappy),
LZ4_ATTRIBUTE => Ok(Self::Lz4),
ZSTD_ATTRIBUTE => Ok(Self::Zstd),
other => Err(KafkaClientError::unsupported(format!(
"unknown Kafka record-batch compression codec {other}"
))),
}
}
pub(crate) fn attribute(self) -> i16 {
match self {
Self::None => 0,
Self::Gzip => GZIP_ATTRIBUTE,
Self::Snappy => SNAPPY_ATTRIBUTE,
Self::Lz4 => LZ4_ATTRIBUTE,
Self::Zstd => ZSTD_ATTRIBUTE,
}
}
pub(crate) fn name(self) -> &'static str {
match self {
Self::None => "none",
Self::Gzip => "gzip",
Self::Snappy => "snappy",
Self::Lz4 => "lz4",
Self::Zstd => "zstd",
}
}
}
pub(crate) struct ZstdEncoderContext {
compressor: ZstdCompressor<'static>,
}
impl ZstdEncoderContext {
pub(crate) fn new() -> KafkaClientResult<Self> {
let mut compressor = ZstdCompressor::new(zstd::DEFAULT_COMPRESSION_LEVEL)
.map_err(|error| compression_error(CompressionCodec::Zstd, error))?;
compressor
.set_parameter(CParameter::ChecksumFlag(false))
.map_err(|error| compression_error(CompressionCodec::Zstd, error))?;
compressor
.set_parameter(CParameter::ContentSizeFlag(false))
.map_err(|error| compression_error(CompressionCodec::Zstd, error))?;
Ok(Self { compressor })
}
}
pub(crate) fn compress_records(
codec: CompressionCodec,
records: Vec<u8>,
zstd_context: &mut ZstdEncoderContext,
) -> KafkaClientResult<Vec<u8>> {
match codec {
CompressionCodec::None => Ok(records),
CompressionCodec::Gzip => compress_gzip(&records),
CompressionCodec::Snappy => compress_snappy(&records),
CompressionCodec::Lz4 => compress_lz4(&records),
CompressionCodec::Zstd => compress_zstd(&records, zstd_context),
}
}
pub(crate) fn decompress_records(
codec: CompressionCodec,
records: &[u8],
) -> KafkaClientResult<Vec<u8>> {
decompress_records_with_limit(codec, records, MAX_DECOMPRESSED_RECORD_BATCH_BYTES)
}
fn compress_gzip(records: &[u8]) -> KafkaClientResult<Vec<u8>> {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder
.write_all(records)
.map_err(|error| compression_error(CompressionCodec::Gzip, error))?;
encoder
.finish()
.map_err(|error| compression_error(CompressionCodec::Gzip, error))
}
fn compress_snappy(records: &[u8]) -> KafkaClientResult<Vec<u8>> {
let mut output = Vec::with_capacity(records.len().saturating_add(XERIAL_SNAPPY_HEADER_LEN));
output.extend_from_slice(XERIAL_SNAPPY_MAGIC);
output.extend_from_slice(&XERIAL_SNAPPY_VERSION.to_be_bytes());
output.extend_from_slice(&XERIAL_SNAPPY_COMPATIBLE_VERSION.to_be_bytes());
let mut encoder = SnappyEncoder::new();
for block in records.chunks(XERIAL_SNAPPY_BLOCK_BYTES) {
let compressed = encoder
.compress_vec(block)
.map_err(|error| compression_error(CompressionCodec::Snappy, error))?;
let block_len = u32::try_from(compressed.len()).map_err(|_| {
KafkaClientError::decompression(
CompressionCodec::Snappy.name(),
"compressed xerial Snappy block exceeds u32",
)
})?;
output.extend_from_slice(&block_len.to_be_bytes());
output.extend_from_slice(&compressed);
}
Ok(output)
}
fn compress_lz4(records: &[u8]) -> KafkaClientResult<Vec<u8>> {
let frame = FrameInfo::new()
.block_size(BlockSize::Max64KB)
.block_mode(BlockMode::Independent);
let mut encoder = FrameEncoder::with_frame_info(frame, Vec::new());
encoder
.write_all(records)
.map_err(|error| compression_error(CompressionCodec::Lz4, error))?;
encoder
.finish()
.map_err(|error| compression_error(CompressionCodec::Lz4, error))
}
fn compress_zstd(records: &[u8], context: &mut ZstdEncoderContext) -> KafkaClientResult<Vec<u8>> {
context
.compressor
.compress(records)
.map_err(|error| compression_error(CompressionCodec::Zstd, error))
}
fn decompress_records_with_limit(
codec: CompressionCodec,
records: &[u8],
limit: usize,
) -> KafkaClientResult<Vec<u8>> {
match codec {
CompressionCodec::None => Ok(records.to_vec()),
CompressionCodec::Gzip => read_bounded(
codec,
MultiGzDecoder::new(Cursor::new(records)),
records.len(),
limit,
),
CompressionCodec::Snappy => decompress_snappy(records, limit),
CompressionCodec::Lz4 => read_bounded(
codec,
FrameDecoder::new(Cursor::new(records)),
records.len(),
limit,
),
CompressionCodec::Zstd => {
let declared_size = zstd::zstd_safe::get_frame_content_size(records)
.map_err(|error| compression_error(codec, error))?;
if declared_size.is_some_and(|size| size > limit as u64) {
return Err(KafkaClientError::decompression_limit(codec.name(), limit));
}
let mut decoder = zstd::stream::read::Decoder::new(Cursor::new(records))
.map_err(|error| compression_error(codec, error))?
.single_frame();
decoder
.window_log_max(ZSTD_WINDOW_LOG_MAX)
.map_err(|error| compression_error(codec, error))?;
let output = read_bounded(codec, &mut decoder, records.len(), limit)?;
let reader = decoder.finish();
let consumed = reader
.get_ref()
.position()
.saturating_sub(reader.buffer().len() as u64);
if consumed != records.len() as u64 {
return Err(KafkaClientError::decompression(
codec.name(),
"trailing bytes after the Zstd frame",
));
}
Ok(output)
}
}
}
fn decompress_snappy(records: &[u8], limit: usize) -> KafkaClientResult<Vec<u8>> {
let mut output = Vec::with_capacity(records.len().saturating_mul(4).min(limit));
let mut decoder = SnappyDecoder::new();
if !records.starts_with(XERIAL_SNAPPY_MAGIC) {
decompress_snappy_block(&mut decoder, records, &mut output, limit)?;
return Ok(output);
}
if records.len() < XERIAL_SNAPPY_HEADER_LEN {
return Err(KafkaClientError::decompression(
CompressionCodec::Snappy.name(),
"truncated xerial Snappy header",
));
}
let version = u32::from_be_bytes(records[8..12].try_into().expect("four-byte version"));
let compatible_version = u32::from_be_bytes(
records[12..16]
.try_into()
.expect("four-byte compatible version"),
);
if version < XERIAL_SNAPPY_VERSION || compatible_version > XERIAL_SNAPPY_COMPATIBLE_VERSION {
return Err(KafkaClientError::decompression(
CompressionCodec::Snappy.name(),
format!(
"unsupported xerial Snappy version {version} (compatible version {compatible_version})"
),
));
}
let mut position = XERIAL_SNAPPY_HEADER_LEN;
while position < records.len() {
if records.len() - position < 4 {
return Err(KafkaClientError::decompression(
CompressionCodec::Snappy.name(),
"truncated xerial Snappy block length",
));
}
let block_len = u32::from_be_bytes(
records[position..position + 4]
.try_into()
.expect("four-byte block length"),
) as usize;
position += 4;
let block_end = position.checked_add(block_len).ok_or_else(|| {
KafkaClientError::decompression(
CompressionCodec::Snappy.name(),
"xerial Snappy block length overflow",
)
})?;
if block_end > records.len() {
return Err(KafkaClientError::decompression(
CompressionCodec::Snappy.name(),
"truncated xerial Snappy block",
));
}
decompress_snappy_block(
&mut decoder,
&records[position..block_end],
&mut output,
limit,
)?;
position = block_end;
}
Ok(output)
}
fn decompress_snappy_block(
decoder: &mut SnappyDecoder,
block: &[u8],
output: &mut Vec<u8>,
limit: usize,
) -> KafkaClientResult<()> {
let expanded_len = decompress_len(block)
.map_err(|error| compression_error(CompressionCodec::Snappy, error))?;
if output.len().saturating_add(expanded_len) > limit {
return Err(KafkaClientError::decompression_limit(
CompressionCodec::Snappy.name(),
limit,
));
}
let output_start = output.len();
output.resize(output_start + expanded_len, 0);
let decoded_len = decoder
.decompress(block, &mut output[output_start..])
.map_err(|error| compression_error(CompressionCodec::Snappy, error))?;
if decoded_len != expanded_len {
return Err(KafkaClientError::decompression(
CompressionCodec::Snappy.name(),
"Snappy block expanded to an unexpected length",
));
}
Ok(())
}
fn read_bounded(
codec: CompressionCodec,
mut reader: impl Read,
compressed_len: usize,
limit: usize,
) -> KafkaClientResult<Vec<u8>> {
let initial_capacity = compressed_len.saturating_mul(4).min(limit);
let mut output = Vec::with_capacity(initial_capacity);
let mut chunk = [0_u8; 16 * 1024];
loop {
let read = reader
.read(&mut chunk)
.map_err(|error| compression_error(codec, error))?;
if read == 0 {
return Ok(output);
}
if output.len().saturating_add(read) > limit {
return Err(KafkaClientError::decompression_limit(codec.name(), limit));
}
output.extend_from_slice(&chunk[..read]);
}
}
fn compression_error(codec: CompressionCodec, error: impl std::fmt::Display) -> KafkaClientError {
KafkaClientError::decompression(codec.name(), error.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
fn compress(codec: CompressionCodec, records: Vec<u8>) -> KafkaClientResult<Vec<u8>> {
let mut zstd_context = ZstdEncoderContext::new().expect("zstd encoder context");
compress_records(codec, records, &mut zstd_context)
}
#[test]
fn lz4_and_zstd_round_trip_standard_frames() {
let records = b"record-data-record-data-record-data";
for codec in [CompressionCodec::Lz4, CompressionCodec::Zstd] {
let encoded = compress(codec, records.to_vec()).expect("compress records");
let decoded = decompress_records(codec, &encoded).expect("decompress records");
assert_eq!(decoded, records, "codec={}", codec.name());
}
}
#[test]
fn gzip_round_trips_standard_frames_including_empty_input() {
for records in [b"".as_slice(), b"record-data-record-data".as_slice()] {
let encoded = compress(CompressionCodec::Gzip, records.to_vec()).expect("gzip");
assert_eq!(&encoded[..2], &[0x1f, 0x8b]);
let decoded =
decompress_records(CompressionCodec::Gzip, &encoded).expect("decompress gzip");
assert_eq!(decoded, records);
}
}
#[test]
fn snappy_round_trips_empty_single_and_multi_block_xerial_frames() {
let cases = [
Vec::new(),
b"record-data-record-data".to_vec(),
(0..XERIAL_SNAPPY_BLOCK_BYTES * 2 + 17)
.map(|index| (index % 251) as u8)
.collect(),
];
for (index, records) in cases.into_iter().enumerate() {
let encoded = compress(CompressionCodec::Snappy, records.clone()).expect("snappy");
assert_eq!(&encoded[..8], XERIAL_SNAPPY_MAGIC);
assert_eq!(
u32::from_be_bytes(encoded[8..12].try_into().expect("version")),
XERIAL_SNAPPY_VERSION
);
assert_eq!(
snappy_block_count(&encoded),
records.len().div_ceil(XERIAL_SNAPPY_BLOCK_BYTES),
"case={index}"
);
let decoded =
decompress_records(CompressionCodec::Snappy, &encoded).expect("decompress snappy");
assert_eq!(decoded, records, "case={index}");
}
}
#[test]
fn snappy_decodes_bounded_raw_librdkafka_frames() {
let records = vec![b'x'; 4_096];
let encoded = SnappyEncoder::new()
.compress_vec(&records)
.expect("raw Snappy compression");
assert_eq!(
decompress_records_with_limit(CompressionCodec::Snappy, &encoded, records.len())
.expect("raw Snappy decompression"),
records
);
let error = decompress_records_with_limit(CompressionCodec::Snappy, &encoded, 1_024)
.expect_err("raw Snappy expansion limit must fail");
assert!(matches!(
error,
KafkaClientError::DecompressionLimitExceeded {
codec: _,
limit: 1_024
}
));
}
fn snappy_block_count(encoded: &[u8]) -> usize {
let mut blocks = 0;
let mut position = XERIAL_SNAPPY_HEADER_LEN;
while position < encoded.len() {
let block_len = u32::from_be_bytes(
encoded[position..position + 4]
.try_into()
.expect("block length"),
) as usize;
position += 4 + block_len;
blocks += 1;
}
assert_eq!(position, encoded.len());
blocks
}
#[test]
fn zstd_context_reuse_produces_independent_valid_frames() {
let mut zstd_context = ZstdEncoderContext::new().expect("zstd encoder context");
let batches: [&[u8]; 3] = [b"first-batch", b"a-different-second-batch", b"third"];
for batch in batches {
let encoded =
compress_records(CompressionCodec::Zstd, batch.to_vec(), &mut zstd_context)
.expect("compress with reused context");
let decoded =
decompress_records(CompressionCodec::Zstd, &encoded).expect("decompress records");
assert_eq!(decoded, batch);
}
}
#[test]
fn decompression_expansion_limit_is_typed() {
let records = vec![b'x'; 4_096];
for codec in [
CompressionCodec::Gzip,
CompressionCodec::Snappy,
CompressionCodec::Lz4,
CompressionCodec::Zstd,
] {
let encoded = compress(codec, records.clone()).expect("compress records");
let error = decompress_records_with_limit(codec, &encoded, 1_024)
.expect_err("expansion limit must fail");
assert!(matches!(
error,
KafkaClientError::DecompressionLimitExceeded {
codec: _,
limit: 1_024
}
));
}
}
#[test]
fn malformed_frames_are_typed_and_never_panic() {
for codec in [CompressionCodec::Lz4, CompressionCodec::Zstd] {
let encoded = compress(codec, b"record-data".to_vec()).expect("compress records");
let malformed = [&encoded[..encoded.len() / 2], &[0xa5; 16][..]];
for bytes in malformed {
let error = decompress_records_with_limit(codec, bytes, 1_024)
.expect_err("malformed frame must fail");
assert!(matches!(error, KafkaClientError::Decompression { .. }));
}
}
let error = decompress_records_with_limit(CompressionCodec::Gzip, &[0xa5; 16], 1_024)
.expect_err("malformed gzip frame must fail");
assert!(matches!(error, KafkaClientError::Decompression { .. }));
let error = decompress_records_with_limit(CompressionCodec::Snappy, &[0xa5; 16], 1_024)
.expect_err("malformed Snappy frame must fail");
assert!(matches!(
error,
KafkaClientError::Decompression { .. }
| KafkaClientError::DecompressionLimitExceeded { .. }
));
let snappy = compress(CompressionCodec::Snappy, b"record-data".to_vec())
.expect("compress snappy records");
let error = decompress_records_with_limit(
CompressionCodec::Snappy,
&snappy[..snappy.len() - 1],
1_024,
)
.expect_err("truncated xerial Snappy block must fail");
assert!(matches!(error, KafkaClientError::Decompression { .. }));
let mut zstd_with_garbage = compress(CompressionCodec::Zstd, b"record-data".to_vec())
.expect("compress zstd records");
zstd_with_garbage.extend_from_slice(b"garbage");
let error =
decompress_records_with_limit(CompressionCodec::Zstd, &zstd_with_garbage, 1_024)
.expect_err("trailing garbage must fail");
assert!(matches!(error, KafkaClientError::Decompression { .. }));
}
#[test]
fn wrong_lz4_content_size_is_typed() {
let frame = FrameInfo::new()
.block_size(BlockSize::Max64KB)
.block_mode(BlockMode::Independent)
.content_size(Some(128));
let mut encoder = FrameEncoder::with_frame_info(frame, Vec::new());
encoder.write_all(b"short").expect("write lz4 input");
encoder
.try_finish()
.expect_err("encoder must reject the wrong declared size");
let mut encoded = encoder.get_ref().clone();
encoded.extend_from_slice(&[0_u8; 4]);
let error = decompress_records_with_limit(CompressionCodec::Lz4, &encoded, 1_024)
.expect_err("decoder must reject the wrong declared size");
assert!(matches!(error, KafkaClientError::Decompression { .. }));
}
}