precomp2 0.2.0

Reversible preprocessing for compressed and container data.
Documentation
use ::base64::Engine as _;
use memchr::memmem;

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

#[derive(Clone, Debug)]
pub(super) struct Base64Block {
  pub(super) header: Vec<u8>,
  pub(super) line_case: u8,
  pub(super) line_lengths: Vec<usize>,
  pub(super) decoded: Vec<u8>,
  pub(super) decoded_len: usize,
}

pub(super) fn parse_base64_block(input: &[u8]) -> Option<Base64Block> {
  let mut header = Vec::new();
  let mut cursor = 0usize;
  if let Some(pos) = memmem::find(input, b"\r\n\r\n") {
    header.extend_from_slice(&input[..pos + 4]);
    cursor = pos + 4;
  }

  let mut line_lengths = Vec::new();
  let mut encoded = Vec::new();
  while cursor < input.len() {
    if input[cursor] == b'-' {
      break;
    }
    let line_end = match memmem::find(&input[cursor..], b"\r\n") {
      Some(v) => cursor + v,
      None => return None,
    };
    let line = &input[cursor..line_end];
    if line.is_empty() {
      return None;
    }
    if line.len() > 255 {
      return None;
    }
    line_lengths.push(line.len());
    encoded.extend_from_slice(line);
    cursor = line_end + 2;
  }
  if line_lengths.is_empty() || cursor != input.len() {
    return None;
  }
  let decoded = BASE64_ENGINE.decode(encoded).ok()?;
  let decoded_len = decoded.len();
  let line_case = if line_lengths.iter().all(|&l| l == line_lengths[0]) {
    0
  } else if line_lengths.len() > 1 && line_lengths[..line_lengths.len() - 1].iter().all(|&l| l == line_lengths[0]) {
    1
  } else {
    2
  };

  Some(Base64Block {
    header,
    line_case,
    line_lengths,
    decoded,
    decoded_len,
  })
}

pub(super) fn encode_base64(input: &[u8], config: &EncodeConfig, depth: u32, budget: &mut EncodeBudget) -> Result<Option<Pcf2Segment>> {
  let Some(block) = parse_base64_block(input) else {
    return Ok(None);
  };
  let (payload_kind, payload) = choose_payload(block.decoded, input.len(), config, depth, budget)?;

  let mut meta = Vec::new();
  meta.push(0);
  meta.push(payload_kind);
  write_varint_vec(block.header.len() as u64, &mut meta);
  meta.extend_from_slice(&block.header);
  meta.push(block.line_case);
  write_varint_vec(block.line_lengths.len() as u64, &mut meta);
  match block.line_case {
    0 => meta.push(block.line_lengths[0] as u8),
    1 => {
      meta.push(block.line_lengths[0] as u8);
      meta.push(*block.line_lengths.last().unwrap() as u8);
    }
    2 => {
      for len in &block.line_lengths {
        meta.push(*len as u8);
      }
    }
    _ => return Ok(None),
  }
  write_varint_vec(block.decoded_len as u64, &mut meta);

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

pub(super) fn decode_base64_segment(
  segment: &Pcf2Segment,
  config: &DecodeConfig,
  depth: u32,
  budget: &mut DecodeBudget,
) -> Result<Vec<u8>> {
  if segment.meta.len() < 2 {
    return Err(Error::InvalidSegment("base64 meta"));
  }
  let meta_version = segment.meta[0];
  let payload_kind = segment.meta[1];
  if meta_version != 0 {
    return Err(Error::InvalidSegment("base64 meta_version"));
  }
  let mut offset = 2usize;
  let header_len = read_varint_slice(&segment.meta, &mut offset)? as usize;
  if offset + header_len > segment.meta.len() {
    return Err(Error::InvalidSegment("base64 header bounds"));
  }
  let header = segment.meta[offset..offset + header_len].to_vec();
  offset += header_len;
  if offset >= segment.meta.len() {
    return Err(Error::InvalidSegment("base64 line_case"));
  }
  let line_case = segment.meta[offset];
  offset += 1;
  let line_count = read_varint_slice(&segment.meta, &mut offset)? as usize;
  if line_count == 0 {
    return Err(Error::InvalidSegment("base64 line_count"));
  }
  let mut line_lengths = Vec::new();
  match line_case {
    0 => {
      if offset >= segment.meta.len() {
        return Err(Error::InvalidSegment("base64 line_len_0"));
      }
      let len = segment.meta[offset] as usize;
      offset += 1;
      line_lengths = vec![len; line_count];
    }
    1 => {
      if offset + 1 >= segment.meta.len() {
        return Err(Error::InvalidSegment("base64 line_len_1"));
      }
      let len = segment.meta[offset] as usize;
      let last = segment.meta[offset + 1] as usize;
      offset += 2;
      line_lengths = vec![len; line_count];
      if let Some(last_len) = line_lengths.last_mut() {
        *last_len = last;
      }
    }
    2 => {
      if offset + line_count > segment.meta.len() {
        return Err(Error::InvalidSegment("base64 line_len_list"));
      }
      for _ in 0..line_count {
        line_lengths.push(segment.meta[offset] as usize);
        offset += 1;
      }
    }
    _ => return Err(Error::InvalidSegment("base64 line_case")),
  }
  let decoded_len = if offset < segment.meta.len() {
    read_varint_slice(&segment.meta, &mut offset)? as usize
  } else {
    segment.data.len()
  };

  let payload = decode_payload(payload_kind, segment.data.clone(), config, depth, budget)?;
  if decoded_len != payload.len() {
    return Err(Error::InvalidSegment("base64 decoded_len mismatch"));
  }

  let encoded = BASE64_ENGINE.encode(payload);
  let mut out = Vec::new();
  out.extend_from_slice(&header);
  let mut cursor = 0usize;
  for len in line_lengths {
    if cursor >= encoded.len() {
      break;
    }
    let end = (cursor + len).min(encoded.len());
    out.extend_from_slice(&encoded.as_bytes()[cursor..end]);
    out.extend_from_slice(b"\r\n");
    cursor = end;
  }
  if cursor != encoded.len() {
    return Err(Error::InvalidSegment("base64 line lengths mismatch"));
  }
  budget.consume(out.len(), config)?;
  Ok(out)
}