use crate::codecs::{
BrotliCodec, DeflateCodec, DeltaCodec, DictionaryCodec, Lz4Codec, RleCodec, SnappyCodec,
ZstdCodec,
};
use crate::error::CompressionError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PipelineStage {
Shuffle {
typesize: u8,
},
Delta,
Rle,
Lz4,
Zstd,
Snappy,
Brotli,
Deflate,
Dictionary,
}
impl PipelineStage {
fn id(self) -> u8 {
match self {
PipelineStage::Shuffle { .. } => 1,
PipelineStage::Delta => 2,
PipelineStage::Rle => 3,
PipelineStage::Lz4 => 4,
PipelineStage::Zstd => 5,
PipelineStage::Snappy => 6,
PipelineStage::Brotli => 7,
PipelineStage::Deflate => 8,
PipelineStage::Dictionary => 9,
}
}
fn typesize_byte(self) -> u8 {
match self {
PipelineStage::Shuffle { typesize } => typesize,
_ => 0,
}
}
fn from_id(id: u8, typesize: u8) -> Result<PipelineStage, CompressionError> {
Ok(match id {
1 => PipelineStage::Shuffle { typesize },
2 => PipelineStage::Delta,
3 => PipelineStage::Rle,
4 => PipelineStage::Lz4,
5 => PipelineStage::Zstd,
6 => PipelineStage::Snappy,
7 => PipelineStage::Brotli,
8 => PipelineStage::Deflate,
9 => PipelineStage::Dictionary,
other => {
return Err(CompressionError::InvalidMetadata(format!(
"unknown pipeline stage id {other} in frame header"
)));
}
})
}
}
const PIPELINE_MAGIC: [u8; 4] = [b'O', b'X', b'P', b'L'];
const PIPELINE_VERSION: u8 = 1;
const HEADER_PREFIX_LEN: usize = 6;
const STAGE_DESCRIPTOR_LEN: usize = 2;
#[derive(Debug, Clone, Default)]
pub struct CodecPipeline {
stages: Vec<PipelineStage>,
}
impl CodecPipeline {
pub fn new() -> Self {
Self::default()
}
pub fn push(mut self, stage: PipelineStage) -> Self {
self.stages.push(stage);
self
}
pub fn stages(&self) -> &[PipelineStage] {
&self.stages
}
pub fn compress(&self, input: &[u8]) -> Result<Vec<u8>, CompressionError> {
let mut output = encode_header(&self.stages);
let mut buffer = input.to_vec();
for &stage in &self.stages {
buffer = apply_stage_forward(stage, &buffer)?;
}
output.extend_from_slice(&buffer);
Ok(output)
}
pub fn decompress(&self, input: &[u8]) -> Result<Vec<u8>, CompressionError> {
let (stages, payload) = decode_header(input)?;
decompress_payload(&stages, payload)
}
pub fn decompress_self_describing(input: &[u8]) -> Result<Vec<u8>, CompressionError> {
let (stages, payload) = decode_header(input)?;
decompress_payload(&stages, payload)
}
}
fn encode_header(stages: &[PipelineStage]) -> Vec<u8> {
let num_stages = stages.len();
let mut header = Vec::with_capacity(HEADER_PREFIX_LEN + STAGE_DESCRIPTOR_LEN * num_stages);
header.extend_from_slice(&PIPELINE_MAGIC);
header.push(PIPELINE_VERSION);
header.push(num_stages as u8);
for &stage in stages {
header.push(stage.id());
header.push(stage.typesize_byte());
}
header
}
fn decode_header(input: &[u8]) -> Result<(Vec<PipelineStage>, &[u8]), CompressionError> {
if input.len() < HEADER_PREFIX_LEN {
return Err(CompressionError::InvalidMetadata(format!(
"pipeline frame truncated: need at least {HEADER_PREFIX_LEN} header bytes, got {}",
input.len()
)));
}
if input[0..4] != PIPELINE_MAGIC {
return Err(CompressionError::InvalidMetadata(
"pipeline frame has bad magic (expected 'OXPL')".to_string(),
));
}
let version = input[4];
if version != PIPELINE_VERSION {
return Err(CompressionError::InvalidMetadata(format!(
"unsupported pipeline frame version {version} (expected {PIPELINE_VERSION})"
)));
}
let num_stages = input[5] as usize;
let header_len = HEADER_PREFIX_LEN + STAGE_DESCRIPTOR_LEN * num_stages;
if input.len() < header_len {
return Err(CompressionError::InvalidMetadata(format!(
"pipeline frame truncated: header declares {num_stages} stage(s) \
needing {header_len} bytes, got {}",
input.len()
)));
}
let mut stages = Vec::with_capacity(num_stages);
for stage_index in 0..num_stages {
let offset = HEADER_PREFIX_LEN + STAGE_DESCRIPTOR_LEN * stage_index;
let id = input[offset];
let typesize = input[offset + 1];
stages.push(PipelineStage::from_id(id, typesize)?);
}
Ok((stages, &input[header_len..]))
}
fn decompress_payload(
stages: &[PipelineStage],
payload: &[u8],
) -> Result<Vec<u8>, CompressionError> {
let mut buffer = payload.to_vec();
for &stage in stages.iter().rev() {
buffer = apply_stage_inverse(stage, &buffer)?;
}
Ok(buffer)
}
fn byte_shuffle(data: &[u8], typesize: usize) -> Vec<u8> {
if typesize <= 1 || data.len() < typesize {
return data.to_vec();
}
let element_count = data.len() / typesize;
let shuffled_len = element_count * typesize;
let mut output = Vec::with_capacity(data.len());
for byte_position in 0..typesize {
for element_index in 0..element_count {
output.push(data[element_index * typesize + byte_position]);
}
}
output.extend_from_slice(&data[shuffled_len..]);
output
}
fn byte_unshuffle(data: &[u8], typesize: usize) -> Vec<u8> {
if typesize <= 1 || data.len() < typesize {
return data.to_vec();
}
let element_count = data.len() / typesize;
let shuffled_len = element_count * typesize;
let mut output = vec![0u8; data.len()];
for byte_position in 0..typesize {
for element_index in 0..element_count {
output[element_index * typesize + byte_position] =
data[byte_position * element_count + element_index];
}
}
output[shuffled_len..].copy_from_slice(&data[shuffled_len..]);
output
}
fn apply_stage_forward(stage: PipelineStage, data: &[u8]) -> Result<Vec<u8>, CompressionError> {
match stage {
PipelineStage::Shuffle { typesize } => Ok(byte_shuffle(data, typesize as usize)),
PipelineStage::Delta => DeltaCodec::default().compress(data),
PipelineStage::Rle => RleCodec::default().compress(data),
PipelineStage::Lz4 => Lz4Codec::default().compress(data),
PipelineStage::Zstd => ZstdCodec::default().compress(data),
PipelineStage::Snappy => SnappyCodec::default().compress(data),
PipelineStage::Brotli => BrotliCodec::default().compress(data),
PipelineStage::Deflate => DeflateCodec::default().compress(data),
PipelineStage::Dictionary => DictionaryCodec::default().compress(data),
}
}
fn apply_stage_inverse(stage: PipelineStage, data: &[u8]) -> Result<Vec<u8>, CompressionError> {
match stage {
PipelineStage::Shuffle { typesize } => Ok(byte_unshuffle(data, typesize as usize)),
PipelineStage::Delta => DeltaCodec::default().decompress(data),
PipelineStage::Rle => RleCodec::default().decompress(data),
PipelineStage::Lz4 => Lz4Codec::default().decompress(data, None),
PipelineStage::Zstd => ZstdCodec::default().decompress(data, None),
PipelineStage::Snappy => SnappyCodec::default().decompress(data),
PipelineStage::Brotli => BrotliCodec::default().decompress(data),
PipelineStage::Deflate => DeflateCodec::default().decompress(data),
PipelineStage::Dictionary => DictionaryCodec::default().decompress(data),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stage_id_round_trips_through_from_id() {
let cases = [
PipelineStage::Shuffle { typesize: 4 },
PipelineStage::Delta,
PipelineStage::Rle,
PipelineStage::Lz4,
PipelineStage::Zstd,
PipelineStage::Snappy,
PipelineStage::Brotli,
PipelineStage::Deflate,
PipelineStage::Dictionary,
];
for stage in cases {
let restored = PipelineStage::from_id(stage.id(), stage.typesize_byte())
.expect("known id must decode");
assert_eq!(stage, restored);
}
}
#[test]
fn from_id_rejects_unknown_id() {
assert!(PipelineStage::from_id(0, 0).is_err());
assert!(PipelineStage::from_id(10, 0).is_err());
assert!(PipelineStage::from_id(255, 0).is_err());
}
#[test]
fn byte_shuffle_round_trips_with_tail() {
let data: Vec<u8> = (0..13u8).collect();
let shuffled = byte_shuffle(&data, 4);
assert_eq!(shuffled.len(), data.len());
let restored = byte_unshuffle(&shuffled, 4);
assert_eq!(restored, data);
}
#[test]
fn byte_shuffle_identity_for_typesize_one() {
let data: Vec<u8> = (0..32u8).collect();
assert_eq!(byte_shuffle(&data, 1), data);
assert_eq!(byte_unshuffle(&data, 1), data);
}
}