use bytes::{BufMut, BytesMut};
use super::{Compression, CompressionError, CompressionErrorKind, Result};
const MAGIC: u32 = 0x184D_2204;
const HEADER_LEN: usize = 7;
const SIZE_PREFIX_LEN: usize = 4;
const MAX_BLOCK_SIZE: usize = 64 * 1024;
const BD_BLOCK_SIZE_64KB: u8 = 4;
const FLG: u8 = 0b0110_0000;
const BD: u8 = BD_BLOCK_SIZE_64KB << 4;
const INCOMPRESSIBLE_BIT: u32 = 0x8000_0000;
const DEFAULT_LEVEL: i32 = 1;
const XXH32_PRIME_1: u32 = 0x9E37_79B1;
const XXH32_PRIME_2: u32 = 0x85EB_CA77;
const XXH32_PRIME_3: u32 = 0xC2B2_AE3D;
const XXH32_PRIME_4: u32 = 0x27D4_EB2F;
const XXH32_PRIME_5: u32 = 0x1656_67B1;
pub fn compress(input: &[u8]) -> Result<Vec<u8>> {
compress_with_level(input, None)
}
pub fn compress_with_level(input: &[u8], level: Option<i32>) -> Result<Vec<u8>> {
let estimated = HEADER_LEN
.saturating_add(input.len())
.saturating_add(input.len() >> 6)
.saturating_add(SIZE_PREFIX_LEN);
let mut out = BytesMut::with_capacity(estimated);
write_frame_header(&mut out);
let mut offset: usize = 0;
while offset < input.len() {
let remaining = input.len().saturating_sub(offset);
let block_size = remaining.min(MAX_BLOCK_SIZE);
let end = offset.saturating_add(block_size);
let block = input
.get(offset..end)
.ok_or_else(|| encode_err("block slice out of bounds".to_owned()))?;
write_block(&mut out, block, level)?;
offset = end;
}
out.put_u32_le(0);
Ok(out.to_vec())
}
pub fn decompress(input: &[u8]) -> Result<Vec<u8>> {
decompress_bounded(input, super::MAX_DECOMPRESSED_LEN)
}
pub fn decompress_bounded(input: &[u8], max_len: usize) -> Result<Vec<u8>> {
let mut offset = read_frame_header(input)?;
let mut output: Vec<u8> = Vec::new();
loop {
let next = offset.saturating_add(SIZE_PREFIX_LEN);
let size_bytes = input
.get(offset..next)
.ok_or_else(|| decode_err("incomplete block size prefix".to_owned()))?;
let raw_arr: [u8; 4] = match size_bytes {
&[a, b, c, d] => [a, b, c, d],
_ => return Err(decode_err("block size prefix wrong length".to_owned())),
};
let raw = u32::from_le_bytes(raw_arr);
offset = next;
if raw == 0 {
break;
}
let is_compressed = (raw & INCOMPRESSIBLE_BIT) == 0;
let block_len_u32 = raw & !INCOMPRESSIBLE_BIT;
let block_len = usize::try_from(block_len_u32)
.map_err(|_| decode_err("block length overflows usize".to_owned()))?;
let block_end = offset
.checked_add(block_len)
.ok_or_else(|| decode_err("block end offset overflows".to_owned()))?;
let block = input.get(offset..block_end).ok_or_else(|| {
decode_err(format!(
"incomplete block: expected {block_len}, got {}",
input.len().saturating_sub(offset)
))
})?;
offset = block_end;
let decompressed;
let bytes: &[u8] = if is_compressed {
decompressed = decompress_block(block)?;
&decompressed
} else {
block
};
if output.len().saturating_add(bytes.len()) > max_len {
return Err(CompressionError::new(
Compression::Lz4,
CompressionErrorKind::DecompressedTooLarge { limit: max_len },
));
}
output.extend_from_slice(bytes);
}
Ok(output)
}
fn write_frame_header(out: &mut BytesMut) {
out.put_u32_le(MAGIC);
out.put_u8(FLG);
out.put_u8(BD);
out.put_u8(header_checksum_byte());
}
fn header_checksum_byte() -> u8 {
let hash = xxh32(&[FLG, BD], 0);
((hash >> 8) & 0xFF) as u8
}
fn read_frame_header(input: &[u8]) -> Result<usize> {
let header = input
.get(..HEADER_LEN)
.ok_or_else(|| decode_err("incomplete frame header".to_owned()))?;
let magic_arr: [u8; 4] = match header.get(..4) {
Some(&[a, b, c, d]) => [a, b, c, d],
_ => return Err(decode_err("frame header missing magic".to_owned())),
};
let magic = u32::from_le_bytes(magic_arr);
if magic != MAGIC {
return Err(decode_err(format!(
"invalid LZ4 frame magic: expected {MAGIC:#010x}, got {magic:#010x}"
)));
}
Ok(HEADER_LEN)
}
fn write_block(out: &mut BytesMut, block: &[u8], level: Option<i32>) -> Result<()> {
let compressed = compress_block(block, level)?;
if compressed.len() < block.len() {
let len = u32::try_from(compressed.len())
.map_err(|_| encode_err("compressed block exceeds u32".to_owned()))?;
out.put_u32_le(len);
out.extend_from_slice(&compressed);
} else {
let len = u32::try_from(block.len())
.map_err(|_| encode_err("uncompressed block exceeds u32".to_owned()))?;
out.put_u32_le(len | INCOMPRESSIBLE_BIT);
out.extend_from_slice(block);
}
Ok(())
}
#[cfg(feature = "lz4-hc")]
fn compress_block(block: &[u8], level: Option<i32>) -> Result<Vec<u8>> {
use lz4::block::CompressionMode;
let raw = level.unwrap_or(DEFAULT_LEVEL);
let mode = if raw <= 2 {
CompressionMode::FAST(raw.max(1))
} else {
CompressionMode::HIGHCOMPRESSION(raw.min(12))
};
lz4::block::compress(block, Some(mode), false)
.map_err(|e: std::io::Error| encode_err(e.to_string()))
}
#[cfg(all(feature = "lz4", not(feature = "lz4-hc")))]
#[expect(
clippy::unnecessary_wraps,
reason = "signature mirrors the fallible C-FFI HC backend"
)]
fn compress_block(block: &[u8], _level: Option<i32>) -> Result<Vec<u8>> {
let _ = DEFAULT_LEVEL;
Ok(lz4_flex::block::compress(block))
}
#[cfg(feature = "lz4-hc")]
fn decompress_block(block: &[u8]) -> Result<Vec<u8>> {
let max_size = i32::try_from(MAX_BLOCK_SIZE)
.map_err(|_| decode_err("MAX_BLOCK_SIZE overflows i32".to_owned()))?;
lz4::block::decompress(block, Some(max_size))
.map_err(|e: std::io::Error| decode_err(e.to_string()))
}
#[cfg(all(feature = "lz4", not(feature = "lz4-hc")))]
fn decompress_block(block: &[u8]) -> Result<Vec<u8>> {
lz4_flex::block::decompress(block, MAX_BLOCK_SIZE).map_err(|e| decode_err(e.to_string()))
}
fn xxh32(input: &[u8], seed: u32) -> u32 {
let len_u32 = u32::try_from(input.len()).unwrap_or(u32::MAX);
let mut hash = seed.wrapping_add(XXH32_PRIME_5).wrapping_add(len_u32);
let mut chunks = input.chunks_exact(4);
for chunk in chunks.by_ref() {
let word = match chunk {
&[b0, b1, b2, b3] => u32::from_le_bytes([b0, b1, b2, b3]),
_ => return hash,
};
hash = hash
.wrapping_add(word.wrapping_mul(XXH32_PRIME_3))
.rotate_left(17)
.wrapping_mul(XXH32_PRIME_4);
}
for &byte in chunks.remainder() {
hash = hash
.wrapping_add(u32::from(byte).wrapping_mul(XXH32_PRIME_5))
.rotate_left(11)
.wrapping_mul(XXH32_PRIME_1);
}
hash ^= hash >> 15;
hash = hash.wrapping_mul(XXH32_PRIME_2);
hash ^= hash >> 13;
hash = hash.wrapping_mul(XXH32_PRIME_3);
hash ^= hash >> 16;
hash
}
const fn encode_err(message: String) -> CompressionError {
CompressionError::new(
Compression::Lz4,
CompressionErrorKind::EncodeFailed { message },
)
}
const fn decode_err(message: String) -> CompressionError {
CompressionError::new(
Compression::Lz4,
CompressionErrorKind::DecodeFailed { message },
)
}
#[cfg(test)]
mod tests {
use super::{
super::CompressionErrorKind, BD, FLG, MAGIC, XXH32_PRIME_5, compress, decompress,
decompress_bounded, xxh32,
};
#[test]
fn decompress_bounded_rejects_a_decompression_bomb() {
let payload = vec![b'x'; 256 * 1024];
let compressed = compress(&payload).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![b'x'; 256 * 1024];
let compressed = compress(&payload).unwrap();
assert_eq!(
decompress_bounded(&compressed, 256 * 1024).unwrap(),
payload
);
}
#[test]
fn frame_starts_with_kafka_lz4_magic() {
let compressed = compress(b"Hello, Kafka!").unwrap();
assert!(compressed.len() >= 7, "compressed output too short");
let magic =
u32::from_le_bytes([compressed[0], compressed[1], compressed[2], compressed[3]]);
assert_eq!(magic, MAGIC, "unexpected LZ4 frame magic");
assert_eq!(compressed[4], FLG);
assert_eq!(compressed[5], BD);
}
#[test]
fn roundtrip_short_payload() {
let payload = b"Hello, Kafka! This is a test message.";
let compressed = compress(payload).unwrap();
let decompressed = decompress(&compressed).unwrap();
assert_eq!(payload.as_slice(), decompressed.as_slice());
}
#[test]
fn roundtrip_empty() {
let compressed = compress(b"").unwrap();
let decompressed = decompress(&compressed).unwrap();
assert!(decompressed.is_empty());
}
#[test]
fn roundtrip_multi_block() {
let payload = vec![b'x'; 256 * 1024];
let compressed = compress(&payload).unwrap();
let decompressed = decompress(&compressed).unwrap();
assert_eq!(payload.as_slice(), decompressed.as_slice());
assert!(compressed.len() < payload.len() / 10);
}
#[test]
fn xxh32_empty_matches_reference() {
let h = xxh32(b"", 0);
assert_eq!(h, 0x02CC_5D05);
}
#[test]
fn xxh32_short_input_matches_reference() {
let h = xxh32(b"a", 0);
assert_eq!(h, 0x550D_7456);
}
#[test]
fn xxh32_uses_prime5_in_tail() {
let byte: u8 = 0xAB;
let mut expected = 0u32.wrapping_add(XXH32_PRIME_5).wrapping_add(1);
expected = expected
.wrapping_add(u32::from(byte).wrapping_mul(XXH32_PRIME_5))
.rotate_left(11)
.wrapping_mul(super::XXH32_PRIME_1);
expected ^= expected >> 15;
expected = expected.wrapping_mul(super::XXH32_PRIME_2);
expected ^= expected >> 13;
expected = expected.wrapping_mul(super::XXH32_PRIME_3);
expected ^= expected >> 16;
assert_eq!(xxh32(&[byte], 0), expected);
}
}