use crate::{
errors::MerkleError,
merkle_tree::{MerklePath, MerkleTreeDigest},
traits::{MerkleParameters, CRH},
};
use snarkvm_utilities::ToBytes;
use std::sync::Arc;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
#[derive(Default)]
pub struct MerkleTree<P: MerkleParameters> {
root: MerkleTreeDigest<P>,
tree: Vec<MerkleTreeDigest<P>>,
hashed_leaves_index: usize,
padding_tree: Vec<(MerkleTreeDigest<P>, MerkleTreeDigest<P>)>,
parameters: Arc<P>,
}
impl<P: MerkleParameters + Send + Sync> MerkleTree<P> {
pub const DEPTH: usize = P::DEPTH;
pub fn new<L: ToBytes + Send + Sync>(parameters: Arc<P>, leaves: &[L]) -> Result<Self, MerkleError> {
let new_time = start_timer!(|| "MerkleTree::new");
let last_level_size = leaves.len().next_power_of_two();
let tree_size = 2 * last_level_size - 1;
let tree_depth = tree_depth(tree_size);
if tree_depth > Self::DEPTH {
return Err(MerkleError::InvalidTreeDepth(tree_depth, Self::DEPTH));
}
let empty_hash = parameters.hash_empty()?;
let mut tree = vec![empty_hash; tree_size];
let mut index = 0;
let mut level_indices = Vec::with_capacity(tree_depth);
for _ in 0..=tree_depth {
level_indices.push(index);
index = left_child(index);
}
let last_level_index = level_indices.pop().unwrap_or(0);
let subsections = Self::hash_row(&*parameters, leaves)?;
let mut subsection_index = 0;
for subsection in subsections.into_iter() {
tree[last_level_index + subsection_index..last_level_index + subsection_index + subsection.len()]
.copy_from_slice(&subsection[..]);
subsection_index += subsection.len();
}
let mut upper_bound = last_level_index;
level_indices.reverse();
for &start_index in &level_indices {
let hashings = (start_index..upper_bound)
.map(|i| (&tree[left_child(i)], &tree[right_child(i)]))
.collect::<Vec<_>>();
let hashes = Self::hash_row(&*parameters, &hashings[..])?;
let mut subsection_index = 0;
for subsection in hashes.into_iter() {
tree[start_index + subsection_index..start_index + subsection_index + subsection.len()]
.copy_from_slice(&subsection[..]);
subsection_index += subsection.len();
}
upper_bound = start_index;
}
let mut current_depth = tree_depth;
let mut padding_tree = Vec::with_capacity((Self::DEPTH).saturating_sub(current_depth + 1));
let mut current_hash = tree[0];
while current_depth < Self::DEPTH {
current_hash = parameters.hash_inner_node(¤t_hash, &empty_hash)?;
if current_depth < Self::DEPTH - 1 {
padding_tree.push((current_hash, empty_hash));
}
current_depth += 1;
}
let root_hash = current_hash;
end_timer!(new_time);
Ok(MerkleTree {
tree,
padding_tree,
hashed_leaves_index: last_level_index,
parameters,
root: root_hash,
})
}
pub fn rebuild<L: ToBytes + Send + Sync>(&self, start_index: usize, new_leaves: &[L]) -> Result<Self, MerkleError> {
let new_time = start_timer!(|| "MerkleTree::rebuild");
let last_level_size = (start_index + new_leaves.len()).next_power_of_two();
let tree_size = 2 * last_level_size - 1;
let tree_depth = tree_depth(tree_size);
if tree_depth > Self::DEPTH {
return Err(MerkleError::InvalidTreeDepth(tree_depth, Self::DEPTH));
}
let empty_hash = self.parameters.hash_empty()?;
let mut tree = vec![empty_hash; tree_size];
let mut index = 0;
let mut level_indices = Vec::with_capacity(tree_depth + 1);
for _ in 0..=tree_depth {
level_indices.push(index);
index = left_child(index);
}
let new_indices = || start_index..start_index + new_leaves.len();
let last_level_index = level_indices.pop().unwrap_or(0);
tree[last_level_index..][..start_index].clone_from_slice(&self.hashed_leaves()[..start_index]);
let subsections = Self::hash_row(&*self.parameters, new_leaves)?;
for (i, subsection) in subsections.into_iter().enumerate() {
tree[last_level_index + start_index + i..last_level_index + start_index + i + subsection.len()]
.copy_from_slice(&subsection[..]);
}
let mut upper_bound = last_level_index;
for start_index in level_indices.into_iter().rev() {
let (parents, children) = tree.split_at_mut(upper_bound);
crate::cfg_iter_mut!(parents[start_index..upper_bound])
.zip(start_index..upper_bound)
.try_for_each(|(parent, current_index)| {
let left_index = left_child(current_index);
let right_index = right_child(current_index);
if new_indices().contains(¤t_index)
|| self.tree.get(left_index) != children.get(left_index - upper_bound)
|| self.tree.get(right_index) != children.get(right_index - upper_bound)
|| new_indices().any(|idx| Ancestors(idx).into_iter().any(|i| i == current_index))
{
*parent = self.parameters.hash_inner_node(
&children[left_index - upper_bound],
&children[right_index - upper_bound],
)?;
} else {
*parent = self.tree[current_index];
}
Ok::<(), MerkleError>(())
})?;
upper_bound = start_index;
}
let mut current_depth = tree_depth;
let mut current_hash = tree[0];
let new_padding_tree = if current_hash == self.tree[0] {
current_hash = self
.parameters
.hash_inner_node(&self.padding_tree.last().unwrap().0, &empty_hash)?;
None
} else {
let mut padding_tree = Vec::with_capacity((Self::DEPTH).saturating_sub(current_depth + 1));
while current_depth < Self::DEPTH {
current_hash = self.parameters.hash_inner_node(¤t_hash, &empty_hash)?;
if current_depth < Self::DEPTH - 1 {
padding_tree.push((current_hash, empty_hash));
}
current_depth += 1;
}
Some(padding_tree)
};
let root_hash = current_hash;
end_timer!(new_time);
Ok(MerkleTree {
root: root_hash,
tree,
hashed_leaves_index: last_level_index,
padding_tree: if let Some(padding_tree) = new_padding_tree {
padding_tree
} else {
self.padding_tree.clone()
},
parameters: self.parameters.clone(),
})
}
#[inline]
pub fn root(&self) -> &<P::H as CRH>::Output {
&self.root
}
#[inline]
pub fn tree(&self) -> &[<P::H as CRH>::Output] {
&self.tree
}
#[inline]
pub fn hashed_leaves(&self) -> &[<P::H as CRH>::Output] {
&self.tree[self.hashed_leaves_index..]
}
pub fn generate_proof<L: ToBytes>(&self, index: usize, leaf: &L) -> Result<MerklePath<P>, MerkleError> {
let prove_time = start_timer!(|| "MerkleTree::generate_proof");
let mut path = vec![];
let leaf_hash = self.parameters.hash_leaf(leaf)?;
let tree_depth = tree_depth(self.tree.len());
let tree_index = convert_index_to_last_level(index, tree_depth);
if tree_index >= self.tree.len() || leaf_hash != self.tree[tree_index] {
return Err(MerkleError::IncorrectLeafIndex(tree_index));
}
let mut current_node = tree_index;
while !is_root(current_node) {
let sibling_node = sibling(current_node).unwrap();
path.push(self.tree[sibling_node]);
current_node = parent(current_node).unwrap();
}
if path.len() > Self::DEPTH {
return Err(MerkleError::InvalidPathLength(path.len(), Self::DEPTH));
}
if path.len() != Self::DEPTH {
let empty_hash = self.parameters.hash_empty()?;
path.push(empty_hash);
for &(ref _hash, ref sibling_hash) in &self.padding_tree {
path.push(*sibling_hash);
}
}
end_timer!(prove_time);
if path.len() != Self::DEPTH {
Err(MerkleError::IncorrectPathLength(path.len()))
} else {
Ok(MerklePath {
parameters: self.parameters.clone(),
path,
leaf_index: index as u64,
})
}
}
fn hash_row<L: ToBytes + Send + Sync>(
parameters: &P,
leaves: &[L],
) -> Result<Vec<Vec<<<P as MerkleParameters>::H as CRH>::Output>>, MerkleError> {
match leaves.len() {
0 => Ok(vec![]),
_ => Ok(vec![
crate::cfg_iter!(leaves)
.map(|leaf| parameters.hash_leaf(&leaf).unwrap())
.collect::<Vec<_>>(),
]),
}
}
}
#[inline]
fn tree_depth(tree_size: usize) -> usize {
fn log2(number: usize) -> usize {
(number as f64).log2() as usize
}
log2(tree_size)
}
#[inline]
fn is_root(index: usize) -> bool {
index == 0
}
#[inline]
fn left_child(index: usize) -> usize {
2 * index + 1
}
#[inline]
fn right_child(index: usize) -> usize {
2 * index + 2
}
#[inline]
fn sibling(index: usize) -> Option<usize> {
if index == 0 {
None
} else if is_left_child(index) {
Some(index + 1)
} else {
Some(index - 1)
}
}
#[inline]
fn is_left_child(index: usize) -> bool {
index % 2 == 1
}
#[inline]
fn parent(index: usize) -> Option<usize> {
if index > 0 { Some((index - 1) >> 1) } else { None }
}
#[inline]
fn convert_index_to_last_level(index: usize, tree_depth: usize) -> usize {
index + (1 << tree_depth) - 1
}
pub struct Ancestors(usize);
impl Iterator for Ancestors {
type Item = usize;
fn next(&mut self) -> Option<usize> {
if let Some(parent) = parent(self.0) {
self.0 = parent;
Some(parent)
} else {
None
}
}
}