assetpack-core 0.3.0

Content-addressed asset packing, chunking, recipe, and SQLite storage primitives.
Documentation
use std::io::{self, BufRead, Read};

use sha3::{Digest, Sha3_256};

#[cfg(feature = "sealed-encryption")]
use super::FrameUnsealer;
use super::{DirectoryRecord, Header, PayloadReader, invalid};
use crate::{
  codec::Codec,
  error::{Error, Result},
  hash::Hash32,
  object::{ObjectSource, VerifiedObject},
};

const STREAM_BUFFER_BYTES: usize = 32 * 1024;
const ZSTD_MAX_WINDOW_LOG: u32 = 22;
const BROTLI_MAX_WINDOW_LOG: u32 = 20;

pub struct SealedPackReader<'a> {
  bytes: &'a [u8],
  pub(super) header: Header,
  root_recipe: Hash32,
  pub(super) records: Vec<DirectoryRecord>,
  #[cfg(feature = "sealed-encryption")]
  unsealer: Option<&'a dyn FrameUnsealer>,
}

impl<'a> SealedPackReader<'a> {
  pub(super) fn from_records(
    bytes: &'a [u8],
    header: Header,
    root_recipe: Hash32,
    records: Vec<DirectoryRecord>,
    #[cfg(feature = "sealed-encryption")] unsealer: Option<&'a dyn FrameUnsealer>,
    #[cfg(not(feature = "sealed-encryption"))] _unsealer: Option<&'a ()>,
  ) -> Result<Self> {
    Ok(Self {
      bytes,
      header,
      root_recipe,
      records,
      #[cfg(feature = "sealed-encryption")]
      unsealer,
    })
  }

  pub fn pack_id(&self) -> Hash32 {
    self.header.pack_id
  }

  pub fn root_recipe(&self) -> Hash32 {
    self.root_recipe
  }

  pub fn object_count(&self) -> usize {
    self.records.len()
  }

  pub fn verify_all_objects(&self) -> Result<()> {
    let mut scratch = [0_u8; 32 * 1024];
    for record in &self.records {
      let mut object = self.open_object(record)?;
      loop {
        let read = object.read(&mut scratch).map_err(error_from_io)?;
        if read == 0 {
          break;
        }
      }
    }
    Ok(())
  }

  fn open_object(&self, record: &DirectoryRecord) -> Result<VerifyingReader<'a>> {
    let source = PayloadReader::new(
      self.bytes,
      &self.header,
      record.clone(),
      #[cfg(feature = "sealed-encryption")]
      self.unsealer,
    )?;
    let decoded: Box<dyn Read + 'a> = match record.codec {
      Codec::Raw => Box::new(source),
      Codec::Zstd => {
        let mut decoder = zstd::stream::read::Decoder::new(source).map_err(|error| Error::Decompress(error.to_string()))?;
        decoder
          .window_log_max(ZSTD_MAX_WINDOW_LOG)
          .map_err(|error| Error::Decompress(error.to_string()))?;
        Box::new(decoder)
      }
      Codec::Brotli => {
        let mut source = io::BufReader::with_capacity(STREAM_BUFFER_BYTES, source);
        let first = source
          .fill_buf()
          .map_err(error_from_io)?
          .first()
          .copied()
          .ok_or_else(|| invalid("empty Brotli payload"))?;
        if brotli_window_bits(first)? > BROTLI_MAX_WINDOW_LOG {
          return Err(invalid("Brotli window exceeds sealed decoder limit"));
        }
        Box::new(brotli::Decompressor::new(source, STREAM_BUFFER_BYTES))
      }
    };
    Ok(VerifyingReader::new(decoded, record.hash, record.decoded_length))
  }
}

impl ObjectSource for SealedPackReader<'_> {
  fn read_object(&self, hash: &Hash32) -> Result<Option<VerifiedObject>> {
    let Some(record) = self
      .records
      .binary_search_by_key(hash, |record| record.hash)
      .ok()
      .and_then(|index| self.records.get(index))
    else {
      return Ok(None);
    };
    let mut reader = self.open_object(record)?;
    let capacity: usize = record
      .decoded_length
      .try_into()
      .map_err(|_| invalid("object length exceeds address space"))?;
    let mut bytes = Vec::with_capacity(capacity);
    reader.read_to_end(&mut bytes).map_err(error_from_io)?;
    Ok(Some(VerifiedObject {
      hash: record.hash,
      kind: record.kind,
      bytes,
    }))
  }
}

fn brotli_window_bits(first: u8) -> Result<u32> {
  if first & 1 == 0 {
    return Ok(16);
  }
  let primary = (first >> 1) & 7;
  if primary != 0 {
    return Ok(17 + u32::from(primary));
  }
  let secondary = (first >> 4) & 7;
  match secondary {
    0 => Ok(17),
    1 => Err(invalid("large-window Brotli is not supported")),
    value => Ok(8 + u32::from(value)),
  }
}

struct VerifyingReader<'a> {
  inner: Option<Box<dyn Read + 'a>>,
  hasher: Sha3_256,
  expected_hash: Hash32,
  expected_length: u64,
  decoded_length: u64,
  verified: bool,
}

impl<'a> VerifyingReader<'a> {
  fn new(inner: Box<dyn Read + 'a>, expected_hash: Hash32, expected_length: u64) -> Self {
    Self {
      inner: Some(inner),
      hasher: Sha3_256::new(),
      expected_hash,
      expected_length,
      decoded_length: 0,
      verified: false,
    }
  }

  fn release_working_set(&mut self) {
    self.inner.take();
  }
}

impl Read for VerifyingReader<'_> {
  fn read(&mut self, output: &mut [u8]) -> io::Result<usize> {
    if self.verified {
      return Ok(0);
    }
    let read = match self.inner.as_mut().expect("unverified reader has an inner stream").read(output) {
      Ok(read) => read,
      Err(error) => {
        self.release_working_set();
        return Err(error);
      }
    };
    if read == 0 {
      let actual = Hash32::new(self.hasher.clone().finalize().into());
      if self.decoded_length != self.expected_length {
        self.release_working_set();
        return Err(io_error(Error::ObjectLengthMismatch {
          expected: self.expected_length,
          actual: self.decoded_length,
        }));
      }
      if actual != self.expected_hash {
        self.release_working_set();
        return Err(io_error(Error::ObjectHashMismatch {
          expected: self.expected_hash,
          actual,
        }));
      }
      self.verified = true;
      self.release_working_set();
      return Ok(0);
    }
    self.hasher.update(&output[..read]);
    self.decoded_length = match self.decoded_length.checked_add(read as u64) {
      Some(length) => length,
      None => {
        self.release_working_set();
        return Err(io::Error::new(io::ErrorKind::InvalidData, "sealed object length overflow"));
      }
    };
    if self.decoded_length > self.expected_length {
      let actual = self.decoded_length;
      self.release_working_set();
      return Err(io_error(Error::ObjectLengthMismatch {
        expected: self.expected_length,
        actual,
      }));
    }
    Ok(read)
  }
}

fn io_error(error: Error) -> io::Error {
  io::Error::new(io::ErrorKind::InvalidData, error)
}

fn error_from_io(error: io::Error) -> Error {
  let kind = error.kind();
  match error.into_inner() {
    Some(inner) => match inner.downcast::<Error>() {
      Ok(error) => *error,
      Err(inner) => Error::Decompress(inner.to_string()),
    },
    None => Error::Decompress(io::Error::from(kind).to_string()),
  }
}

#[cfg(test)]
mod tests {
  use super::*;

  #[test]
  fn brotli_window_bits_decodes_rfc7932_first_byte() {
    for (first, expected) in [
      (0x00, 16),
      (0x02, 16),
      (0xfe, 16),
      (0x01, 17),
      (0x03, 18),
      (0x05, 19),
      (0x0b, 22),
      (0x0f, 24),
      (0x21, 10),
      (0x71, 15),
    ] {
      assert_eq!(brotli_window_bits(first).unwrap(), expected, "first byte {first:#04x}");
    }
    for first in [0x11, 0x91] {
      let error = brotli_window_bits(first).unwrap_err();
      assert!(error.to_string().contains("large-window"), "first byte {first:#04x}: {error}");
    }
  }
}