use alloc::{
string::{String, ToString},
sync::Arc,
vec::Vec,
};
use core::{iter::FusedIterator, ops::Index, slice};
use super::{
super::{InnerNodeInfo, MerklePath},
MmrDelta, MmrError, MmrPath, MmrPeaks, MmrProof,
forest::{Forest, TreeSizeIterator},
nodes_from_mask,
};
use crate::{
Word,
hash::poseidon2::Poseidon2,
utils::{ByteReader, ByteWriter, Deserializable, DeserializationError, Serializable},
};
const NODE_CHUNK_CAPACITY: usize = 1024;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub(super) struct NodeStore {
chunks: Vec<Arc<Vec<Word>>>,
}
impl NodeStore {
pub fn new() -> Self {
Self { chunks: Vec::new() }
}
pub fn len(&self) -> usize {
match self.chunks.last() {
Some(last) => (self.chunks.len() - 1) * NODE_CHUNK_CAPACITY + last.len(),
None => 0,
}
}
pub fn push(&mut self, node: Word) {
match self.chunks.last_mut() {
Some(last) if last.len() < NODE_CHUNK_CAPACITY => {
if Arc::get_mut(last).is_none() {
let mut copy = Vec::with_capacity(NODE_CHUNK_CAPACITY);
copy.extend_from_slice(last);
*last = Arc::new(copy);
}
Arc::get_mut(last).expect("chunk is unique").push(node);
},
_ => {
let mut chunk = Vec::with_capacity(NODE_CHUNK_CAPACITY);
chunk.push(node);
self.chunks.push(Arc::new(chunk));
},
}
}
pub fn iter(&self) -> MmrNodeIter<'_> {
self.iter_from(0)
}
pub fn iter_from(&self, start: usize) -> MmrNodeIter<'_> {
let first_chunk = start / NODE_CHUNK_CAPACITY;
let offset = start % NODE_CHUNK_CAPACITY;
match self.chunks.get(first_chunk) {
Some(chunk) => MmrNodeIter {
current: chunk[offset.min(chunk.len())..].iter(),
chunks: &self.chunks[first_chunk + 1..],
},
None => MmrNodeIter { current: [].iter(), chunks: &[] },
}
}
}
impl Index<usize> for NodeStore {
type Output = Word;
fn index(&self, index: usize) -> &Word {
&self.chunks[index / NODE_CHUNK_CAPACITY][index % NODE_CHUNK_CAPACITY]
}
}
impl FromIterator<Word> for NodeStore {
fn from_iter<T: IntoIterator<Item = Word>>(iter: T) -> Self {
let mut store = Self::new();
for node in iter {
store.push(node);
}
store
}
}
impl PartialEq<&[Word]> for NodeStore {
fn eq(&self, other: &&[Word]) -> bool {
self.len() == other.len() && self.iter().eq(other.iter())
}
}
#[derive(Debug, Clone)]
pub struct Mmr {
pub(super) forest: Forest,
pub(super) nodes: NodeStore,
}
impl Default for Mmr {
fn default() -> Self {
Self::new()
}
}
impl Mmr {
pub fn new() -> Mmr {
Mmr {
forest: Forest::empty(),
nodes: NodeStore::new(),
}
}
pub fn try_from_iter<T: IntoIterator<Item = Word>>(values: T) -> Result<Self, MmrError> {
Self::try_from_iter_with_limit(values, Forest::MAX_LEAVES)
}
pub fn from_nodes_unchecked(
forest: Forest,
nodes: impl IntoIterator<Item = Word>,
) -> Result<Self, MmrError> {
Self::from_store(forest, nodes.into_iter().collect())
}
fn from_store(forest: Forest, nodes: NodeStore) -> Result<Self, MmrError> {
if nodes.len() != forest.num_nodes() {
return Err(MmrError::InvalidNodeCount {
expected: forest.num_nodes(),
actual: nodes.len(),
});
}
Ok(Self { forest, nodes })
}
pub(crate) fn try_from_iter_with_limit<T: IntoIterator<Item = Word>>(
values: T,
max_leaves: usize,
) -> Result<Self, MmrError> {
let mut mmr = Mmr::new();
let iter = values.into_iter();
let (lower, _) = iter.size_hint();
if lower > max_leaves {
return Err(MmrError::ForestSizeExceeded { requested: lower, max: max_leaves });
}
let mut count = 0usize;
for v in iter {
count += 1;
if count > max_leaves {
return Err(MmrError::ForestSizeExceeded { requested: count, max: max_leaves });
}
mmr.add(v)?;
}
Ok(mmr)
}
fn from_serialized_parts(forest: Forest, nodes: NodeStore) -> Result<Self, String> {
let mmr = Self::from_store(forest, nodes).map_err(|err| err.to_string())?;
if mmr
.inner_nodes()
.any(|node| node.value != Poseidon2::merge(&[node.left, node.right]))
{
return Err("Mmr contains a parent node inconsistent with its children".into());
}
Ok(mmr)
}
pub const fn forest(&self) -> Forest {
self.forest
}
pub fn nodes_from(&self, start: usize) -> impl ExactSizeIterator<Item = &Word> + Clone {
self.nodes.iter_from(start)
}
pub fn open(&self, pos: usize) -> Result<MmrProof, MmrError> {
self.open_at(pos, self.forest)
}
pub fn open_at(&self, pos: usize, forest: Forest) -> Result<MmrProof, MmrError> {
if forest > self.forest {
return Err(MmrError::ForestOutOfBounds(forest.num_leaves(), self.forest.num_leaves()));
}
let (leaf, path) = self.collect_merkle_path_and_value(pos, forest)?;
let path = MmrPath::new(forest, pos, MerklePath::new(path));
Ok(MmrProof::new(path, leaf))
}
pub fn get(&self, pos: usize) -> Result<Word, MmrError> {
let (value, _) = self.collect_merkle_path_and_value(pos, self.forest)?;
Ok(value)
}
pub fn add(&mut self, el: Word) -> Result<(), MmrError> {
let old_forest = self.forest;
self.forest.append_leaf()?;
self.nodes.push(el);
let mut left_offset = self.nodes.len().saturating_sub(2);
let mut right = el;
let mut left_tree = 1usize;
while (old_forest.num_leaves() & left_tree) != 0 {
right = Poseidon2::merge(&[self.nodes[left_offset], right]);
self.nodes.push(right);
debug_assert!(left_tree <= Forest::MAX_LEAVES);
let left_nodes = left_tree * 2 - 1;
left_offset = left_offset.saturating_sub(left_nodes);
match left_tree.checked_shl(1) {
Some(next) => left_tree = next,
None => break,
}
}
Ok(())
}
pub fn peaks(&self) -> MmrPeaks {
self.peaks_at(self.forest).expect("failed to get peaks at current forest")
}
pub fn peaks_at(&self, forest: Forest) -> Result<MmrPeaks, MmrError> {
if forest > self.forest {
return Err(MmrError::ForestOutOfBounds(forest.num_leaves(), self.forest.num_leaves()));
}
let peaks: Vec<Word> = TreeSizeIterator::new(forest)
.rev()
.map(Forest::num_nodes)
.scan(0, |offset, el| {
*offset += el;
Some(*offset)
})
.map(|offset| self.nodes[offset - 1])
.collect();
let peaks = MmrPeaks::new(forest, peaks)?;
Ok(peaks)
}
pub fn get_delta(&self, from_forest: Forest, to_forest: Forest) -> Result<MmrDelta, MmrError> {
if to_forest > self.forest {
return Err(MmrError::ForestOutOfBounds(
to_forest.num_leaves(),
self.forest.num_leaves(),
));
}
if from_forest > to_forest {
return Err(MmrError::ForestOutOfBounds(
from_forest.num_leaves(),
to_forest.num_leaves(),
));
}
if from_forest == to_forest {
return Ok(MmrDelta { forest: to_forest, data: Vec::new() });
}
let mut result = Vec::new();
let candidate_mask = to_forest.num_leaves() ^ from_forest.num_leaves();
let mut new_high = super::forest::largest_tree_from_mask(candidate_mask);
let mut merges = from_forest & new_high.all_smaller_trees_unchecked();
let common_trees = from_forest ^ merges;
if !merges.is_empty() {
let mut target = merges.smallest_tree_unchecked();
while target < new_high {
let known_mask =
common_trees.num_leaves() | merges.num_leaves() | target.num_leaves();
let known = nodes_from_mask(known_mask);
let sibling = target.num_nodes();
result.push(self.nodes[known + sibling - 1]);
target = target.next_larger_tree()?;
while !(merges & target).is_empty() {
target = target.next_larger_tree()?;
}
merges ^= merges & target.all_smaller_trees_unchecked();
}
} else {
new_high = Forest::empty();
}
let mut new_peaks = to_forest ^ common_trees ^ new_high;
let old_peaks = to_forest ^ new_peaks;
let mut offset = old_peaks.num_nodes();
while !new_peaks.is_empty() {
let target = new_peaks.largest_tree_unchecked();
offset += target.num_nodes();
result.push(self.nodes[offset - 1]);
new_peaks ^= target;
}
Ok(MmrDelta { forest: to_forest, data: result })
}
pub fn inner_nodes(&self) -> MmrNodes<'_> {
MmrNodes {
mmr: self,
forest: 0,
last_right: 0,
index: 0,
}
}
fn collect_merkle_path_and_value(
&self,
leaf_idx: usize,
forest: Forest,
) -> Result<(Word, Vec<Word>), MmrError> {
let tree_bit = forest
.leaf_to_corresponding_tree(leaf_idx)
.ok_or(MmrError::PositionNotFound(leaf_idx))?;
let forest_before = forest.trees_larger_than(tree_bit);
let index_offset = forest_before.num_nodes();
let relative_pos = leaf_idx - forest_before.num_leaves();
let tree_depth = (tree_bit + 1) as usize;
let mut path = Vec::with_capacity(tree_depth);
let mut forest_target: usize = 1usize << tree_bit;
let mut index = nodes_from_mask(forest_target) - 1;
while forest_target > 1 {
forest_target >>= 1;
let right_offset = index - 1;
let left_offset = right_offset - nodes_from_mask(forest_target);
let left_or_right = relative_pos & forest_target;
let sibling = if left_or_right != 0 {
index = right_offset;
self.nodes[index_offset + left_offset]
} else {
index = left_offset;
self.nodes[index_offset + right_offset]
};
path.push(sibling);
}
debug_assert!(path.len() == tree_depth - 1);
path.reverse();
let value = self.nodes[index_offset + index];
Ok((value, path))
}
}
impl Serializable for Mmr {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
self.forest.write_into(target);
target.write_usize(self.nodes.len());
target.write_many(self.nodes.iter());
}
}
impl Deserializable for Mmr {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
let forest = Forest::read_from(source)?;
let count = source.read_usize()?;
if count != forest.num_nodes() {
return Err(DeserializationError::InvalidValue(
MmrError::InvalidNodeCount {
expected: forest.num_nodes(),
actual: count,
}
.to_string(),
));
}
let nodes = source.read_many_iter(count)?.collect::<Result<NodeStore, _>>()?;
Self::from_serialized_parts(forest, nodes).map_err(DeserializationError::InvalidValue)
}
}
#[derive(Clone, Debug)]
pub(super) struct MmrNodeIter<'a> {
current: slice::Iter<'a, Word>,
chunks: &'a [Arc<Vec<Word>>],
}
impl<'a> MmrNodeIter<'a> {
fn advance_chunk(&mut self) -> Option<()> {
match self.chunks.split_first() {
Some((chunk, rest)) => {
self.current = chunk.iter();
self.chunks = rest;
Some(())
},
None => {
self.current = [].iter();
None
},
}
}
}
impl<'a> Iterator for MmrNodeIter<'a> {
type Item = &'a Word;
fn next(&mut self) -> Option<&'a Word> {
loop {
if let Some(node) = self.current.next() {
return Some(node);
}
self.advance_chunk()?;
}
}
fn nth(&mut self, mut n: usize) -> Option<&'a Word> {
loop {
let len = self.current.len();
if n < len {
return self.current.nth(n);
}
n -= len;
self.advance_chunk()?;
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
let len = self.len();
(len, Some(len))
}
fn count(self) -> usize {
self.len()
}
}
impl ExactSizeIterator for MmrNodeIter<'_> {
fn len(&self) -> usize {
self.current.len()
+ match self.chunks.split_last() {
Some((last, full)) => full.len() * NODE_CHUNK_CAPACITY + last.len(),
None => 0,
}
}
}
impl FusedIterator for MmrNodeIter<'_> {}
pub struct MmrNodes<'a> {
mmr: &'a Mmr,
forest: usize,
last_right: usize,
index: usize,
}
impl Iterator for MmrNodes<'_> {
type Item = InnerNodeInfo;
fn next(&mut self) -> Option<Self::Item> {
debug_assert!(self.last_right.count_ones() <= 1, "last_right tracks zero or one element");
let target = self.mmr.forest.without_single_leaf().num_leaves();
if self.forest < target {
if self.last_right == 0 {
debug_assert!(self.last_right == 0, "left must be before right");
self.forest |= 1;
self.index += 1;
debug_assert!((self.forest & 1) == 1, "right must be after left");
self.last_right |= 1;
self.index += 1;
};
debug_assert!(
self.forest & self.last_right != 0,
"parent requires both a left and right",
);
let right_nodes = Forest::new(self.last_right).unwrap().num_nodes();
let parent = self.last_right << 1;
self.forest ^= self.last_right;
if self.forest & parent == 0 {
debug_assert!(self.forest & 1 == 0, "next iteration yields a left leaf");
self.last_right = 0;
self.forest ^= parent;
} else {
self.last_right = parent;
}
let value = self.mmr.nodes[self.index];
let right = self.mmr.nodes[self.index - 1];
let left = self.mmr.nodes[self.index - 1 - right_nodes];
self.index += 1;
let node = InnerNodeInfo { value, left, right };
Some(node)
} else {
None
}
}
}
#[cfg(test)]
mod tests {
use alloc::{sync::Arc, vec::Vec};
use super::{super::nodes_from_mask, NODE_CHUNK_CAPACITY};
use crate::{
Felt, Word, ZERO,
merkle::mmr::{Forest, Mmr, MmrError},
utils::{Deserializable, DeserializationError, Serializable},
};
fn leaves(count: u64) -> impl Iterator<Item = Word> {
(0..count).map(|value| Word::new([ZERO, ZERO, ZERO, Felt::new_unchecked(value)]))
}
#[test]
fn test_serialization() {
let nodes = (0u64..128u64)
.map(|value| Word::new([ZERO, ZERO, ZERO, Felt::new_unchecked(value)]))
.collect::<Vec<_>>();
let mmr = Mmr::try_from_iter(nodes).unwrap();
let serialized = mmr.to_bytes();
let deserialized = Mmr::read_from_bytes(&serialized).unwrap();
assert_eq!(mmr.forest, deserialized.forest);
assert_eq!(mmr.nodes, deserialized.nodes);
}
#[test]
fn test_deserialization_rejects_large_forest() {
let mut bytes = (Forest::MAX_LEAVES + 1).to_bytes();
bytes.extend_from_slice(&0usize.to_bytes());
let result = Mmr::read_from_bytes(&bytes);
assert!(matches!(result, Err(DeserializationError::InvalidValue(_))));
}
#[test]
fn test_serialization_matches_vec_format() {
let num_leaves = NODE_CHUNK_CAPACITY as u64 + NODE_CHUNK_CAPACITY as u64 / 2;
let mmr = Mmr::try_from_iter(leaves(num_leaves)).unwrap();
assert!(mmr.nodes.len() > 2 * NODE_CHUNK_CAPACITY);
let mut expected = mmr.forest.to_bytes();
let nodes_vec: Vec<Word> = mmr.nodes.iter().copied().collect();
expected.extend_from_slice(&nodes_vec.to_bytes());
assert_eq!(mmr.to_bytes(), expected);
let deserialized = Mmr::read_from_bytes(&expected).unwrap();
assert_eq!(mmr.forest, deserialized.forest);
assert_eq!(mmr.nodes, deserialized.nodes);
}
#[test]
fn test_deserialization_rejects_node_count_mismatch() {
let mmr = Mmr::try_from_iter(leaves(8)).unwrap();
let mut bytes = mmr.forest.to_bytes();
let mut nodes_vec: Vec<Word> = mmr.nodes.iter().copied().collect();
nodes_vec.pop();
bytes.extend_from_slice(&nodes_vec.to_bytes());
let result = Mmr::read_from_bytes(&bytes);
assert!(matches!(result, Err(DeserializationError::InvalidValue(_))));
}
#[test]
fn test_clone_stays_frozen_after_push() {
let num_leaves = 2 * NODE_CHUNK_CAPACITY as u64;
let mut mmr = Mmr::try_from_iter(leaves(num_leaves)).unwrap();
let clone = mmr.clone();
for leaf in leaves(3 * NODE_CHUNK_CAPACITY as u64).skip(num_leaves as usize) {
mmr.add(leaf).unwrap();
}
let reference = Mmr::try_from_iter(leaves(num_leaves)).unwrap();
assert_eq!(clone.forest, reference.forest);
assert_eq!(clone.nodes, reference.nodes);
assert_eq!(clone.peaks(), reference.peaks());
let reference = Mmr::try_from_iter(leaves(3 * NODE_CHUNK_CAPACITY as u64)).unwrap();
assert_eq!(mmr.forest, reference.forest);
assert_eq!(mmr.nodes, reference.nodes);
assert_eq!(mmr.peaks(), reference.peaks());
}
#[test]
fn test_clone_shares_full_chunks() {
let mut mmr = Mmr::try_from_iter(leaves(NODE_CHUNK_CAPACITY as u64)).unwrap();
let clone = mmr.clone();
mmr.add(Word::empty()).unwrap();
let orig_chunks = &mmr.nodes.chunks;
let clone_chunks = &clone.nodes.chunks;
assert_eq!(orig_chunks.len(), clone_chunks.len());
for (orig, cloned) in orig_chunks.iter().zip(clone_chunks).take(orig_chunks.len() - 1) {
assert!(Arc::ptr_eq(orig, cloned));
}
assert!(!Arc::ptr_eq(orig_chunks.last().unwrap(), clone_chunks.last().unwrap()));
}
#[test]
fn test_nodes_from() {
let mmr = Mmr::try_from_iter(leaves(2 * NODE_CHUNK_CAPACITY as u64)).unwrap();
let num_nodes = mmr.forest().num_nodes();
let all: Vec<Word> = mmr.nodes_from(0).copied().collect();
assert_eq!(all.len(), num_nodes);
for start in [
0,
1,
NODE_CHUNK_CAPACITY - 1,
NODE_CHUNK_CAPACITY,
NODE_CHUNK_CAPACITY + 1,
num_nodes - 1,
num_nodes,
num_nodes + 1,
] {
let suffix: Vec<Word> = mmr.nodes_from(start).copied().collect();
assert_eq!(suffix, all[start.min(num_nodes)..]);
}
}
#[test]
fn test_node_iter_skip_matches_nodes_from() {
let mmr = Mmr::try_from_iter(leaves(2 * NODE_CHUNK_CAPACITY as u64)).unwrap();
let num_nodes = mmr.forest().num_nodes();
let all: Vec<Word> = mmr.nodes_from(0).copied().collect();
for start in [
0,
1,
NODE_CHUNK_CAPACITY - 1,
NODE_CHUNK_CAPACITY,
NODE_CHUNK_CAPACITY + 1,
num_nodes - 1,
num_nodes,
num_nodes + 1,
] {
let skipped: Vec<Word> = mmr.nodes_from(0).skip(start).copied().collect();
assert_eq!(skipped, all[start.min(num_nodes)..]);
}
let mut iter = mmr.nodes_from(0);
assert_eq!(iter.nth(NODE_CHUNK_CAPACITY + 1), Some(&all[NODE_CHUNK_CAPACITY + 1]));
assert_eq!(iter.next(), Some(&all[NODE_CHUNK_CAPACITY + 2]));
let mut iter = mmr.nodes_from(0);
assert_eq!(iter.nth(num_nodes), None);
assert_eq!(iter.next(), None);
}
#[test]
fn test_node_iter_len() {
let mmr = Mmr::try_from_iter(leaves(2 * NODE_CHUNK_CAPACITY as u64)).unwrap();
let num_nodes = mmr.forest().num_nodes();
for start in [0, 1, NODE_CHUNK_CAPACITY, num_nodes - 1, num_nodes, num_nodes + 1] {
let iter = mmr.nodes_from(start);
assert_eq!(iter.len(), num_nodes.saturating_sub(start));
assert_eq!(iter.size_hint(), (iter.len(), Some(iter.len())));
}
let mut iter = mmr.nodes_from(0);
iter.next();
assert_eq!(iter.len(), num_nodes - 1);
iter.nth(NODE_CHUNK_CAPACITY);
assert_eq!(iter.len(), num_nodes - NODE_CHUNK_CAPACITY - 2);
assert_eq!(iter.count(), num_nodes - NODE_CHUNK_CAPACITY - 2);
}
#[test]
fn test_nodes_from_incremental_persistence() {
let initial_leaves = NODE_CHUNK_CAPACITY as u64 / 2;
let final_leaves = 2 * NODE_CHUNK_CAPACITY as u64;
let mut mmr = Mmr::try_from_iter(leaves(initial_leaves)).unwrap();
let mut persisted: Vec<Word> = mmr.nodes_from(0).copied().collect();
for leaf in leaves(final_leaves).skip(initial_leaves as usize) {
mmr.add(leaf).unwrap();
}
persisted.extend(mmr.nodes_from(persisted.len()).copied());
assert_eq!(persisted.len(), mmr.forest().num_nodes());
assert!(mmr.nodes == persisted.as_slice());
}
#[test]
fn test_from_nodes_unchecked_round_trip() {
let multi_chunk = NODE_CHUNK_CAPACITY as u64 + NODE_CHUNK_CAPACITY as u64 / 2;
for num_leaves in [0, 1, 8, multi_chunk] {
let mmr = Mmr::try_from_iter(leaves(num_leaves)).unwrap();
let rebuilt =
Mmr::from_nodes_unchecked(mmr.forest(), mmr.nodes_from(0).copied()).unwrap();
assert_eq!(mmr.forest, rebuilt.forest);
assert_eq!(mmr.nodes, rebuilt.nodes);
assert_eq!(mmr.peaks(), rebuilt.peaks());
let peaks = rebuilt.peaks();
for pos in [0, num_leaves.saturating_sub(1) as usize] {
if num_leaves > 0 {
let proof = rebuilt.open(pos).unwrap();
let leaf = rebuilt.get(pos).unwrap();
peaks.verify(leaf, proof).unwrap();
}
}
}
}
#[test]
fn test_from_nodes_unchecked_rejects_count_mismatch() {
let mmr = Mmr::try_from_iter(leaves(8)).unwrap();
let nodes: Vec<Word> = mmr.nodes_from(0).copied().collect();
let too_few =
Mmr::from_nodes_unchecked(mmr.forest(), nodes.iter().copied().take(nodes.len() - 1));
assert!(matches!(
too_few,
Err(MmrError::InvalidNodeCount { expected, actual })
if expected == nodes.len() && actual == nodes.len() - 1
));
let too_many =
Mmr::from_nodes_unchecked(mmr.forest(), nodes.iter().copied().chain([Word::empty()]));
assert!(matches!(
too_many,
Err(MmrError::InvalidNodeCount { expected, actual })
if expected == nodes.len() && actual == nodes.len() + 1
));
}
#[test]
fn test_nodes_from_mask_at_max_leaves() {
let expected = (Forest::MAX_LEAVES as u128)
.saturating_mul(2)
.saturating_sub(Forest::MAX_LEAVES.count_ones() as u128);
assert!(expected <= usize::MAX as u128);
assert_eq!(nodes_from_mask(Forest::MAX_LEAVES), expected as usize);
}
}