use alloc::vec::Vec;
use xxhash_rust::xxh3::{Xxh3, xxh3_64};
use crate::error::Error;
pub const MAGIC: u32 = 0x504C_474D;
pub const FORMAT_VERSION: u16 = 1;
pub const FLAG_VECTORS: u16 = 1;
const HEADER: usize = 64;
const ENTRY: usize = 32;
const ALIGN: usize = 64;
const FILE_HASH_AT: usize = 20;
fn align_up(v: u64) -> u64 {
v.div_ceil(ALIGN as u64) * ALIGN as u64
}
fn file_hash(bytes: &[u8]) -> u64 {
let mut h = Xxh3::new();
h.update(&bytes[..FILE_HASH_AT]);
h.update(&[0u8; 8]);
h.update(&bytes[FILE_HASH_AT + 8..]);
h.digest()
}
#[derive(Debug, Default)]
pub struct SnapshotWriter {
sections: Vec<(u16, Vec<u8>)>,
}
impl SnapshotWriter {
pub fn new() -> Self {
Self::default()
}
pub fn section(&mut self, kind: u16, bytes: Vec<u8>) -> Result<(), Error> {
if self.sections.iter().any(|&(k, _)| k == kind) {
return Err(Error::Corrupt("duplicate section kind"));
}
self.sections.push((kind, bytes));
Ok(())
}
pub fn finish(self, config: &[u8], flags: u16, created_at: u64, engine_ver: &str) -> Vec<u8> {
let config_end = HEADER as u64 + config.len() as u64;
let table_start = align_up(config_end);
let table_end = table_start + (self.sections.len() * ENTRY) as u64;
let mut offsets = Vec::with_capacity(self.sections.len());
let mut cursor = align_up(table_end);
for (_, bytes) in &self.sections {
offsets.push(cursor);
cursor = align_up(cursor + bytes.len() as u64);
}
let file_len = cursor as usize;
let mut out = alloc::vec![0u8; file_len];
out[0..4].copy_from_slice(&MAGIC.to_le_bytes());
out[4..6].copy_from_slice(&FORMAT_VERSION.to_le_bytes());
out[6..8].copy_from_slice(&flags.to_le_bytes());
out[8..10].copy_from_slice(&(self.sections.len() as u16).to_le_bytes());
out[16..20].copy_from_slice(&(config.len() as u32).to_le_bytes());
out[28..36].copy_from_slice(&created_at.to_le_bytes());
let ver = engine_ver.as_bytes();
let ver_len = ver.len().min(24);
out[36..36 + ver_len].copy_from_slice(&ver[..ver_len]);
out[HEADER..HEADER + config.len()].copy_from_slice(config);
for (i, (kind, bytes)) in self.sections.iter().enumerate() {
let at = table_start as usize + i * ENTRY;
out[at..at + 2].copy_from_slice(&kind.to_le_bytes());
out[at + 2..at + 4].copy_from_slice(&(ALIGN as u16).to_le_bytes());
out[at + 8..at + 16].copy_from_slice(&offsets[i].to_le_bytes());
out[at + 16..at + 24].copy_from_slice(&(bytes.len() as u64).to_le_bytes());
out[at + 24..at + 32].copy_from_slice(&xxh3_64(bytes).to_le_bytes());
let start = offsets[i] as usize;
out[start..start + bytes.len()].copy_from_slice(bytes);
}
let hash = file_hash(&out);
out[FILE_HASH_AT..FILE_HASH_AT + 8].copy_from_slice(&hash.to_le_bytes());
out
}
}
#[derive(Debug, Clone, Copy)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct SectionMeta {
pub kind: u16,
pub len: u64,
pub hash: u64,
}
#[derive(Debug)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Prefix {
pub bytes: Vec<u8>,
pub offsets: Vec<u64>,
pub file_len: u64,
}
pub trait SnapshotSink {
fn write(&mut self, bytes: &[u8]) -> Result<(), Error>;
fn patch(&mut self, at: u64, bytes: &[u8]) -> Result<(), Error>;
}
impl SnapshotSink for &mut Vec<u8> {
fn write(&mut self, bytes: &[u8]) -> Result<(), Error> {
self.extend_from_slice(bytes);
Ok(())
}
fn patch(&mut self, at: u64, bytes: &[u8]) -> Result<(), Error> {
let at = at as usize;
self[at..at + bytes.len()].copy_from_slice(bytes);
Ok(())
}
}
pub const FILE_HASH_OFFSET: u64 = FILE_HASH_AT as u64;
pub fn build_prefix(
config: &[u8],
flags: u16,
created_at: u64,
engine_ver: &str,
metas: &[SectionMeta],
) -> Prefix {
let count = metas.len();
let config_end = HEADER as u64 + config.len() as u64;
let table_start = align_up(config_end);
let table_end = table_start + (count * ENTRY) as u64;
let mut offsets = Vec::with_capacity(count);
let mut cursor = align_up(table_end);
for m in metas {
offsets.push(cursor);
cursor = align_up(cursor + m.len);
}
let file_len = cursor;
let prefix_len = align_up(table_end) as usize;
let mut out = alloc::vec![0u8; prefix_len];
out[0..4].copy_from_slice(&MAGIC.to_le_bytes());
out[4..6].copy_from_slice(&FORMAT_VERSION.to_le_bytes());
out[6..8].copy_from_slice(&flags.to_le_bytes());
out[8..10].copy_from_slice(&(count as u16).to_le_bytes());
out[16..20].copy_from_slice(&(config.len() as u32).to_le_bytes());
out[28..36].copy_from_slice(&created_at.to_le_bytes());
let ver = engine_ver.as_bytes();
let ver_len = ver.len().min(24);
out[36..36 + ver_len].copy_from_slice(&ver[..ver_len]);
out[HEADER..HEADER + config.len()].copy_from_slice(config);
for (i, m) in metas.iter().enumerate() {
let at = table_start as usize + i * ENTRY;
out[at..at + 2].copy_from_slice(&m.kind.to_le_bytes());
out[at + 2..at + 4].copy_from_slice(&(ALIGN as u16).to_le_bytes());
out[at + 8..at + 16].copy_from_slice(&offsets[i].to_le_bytes());
out[at + 16..at + 24].copy_from_slice(&m.len.to_le_bytes());
out[at + 24..at + 32].copy_from_slice(&m.hash.to_le_bytes());
}
Prefix {
bytes: out,
offsets,
file_len,
}
}
pub fn pad_len(offset: u64, len: u64) -> usize {
(align_up(offset + len) - (offset + len)) as usize
}
#[derive(Debug)]
pub struct Snapshot<'a> {
bytes: &'a [u8],
pub flags: u16,
pub created_at: u64,
config_len: usize,
sections: Vec<(u16, usize, usize, u64)>,
engine_ver_len: usize,
}
impl<'a> Snapshot<'a> {
pub fn parse(bytes: &'a [u8]) -> Result<Self, Error> {
if bytes.len() < HEADER {
return Err(Error::Corrupt("snapshot shorter than its header"));
}
if !bytes.len().is_multiple_of(ALIGN) {
return Err(Error::Corrupt("snapshot length is not 64-byte aligned"));
}
if u32::from_le_bytes(bytes[0..4].try_into().unwrap()) != MAGIC {
return Err(Error::Corrupt("bad magic"));
}
let version = u16::from_le_bytes(bytes[4..6].try_into().unwrap());
if version != FORMAT_VERSION {
return Err(Error::UnsupportedVersion(version));
}
let flags = u16::from_le_bytes(bytes[6..8].try_into().unwrap());
if flags & !FLAG_VECTORS != 0 {
return Err(Error::Corrupt("unknown flag bits set"));
}
let section_cnt = u16::from_le_bytes(bytes[8..10].try_into().unwrap()) as usize;
if bytes[10..16] != [0u8; 6] || bytes[60..64] != [0u8; 4] {
return Err(Error::Corrupt("reserved header bytes must be zero"));
}
let config_len = u32::from_le_bytes(bytes[16..20].try_into().unwrap()) as usize;
let created_at = u64::from_le_bytes(bytes[28..36].try_into().unwrap());
let ver_bytes = &bytes[36..60];
let engine_ver_len = ver_bytes.iter().position(|&b| b == 0).unwrap_or(24);
if ver_bytes[engine_ver_len..].iter().any(|&b| b != 0) {
return Err(Error::Corrupt("engine version is not zero-terminated"));
}
if core::str::from_utf8(&ver_bytes[..engine_ver_len]).is_err() {
return Err(Error::Corrupt("engine version is not UTF-8"));
}
let file_len = bytes.len() as u64;
let config_end = HEADER as u64 + config_len as u64;
let table_start = align_up(config_end);
let table_end = table_start + (section_cnt * ENTRY) as u64;
if table_end > file_len {
return Err(Error::Corrupt("section table out of bounds"));
}
if bytes[config_end as usize..table_start as usize]
.iter()
.any(|&b| b != 0)
{
return Err(Error::Corrupt("nonzero padding after the config block"));
}
let mut sections = Vec::with_capacity(section_cnt);
let mut expected = align_up(table_end);
if bytes[table_end as usize..expected as usize]
.iter()
.any(|&b| b != 0)
{
return Err(Error::Corrupt("nonzero padding after the section table"));
}
for i in 0..section_cnt {
let at = table_start as usize + i * ENTRY;
let kind = u16::from_le_bytes(bytes[at..at + 2].try_into().unwrap());
let align = u16::from_le_bytes(bytes[at + 2..at + 4].try_into().unwrap());
if align as usize != ALIGN {
return Err(Error::Corrupt("section alignment must be 64"));
}
if bytes[at + 4..at + 8] != [0u8; 4] {
return Err(Error::Corrupt("reserved section bytes must be zero"));
}
let offset = u64::from_le_bytes(bytes[at + 8..at + 16].try_into().unwrap());
let len = u64::from_le_bytes(bytes[at + 16..at + 24].try_into().unwrap());
let want = u64::from_le_bytes(bytes[at + 24..at + 32].try_into().unwrap());
if offset != expected {
return Err(Error::Corrupt("sections must be contiguous in table order"));
}
let end = offset
.checked_add(len)
.ok_or(Error::Corrupt("section length overflow"))?;
if end > file_len {
return Err(Error::Corrupt("section out of bounds"));
}
if sections.iter().any(|&(k, _, _, _)| k == kind) {
return Err(Error::Corrupt("duplicate section kind"));
}
expected = align_up(end);
if bytes[end as usize..expected.min(file_len) as usize]
.iter()
.any(|&b| b != 0)
{
return Err(Error::Corrupt("nonzero padding after a section"));
}
sections.push((kind, offset as usize, len as usize, want));
}
if expected != file_len {
return Err(Error::Corrupt("trailing bytes after the last section"));
}
Ok(Self {
bytes,
flags,
created_at,
config_len,
sections,
engine_ver_len,
})
}
pub fn config(&self) -> &'a [u8] {
&self.bytes[HEADER..HEADER + self.config_len]
}
pub fn section(&self, kind: u16) -> Option<&'a [u8]> {
self.sections
.iter()
.find(|&&(k, _, _, _)| k == kind)
.map(|&(_, start, len, _)| &self.bytes[start..start + len])
}
pub fn engine_ver(&self) -> &'a str {
core::str::from_utf8(&self.bytes[36..36 + self.engine_ver_len])
.expect("validated during parse")
}
pub fn scrub(&self) -> ScrubCursor<'a> {
self.scrub_with_budget(DEFAULT_SCRUB_BUDGET)
}
pub fn scrub_with_budget(&self, budget: usize) -> ScrubCursor<'a> {
ScrubCursor {
bytes: self.bytes,
sections: self
.sections
.iter()
.map(|&(_, s, l, w)| (s, l, w))
.collect(),
budget: budget.max(1),
pos: 0,
sec: 0,
sec_hash: Xxh3::new(),
file: Xxh3::new(),
done: false,
}
}
}
pub const DEFAULT_SCRUB_BUDGET: usize = 1 << 20;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ScrubProgress {
pub done_bytes: u64,
pub total_bytes: u64,
}
pub struct ScrubCursor<'a> {
bytes: &'a [u8],
sections: Vec<(usize, usize, u64)>,
budget: usize,
pos: usize,
sec: usize,
sec_hash: Xxh3,
file: Xxh3,
done: bool,
}
impl ScrubCursor<'_> {
fn feed_file(&mut self, from: usize, to: usize) {
let before_end = to.min(FILE_HASH_AT);
if from < before_end {
self.file.update(&self.bytes[from..before_end]);
}
let z_start = from.max(FILE_HASH_AT);
let z_end = to.min(FILE_HASH_AT + 8);
if z_start < z_end {
self.file.update(&[0u8; 8][..z_end - z_start]);
}
let after_start = from.max(FILE_HASH_AT + 8);
if after_start < to {
self.file.update(&self.bytes[after_start..to]);
}
}
fn complete_section(&mut self, want: u64) -> Result<(), Error> {
if self.sec_hash.digest() != want {
self.done = true;
return Err(Error::Corrupt("section checksum mismatch"));
}
self.sec += 1;
self.sec_hash = Xxh3::new();
Ok(())
}
}
impl Iterator for ScrubCursor<'_> {
type Item = Result<ScrubProgress, Error>;
fn next(&mut self) -> Option<Self::Item> {
if self.done {
return None;
}
let n = self.sections.len();
let file_len = self.bytes.len();
let mut budget_left = self.budget;
while budget_left > 0 && self.pos < file_len {
while self.sec < n {
let (start, len, want) = self.sections[self.sec];
if self.pos == start && len == 0 {
if let Err(e) = self.complete_section(want) {
return Some(Err(e));
}
} else {
break;
}
}
let (in_body, boundary) = match self.sections.get(self.sec) {
Some(&(start, len, _)) if self.pos >= start => (true, start + len),
Some(&(start, _, _)) => (false, start),
None => (false, file_len),
};
let end = (self.pos + budget_left).min(boundary);
self.feed_file(self.pos, end);
if in_body {
self.sec_hash.update(&self.bytes[self.pos..end]);
}
budget_left -= end - self.pos;
self.pos = end;
if in_body && self.pos == boundary {
let want = self.sections[self.sec].2;
if let Err(e) = self.complete_section(want) {
return Some(Err(e));
}
}
}
if self.pos == file_len {
while self.sec < n {
let want = self.sections[self.sec].2;
if let Err(e) = self.complete_section(want) {
return Some(Err(e));
}
}
self.done = true;
let stored = u64::from_le_bytes(
self.bytes[FILE_HASH_AT..FILE_HASH_AT + 8]
.try_into()
.unwrap(),
);
if stored != 0 && self.file.digest() != stored {
return Some(Err(Error::Corrupt("file checksum mismatch")));
}
}
Some(Ok(ScrubProgress {
done_bytes: self.pos as u64,
total_bytes: file_len as u64,
}))
}
}