use std::collections::BTreeSet;
use bytes::{Buf, BufMut};
use commonware_codec::{EncodeSize, Error as CodecError, Read, ReadExt, ReadRangeExt, Write};
use crate::hash::{HASH_LEN, Hash, Hasher, domain_digest, hash};
use crate::object::{ChunkedBlob, Tree, TreeEntry};
const CHUNKED_TYPE_DOMAIN: &[u8] = b"mkit.chunked\x00";
const TREE_TYPE_DOMAIN: &[u8] = b"mkit.tree\x00";
const CBLOB_META_DOMAIN: &[u8] = b"mkit-cblob-meta-v1";
const TREE_ENTRY_DOMAIN: &[u8] = b"mkit-tree-entry-v1";
pub const MAX_LEVELS: usize = u32::BITS as usize;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ObjectKind {
Tree,
ChunkedBlob,
}
impl ObjectKind {
fn type_domain(self) -> &'static [u8] {
match self {
Self::Tree => TREE_TYPE_DOMAIN,
Self::ChunkedBlob => CHUNKED_TYPE_DOMAIN,
}
}
}
#[must_use]
pub fn wrap_id(kind: ObjectKind, inner_root: &Hash) -> Hash {
domain_digest(kind.type_domain(), inner_root)
}
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum MerkleError {
#[error("merkle position {0} is out of range")]
PositionOutOfRange(u32),
#[error("merkle position {0} is duplicated")]
DuplicatePosition(u32),
#[error("no merkle positions given")]
NoPositions,
#[error("merkle range start {start} is greater than end {end}")]
InvalidRange {
start: u32,
end: u32,
},
#[error("merkle proof is malformed")]
MalformedProof,
#[error("merkle proof is unaligned with the requested position(s)")]
UnalignedProof,
#[error("merkle proof verification failed")]
VerificationFailed,
}
impl From<CodecError> for MerkleError {
fn from(_: CodecError) -> Self {
Self::MalformedProof
}
}
fn u32_of(n: usize) -> u32 {
u32::try_from(n).expect("merkle leaf/index count fits u32 (objects capped at 1M)")
}
fn h2(a: &[u8], b: &[u8]) -> Hash {
let mut h = Hasher::new();
h.update(a).update(b);
h.finalize()
}
fn position_leaf(index: u32, leaf: &Hash) -> Hash {
h2(&index.to_be_bytes(), leaf)
}
fn levels_in_tree(leaf_count: u32) -> usize {
(u32::BITS - leaf_count.saturating_sub(1).leading_zeros() + 1) as usize
}
struct BmtTree {
leaf_count: u32,
empty: bool,
levels: Vec<Vec<Hash>>,
root: Hash,
}
fn build_bmt(leaves: &[Hash]) -> BmtTree {
let leaf_count = u32_of(leaves.len());
let empty = leaves.is_empty();
let position_hashed: Vec<Hash> = if empty {
vec![hash(b"")]
} else {
leaves
.iter()
.enumerate()
.map(|(i, l)| position_leaf(u32_of(i), l))
.collect()
};
let mut levels = vec![position_hashed];
while levels.last().expect("levels is never empty").len() > 1 {
let cur = levels.last().expect("levels is never empty");
let mut next = Vec::with_capacity(cur.len().div_ceil(2));
for pair in cur.chunks(2) {
let right = if pair.len() == 2 { &pair[1] } else { &pair[0] };
next.push(h2(&pair[0], right));
}
levels.push(next);
}
let tree_root = levels.last().expect("levels is never empty")[0];
let root = h2(&leaf_count.to_be_bytes(), &tree_root);
BmtTree {
leaf_count,
empty,
levels,
root,
}
}
fn siblings_required_for_multi_proof(
leaf_count: u32,
positions: impl IntoIterator<Item = u32>,
) -> Result<BTreeSet<(usize, usize)>, MerkleError> {
let mut current = BTreeSet::new();
for pos in positions {
if pos >= leaf_count {
return Err(MerkleError::PositionOutOfRange(pos));
}
if !current.insert(pos as usize) {
return Err(MerkleError::DuplicatePosition(pos));
}
}
if current.is_empty() {
return Err(MerkleError::NoPositions);
}
let mut sibling_positions = BTreeSet::new();
let levels_count = levels_in_tree(leaf_count);
let mut level_size = leaf_count as usize;
for level in 0..levels_count.saturating_sub(1) {
for &index in ¤t {
let sibling_index = if index.is_multiple_of(2) {
if index + 1 < level_size {
index + 1
} else {
index
}
} else {
index - 1
};
if sibling_index != index && !current.contains(&sibling_index) {
sibling_positions.insert((level, sibling_index));
}
}
current = current.iter().map(|idx| idx / 2).collect();
level_size = level_size.div_ceil(2);
}
Ok(sibling_positions)
}
fn siblings_required_for_range_proof(
leaf_count: u32,
start: u32,
end: u32,
) -> Result<BTreeSet<(usize, usize)>, MerkleError> {
if leaf_count == 0 {
return Err(MerkleError::NoPositions);
}
if start > end {
return Err(MerkleError::InvalidRange { start, end });
}
if start >= leaf_count {
return Err(MerkleError::PositionOutOfRange(start));
}
if end >= leaf_count {
return Err(MerkleError::PositionOutOfRange(end));
}
let mut sibling_positions = BTreeSet::new();
let levels_count = levels_in_tree(leaf_count);
let mut level_start = start as usize;
let mut level_end = end as usize;
let mut level_size = leaf_count as usize;
for level in 0..levels_count.saturating_sub(1) {
if !level_start.is_multiple_of(2) {
sibling_positions.insert((level, level_start - 1));
}
if level_end.is_multiple_of(2) {
let right = level_end + 1;
if right < level_size {
sibling_positions.insert((level, right));
}
}
level_start /= 2;
level_end /= 2;
level_size = level_size.div_ceil(2);
}
Ok(sibling_positions)
}
impl BmtTree {
fn proof(&self, position: u32) -> Result<Proof, MerkleError> {
self.multi_proof(core::iter::once(position))
}
fn range_proof(&self, start: u32, end: u32) -> Result<Proof, MerkleError> {
if self.empty {
return Err(MerkleError::PositionOutOfRange(start));
}
if start > end {
return Err(MerkleError::InvalidRange { start, end });
}
let leaf_count = self.leaf_count;
if start >= leaf_count {
return Err(MerkleError::PositionOutOfRange(start));
}
if end >= leaf_count {
return Err(MerkleError::PositionOutOfRange(end));
}
let sibling_positions = siblings_required_for_range_proof(leaf_count, start, end)?;
let siblings = sibling_positions
.iter()
.map(|&(level, index)| self.levels[level][index])
.collect();
Ok(Proof {
leaf_count,
siblings,
})
}
fn multi_proof(&self, positions: impl IntoIterator<Item = u32>) -> Result<Proof, MerkleError> {
let mut positions = positions.into_iter().peekable();
let first = *positions.peek().ok_or(MerkleError::NoPositions)?;
if self.empty {
return Err(MerkleError::PositionOutOfRange(first));
}
let leaf_count = self.leaf_count;
let sibling_positions = siblings_required_for_multi_proof(leaf_count, positions)?;
let siblings = sibling_positions
.iter()
.map(|&(level, index)| self.levels[level][index])
.collect();
Ok(Proof {
leaf_count,
siblings,
})
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Proof {
pub leaf_count: u32,
pub siblings: Vec<Hash>,
}
impl Write for Proof {
fn write(&self, writer: &mut impl BufMut) {
self.leaf_count.write(writer);
self.siblings.write(writer);
}
}
impl EncodeSize for Proof {
fn encode_size(&self) -> usize {
self.leaf_count.encode_size() + self.siblings.encode_size()
}
}
impl Read for Proof {
type Cfg = usize;
fn read_cfg(reader: &mut impl Buf, max_items: &Self::Cfg) -> Result<Self, CodecError> {
let leaf_count = u32::read(reader)?;
let max_siblings = max_items.saturating_mul(MAX_LEVELS);
let siblings = Vec::<Hash>::read_range(reader, ..=max_siblings)?;
Ok(Self {
leaf_count,
siblings,
})
}
}
impl Proof {
#[must_use]
pub fn encode(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(self.encode_size());
self.write(&mut out);
out
}
pub fn decode(bytes: &[u8], max_items: usize) -> Result<Self, MerkleError> {
let mut buf = bytes;
let proof = Self::read_cfg(&mut buf, &max_items)?;
if buf.has_remaining() {
return Err(MerkleError::MalformedProof);
}
Ok(proof)
}
pub(crate) fn reconstruct_element_root(
&self,
leaf: &Hash,
mut position: u32,
) -> Result<Hash, MerkleError> {
if position >= self.leaf_count {
return Err(MerkleError::PositionOutOfRange(position));
}
let mut computed = position_leaf(position, leaf);
let mut level_size = self.leaf_count as usize;
let mut sibling_iter = self.siblings.iter();
while level_size > 1 {
let is_last_odd = position.is_multiple_of(2) && position as usize + 1 >= level_size;
let (left, right) = if is_last_odd {
(computed, computed)
} else if position.is_multiple_of(2) {
let sib = *sibling_iter.next().ok_or(MerkleError::UnalignedProof)?;
(computed, sib)
} else {
let sib = *sibling_iter.next().ok_or(MerkleError::UnalignedProof)?;
(sib, computed)
};
computed = h2(&left, &right);
position /= 2;
level_size = level_size.div_ceil(2);
}
if sibling_iter.next().is_some() {
return Err(MerkleError::UnalignedProof);
}
Ok(h2(&self.leaf_count.to_be_bytes(), &computed))
}
pub(crate) fn reconstruct_multi_root(
&self,
elements: &[(Hash, u32)],
) -> Result<Hash, MerkleError> {
if elements.is_empty() {
return Err(MerkleError::NoPositions);
}
for (_, position) in elements {
if *position >= self.leaf_count {
return Err(MerkleError::PositionOutOfRange(*position));
}
}
let mut sorted: Vec<(u32, Hash)> = elements
.iter()
.map(|(leaf, pos)| (*pos, position_leaf(*pos, leaf)))
.collect();
sorted.sort_unstable_by_key(|(pos, _)| *pos);
for i in 1..sorted.len() {
if sorted[i - 1].0 == sorted[i].0 {
return Err(MerkleError::DuplicatePosition(sorted[i].0));
}
}
let levels = levels_in_tree(self.leaf_count);
let mut level_size = self.leaf_count;
let mut sibling_iter = self.siblings.iter();
let mut current = sorted;
for _ in 0..levels.saturating_sub(1) {
let mut next_level: Vec<(u32, Hash)> = Vec::with_capacity(current.len().div_ceil(2));
let mut idx = 0;
while idx < current.len() {
let (pos, digest) = current[idx];
let parent_pos = pos / 2;
let (left, right) = if pos.is_multiple_of(2) {
let left = digest;
let right = if idx + 1 < current.len() && current[idx + 1].0 == pos + 1 {
idx += 1;
current[idx].1
} else if pos + 1 >= level_size {
left
} else {
*sibling_iter.next().ok_or(MerkleError::UnalignedProof)?
};
(left, right)
} else {
let right = digest;
let left = *sibling_iter.next().ok_or(MerkleError::UnalignedProof)?;
(left, right)
};
next_level.push((parent_pos, h2(&left, &right)));
idx += 1;
}
current = next_level;
level_size = level_size.div_ceil(2);
}
if sibling_iter.next().is_some() {
return Err(MerkleError::UnalignedProof);
}
if current.len() != 1 {
return Err(MerkleError::UnalignedProof);
}
Ok(h2(&self.leaf_count.to_be_bytes(), ¤t[0].1))
}
fn reconstruct_range_root(&self, position: u32, leaves: &[Hash]) -> Result<Hash, MerkleError> {
if leaves.is_empty() && position != 0 {
return Err(MerkleError::PositionOutOfRange(position));
}
if !leaves.is_empty() {
let leaves_len = u32_of(leaves.len());
let end = position
.checked_add(leaves_len - 1)
.ok_or(MerkleError::PositionOutOfRange(position))?;
if end >= self.leaf_count {
return Err(MerkleError::PositionOutOfRange(end));
}
}
let elements: Vec<(Hash, u32)> = leaves
.iter()
.enumerate()
.map(|(i, l)| (*l, position + u32_of(i)))
.collect();
self.reconstruct_multi_root(&elements)
}
#[cfg(test)]
pub(crate) fn verify_element_inclusion(
&self,
leaf: &Hash,
position: u32,
inner_root: &Hash,
) -> Result<(), MerkleError> {
let got = self.reconstruct_element_root(leaf, position)?;
if &got == inner_root {
Ok(())
} else {
Err(MerkleError::VerificationFailed)
}
}
#[cfg(test)]
pub(crate) fn verify_multi_inclusion(
&self,
elements: &[(Hash, u32)],
inner_root: &Hash,
) -> Result<(), MerkleError> {
let got = self.reconstruct_multi_root(elements)?;
if &got == inner_root {
Ok(())
} else {
Err(MerkleError::VerificationFailed)
}
}
#[cfg(test)]
pub(crate) fn verify_range_inclusion(
&self,
position: u32,
leaves: &[Hash],
inner_root: &Hash,
) -> Result<(), MerkleError> {
let got = self.reconstruct_range_root(position, leaves)?;
if &got == inner_root {
Ok(())
} else {
Err(MerkleError::VerificationFailed)
}
}
}
pub(crate) fn chunked_meta_leaf_raw(total_size: u64, chunk_size: u32) -> Hash {
let mut body = [0u8; 12];
body[..8].copy_from_slice(&total_size.to_le_bytes());
body[8..].copy_from_slice(&chunk_size.to_le_bytes());
domain_digest(CBLOB_META_DOMAIN, &body)
}
fn chunked_meta_leaf(cb: &ChunkedBlob) -> Hash {
chunked_meta_leaf_raw(cb.total_size, cb.chunk_size)
}
pub(crate) fn tree_entry_leaf(e: &TreeEntry) -> Hash {
let mut body = Vec::with_capacity(4 + e.name.len() + 1 + HASH_LEN);
body.extend_from_slice(&u32_of(e.name.len()).to_le_bytes());
body.extend_from_slice(&e.name);
body.push(e.mode as u8);
body.extend_from_slice(&e.object_hash);
domain_digest(TREE_ENTRY_DOMAIN, &body)
}
fn chunked_leaves(cb: &ChunkedBlob) -> Vec<Hash> {
let mut leaves = Vec::with_capacity(1 + cb.chunks.len());
leaves.push(chunked_meta_leaf(cb));
leaves.extend_from_slice(&cb.chunks);
leaves
}
fn tree_leaves(tree: &Tree) -> Vec<Hash> {
tree.entries.iter().map(tree_entry_leaf).collect()
}
#[must_use]
pub fn chunked_inner_root(cb: &ChunkedBlob) -> Hash {
build_bmt(&chunked_leaves(cb)).root
}
#[must_use]
pub fn tree_inner_root(tree: &Tree) -> Hash {
build_bmt(&tree_leaves(tree)).root
}
#[must_use]
pub fn compute_chunked_id(cb: &ChunkedBlob) -> Hash {
wrap_id(ObjectKind::ChunkedBlob, &chunked_inner_root(cb))
}
#[must_use]
pub fn compute_tree_id(tree: &Tree) -> Hash {
wrap_id(ObjectKind::Tree, &tree_inner_root(tree))
}
pub const TREE_EMPTY_ID: Hash = [
0x1a, 0xb8, 0xd0, 0x78, 0x8b, 0x29, 0xfe, 0x59, 0x92, 0x01, 0x1e, 0x64, 0xd6, 0xc9, 0x22, 0xec,
0x93, 0xf4, 0x24, 0x8b, 0x37, 0x55, 0xb9, 0x2b, 0x15, 0xb0, 0x7e, 0x66, 0x4c, 0xb1, 0x56, 0x52,
];
#[must_use]
pub fn chunk_position(cb: &ChunkedBlob, chunk_hash: &Hash) -> Option<u32> {
cb.chunks
.iter()
.position(|c| c == chunk_hash)
.map(|i| u32_of(i + 1))
}
#[must_use]
pub fn tree_entry_position(tree: &Tree, name: &[u8]) -> Option<u32> {
tree.entries.iter().position(|e| e.name == name).map(u32_of)
}
pub(crate) fn proof_encoded_size(
leaf_count: u32,
positions: impl IntoIterator<Item = u32>,
) -> Result<usize, MerkleError> {
let siblings = siblings_required_for_multi_proof(leaf_count, positions)?.len();
let mut count = siblings;
let mut prefix = 1;
while count >= 128 {
prefix += 1;
count >>= 7;
}
Ok(4 + prefix + siblings * 32)
}
pub fn build_chunk_proof(cb: &ChunkedBlob, position: u32) -> Result<Proof, MerkleError> {
build_bmt(&chunked_leaves(cb)).proof(position)
}
pub fn build_chunks_range_proof(
cb: &ChunkedBlob,
start: u32,
end: u32,
) -> Result<Proof, MerkleError> {
build_bmt(&chunked_leaves(cb)).range_proof(start, end)
}
pub fn build_chunks_multi_proof(
cb: &ChunkedBlob,
positions: impl IntoIterator<Item = u32>,
) -> Result<Proof, MerkleError> {
build_bmt(&chunked_leaves(cb)).multi_proof(positions)
}
pub fn build_tree_entry_proof(tree: &Tree, position: u32) -> Result<Proof, MerkleError> {
build_bmt(&tree_leaves(tree)).proof(position)
}
pub fn build_tree_entries_range_proof(
tree: &Tree,
start: u32,
end: u32,
) -> Result<Proof, MerkleError> {
build_bmt(&tree_leaves(tree)).range_proof(start, end)
}
pub fn build_tree_entries_multi_proof(
tree: &Tree,
positions: impl IntoIterator<Item = u32>,
) -> Result<Proof, MerkleError> {
build_bmt(&tree_leaves(tree)).multi_proof(positions)
}
pub fn verify_tree_entry(
tree_id: &Hash,
entry: &TreeEntry,
position: u32,
proof: &Proof,
) -> Result<(), MerkleError> {
let root = proof.reconstruct_element_root(&tree_entry_leaf(entry), position)?;
check_wrapped(ObjectKind::Tree, &root, tree_id)
}
pub fn verify_tree_entries_range(
tree_id: &Hash,
start: u32,
entries: &[TreeEntry],
proof: &Proof,
) -> Result<(), MerkleError> {
let leaves: Vec<Hash> = entries.iter().map(tree_entry_leaf).collect();
let root = proof.reconstruct_range_root(start, &leaves)?;
check_wrapped(ObjectKind::Tree, &root, tree_id)
}
pub fn verify_tree_entries_multi(
tree_id: &Hash,
entries: &[(TreeEntry, u32)],
proof: &Proof,
) -> Result<(), MerkleError> {
let elements: Vec<(Hash, u32)> = entries
.iter()
.map(|(e, pos)| (tree_entry_leaf(e), *pos))
.collect();
let root = proof.reconstruct_multi_root(&elements)?;
check_wrapped(ObjectKind::Tree, &root, tree_id)
}
pub fn verify_chunk(
chunked_id: &Hash,
chunk_hash: &Hash,
position: u32,
proof: &Proof,
) -> Result<(), MerkleError> {
if position == 0 {
return Err(MerkleError::PositionOutOfRange(0));
}
let root = proof.reconstruct_element_root(chunk_hash, position)?;
check_wrapped(ObjectKind::ChunkedBlob, &root, chunked_id)
}
pub fn verify_chunks_range(
chunked_id: &Hash,
start: u32,
chunk_hashes: &[Hash],
proof: &Proof,
) -> Result<(), MerkleError> {
if start == 0 {
return Err(MerkleError::PositionOutOfRange(0));
}
let root = proof.reconstruct_range_root(start, chunk_hashes)?;
check_wrapped(ObjectKind::ChunkedBlob, &root, chunked_id)
}
pub fn verify_chunks_multi(
chunked_id: &Hash,
chunks: &[(Hash, u32)],
proof: &Proof,
) -> Result<(), MerkleError> {
if chunks.iter().any(|(_, pos)| *pos == 0) {
return Err(MerkleError::PositionOutOfRange(0));
}
let elements: Vec<(Hash, u32)> = chunks.to_vec();
let root = proof.reconstruct_multi_root(&elements)?;
check_wrapped(ObjectKind::ChunkedBlob, &root, chunked_id)
}
pub(crate) fn verify_chunk_with_meta_leaf(
chunked_id: &Hash,
total_size: u64,
chunk_size: u32,
chunk_hash: &Hash,
chunk_position: u32,
proof: &Proof,
) -> Result<(), MerkleError> {
if chunk_position == 0 {
return Err(MerkleError::PositionOutOfRange(0));
}
let meta_leaf = chunked_meta_leaf_raw(total_size, chunk_size);
let elements = [(meta_leaf, 0u32), (*chunk_hash, chunk_position)];
let root = proof.reconstruct_multi_root(&elements)?;
check_wrapped(ObjectKind::ChunkedBlob, &root, chunked_id)
}
fn check_wrapped(kind: ObjectKind, root: &Hash, expected_id: &Hash) -> Result<(), MerkleError> {
if &wrap_id(kind, root) == expected_id {
Ok(())
} else {
Err(MerkleError::VerificationFailed)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::object::EntryMode;
fn cb(total: u64, chunk_size: u32, chunks: &[u8]) -> ChunkedBlob {
ChunkedBlob {
total_size: total,
chunk_size,
chunks: chunks.iter().map(|b| [*b; 32]).collect(),
}
}
fn entry(name: &[u8], mode: EntryMode, h: u8) -> TreeEntry {
TreeEntry {
name: name.to_vec(),
mode,
object_hash: [h; 32],
}
}
fn tree(entries: Vec<TreeEntry>) -> Tree {
Tree { entries }
}
#[test]
fn id_changes_when_a_leaf_changes() {
let a = cb(100, 0, &[1, 2, 3]);
let b = cb(100, 0, &[1, 2, 4]);
assert_ne!(compute_chunked_id(&a), compute_chunked_id(&b));
}
#[test]
fn id_changes_when_leaf_count_changes() {
let a = cb(100, 0, &[1, 2, 3]);
let b = cb(100, 0, &[1, 2, 3, 3]); assert_ne!(compute_chunked_id(&a), compute_chunked_id(&b));
}
#[test]
fn chunked_id_changes_when_metadata_changes() {
let a = cb(100, 0, &[1, 2, 3]);
let b = cb(101, 0, &[1, 2, 3]);
let c = cb(100, 64, &[1, 2, 3]);
assert_ne!(compute_chunked_id(&a), compute_chunked_id(&b));
assert_ne!(compute_chunked_id(&a), compute_chunked_id(&c));
}
#[test]
fn tree_ordering_matters() {
let a = tree(vec![
entry(b"a", EntryMode::Blob, 1),
entry(b"b", EntryMode::Blob, 2),
]);
let b = tree(vec![
entry(b"b", EntryMode::Blob, 2),
entry(b"a", EntryMode::Blob, 1),
]);
assert_ne!(compute_tree_id(&a), compute_tree_id(&b));
}
#[test]
fn empty_tree_id_matches_constant() {
let got = compute_tree_id(&tree(vec![]));
assert_eq!(got, TREE_EMPTY_ID, "update TREE_EMPTY_ID to {got:02x?}");
}
#[test]
fn type_binding_no_cross_collisions() {
let empty_tree = compute_tree_id(&tree(vec![]));
let empty_cblob = compute_chunked_id(&cb(0, 0, &[]));
assert_ne!(empty_tree, empty_cblob);
let t = tree(vec![entry(b"x", EntryMode::Blob, 9)]);
let c = cb(10, 0, &[9]);
assert_ne!(compute_tree_id(&t), compute_chunked_id(&c));
}
#[test]
fn id_ne_flat_blake3_of_serialized_bytes() {
let c = cb(100, 0, &[1, 2, 3]);
let serialized =
crate::serialize::serialize(&crate::object::Object::ChunkedBlob(c.clone())).unwrap();
assert_ne!(compute_chunked_id(&c), crate::hash::hash(&serialized));
}
#[test]
fn chunk_position_offsets_by_one() {
let c = cb(100, 0, &[7, 8, 9]);
assert_eq!(chunk_position(&c, &[7; 32]), Some(1));
assert_eq!(chunk_position(&c, &[9; 32]), Some(3));
assert_eq!(chunk_position(&c, &[0; 32]), None);
}
#[test]
fn chunk_inclusion_proof_round_trips() {
let c = cb(100, 0, &[10, 20, 30, 40]);
let id = compute_chunked_id(&c);
for (idx, byte) in [(0usize, 10u8), (2, 30), (3, 40)] {
let pos = chunk_position(&c, &[byte; 32]).unwrap();
assert_eq!(pos, u32_of(idx) + 1);
let proof = build_chunk_proof(&c, pos).unwrap();
verify_chunk(&id, &[byte; 32], pos, &proof).unwrap();
assert!(verify_chunk(&id, &[0xFF; 32], pos, &proof).is_err());
}
}
#[test]
fn verify_chunk_rejects_meta_leaf() {
let c = cb(100, 0, &[10, 20, 30]);
let id = compute_chunked_id(&c);
let meta_leaf = chunked_meta_leaf(&c);
let proof = build_chunk_proof(&c, 0).unwrap();
proof
.verify_element_inclusion(&meta_leaf, 0, &chunked_inner_root(&c))
.expect("meta leaf is a valid BMT element at position 0");
assert_eq!(
verify_chunk(&id, &meta_leaf, 0, &proof),
Err(MerkleError::PositionOutOfRange(0))
);
}
#[test]
fn verify_chunk_with_meta_leaf_round_trips_and_rejects_forgery() {
let c = cb(100, 0, &[10, 20, 30]);
let id = compute_chunked_id(&c);
let pos = chunk_position(&c, &[20; 32]).unwrap(); let proof = build_chunks_multi_proof(&c, [0, pos]).unwrap();
verify_chunk_with_meta_leaf(&id, 100, 0, &[20; 32], pos, &proof).unwrap();
assert!(verify_chunk_with_meta_leaf(&id, 101, 0, &[20; 32], pos, &proof).is_err());
assert!(verify_chunk_with_meta_leaf(&id, 100, 64, &[20; 32], pos, &proof).is_err());
assert!(verify_chunk_with_meta_leaf(&id, 100, 0, &[0xFF; 32], pos, &proof).is_err());
assert_eq!(
verify_chunk_with_meta_leaf(&id, 100, 0, &[10; 32], 0, &proof),
Err(MerkleError::PositionOutOfRange(0))
);
}
#[test]
fn tree_inclusion_proof_round_trips() {
let t = tree(vec![
entry(b"a", EntryMode::Blob, 1),
entry(b"b", EntryMode::Tree, 2),
entry(b"c", EntryMode::Executable, 3),
]);
let id = compute_tree_id(&t);
let pos = tree_entry_position(&t, b"b").unwrap();
assert_eq!(pos, 1);
let proof = build_tree_entry_proof(&t, pos).unwrap();
verify_tree_entry(&id, &t.entries[1], pos, &proof).unwrap();
let wrong = entry(b"b", EntryMode::Blob, 2);
assert!(verify_tree_entry(&id, &wrong, pos, &proof).is_err());
}
#[test]
fn range_and_multi_proofs_round_trip() {
let t = tree(
(0..9)
.map(|i| entry(&[b'a' + i], EntryMode::Blob, i))
.collect(),
);
let id = compute_tree_id(&t);
let range_proof = build_tree_entries_range_proof(&t, 2, 4).unwrap();
verify_tree_entries_range(&id, 2, &t.entries[2..=4], &range_proof).unwrap();
assert!(verify_tree_entries_range(&id, 2, &t.entries[2..4], &range_proof).is_err());
let multi_proof = build_tree_entries_multi_proof(&t, [0, 4, 8]).unwrap();
let elements = [
(t.entries[0].clone(), 0),
(t.entries[4].clone(), 4),
(t.entries[8].clone(), 8),
];
verify_tree_entries_multi(&id, &elements, &multi_proof).unwrap();
let wrong_elements = [
(t.entries[0].clone(), 0),
(t.entries[4].clone(), 5), (t.entries[8].clone(), 8),
];
assert!(verify_tree_entries_multi(&id, &wrong_elements, &multi_proof).is_err());
}
#[test]
fn chunk_range_and_multi_proofs_reject_position_zero() {
let c = cb(100, 0, &[1, 2, 3]);
let id = compute_chunked_id(&c);
let proof = build_chunks_range_proof(&c, 0, 1).unwrap();
assert_eq!(
verify_chunks_range(&id, 0, &c.chunks[..2], &proof),
Err(MerkleError::PositionOutOfRange(0))
);
let multi = build_chunks_multi_proof(&c, [0, 2]).unwrap();
assert_eq!(
verify_chunks_multi(&id, &[(c.chunks[0], 0), (c.chunks[1], 2)], &multi),
Err(MerkleError::PositionOutOfRange(0))
);
}
#[test]
fn single_leaf_tree_proof_has_no_siblings() {
let t = tree(vec![entry(b"only", EntryMode::Blob, 1)]);
let id = compute_tree_id(&t);
let proof = build_tree_entry_proof(&t, 0).unwrap();
assert!(
proof.siblings.is_empty(),
"single-leaf tree proof must have zero siblings"
);
verify_tree_entry(&id, &t.entries[0], 0, &proof).unwrap();
}
#[test]
fn verify_rejects_empty_element_set() {
assert_eq!(
verify_tree_entries_range(&TREE_EMPTY_ID, 0, &[], &Proof::default()),
Err(MerkleError::NoPositions)
);
assert_eq!(
verify_tree_entries_multi(&TREE_EMPTY_ID, &[], &Proof::default()),
Err(MerkleError::NoPositions)
);
let t = tree(vec![entry(b"a", EntryMode::Blob, 1)]);
let id = compute_tree_id(&t);
let real_proof = build_tree_entry_proof(&t, 0).unwrap();
assert_eq!(
verify_tree_entries_range(&id, 0, &[], &real_proof),
Err(MerkleError::NoPositions)
);
assert_eq!(
verify_tree_entries_multi(&id, &[], &real_proof),
Err(MerkleError::NoPositions)
);
}
#[test]
fn empty_tree_builders_refuse_position_zero() {
let empty = tree(vec![]);
assert_eq!(
build_tree_entries_range_proof(&empty, 0, 0),
Err(MerkleError::PositionOutOfRange(0))
);
assert_eq!(
build_tree_entry_proof(&empty, 0),
Err(MerkleError::PositionOutOfRange(0))
);
assert_eq!(
build_tree_entries_multi_proof(&empty, [0]),
Err(MerkleError::PositionOutOfRange(0))
);
}
#[test]
fn out_of_range_position_rejected() {
let c = cb(10, 0, &[1]);
assert_eq!(
build_chunk_proof(&c, 2),
Err(MerkleError::PositionOutOfRange(2))
);
}
#[test]
fn proof_round_trips_through_encode_decode() {
let t = tree(
(0..7)
.map(|i| entry(&[b'a' + i], EntryMode::Blob, i))
.collect(),
);
let proof = build_tree_entry_proof(&t, 6).unwrap();
let bytes = proof.encode();
let decoded = Proof::decode(&bytes, 1).unwrap();
assert_eq!(proof, decoded);
let mut truncated_extra = bytes.clone();
truncated_extra.push(0);
assert_eq!(
Proof::decode(&truncated_extra, 1),
Err(MerkleError::MalformedProof)
);
assert_eq!(
Proof::decode(&bytes[..bytes.len() - 1], 1),
Err(MerkleError::MalformedProof)
);
assert_eq!(Proof::decode(&bytes, 0), Err(MerkleError::MalformedProof));
}
#[test]
fn vendored_root_matches_commonware() {
use commonware_cryptography::blake3::{Blake3, Digest};
use commonware_storage::bmt::Builder;
for n in [1usize, 2, 3, 4, 5, 7, 8, 9, 16, 33] {
let leaves: Vec<Hash> = (0..n)
.map(|i| hash(&[u8::try_from(i % 256).unwrap(); 4]))
.collect();
let mut builder = Builder::<Blake3>::new(n);
for l in &leaves {
builder.add(&Digest(*l));
}
let cw_root = builder.build().root().0;
let ours = build_bmt(&leaves).root;
assert_eq!(ours, cw_root, "vendored BMT root diverged at n={n}");
}
}
fn tree_and_positions()
-> impl proptest::strategy::Strategy<Value = (usize, u32, (u32, u32), Vec<u32>)> {
use proptest::prelude::*;
let boundary_counts: Vec<usize> = [1usize, 2, 4, 8, 16, 32, 64, 128, 256]
.into_iter()
.flat_map(|p| [p.saturating_sub(1).max(1), p, p + 1])
.chain([
3usize, 5, 7, 9, 15, 17, 31, 33, 63, 65, 127, 129, 255, 257, 300,
])
.collect();
prop_oneof![
3 => 1usize..=300,
2 => proptest::sample::select(boundary_counts),
]
.prop_flat_map(|n| {
let n_u32 =
u32::try_from(n).expect("n is bounded well under u32::MAX by the strategy above");
(
Just(n),
0..n_u32,
(0..n_u32).prop_flat_map(move |s| (Just(s), s..n_u32)),
proptest::collection::vec(0..n_u32, 1..=n.min(6)).prop_map(|mut v| {
v.sort_unstable();
v.dedup();
v
}),
)
})
}
proptest::proptest! {
#![proptest_config(proptest::prelude::ProptestConfig::with_cases(400))]
#[test]
fn proofs_match_commonware((n, pos, (start, end), positions) in tree_and_positions()) {
use commonware_cryptography::blake3::{Blake3, Digest};
use commonware_storage::bmt::Builder;
let leaves: Vec<Hash> = (0..n).map(|i| hash(&(i as u64).to_le_bytes())).collect();
let ours = build_bmt(&leaves);
let mut cw_builder = Builder::<Blake3>::new(n);
for l in &leaves {
cw_builder.add(&Digest(*l));
}
let cw_tree = cw_builder.build();
let cw_root = cw_tree.root();
assert_eq!(ours.root, cw_root.0, "root mismatch at n={n}");
let our_proof = ours.proof(pos).unwrap();
let cw_proof = cw_tree.proof(pos).unwrap();
assert_eq!(
our_proof.encode(),
commonware_codec::Encode::encode(&cw_proof).to_vec(),
"single-proof bytes diverged at n={n} pos={pos}"
);
cw_proof
.verify_element_inclusion::<Blake3>(&Digest(leaves[pos as usize]), pos, &cw_root)
.expect("upstream must accept its own proof");
our_proof
.verify_element_inclusion(&leaves[pos as usize], pos, &ours.root)
.expect("ours must accept its own proof");
let cw_bytes = commonware_codec::Encode::encode(&cw_proof).to_vec();
let our_decoded = Proof::decode(&cw_bytes, 1).unwrap();
our_decoded
.verify_element_inclusion(&leaves[pos as usize], pos, &ours.root)
.expect("ours must accept upstream's proof bytes");
let our_bytes = our_proof.encode();
let mut our_bytes_buf: &[u8] = &our_bytes;
let cw_decoded =
<commonware_storage::bmt::Proof<Digest> as commonware_codec::Read>::read_cfg(
&mut our_bytes_buf,
&1usize,
)
.unwrap();
cw_decoded
.verify_element_inclusion::<Blake3>(&Digest(leaves[pos as usize]), pos, &cw_root)
.expect("upstream must accept our proof bytes");
if !our_proof.siblings.is_empty() {
let mut mutated = our_proof.clone();
mutated.siblings[0][0] ^= 0x01;
assert!(
mutated
.verify_element_inclusion(&leaves[pos as usize], pos, &ours.root)
.is_err()
);
let mut dropped = our_proof.clone();
dropped.siblings.pop();
assert!(
dropped
.verify_element_inclusion(&leaves[pos as usize], pos, &ours.root)
.is_err()
);
let mut extra = our_proof.clone();
extra.siblings.push(hash(b"extra"));
assert!(
extra
.verify_element_inclusion(&leaves[pos as usize], pos, &ours.root)
.is_err()
);
}
let mut wrong_count = our_proof.clone();
wrong_count.leaf_count = wrong_count.leaf_count.wrapping_add(1);
assert!(
wrong_count
.verify_element_inclusion(&leaves[pos as usize], pos, &ours.root)
.is_err()
);
if n >= 2 {
let our_range = ours.range_proof(start, end).unwrap();
let cw_range = cw_tree.range_proof(start, end).unwrap();
assert_eq!(
our_range.encode(),
commonware_codec::Encode::encode(&cw_range).to_vec(),
"range-proof bytes diverged at n={n} start={start} end={end}"
);
let range_leaves: Vec<Hash> = leaves[start as usize..=end as usize].to_vec();
let cw_range_leaves: Vec<Digest> =
range_leaves.iter().map(|h| Digest(*h)).collect();
cw_range
.verify_range_inclusion::<Blake3>(start, &cw_range_leaves, &cw_root)
.expect("upstream must accept its own range proof");
our_range
.verify_range_inclusion(start, &range_leaves, &ours.root)
.expect("ours must accept its own range proof");
if !positions.is_empty() {
let our_multi = ours.multi_proof(positions.iter().copied()).unwrap();
let cw_multi = cw_tree.multi_proof(positions.iter().copied()).unwrap();
assert_eq!(
our_multi.encode(),
commonware_codec::Encode::encode(&cw_multi).to_vec(),
"multi-proof bytes diverged at n={n} positions={positions:?}"
);
let elements: Vec<(Hash, u32)> =
positions.iter().map(|&p| (leaves[p as usize], p)).collect();
let cw_elements: Vec<(Digest, u32)> = positions
.iter()
.map(|&p| (Digest(leaves[p as usize]), p))
.collect();
cw_multi
.verify_multi_inclusion::<Blake3>(&cw_elements, &cw_root)
.expect("upstream must accept its own multi proof");
our_multi
.verify_multi_inclusion(&elements, &ours.root)
.expect("ours must accept its own multi proof");
}
}
}
}
}
#[cfg(kani)]
mod kani_proofs {
use super::*;
fn toy(a: &[u8], b: &[u8]) -> Hash {
fn first16(x: &[u8]) -> u128 {
let mut w = [0u8; 16];
let n = x.len().min(16);
w[..n].copy_from_slice(&x[..n]);
u128::from_le_bytes(w)
}
let (ha, hb) = (first16(a), first16(b));
let lo = ha ^ hb.rotate_left(8) ^ a.len() as u128;
let hi = hb ^ ha.rotate_left(16) ^ ((b.len() as u128) << 64);
let mut out = [0u8; HASH_LEN];
out[..16].copy_from_slice(&lo.to_le_bytes());
out[16..].copy_from_slice(&hi.to_le_bytes());
out
}
fn toy_h2(a: &[u8], b: &[u8]) -> Hash {
toy(a, b)
}
fn toy_domain_digest(domain: &[u8], body: &[u8]) -> Hash {
toy(body, domain)
}
fn toy_hash(data: &[u8]) -> Hash {
toy(data, &[])
}
fn spec_sibling_count(leaf_count: u32, mut position: u32) -> usize {
let mut level_size = u64::from(leaf_count);
let mut need = 0;
while level_size > 1 {
if !(position % 2 == 0 && u64::from(position) + 1 >= level_size) {
need += 1;
}
position /= 2;
level_size = level_size.div_ceil(2);
}
need
}
fn short_at<const N: usize>() {
let buf: [u8; N] = kani::any();
let b: &[u8] = &buf;
if let Ok(p) = Proof::decode(b, 1) {
assert!(N == 5 && b[4] == 0 && p.siblings.is_empty());
let lc: [u8; 4] = b[..4].try_into().expect("4 bytes");
assert_eq!(p.leaf_count, u32::from_be_bytes(lc));
kani::cover!(true, "ok_no_siblings");
}
}
#[kani::proof]
#[kani::unwind(4)]
fn merkle_proof_decode_no_panic() {
short_at::<0>();
short_at::<1>();
short_at::<2>();
short_at::<3>();
short_at::<4>();
}
#[kani::proof]
#[kani::unwind(4)]
fn merkle_proof_decode_empty_proof() {
short_at::<5>();
}
#[kani::proof]
#[kani::unwind(6)]
fn merkle_proof_decode_one_sibling() {
let mut buf: [u8; 37] = kani::any();
buf[4] = 1;
let p = Proof::decode(&buf, 1).expect("one-sibling proof decodes");
assert_eq!(
p.leaf_count,
u32::from_be_bytes([buf[0], buf[1], buf[2], buf[3]])
);
assert_eq!(p.siblings.len(), 1);
let d = &p.siblings[0];
let w = |x: &[u8], i: usize| u64::from_le_bytes(x[i..i + 8].try_into().expect("8"));
assert!((0..4).all(|k| w(d, 8 * k) == w(&buf[5..], 8 * k)));
}
fn verify_with<const S: usize>(max_leaves: u32) -> bool {
let proof = Proof {
leaf_count: kani::any_where(|&n| n <= max_leaves),
siblings: kani::any::<[Hash; S]>().to_vec(),
};
let position: u32 = kani::any();
let leaf: Hash = kani::any();
let id: Hash = kani::any();
let folded = proof.reconstruct_element_root(&leaf, position);
let shape_ok = position < proof.leaf_count
&& proof.siblings.len() == spec_sibling_count(proof.leaf_count, position);
assert_eq!(folded.is_ok(), shape_ok);
let chunk = verify_chunk(&id, &leaf, position, &proof);
if position == 0 {
assert_eq!(chunk, Err(MerkleError::PositionOutOfRange(0)));
}
if chunk.is_ok() {
assert!(shape_ok);
}
let entry = TreeEntry {
name: vec![b'a'],
mode: crate::object::EntryMode::Blob,
object_hash: leaf,
};
let _ = verify_tree_entry(&id, &entry, position, &proof);
chunk.is_ok()
}
#[kani::proof]
#[kani::stub(h2, toy_h2)]
#[kani::stub(crate::hash::domain_digest, toy_domain_digest)]
#[kani::stub(crate::hash::hash, toy_hash)]
#[kani::unwind(34)]
fn merkle_verify_s0() {
verify_with::<0>(u32::MAX);
}
#[kani::proof]
#[kani::stub(h2, toy_h2)]
#[kani::stub(crate::hash::domain_digest, toy_domain_digest)]
#[kani::stub(crate::hash::hash, toy_hash)]
#[kani::unwind(5)]
fn merkle_verify_s1() {
kani::cover!(verify_with::<1>(8), "accepts_one_sibling_proof");
}
#[kani::proof]
#[kani::stub(h2, toy_h2)]
#[kani::stub(crate::hash::domain_digest, toy_domain_digest)]
#[kani::stub(crate::hash::hash, toy_hash)]
#[kani::unwind(5)]
fn merkle_verify_s2() {
kani::cover!(verify_with::<2>(8), "accepts_two_sibling_proof");
}
fn spec_proof(leaves: &[Hash], pos: usize) -> (Hash, Vec<Hash>) {
#[allow(clippy::cast_possible_truncation)]
let mut level: Vec<Hash> = leaves
.iter()
.enumerate()
.map(|(i, l)| toy_h2(&(i as u32).to_be_bytes(), l))
.collect();
let mut p = pos;
let mut siblings = Vec::new();
while level.len() > 1 {
if p % 2 == 1 {
siblings.push(level[p - 1]);
} else if p + 1 < level.len() {
siblings.push(level[p + 1]);
}
let mut next = Vec::new();
let mut i = 0;
while i < level.len() {
let right = if i + 1 < level.len() {
level[i + 1]
} else {
level[i]
};
next.push(toy_h2(&level[i], &right));
i += 2;
}
level = next;
p /= 2;
}
#[allow(clippy::cast_possible_truncation)]
let root = toy_h2(&(leaves.len() as u32).to_be_bytes(), &level[0]);
(root, siblings)
}
fn chunked_fixture<const N: usize>() -> (ChunkedBlob, Vec<Hash>) {
let cb = ChunkedBlob {
total_size: kani::any(),
chunk_size: kani::any(),
chunks: kani::any::<[Hash; N]>().to_vec(),
};
let mut leaves = vec![chunked_meta_leaf_raw(cb.total_size, cb.chunk_size)];
leaves.extend_from_slice(&cb.chunks);
(cb, leaves)
}
fn chunk_rt<const N: usize>(pos: u32) {
let (cb, leaves) = chunked_fixture::<N>();
let (root, siblings) = spec_proof(&leaves, pos as usize);
let id = wrap_id(ObjectKind::ChunkedBlob, &root);
#[allow(clippy::cast_possible_truncation)]
let proof = Proof {
leaf_count: leaves.len() as u32,
siblings,
};
assert_eq!(
proof.siblings.len(),
spec_sibling_count(proof.leaf_count, pos)
);
assert_eq!(
verify_chunk(&id, &cb.chunks[(pos - 1) as usize], pos, &proof),
Ok(())
);
}
#[kani::proof]
#[kani::stub(h2, toy_h2)]
#[kani::stub(crate::hash::domain_digest, toy_domain_digest)]
#[kani::stub(crate::hash::hash, toy_hash)]
#[kani::unwind(5)]
fn merkle_roundtrip_one_chunk() {
chunk_rt::<1>(1);
}
#[kani::proof]
#[kani::stub(h2, toy_h2)]
#[kani::stub(crate::hash::domain_digest, toy_domain_digest)]
#[kani::stub(crate::hash::hash, toy_hash)]
#[kani::unwind(5)]
fn merkle_roundtrip_two_chunks() {
chunk_rt::<2>(1);
chunk_rt::<2>(2);
}
#[kani::proof]
#[kani::stub(h2, toy_h2)]
#[kani::stub(crate::hash::domain_digest, toy_domain_digest)]
#[kani::stub(crate::hash::hash, toy_hash)]
#[kani::unwind(3)]
fn merkle_builder_empty_tree_refuses() {
let empty = Tree {
entries: Vec::new(),
};
let (p, start, end): (u32, u32, u32) = (kani::any(), kani::any(), kani::any());
assert!(build_tree_entry_proof(&empty, p).is_err());
assert!(build_tree_entries_multi_proof(&empty, [p]).is_err());
let range = build_tree_entries_range_proof(&empty, start, end);
kani::cover!(range.is_err(), "range_refused");
assert!(
range.is_err(),
"empty-Tree range proof must be refused (SPEC-MERKLE-OBJECTS §5.4)"
);
}
#[kani::proof]
#[kani::stub(h2, toy_h2)]
#[kani::stub(crate::hash::domain_digest, toy_domain_digest)]
#[kani::stub(crate::hash::hash, toy_hash)]
#[kani::unwind(6)]
#[kani::should_panic]
fn merkle_canary_tampered_leaf_verifies() {
let (_cb, leaves) = chunked_fixture::<2>();
let (root, siblings) = spec_proof(&leaves, 1);
let id = wrap_id(ObjectKind::ChunkedBlob, &root);
let proof = Proof {
leaf_count: 3,
siblings,
};
let forged: Hash = kani::any();
assert!(verify_chunk(&id, &forged, 1, &proof).is_ok());
}
}