use incrementalmerkletree::{
frontier::{Frontier, FrontierError},
Hashable, Level, Position,
};
use rayon::prelude::*;
use std::{error::Error, fmt};
type CompleteSubtreeRoots<H> = Vec<Option<H>>;
struct TreeCapacity<const DEPTH: u8>;
impl<const DEPTH: u8> TreeCapacity<DEPTH> {
const MAX_LEAVES: u64 = 1u64 << DEPTH;
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum BatchFrontierError {
Frontier(FrontierError),
BatchSpansMultipleSubtrees,
}
impl fmt::Display for BatchFrontierError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
BatchFrontierError::Frontier(error) => {
write!(f, "frontier reconstruction error: {error:?}")
}
BatchFrontierError::BatchSpansMultipleSubtrees => {
write!(f, "batch spans more than one tracked subtree boundary")
}
}
}
}
impl Error for BatchFrontierError {}
impl From<FrontierError> for BatchFrontierError {
fn from(error: FrontierError) -> Self {
BatchFrontierError::Frontier(error)
}
}
fn merge_complete_subtree<H: Hashable + Clone>(
slots: &mut CompleteSubtreeRoots<H>,
level: usize,
node: H,
) {
let mut idx = level;
let mut carry = node;
loop {
match slots[idx].take() {
None => {
slots[idx] = Some(carry);
break;
}
Some(existing) => {
carry = H::combine(Level::from(idx as u8), &existing, &carry);
idx += 1;
}
}
}
}
fn perfect_subtree_root<H: Hashable + Clone + Send + Sync>(leaves: &[H]) -> H {
debug_assert!(leaves.len().is_power_of_two());
if leaves.len() == 1 {
return leaves[0].clone();
}
let half = leaves.len() / 2;
let child_level = Level::from(half.trailing_zeros() as u8);
let (left, right) = leaves.split_at(half);
let (l, r) = rayon::join(
|| perfect_subtree_root(left),
|| perfect_subtree_root(right),
);
H::combine(child_level, &l, &r)
}
fn contains_complete_subtree(position: Position, level: u32) -> bool {
u64::from(position) & (1u64 << level) != 0
}
fn frontier_complete_subtree_roots<H, const DEPTH: u8>(
frontier: &Frontier<H, DEPTH>,
) -> (CompleteSubtreeRoots<H>, u64)
where
H: Hashable + Clone,
{
let Some(frontier) = frontier.value() else {
return (vec![None; usize::from(DEPTH)], 0);
};
let position = frontier.position();
let mut slots = vec![None; usize::from(DEPTH)];
let mut sibling_roots = frontier.ommers().iter().cloned();
for level in 0..u64::BITS {
if contains_complete_subtree(position, level) {
slots[level as usize] = Some(sibling_roots.next().expect("sibling root per set bit"));
}
}
merge_complete_subtree(&mut slots, 0, frontier.leaf().clone());
(slots, u64::from(position) + 1)
}
fn complete_subtree_chunks<H>(start_position: u64, leaves: &[H]) -> Vec<(usize, &[H])> {
let mut chunks = Vec::new();
let mut global_pos = start_position;
let mut leaf_offset = 0usize;
let end_position = start_position + leaves.len() as u64;
while global_pos < end_position {
let leaves_left = end_position - global_pos;
let max_available_level = u64::BITS - 1 - leaves_left.leading_zeros();
let max_aligned_level = if global_pos == 0 {
max_available_level
} else {
global_pos.trailing_zeros()
};
let level = max_aligned_level.min(max_available_level) as usize;
let chunk_len = 1usize << level;
chunks.push((level, &leaves[leaf_offset..leaf_offset + chunk_len]));
leaf_offset += chunk_len;
global_pos += chunk_len as u64;
}
chunks
}
pub(crate) fn parallel_append<H, const DEPTH: u8>(
frontier: Frontier<H, DEPTH>,
mut new_leaves: Vec<H>,
) -> Result<Frontier<H, DEPTH>, FrontierError>
where
H: Hashable + Clone + Send + Sync,
{
if new_leaves.is_empty() {
return Ok(frontier);
}
let (mut complete_subtree_roots, next_leaf_position) =
frontier_complete_subtree_roots(&frontier);
let new_tip_leaf = new_leaves
.pop()
.expect("new_leaves is not empty because it was checked above");
let leaves_to_merge = new_leaves;
let chunks = complete_subtree_chunks(next_leaf_position, &leaves_to_merge);
let new_subtree_roots: Vec<(usize, H)> = chunks
.into_par_iter()
.map(|(level, leaves)| (level, perfect_subtree_root(leaves)))
.collect();
for (level, root) in new_subtree_roots {
merge_complete_subtree(&mut complete_subtree_roots, level, root);
}
let new_tip_position = next_leaf_position + leaves_to_merge.len() as u64;
let complete_subtree_roots = complete_subtree_roots.into_iter().flatten().collect();
Frontier::from_parts(
Position::from(new_tip_position),
new_tip_leaf,
complete_subtree_roots,
)
}
pub fn append_batch_with_subtree<H, const DEPTH: u8>(
frontier: Frontier<H, DEPTH>,
nodes: Vec<H>,
) -> Result<(Frontier<H, DEPTH>, Option<(u64, H)>), BatchFrontierError>
where
H: Hashable + Clone + Send + Sync,
{
use crate::subtree::TRACKED_SUBTREE_HEIGHT;
if nodes.is_empty() {
return Ok((frontier, None));
}
let old_size = frontier.tree_size();
let new_size = old_size + nodes.len() as u64;
if new_size > TreeCapacity::<DEPTH>::MAX_LEAVES {
return Err(FrontierError::MaxDepthExceeded {
depth: DEPTH.saturating_add(1),
}
.into());
}
let subtree_size = 1u64 << TRACKED_SUBTREE_HEIGHT;
let boundary = (old_size / subtree_size)
.checked_add(1)
.and_then(|n| n.checked_mul(subtree_size));
if boundary
.and_then(|b| b.checked_add(subtree_size))
.is_some_and(|second_boundary| second_boundary <= new_size)
{
return Err(BatchFrontierError::BatchSpansMultipleSubtrees);
}
if boundary.is_some_and(|b| b <= new_size) {
let boundary = boundary.expect("checked above");
let head_len = (boundary - old_size) as usize;
let mut head = nodes;
let tail = head.split_off(head_len);
let f1 = parallel_append(frontier, head)?;
let index_value = (boundary >> TRACKED_SUBTREE_HEIGHT) - 1;
let root = f1
.value()
.expect("just appended at least one leaf")
.root(Some(Level::from(TRACKED_SUBTREE_HEIGHT)));
let f2 = parallel_append(f1, tail)?;
Ok((f2, Some((index_value, root))))
} else {
let f = parallel_append(frontier, nodes)?;
Ok((f, None))
}
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
const DEPTH: u8 = 32;
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
struct TestNode(u64);
fn mix3(level: u64, a: u64, b: u64) -> u64 {
const FNV_PRIME: u64 = 0x00000100000001B3;
const FNV_OFFSET: u64 = 0xcbf29ce484222325;
let mut h = FNV_OFFSET;
h ^= level;
h = h.wrapping_mul(FNV_PRIME);
h ^= a;
h = h.wrapping_mul(FNV_PRIME);
h ^= b;
h = h.wrapping_mul(FNV_PRIME);
h
}
impl Hashable for TestNode {
fn empty_leaf() -> Self {
Self(0)
}
fn combine(level: Level, a: &Self, b: &Self) -> Self {
Self(mix3(u8::from(level) as u64, a.0, b.0))
}
}
fn sequential_append<const DEPTH: u8>(
start: Frontier<TestNode, DEPTH>,
leaves: &[TestNode],
) -> Frontier<TestNode, DEPTH> {
let mut f = start;
for leaf in leaves {
assert!(f.append(*leaf), "test trees never overflow");
}
f
}
fn build_frontier<const DEPTH: u8>(prefix: &[TestNode]) -> Frontier<TestNode, DEPTH> {
let mut f = Frontier::<TestNode, DEPTH>::empty();
for leaf in prefix {
assert!(f.append(*leaf));
}
f
}
fn chunk_levels_and_values(start_position: u64, leaves: &[u64]) -> Vec<(usize, Vec<u64>)> {
complete_subtree_chunks(start_position, leaves)
.into_iter()
.map(|(level, chunk)| (level, chunk.to_vec()))
.collect()
}
#[test]
fn frontier_complete_subtree_roots_empty_frontier() {
let empty = Frontier::<TestNode, DEPTH>::empty();
let (complete_subtree_roots, next_leaf_position) = frontier_complete_subtree_roots(&empty);
assert_eq!(next_leaf_position, 0);
assert_eq!(
complete_subtree_roots,
vec![None; usize::from(DEPTH)],
"empty frontier has no complete subtree roots"
);
}
#[test]
fn complete_subtree_chunks_match_expected_decompositions() {
let cases = [
("empty at zero", 0, vec![], vec![]),
("empty after nonzero position", 17, vec![], vec![]),
(
"start at zero",
0,
vec![0, 1, 2, 3, 4, 5, 6],
vec![(2, vec![0, 1, 2, 3]), (1, vec![4, 5]), (0, vec![6])],
),
(
"aligned start",
8,
vec![10, 11, 12, 13, 14, 15, 16, 17],
vec![(3, vec![10, 11, 12, 13, 14, 15, 16, 17])],
),
(
"unaligned start",
6,
vec![100, 101, 102, 103, 104, 105, 106],
vec![
(1, vec![100, 101]),
(2, vec![102, 103, 104, 105]),
(0, vec![106]),
],
),
(
"preserve global alignment",
5,
vec![20, 21, 22, 23, 24, 25],
vec![
(0, vec![20]),
(1, vec![21, 22]),
(1, vec![23, 24]),
(0, vec![25]),
],
),
];
for (name, start_position, leaves, expected) in cases {
assert_eq!(
chunk_levels_and_values(start_position, &leaves),
expected,
"{name}"
);
}
}
#[test]
fn complete_subtree_roots_flatten_to_frontier_order() {
let interesting_prefix_lengths = [
0usize, 1, 2, 3, 4, 5, 7, 8, 9, 15, 16, 17, 31, 32, 33, 63, 64, 65,
];
for prefix_len in interesting_prefix_lengths {
let prefix_len = u64::try_from(prefix_len).expect("test prefix length fits in u64");
let prefix: Vec<TestNode> = (0..prefix_len).map(TestNode).collect();
let start = build_frontier::<DEPTH>(&prefix);
let new_tip_leaf = TestNode(10_000 + prefix_len);
let (complete_subtree_roots, next_leaf_position) =
frontier_complete_subtree_roots(&start);
assert_eq!(
next_leaf_position, prefix_len,
"next leaf position mismatch for prefix length {prefix_len}"
);
let complete_subtree_roots: Vec<TestNode> =
complete_subtree_roots.into_iter().flatten().collect();
let reconstructed = Frontier::<TestNode, DEPTH>::from_parts(
Position::from(next_leaf_position),
new_tip_leaf,
complete_subtree_roots,
)
.expect("test frontier reconstruction succeeds");
let sequential = sequential_append::<DEPTH>(start, &[new_tip_leaf]);
assert_eq!(
sequential.value().map(|f| f.clone().into_parts()),
reconstructed.value().map(|f| f.clone().into_parts()),
"flattened complete_subtree_roots must be ordered for Frontier::from_parts at prefix length {prefix_len}"
);
}
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(2000))]
#[test]
fn parallel_matches_sequential(
prefix_len in 0usize..300,
batch in proptest::collection::vec(any::<u64>().prop_map(TestNode), 0..300),
) {
let prefix: Vec<TestNode> = (0..prefix_len as u64).map(TestNode).collect();
let start = build_frontier::<DEPTH>(&prefix);
let seq = sequential_append::<DEPTH>(start.clone(), &batch);
let par = parallel_append(start, batch.clone()).expect("no overflow in tests");
prop_assert_eq!(seq.root(), par.root(), "root mismatch");
prop_assert_eq!(
seq.value().map(|f| f.clone().into_parts()),
par.value().map(|f| f.clone().into_parts()),
"frontier parts mismatch"
);
}
}
#[test]
fn exhaustive_small() {
for prefix_len in 0u64..40 {
let prefix: Vec<TestNode> = (0..prefix_len).map(TestNode).collect();
let start = build_frontier::<DEPTH>(&prefix);
for batch_len in 0u64..40 {
let batch: Vec<TestNode> = (1000..1000 + batch_len).map(TestNode).collect();
let seq = sequential_append::<DEPTH>(start.clone(), &batch);
let par = parallel_append(start.clone(), batch).expect("no overflow");
assert_eq!(
seq.root(),
par.root(),
"root mismatch p={prefix_len} b={batch_len}"
);
assert_eq!(
seq.value().map(|f| f.clone().into_parts()),
par.value().map(|f| f.clone().into_parts()),
"parts mismatch p={prefix_len} b={batch_len}"
);
}
}
}
#[test]
fn overflow_is_reported() {
const SMALL_DEPTH: u8 = 3;
let prefix: Vec<TestNode> = (0..7).map(TestNode).collect();
let start = build_frontier::<SMALL_DEPTH>(&prefix);
let exact_capacity_batch = [TestNode(100)];
let seq = sequential_append::<SMALL_DEPTH>(start.clone(), &exact_capacity_batch);
let par = parallel_append(start.clone(), exact_capacity_batch.to_vec())
.expect("one remaining leaf fits");
assert_eq!(seq.root(), par.root(), "root mismatch at exact capacity");
assert_eq!(
seq.value().map(|f| f.clone().into_parts()),
par.value().map(|f| f.clone().into_parts()),
"parts mismatch at exact capacity"
);
let empty_append = parallel_append(par.clone(), Vec::new()).expect("empty append succeeds");
assert_eq!(
par.value().map(|f| f.clone().into_parts()),
empty_append.value().map(|f| f.clone().into_parts()),
"empty append changed a full frontier"
);
let full_tree_overflow = append_batch_with_subtree(par, vec![TestNode(101)]);
assert!(
full_tree_overflow.is_err(),
"appending to a full tree overflows"
);
let partial_batch_overflow =
append_batch_with_subtree(start, vec![TestNode(100), TestNode(101)]);
assert!(
partial_batch_overflow.is_err(),
"batch crossing tree capacity overflows"
);
}
#[test]
fn append_batch_errors_on_multiple_subtree_boundaries() {
use crate::subtree::TRACKED_SUBTREE_HEIGHT;
let start = Frontier::<TestNode, DEPTH>::empty();
let subtree_size = 1usize << TRACKED_SUBTREE_HEIGHT;
let batch = vec![TestNode(0); subtree_size * 2];
let result = append_batch_with_subtree(start, batch);
assert_eq!(result, Err(BatchFrontierError::BatchSpansMultipleSubtrees));
}
#[test]
fn matches_sequential_at_alignment_boundaries() {
let interesting_prefix_lengths = [
0usize, 1, 2, 3, 7, 8, 9, 15, 16, 17, 255, 256, 257, 65_535, 65_536, 65_537,
];
let interesting_batch_lengths = [0usize, 1, 2, 3, 4, 5, 31, 32, 33];
let max_prefix_len = *interesting_prefix_lengths
.last()
.expect("interesting prefixes are non-empty");
let mut frontier = Frontier::<TestNode, DEPTH>::empty();
let mut snapshots = Vec::new();
for prefix_len in 0..=max_prefix_len {
if interesting_prefix_lengths.contains(&prefix_len) {
snapshots.push((prefix_len, frontier.clone()));
}
if prefix_len < max_prefix_len {
assert!(frontier.append(TestNode(
u64::try_from(prefix_len).expect("test prefix length fits in u64")
)));
}
}
for (prefix_len, start) in snapshots {
for batch_len in interesting_batch_lengths {
let prefix_len = u64::try_from(prefix_len).expect("test prefix length fits in u64");
let batch: Vec<TestNode> = (0..batch_len)
.map(|leaf| {
TestNode(
1_000_000
+ prefix_len
+ u64::try_from(leaf).expect("test batch length fits in u64"),
)
})
.collect();
let seq = sequential_append::<DEPTH>(start.clone(), &batch);
let par = parallel_append(start.clone(), batch).expect("no overflow");
assert_eq!(
seq.root(),
par.root(),
"root mismatch p={prefix_len} b={batch_len}"
);
assert_eq!(
seq.value().map(|f| f.clone().into_parts()),
par.value().map(|f| f.clone().into_parts()),
"parts mismatch p={prefix_len} b={batch_len}"
);
}
}
}
#[test]
fn merge_complete_subtree_carry_chain() {
let leaves: Vec<TestNode> = (0u64..8).map(TestNode).collect();
let frontier = build_frontier::<DEPTH>(&leaves);
let (expected_slots, expected_next) = frontier_complete_subtree_roots(&frontier);
let mut slots: CompleteSubtreeRoots<TestNode> = vec![None; usize::from(DEPTH)];
for leaf in &leaves {
merge_complete_subtree(&mut slots, 0, *leaf);
}
assert_eq!(
slots, expected_slots,
"slot state after merging 8 leaves must match frontier expansion"
);
assert_eq!(
expected_next, 8,
"frontier covering 8 leaves has next position 8"
);
}
#[test]
fn perfect_subtree_root_matches_sequential() {
for log2_len in 0usize..=4 {
let len = 1usize << log2_len;
let leaves: Vec<TestNode> = (0..len as u64).map(TestNode).collect();
let frontier = build_frontier::<DEPTH>(&leaves);
let (slots, _) = frontier_complete_subtree_roots(&frontier);
let sequential_root =
slots[log2_len].expect("complete 2^k subtree fills exactly slot k after expansion");
let parallel_root = perfect_subtree_root(&leaves);
assert_eq!(
sequential_root, parallel_root,
"perfect_subtree_root mismatch for 2^{log2_len} leaves"
);
}
}
}