use omnizip_filters::Filter;
use crate::codec::{compress, decompress, Codec, CoreError};
pub fn filter_then_compress<F: Filter>(
plaintext: &[u8],
filter: &F,
inner_codec: u8,
) -> Result<Vec<u8>, CoreError> {
let filtered = filter.encode(plaintext);
let inner = compress(inner_codec, &filtered)?;
let filtered_len = u32::try_from(filtered.len()).unwrap_or(u32::MAX);
let mut out = Vec::with_capacity(4 + inner.len());
out.extend_from_slice(&filtered_len.to_le_bytes());
out.extend_from_slice(&inner);
Ok(out)
}
pub fn decompress_then_filter<F: Filter>(
compressed: &[u8],
filter: &F,
inner_codec: u8,
error_label: &str,
) -> Result<Vec<u8>, CoreError> {
if compressed.len() < 4 {
return Err(CoreError::Corrupt {
reason: format!("{error_label}: input too short for length prefix"),
});
}
let filtered_len =
u32::from_le_bytes([compressed[0], compressed[1], compressed[2], compressed[3]]);
let inner_bytes = &compressed[4..];
let filtered = decompress(inner_codec, inner_bytes, filtered_len)?;
Ok(filter.decode(&filtered))
}
pub struct FilterCodecComposite<F: Filter> {
filter: F,
inner: u8,
id: u8,
name: &'static str,
min_compress_size: usize,
}
impl<F: Filter> FilterCodecComposite<F> {
#[must_use]
pub const fn new(
filter: F,
inner: u8,
id: u8,
name: &'static str,
min_compress_size: usize,
) -> Self {
Self {
filter,
inner,
id,
name,
min_compress_size,
}
}
}
impl<F: Filter> Codec for FilterCodecComposite<F> {
fn id(&self) -> u8 {
self.id
}
fn name(&self) -> &'static str {
self.name
}
fn min_compress_size(&self) -> usize {
self.min_compress_size
}
fn compress(&self, plaintext: &[u8]) -> Result<Vec<u8>, CoreError> {
filter_then_compress(plaintext, &self.filter, self.inner)
}
fn decompress(&self, compressed: &[u8], _expected_len: u32) -> Result<Vec<u8>, CoreError> {
decompress_then_filter(compressed, &self.filter, self.inner, self.name)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::codec::Codec;
use omnizip_filters::shuffle::ByteShuffle;
#[test]
fn round_trips_through_any_filter() {
let plaintext = b"hello world. hello world. hello world.".repeat(10);
let filter = ByteShuffle::new(4);
let compressed =
filter_then_compress(&plaintext, &filter, crate::codec::CODEC_LZ4).expect("compress");
let recovered =
decompress_then_filter(&compressed, &filter, crate::codec::CODEC_LZ4, "test")
.expect("decompress");
assert_eq!(recovered, plaintext);
}
#[test]
fn all_registered_composites_round_trip_through_the_pipeline() {
let composite_ids = [
crate::codec::CODEC_BLOSC2_SHUFFLE_LZ4, crate::codec::CODEC_SHUFFLE_ZSTD, crate::codec::CODEC_BITSHUFFLE_LZ4, crate::codec::CODEC_BCJ_X86_LZ4, crate::codec::CODEC_BCJ_X86_ZSTD, crate::codec::CODEC_BCJ_ARM64_LZ4, crate::codec::CODEC_BCJ_ARM64_ZSTD, ];
let registry = crate::codec::default_registry();
let mut payload = Vec::with_capacity(8 * 1024);
let mut v = 0.25f32;
for i in 0..2048 {
v = (v + i as f32 * 0.001).sin();
payload.extend_from_slice(&v.to_le_bytes());
}
for id in composite_ids {
let compressed = registry
.compress(id, &payload)
.unwrap_or_else(|e| panic!("0x{id:02X} compress: {e}"));
let recovered = registry
.decompress(id, &compressed, payload.len() as u32)
.unwrap_or_else(|e| panic!("0x{id:02X} decompress: {e}"));
assert_eq!(recovered, payload, "0x{id:02X} round-trip");
}
}
#[test]
fn composite_wire_format_is_the_shared_pipeline() {
let codec = FilterCodecComposite::new(
ByteShuffle::new(4),
crate::codec::CODEC_LZ4,
0x0A,
"shuffle+lz4-test",
512,
);
let payload: Vec<u8> = (0..4096u32).flat_map(|i| i.to_le_bytes()).collect();
let via_codec = codec.compress(&payload).expect("codec path");
let via_pipeline =
filter_then_compress(&payload, &ByteShuffle::new(4), crate::codec::CODEC_LZ4)
.expect("pipeline path");
assert_eq!(via_codec, via_pipeline, "codec output == pipeline output");
}
}