assetpack-core 0.3.0

Content-addressed asset packing, chunking, recipe, and SQLite storage primitives.
Documentation
use zstd::stream::raw::{DParameter, Decoder, InBuffer, Operation, OutBuffer};

use super::{
  DIRECTORY_RECORD_BYTES, DirectoryRecord, HEADER_BYTES, Header, MAX_DECODED_OBJECT_BYTES, MAX_INDEX_DECODED_BYTES, MAX_OBJECT_COUNT,
  MAX_STORED_OBJECT_BYTES, checked_range, codec_from_id, frame_count, invalid, logical_pack_id_records,
};
use crate::{Hash32, ObjectKind, Result};

const ZSTD_MAGIC: [u8; 4] = [0x28, 0xb5, 0x2f, 0xfd];
const INDEX_WINDOW_LOG: u32 = 20;

pub(super) fn encode_decoded_index(root: Hash32, records: &[DirectoryRecord]) -> Vec<u8> {
  let mut decoded = Vec::with_capacity(32 + records.len() * DIRECTORY_RECORD_BYTES);
  decoded.extend_from_slice(root.as_bytes());
  for record in records {
    decoded.extend_from_slice(record.hash.as_bytes());
    decoded.push(record.kind as u8);
    decoded.push(super::codec_id(record.codec));
    decoded.extend_from_slice(&record.decoded_length.to_le_bytes());
    decoded.extend_from_slice(&record.stored_length.to_le_bytes());
  }
  decoded
}

pub(super) fn compress_index(decoded: &[u8]) -> Result<Vec<u8>> {
  if decoded.len() > MAX_INDEX_DECODED_BYTES {
    return Err(invalid("decoded index exceeds hard limit"));
  }
  let mut encoder = zstd::stream::write::Encoder::new(Vec::new(), 19).map_err(|error| invalid(error.to_string()))?;
  encoder
    .set_parameter(zstd::zstd_safe::CParameter::WindowLog(INDEX_WINDOW_LOG))
    .map_err(|error| invalid(error.to_string()))?;
  encoder
    .set_parameter(zstd::zstd_safe::CParameter::NbWorkers(0))
    .map_err(|error| invalid(error.to_string()))?;
  encoder
    .set_parameter(zstd::zstd_safe::CParameter::EnableLongDistanceMatching(false))
    .map_err(|error| invalid(error.to_string()))?;
  encoder.include_checksum(false).map_err(|error| invalid(error.to_string()))?;
  encoder.include_dictid(false).map_err(|error| invalid(error.to_string()))?;
  encoder.include_contentsize(false).map_err(|error| invalid(error.to_string()))?;
  std::io::Write::write_all(&mut encoder, decoded).map_err(|error| invalid(error.to_string()))?;
  let compressed = encoder.finish().map_err(|error| invalid(error.to_string()))?;
  if compressed.len() > super::MAX_INDEX_COMPRESSED_BYTES {
    return Err(invalid("compressed index exceeds hard limit"));
  }
  Ok(compressed)
}

pub(super) fn decode_index(header: &Header, compressed: &[u8], pack_len: usize) -> Result<(Hash32, Vec<DirectoryRecord>)> {
  if compressed.get(..4) != Some(&ZSTD_MAGIC) {
    return Err(invalid("index is not a standard Zstandard frame"));
  }
  let frame_size = zstd::zstd_safe::find_frame_compressed_size(compressed).map_err(|error| invalid(error.to_string()))?;
  if frame_size != compressed.len() {
    return Err(invalid("index must contain exactly one Zstandard frame"));
  }
  let mut decoder = Decoder::new().map_err(|error| invalid(error.to_string()))?;
  decoder
    .set_parameter(DParameter::WindowLogMax(INDEX_WINDOW_LOG))
    .map_err(|error| invalid(error.to_string()))?;
  let mut input = InBuffer::around(compressed);
  let mut parser = IndexParser::new(header, pack_len);
  loop {
    let mut scratch = [0_u8; 8192];
    let mut output = OutBuffer::around(&mut scratch[..]);
    let remaining = decoder.run(&mut input, &mut output).map_err(|error| invalid(error.to_string()))?;
    let produced = output.pos();
    parser.push(&scratch[..produced])?;
    if remaining == 0 {
      if input.pos() != compressed.len() {
        return Err(invalid("index decoder did not consume the complete frame"));
      }
      break;
    }
    if produced == 0 && input.pos() == compressed.len() {
      return Err(invalid("truncated Zstandard index"));
    }
  }
  parser.finish()
}

struct IndexParser<'a> {
  header: &'a Header,
  pack_len: usize,
  decoded_bytes: usize,
  root: [u8; 32],
  root_bytes: usize,
  record: [u8; DIRECTORY_RECORD_BYTES],
  record_bytes: usize,
  records: Vec<DirectoryRecord>,
}

impl<'a> IndexParser<'a> {
  fn new(header: &'a Header, pack_len: usize) -> Self {
    Self {
      header,
      pack_len,
      decoded_bytes: 0,
      root: [0; 32],
      root_bytes: 0,
      record: [0; DIRECTORY_RECORD_BYTES],
      record_bytes: 0,
      records: Vec::new(),
    }
  }

  fn push(&mut self, mut bytes: &[u8]) -> Result<()> {
    self.decoded_bytes = self
      .decoded_bytes
      .checked_add(bytes.len())
      .filter(|length| *length <= MAX_INDEX_DECODED_BYTES)
      .ok_or_else(|| invalid("decoded index exceeds hard limit"))?;
    if self.root_bytes < self.root.len() {
      let count = bytes.len().min(self.root.len() - self.root_bytes);
      self.root[self.root_bytes..self.root_bytes + count].copy_from_slice(&bytes[..count]);
      self.root_bytes += count;
      bytes = &bytes[count..];
    }
    while !bytes.is_empty() {
      let count = bytes.len().min(DIRECTORY_RECORD_BYTES - self.record_bytes);
      self.record[self.record_bytes..self.record_bytes + count].copy_from_slice(&bytes[..count]);
      self.record_bytes += count;
      bytes = &bytes[count..];
      if self.record_bytes == DIRECTORY_RECORD_BYTES {
        self.emit_record()?;
        self.record_bytes = 0;
      }
    }
    Ok(())
  }

  fn emit_record(&mut self) -> Result<()> {
    if self.records.len() as u64 >= MAX_OBJECT_COUNT {
      return Err(invalid("object count exceeds hard limit"));
    }
    let record = DirectoryRecord {
      hash: Hash32::from_bytes(&self.record[..32])?,
      kind: ObjectKind::from_u8(self.record[32]).ok_or_else(|| invalid("unknown object kind"))?,
      codec: codec_from_id(self.record[33])?,
      decoded_length: u64::from_le_bytes(self.record[34..42].try_into().unwrap()),
      stored_length: u64::from_le_bytes(self.record[42..50].try_into().unwrap()),
      payload_offset: 0,
    };
    if record.decoded_length > MAX_DECODED_OBJECT_BYTES || record.stored_length > MAX_STORED_OBJECT_BYTES {
      return Err(invalid("non-canonical directory record"));
    }
    if self.records.last().is_some_and(|previous| previous.hash >= record.hash) {
      return Err(invalid("directory is not strictly hash sorted"));
    }
    self.records.push(record);
    Ok(())
  }

  fn finish(mut self) -> Result<(Hash32, Vec<DirectoryRecord>)> {
    if self.root_bytes != self.root.len() {
      return Err(invalid("decoded index is shorter than root hash"));
    }
    if self.record_bytes != 0 {
      return Err(invalid("decoded index ends with a partial directory record"));
    }
    if self.records.is_empty() {
      return Err(invalid("decoded index has no directory records"));
    }
    let root = Hash32::new(self.root);
    let root_record = self
      .records
      .binary_search_by_key(&root, |record| record.hash)
      .ok()
      .and_then(|index| self.records.get(index))
      .ok_or_else(|| invalid("root recipe is absent from directory"))?;
    if root_record.kind != ObjectKind::Recipe {
      return Err(invalid("root directory object is not a recipe"));
    }
    if logical_pack_id_records(root, &self.records) != self.header.pack_id {
      return Err(invalid("logical pack identity mismatch"));
    }
    let mut cursor = (HEADER_BYTES as u64)
      .checked_add(self.header.index_stored_length)
      .ok_or_else(|| invalid("payload cursor overflow"))?;
    for record in &mut self.records {
      record.payload_offset = cursor;
      let tag_bytes = if self.header.encrypted() {
        frame_count(record.stored_length)
          .checked_mul(16)
          .ok_or_else(|| invalid("payload tag length overflow"))?
      } else {
        0
      };
      let physical = record
        .stored_length
        .checked_add(tag_bytes)
        .ok_or_else(|| invalid("payload length overflow"))?;
      checked_range(self.pack_len, cursor, physical, "payload")?;
      cursor = cursor.checked_add(physical).ok_or_else(|| invalid("payload cursor overflow"))?;
    }
    if cursor != self.pack_len as u64 {
      return Err(invalid("non-canonical trailing bytes"));
    }
    Ok((root, self.records))
  }
}

#[cfg(test)]
mod tests {
  use std::io::Write;

  use super::*;
  use crate::Codec;

  fn record(hash: Hash32, kind: ObjectKind, decoded_length: u64, stored_length: u64) -> DirectoryRecord {
    DirectoryRecord {
      hash,
      kind,
      codec: Codec::Raw,
      decoded_length,
      stored_length,
      payload_offset: 0,
    }
  }

  fn root_and_records(count: usize) -> (Hash32, Vec<DirectoryRecord>) {
    let mut records = Vec::with_capacity(count);
    for value in 0..count as u64 {
      let mut bytes = [0_u8; 32];
      bytes[24..].copy_from_slice(&value.to_be_bytes());
      records.push(record(
        Hash32::new(bytes),
        if value == 0 { ObjectKind::Recipe } else { ObjectKind::Chunk },
        0,
        0,
      ));
    }
    (records[0].hash, records)
  }

  fn header(root: Hash32, records: &[DirectoryRecord], index_stored_length: u64) -> Header {
    Header {
      tag: super::super::DEFAULT_FORMAT_TAG,
      crypto_suite: super::super::CRYPTO_PLAIN_SHA3,
      index_stored_length,
      pack_id: super::super::logical_pack_id_records(root, records),
      pack_salt: [0; 32],
      key_slot: [0; 16],
      index_digest: Hash32::new([0; 32]),
    }
  }

  fn compress_unchecked(decoded: &[u8], window_log: u32) -> Vec<u8> {
    let mut encoder = zstd::stream::write::Encoder::new(Vec::new(), 1).unwrap();
    encoder.set_parameter(zstd::zstd_safe::CParameter::WindowLog(window_log)).unwrap();
    encoder.include_checksum(false).unwrap();
    encoder.include_dictid(false).unwrap();
    encoder.include_contentsize(false).unwrap();
    encoder.write_all(decoded).unwrap();
    encoder.finish().unwrap()
  }

  fn rejection(result: Result<(Hash32, Vec<DirectoryRecord>)>) -> crate::Error {
    match result {
      Ok(_) => panic!("index was accepted"),
      Err(error) => error,
    }
  }

  #[test]
  fn compact_directory_record_is_byte_exact() {
    assert_eq!(DIRECTORY_RECORD_BYTES, 50);
    let hash = Hash32::new([0x55; 32]);
    let record = DirectoryRecord {
      hash,
      kind: ObjectKind::Recipe,
      codec: Codec::Brotli,
      decoded_length: 0x1112_1314_1516_1718,
      stored_length: 0x2122_2324_2526_2728,
      payload_offset: 999,
    };
    let decoded = encode_decoded_index(hash, &[record]);
    let raw = &decoded[32..];
    assert_eq!(&raw[..32], &[0x55; 32]);
    assert_eq!(raw[32], ObjectKind::Recipe as u8);
    assert_eq!(raw[33], 2);
    assert_eq!(&raw[34..42], &0x1112_1314_1516_1718_u64.to_le_bytes());
    assert_eq!(&raw[42..50], &0x2122_2324_2526_2728_u64.to_le_bytes());
  }

  #[test]
  fn fixed_encoder_rejects_decoded_input_over_hard_limit() {
    let error = compress_index(&vec![0; MAX_INDEX_DECODED_BYTES + 1]).unwrap_err();
    assert!(error.to_string().contains("decoded index exceeds hard limit"));
  }

  #[test]
  fn fixed_encoder_rejects_compressed_output_over_hard_limit() {
    let mut state = 0x1234_5678_u32;
    let decoded = (0..super::super::MAX_INDEX_COMPRESSED_BYTES + 64 * 1024)
      .map(|_| {
        state ^= state << 13;
        state ^= state >> 17;
        state ^= state << 5;
        state as u8
      })
      .collect::<Vec<_>>();
    let error = compress_index(&decoded).unwrap_err();
    assert!(error.to_string().contains("compressed index exceeds hard limit"));
  }

  #[test]
  fn decoder_rejects_short_nonstandard_multiple_trailing_and_truncated_frames() {
    let (root, records) = root_and_records(1);
    let decoded = encode_decoded_index(root, &records);
    let stored = compress_index(&decoded).unwrap();
    let cases = [
      compress_unchecked(&decoded[..31], INDEX_WINDOW_LOG),
      vec![0x50, 0x2a, 0x4d, 0x18, 0, 0, 0, 0],
      [stored.as_slice(), stored.as_slice()].concat(),
      [stored.as_slice(), &[0]].concat(),
      stored[..stored.len() - 1].to_vec(),
      [vec![0xff], stored[1..].to_vec()].concat(),
    ];
    for (case, replacement) in cases.into_iter().enumerate() {
      let header = header(root, &records, replacement.len() as u64);
      let result = decode_index(&header, &replacement, HEADER_BYTES + replacement.len());
      assert!(result.is_err(), "invalid frame case {case} was accepted");
    }
  }

  #[test]
  fn decoder_rejects_window_over_limit_before_directory_parse() {
    let oversized = compress_unchecked(&vec![0; 2 * 1024 * 1024], 21);
    let (root, records) = root_and_records(1);
    let header = header(root, &records, oversized.len() as u64);
    let error = rejection(decode_index(&header, &oversized, HEADER_BYTES + oversized.len()));
    assert!(error.to_string().contains("Frame requires too much memory for decoding"), "{error}");
  }

  #[test]
  fn decoder_rejects_semantically_valid_compression_bomb_at_decoded_limit() {
    let record_count = (MAX_INDEX_DECODED_BYTES - 32) / DIRECTORY_RECORD_BYTES + 1;
    let (root, records) = root_and_records(record_count);
    let decoded = encode_decoded_index(root, &records);
    assert!(decoded.len() > MAX_INDEX_DECODED_BYTES);
    let bomb = compress_unchecked(&decoded, INDEX_WINDOW_LOG);
    assert!(bomb.len() <= super::super::MAX_INDEX_COMPRESSED_BYTES);
    let header = header(root, &records, bomb.len() as u64);
    let error = rejection(decode_index(&header, &bomb, HEADER_BYTES + bomb.len()));
    assert!(error.to_string().contains("decoded index exceeds hard limit"), "{error}");
  }

  #[test]
  fn decoder_rejects_payload_cursor_overflow() {
    let (root, records) = root_and_records(1);
    let decoded = encode_decoded_index(root, &records);
    let stored = compress_index(&decoded).unwrap();
    let header = header(root, &records, u64::MAX);
    let error = rejection(decode_index(&header, &stored, HEADER_BYTES + stored.len()));
    assert!(error.to_string().contains("payload cursor overflow"), "{error}");
  }

  #[cfg(feature = "sealed-encryption")]
  #[test]
  fn encrypted_odd_length_payloads_have_consecutive_exact_cursors() {
    let (root, mut records) = root_and_records(4);
    for (record, stored_length) in records.iter_mut().zip([1, 3, 5, 7]) {
      record.stored_length = stored_length;
    }
    let decoded = encode_decoded_index(root, &records);
    let stored = compress_index(&decoded).unwrap();
    let mut header = header(root, &records, stored.len() as u64 + 16);
    header.crypto_suite = super::super::CRYPTO_CHACHA20_POLY1305_FRAMES;
    let payload_bytes = records.iter().map(|record| record.stored_length + 16).sum::<u64>();
    let pack_len = HEADER_BYTES + header.index_stored_length as usize + payload_bytes as usize;
    let (_, decoded_records) = decode_index(&header, &stored, pack_len).unwrap();
    let mut cursor = HEADER_BYTES as u64 + header.index_stored_length;
    for record in decoded_records {
      assert_eq!(record.payload_offset, cursor);
      cursor += record.stored_length + 16;
    }
    assert_eq!(cursor, pack_len as u64);
  }
}