Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 88 additions & 25 deletions src/snapshot/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,8 @@ use std::{
};

const MAGIC: &[u8; 8] = b"DUASNAP\0";
const VERSION: u16 = 1;
const VERSION: u16 = 2;
const VERSION_1: u16 = 1;
const HEADER_LEN: usize = 12;
const MAX_NAME_LEN: usize = 1024 * 1024;
const MAX_RECORD_LEN: usize = 2 * 1024 * 1024;
Expand Down Expand Up @@ -79,10 +80,12 @@ enum SnapshotReader<R> {

struct Decoder<R> {
reader: HashingReader<BufReader<SnapshotReader<R>>>,
allocation: Option<(usize, usize)>,
open_nodes: Vec<OpenNode>,
sibling_names: Vec<Vec<u8>>,
record: Vec<u8>,
node_count: u64,
name_bytes: u64,
total_size: u128,
total_entries: u64,
summary: Option<DecodeSummary>,
Expand Down Expand Up @@ -285,7 +288,7 @@ impl<R: Read + Seek> Replay<R> {
}
}

/// Write `traversal` as a deterministic version-1 snapshot, optionally compressed with zlib.
/// Write `traversal` as a deterministic version-2 snapshot, optionally compressed with zlib.
///
/// `roots` must contain the traversal's top-level nodes in original input order.
pub fn write(
Expand All @@ -305,13 +308,6 @@ pub fn write(
}

fn write_raw(writer: impl Write, traversal: &Traversal, roots: &[TreeIndex]) -> Result<()> {
let mut writer = gix::hash::io::Write::new(BufWriter::new(writer), gix::hash::Kind::Sha256);
let mut header = [0; HEADER_LEN];
header[..MAGIC.len()].copy_from_slice(MAGIC);
header[8..10].copy_from_slice(&VERSION.to_le_bytes());
header[10] = PATH_ENCODING;
writer.write_all(&header)?;

let mut seen_roots = HashSet::with_capacity(roots.len());
for &root in roots {
if !seen_roots.insert(root) {
Expand All @@ -320,6 +316,32 @@ fn write_raw(writer: impl Write, traversal: &Traversal, roots: &[TreeIndex]) ->
}
drop(seen_roots);

let mut stack = roots.to_vec();
let mut stored_node_count = 0u64;
let mut stored_name_bytes = 0u64;
while let Some(index) = stack.pop() {
let name = traversal
.tree
.native_name(index)
.ok_or_else(|| anyhow!("snapshot node {index:?} does not exist"))?;
stored_node_count = stored_node_count
.checked_add(1)
.context("snapshot contains too many nodes")?;
stored_name_bytes = stored_name_bytes
.checked_add(u64::try_from(name.len()).context("snapshot names are too long")?)
.context("snapshot names are too long")?;
stack.extend(traversal.tree.children(index));
}

let mut writer = gix::hash::io::Write::new(BufWriter::new(writer), gix::hash::Kind::Sha256);
let mut header = [0; HEADER_LEN];
header[..MAGIC.len()].copy_from_slice(MAGIC);
header[8..10].copy_from_slice(&VERSION.to_le_bytes());
header[10] = PATH_ENCODING;
writer.write_all(&header)?;
write_uleb128(&mut writer, u128::from(stored_node_count))?;
write_uleb128(&mut writer, u128::from(stored_name_bytes))?;

let mut stack = Vec::new();
let mut record = Vec::new();
let mut node_count = 0u64;
Expand Down Expand Up @@ -414,13 +436,20 @@ fn write_raw(writer: impl Write, traversal: &Traversal, roots: &[TreeIndex]) ->
Ok(())
}

/// Read and fully verify a version-1 snapshot before returning its traversal.
/// Read and fully verify a version-1 or version-2 snapshot before returning its traversal.
pub fn read(reader: impl Read) -> Result<Snapshot> {
let mut decoder = Decoder::new(reader)?;
let mut traversal = Traversal::new();
traversal.cost = Some(Duration::ZERO);
if let Some((node_count, name_bytes)) = decoder.allocation {
traversal
.tree
.try_reserve_exact(node_count, name_bytes)
.map_err(|err| anyhow!("could not preallocate snapshot tree: {err}"))?;
}
let mut parents = Vec::new();
let mut roots = Vec::new();
let summary = decode(reader, |entry| {
while let Some(entry) = decoder.next_entry(None)? {
parents.truncate(entry.depth);
let parent = parents.last().copied().unwrap_or(traversal.root_index);
let node = traversal
Expand All @@ -437,8 +466,10 @@ pub fn read(reader: impl Read) -> Result<Snapshot> {
.try_reserve(1)
.context("could not grow snapshot ancestor stack")?;
parents.push(node);
Ok(())
})?;
}
let summary = decoder
.summary
.context("snapshot ended without a summary")?;

traversal
.tree
Expand All @@ -450,17 +481,6 @@ pub fn read(reader: impl Read) -> Result<Snapshot> {
Ok(Snapshot { traversal, roots })
}

fn decode(
reader: impl Read,
mut on_entry: impl for<'entry> FnMut(DecodedEntry<'entry>) -> Result<()>,
) -> Result<DecodeSummary> {
let mut decoder = Decoder::new(reader)?;
while let Some(entry) = decoder.next_entry(None)? {
on_entry(entry)?;
}
Ok(decoder.summary.expect("end of snapshot stores its summary"))
}

impl<R: Read> Decoder<R> {
fn new(reader: R) -> Result<Self> {
let mut reader = HashingReader::new(BufReader::new(SnapshotReader::new(reader)?));
Expand All @@ -470,7 +490,7 @@ impl<R: Read> Decoder<R> {
bail!("invalid snapshot at byte 0: bad magic");
}
let version = u16::from_le_bytes([header[8], header[9]]);
if version != VERSION {
if version != VERSION_1 && version != VERSION {
bail!("invalid snapshot at byte 8: unsupported version {version}");
}
if header[10] != PATH_ENCODING {
Expand All @@ -483,12 +503,41 @@ impl<R: Read> Decoder<R> {
bail!("invalid snapshot at byte 11: unknown header flags");
}

let allocation = match version {
VERSION_1 => None,
VERSION => {
let node_count_offset = reader.offset;
let node_count = read_u64(&mut reader).map_err(|err| {
anyhow!("invalid snapshot integer at byte {node_count_offset}: {err}")
})?;
if node_count >= u64::from(u32::MAX) {
bail!("snapshot exceeds the tree node-index limit");
}
let name_bytes_offset = reader.offset;
let name_bytes = read_u64(&mut reader).map_err(|err| {
anyhow!("invalid snapshot integer at byte {name_bytes_offset}: {err}")
})?;
if name_bytes > u64::from(u32::MAX) {
bail!("snapshot exceeds the tree name-storage limit");
}
Some((
usize::try_from(node_count)
.context("snapshot node count exceeds this address space")?,
usize::try_from(name_bytes)
.context("snapshot name storage exceeds this address space")?,
))
}
_ => unreachable!("snapshot version was checked above"),
};

Ok(Self {
reader,
allocation,
open_nodes: Vec::new(),
sibling_names: Vec::new(),
record: Vec::new(),
node_count: 0,
name_bytes: 0,
total_size: 0,
total_entries: 0,
summary: None,
Expand Down Expand Up @@ -546,6 +595,15 @@ impl<R: Read> Decoder<R> {
if node_id >= u64::from(u32::MAX) {
bail!("snapshot exceeds the tree node-index limit");
}
self.name_bytes = self
.name_bytes
.checked_add(u64::try_from(native_name.len()).context("snapshot names are too long")?)
.context("snapshot names are too long")?;
if self.allocation.is_some_and(|(node_count, name_bytes)| {
node_id > node_count as u64 || self.name_bytes > name_bytes as u64
}) {
bail!("snapshot exceeds its declared storage requirements");
}

let sibling_ordinal = if parent_id == 0 {
self.open_nodes.clear();
Expand Down Expand Up @@ -635,6 +693,11 @@ impl<R: Read> Decoder<R> {
self.node_count
);
}
if self.allocation.is_some_and(|(node_count, name_bytes)| {
self.node_count != node_count as u64 || self.name_bytes != name_bytes as u64
}) {
bail!("snapshot does not match its declared storage requirements");
}

let actual = self.reader.hash.clone().try_finalize()?;
let mut expected = [0; DIGEST_LEN];
Expand Down
21 changes: 20 additions & 1 deletion src/snapshot/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ fn replay_rejects_source_changes_after_validation() {
let mut replay = Replay::new(Cursor::new(bytes.clone())).unwrap();

let mut changed = bytes[..bytes.len() - DIGEST_LEN].to_vec();
changed[16] = b'b';
changed[18] = b'b';
*replay.reader.get_mut() = rehash(changed);

assert!(
Expand Down Expand Up @@ -251,6 +251,25 @@ fn rejects_header_record_parent_and_name_errors() {
);
}

#[test]
fn rejects_incorrect_v2_storage_requirements() {
let mut traversal = Traversal::new();
let root_index = traversal.root_index;
let root = add(&mut traversal, root_index, "a", EntryData::default());
let bytes = encoded(&traversal, &[root]);

for offset in [HEADER_LEN, HEADER_LEN + 1] {
let mut changed = bytes[..bytes.len() - DIGEST_LEN].to_vec();
changed[offset] += 1;
let err = read(Cursor::new(rehash(changed))).unwrap_err();
assert!(
err.to_string()
.contains("does not match its declared storage requirements"),
"{err:#}"
);
}
}

#[test]
fn rejects_mismatched_record_length_footer_and_timestamp() {
let root = record(
Expand Down
23 changes: 23 additions & 0 deletions src/traverse.rs
Original file line number Diff line number Diff line change
Expand Up @@ -267,6 +267,29 @@ impl Tree {
}
}

pub(crate) fn try_reserve_exact(
&mut self,
additional_nodes: usize,
additional_name_bytes: usize,
) -> Result<(), TreeError> {
let node_count = self
.nodes
.len()
.checked_add(additional_nodes)
.ok_or(TreeError::Capacity)?;
let name_bytes = self
.names
.len()
.checked_add(additional_name_bytes)
.ok_or(TreeError::Capacity)?;
if node_count > u32::MAX as usize || name_bytes > u32::MAX as usize {
return Err(TreeError::Capacity);
}
self.nodes.try_reserve_exact(additional_nodes)?;
self.names.try_reserve_exact(additional_name_bytes)?;
Ok(())
}

/// Add a parentless node.
///
/// # Panics
Expand Down
27 changes: 23 additions & 4 deletions tests/snapshot.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,21 @@ use std::{
};

#[cfg(unix)]
const ONE_FILE: &[u8] = &[
const ONE_FILE_V1: &[u8] = &[
0x44, 0x55, 0x41, 0x53, 0x4e, 0x41, 0x50, 0x00, 0x01, 0x00, 0x00, 0x00, 0x07, 0x01, 0x00, 0x01,
0x61, 0x01, 0x00, 0x00, 0x00, 0x01, 0x69, 0x58, 0x7d, 0x2b, 0xa3, 0xd1, 0xef, 0x5b, 0x1d, 0x3d,
0x5b, 0xef, 0xc7, 0x03, 0x6e, 0x80, 0xfb, 0x3a, 0xf4, 0x1a, 0xea, 0xa0, 0xd0, 0x8f, 0x12, 0xcf,
0x15, 0x79, 0xc3, 0xb1, 0xa4, 0x7b,
];

#[cfg(unix)]
const ONE_FILE_V2: &[u8] = &[
0x44, 0x55, 0x41, 0x53, 0x4e, 0x41, 0x50, 0x00, 0x02, 0x00, 0x00, 0x00, 0x01, 0x01, 0x07, 0x01,
0x00, 0x01, 0x61, 0x01, 0x00, 0x00, 0x00, 0x01, 0x07, 0x2c, 0x93, 0x13, 0x2c, 0xe6, 0xc5, 0xed,
0x37, 0x7b, 0xa8, 0xe3, 0x21, 0x9b, 0x9d, 0x50, 0x97, 0xd4, 0x4a, 0xa6, 0x96, 0x1a, 0xfd, 0x26,
0x57, 0x54, 0xb8, 0x59, 0x5f, 0x0a, 0xad, 0x61,
];

fn add(
traversal: &mut Traversal,
parent: TreeIndex,
Expand Down Expand Up @@ -52,7 +60,7 @@ fn one_file_matches_the_golden_stream() {
);

let bytes = encoded(&traversal, &[root]);
assert_eq!(bytes, ONE_FILE);
assert_eq!(bytes, ONE_FILE_V2);

let snapshot = read(Cursor::new(bytes)).unwrap();
assert_eq!(snapshot.roots.len(), 1);
Expand All @@ -79,7 +87,18 @@ fn one_file_matches_the_golden_stream() {
Some(1)
);
assert_eq!(snapshot.traversal.cost, Some(Duration::ZERO));
assert_eq!(encoded(&snapshot.traversal, &snapshot.roots), ONE_FILE);
assert_eq!(encoded(&snapshot.traversal, &snapshot.roots), ONE_FILE_V2);
}

#[cfg(unix)]
#[test]
fn version_1_remains_readable() {
let snapshot = read(Cursor::new(ONE_FILE_V1)).unwrap();
assert_eq!(
snapshot.traversal.tree.name(snapshot.roots[0]).unwrap(),
Path::new("a")
);
assert_eq!(encoded(&snapshot.traversal, &snapshot.roots), ONE_FILE_V2);
}

#[test]
Expand Down Expand Up @@ -316,7 +335,7 @@ fn rejects_corruption_bad_digest_truncation_and_trailing_data() {
let bytes = encoded(&traversal, &[root]);

let mut corrupt = bytes.clone();
corrupt[16] = b'b';
corrupt[18] = b'b';
assert!(
read(Cursor::new(corrupt))
.unwrap_err()
Expand Down
Loading