precomp2 0.2.0

Reversible preprocessing for compressed and container data.
Documentation
use std::io::{Read, Write};

use bzip2::{Compression, read::BzDecoder, write::BzEncoder};

use super::{
  DecodeConfig, EncodeConfig, PenaltyEntry,
  budget::{DecodeBudget, EncodeBudget},
  payload::{choose_payload, decode_payload},
  varint::{read_varint_slice, write_varint_vec},
};
use crate::{
  error::{Error, Result},
  pcf2::{Pcf2Segment, SegmentKind},
};

pub(super) fn encode_bzip2(input: &[u8], config: &EncodeConfig, depth: u32, budget: &mut EncodeBudget) -> Result<Option<Pcf2Segment>> {
  if input.len() < 4 || &input[..3] != b"BZh" {
    return Ok(None);
  }
  let level = input[3];
  if !(b'1'..=b'9').contains(&level) {
    return Ok(None);
  }
  let level_num = (level - b'0') as u32;
  let plain = match decode_bzip2_raw(input) {
    Ok(v) => v,
    Err(_) => return Ok(None),
  };
  let (payload_kind, payload) = choose_payload(plain.clone(), input.len(), config, depth, budget)?;
  let recompressed = encode_bzip2_raw(&plain, level_num)?;
  if recompressed.len() != input.len() {
    return Ok(None);
  }
  let penalty = compute_penalty(input, &recompressed);

  let mut meta = Vec::new();
  meta.push(0);
  meta.push(payload_kind);
  meta.push(level);
  write_varint_vec(penalty.len() as u64, &mut meta);
  for entry in &penalty {
    meta.extend_from_slice(&entry.pos.to_be_bytes());
    meta.push(entry.value);
  }
  write_varint_vec(plain.len() as u64, &mut meta);

  Ok(Some(Pcf2Segment {
    kind: SegmentKind::Bzip2 as u8,
    flags: 0,
    orig_len: input.len() as u64,
    meta,
    data: payload,
  }))
}

pub(super) fn decode_bzip2_segment(segment: &Pcf2Segment, config: &DecodeConfig, depth: u32, budget: &mut DecodeBudget) -> Result<Vec<u8>> {
  if segment.meta.len() < 3 {
    return Err(Error::InvalidSegment("bzip2 meta"));
  }
  let meta_version = segment.meta[0];
  let payload_kind = segment.meta[1];
  let level = segment.meta[2];
  if meta_version != 0 {
    return Err(Error::InvalidSegment("bzip2 meta_version"));
  }
  if !(b'1'..=b'9').contains(&level) {
    return Err(Error::InvalidSegment("bzip2 level"));
  }
  let mut offset = 3usize;
  let penalty_count = read_varint_slice(&segment.meta, &mut offset)? as usize;
  let mut penalty = Vec::with_capacity(penalty_count);
  for _ in 0..penalty_count {
    if offset + 5 > segment.meta.len() {
      return Err(Error::InvalidSegment("bzip2 penalty bounds"));
    }
    let pos = u32::from_be_bytes([
      segment.meta[offset],
      segment.meta[offset + 1],
      segment.meta[offset + 2],
      segment.meta[offset + 3],
    ]);
    let value = segment.meta[offset + 4];
    offset += 5;
    penalty.push(PenaltyEntry { pos, value });
  }
  let decoded_len = if offset < segment.meta.len() {
    Some(read_varint_slice(&segment.meta, &mut offset)? as usize)
  } else {
    None
  };

  let payload = decode_payload(payload_kind, segment.data.clone(), config, depth, budget)?;
  if let Some(expected) = decoded_len
    && expected != payload.len()
  {
    return Err(Error::InvalidSegment("bzip2 decoded_len mismatch"));
  }
  let level_num = (level - b'0') as u32;
  let mut recompressed = encode_bzip2_raw(&payload, level_num)?;
  for entry in penalty {
    let pos = entry.pos as usize;
    if pos >= recompressed.len() {
      return Err(Error::InvalidSegment("bzip2 penalty out of range"));
    }
    recompressed[pos] = entry.value;
  }
  budget.consume(recompressed.len(), config)?;
  Ok(recompressed)
}

fn decode_bzip2_raw(input: &[u8]) -> Result<Vec<u8>> {
  let mut decoder = BzDecoder::new(input);
  let mut out = Vec::new();
  decoder.read_to_end(&mut out)?;
  Ok(out)
}

fn encode_bzip2_raw(input: &[u8], level: u32) -> Result<Vec<u8>> {
  let mut encoder = BzEncoder::new(Vec::new(), Compression::new(level));
  encoder.write_all(input)?;
  Ok(encoder.finish()?)
}

fn compute_penalty(original: &[u8], recompressed: &[u8]) -> Vec<PenaltyEntry> {
  let mut penalty = Vec::new();
  for (pos, (&a, &b)) in original.iter().zip(recompressed).enumerate() {
    if a != b {
      penalty.push(PenaltyEntry { pos: pos as u32, value: a });
    }
  }
  penalty
}