use std::{cmp::max, collections::HashMap, fmt::Debug, str::FromStr};
use rayon::iter::{IntoParallelIterator, ParallelIterator};
use super::{
error::{FromConfigError, ZerokitMerkleTreeError},
merkle_tree::{FrOf, Hasher, ZerokitMerkleProof, ZerokitMerkleTree, MIN_PARALLEL_NODES},
override_range_validation::{validate_override_range_inputs, EmptyIndicesPolicy},
};
#[derive(Clone, PartialEq, Eq, Debug)]
pub struct OptimalMerkleTree<H>
where
H: Hasher,
{
depth: usize,
cached_nodes: Vec<H::Fr>,
nodes: HashMap<(usize, usize), H::Fr>,
cached_leaves_indices: Vec<u8>,
next_index: usize,
metadata: Vec<u8>,
}
#[derive(Clone, PartialEq, Eq)]
pub struct OptimalMerkleProof<H: Hasher>(pub Vec<(H::Fr, u8)>);
#[derive(Default)]
pub struct OptimalMerkleConfig(());
impl FromStr for OptimalMerkleConfig {
type Err = FromConfigError;
fn from_str(_s: &str) -> Result<Self, Self::Err> {
Ok(OptimalMerkleConfig::default())
}
}
impl<H: Hasher> ZerokitMerkleTree for OptimalMerkleTree<H>
where
H: Hasher,
{
type Proof = OptimalMerkleProof<H>;
type Hasher = H;
type Config = OptimalMerkleConfig;
fn default(depth: usize) -> Result<Self, ZerokitMerkleTreeError> {
OptimalMerkleTree::<H>::new(depth, H::default_leaf(), Self::Config::default())
}
fn new(
depth: usize,
default_leaf: H::Fr,
_config: Self::Config,
) -> Result<Self, ZerokitMerkleTreeError> {
if depth >= usize::BITS as usize {
return Err(ZerokitMerkleTreeError::InvalidDepth);
}
let mut cached_nodes: Vec<H::Fr> = Vec::with_capacity(depth + 1);
cached_nodes.push(default_leaf);
for i in 0..depth {
cached_nodes.push(H::hash(&[cached_nodes[i]; 2]).map_err(Into::into)?);
}
cached_nodes.reverse();
Ok(OptimalMerkleTree {
depth,
cached_nodes,
nodes: HashMap::with_capacity(1 << depth),
cached_leaves_indices: vec![0; 1 << depth],
next_index: 0,
metadata: Vec::new(),
})
}
fn close_db_connection(&mut self) -> Result<(), ZerokitMerkleTreeError> {
Ok(())
}
fn depth(&self) -> usize {
self.depth
}
fn capacity(&self) -> usize {
1 << self.depth
}
fn leaves_set(&self) -> usize {
self.next_index
}
fn root(&self) -> H::Fr {
self.get_node(0, 0)
}
fn set(&mut self, index: usize, leaf: H::Fr) -> Result<(), ZerokitMerkleTreeError> {
if index >= self.capacity() {
return Err(ZerokitMerkleTreeError::InvalidLeaf);
}
self.nodes.insert((self.depth, index), leaf);
self.update_hashes(index, 1)?;
self.next_index = max(self.next_index, index + 1);
self.cached_leaves_indices[index] = 1;
Ok(())
}
fn get(&self, index: usize) -> Result<H::Fr, ZerokitMerkleTreeError> {
if index >= self.capacity() {
return Err(ZerokitMerkleTreeError::InvalidLeaf);
}
Ok(self.get_node(self.depth, index))
}
fn get_subtree_root(&self, n: usize, index: usize) -> Result<H::Fr, ZerokitMerkleTreeError> {
if n > self.depth() {
return Err(ZerokitMerkleTreeError::InvalidLevel);
}
if index >= self.capacity() {
return Err(ZerokitMerkleTreeError::InvalidLeaf);
}
if n == 0 {
Ok(self.root())
} else if n == self.depth {
self.get(index)
} else {
Ok(self.get_node(n, index >> (self.depth - n)))
}
}
fn get_empty_leaves_indices(&self) -> Vec<usize> {
self.cached_leaves_indices
.iter()
.take(self.next_index)
.enumerate()
.filter(|&(_, &v)| v == 0u8)
.map(|(idx, _)| idx)
.collect()
}
fn set_range<I: ExactSizeIterator<Item = H::Fr>>(
&mut self,
start: usize,
leaves: I,
) -> Result<(), ZerokitMerkleTreeError> {
let leaves_len = leaves.len();
let end = start
.checked_add(leaves_len)
.ok_or(ZerokitMerkleTreeError::TooManySet)?;
if end > self.capacity() {
return Err(ZerokitMerkleTreeError::TooManySet);
}
for (i, leaf) in leaves.enumerate() {
self.nodes.insert((self.depth, start + i), leaf);
self.cached_leaves_indices[start + i] = 1;
}
self.update_hashes(start, leaves_len)?;
self.next_index = max(self.next_index, start + leaves_len);
Ok(())
}
fn override_range<I, J>(
&mut self,
start: usize,
leaves: I,
indices: J,
) -> Result<(), ZerokitMerkleTreeError>
where
I: ExactSizeIterator<Item = FrOf<Self::Hasher>>,
J: ExactSizeIterator<Item = usize>,
{
let leaves_vec = leaves.into_iter().collect::<Vec<_>>();
let validated = validate_override_range_inputs(
start,
leaves_vec.len(),
indices.into_iter().collect::<Vec<_>>(),
self.capacity(),
EmptyIndicesPolicy::Reject,
)?;
let indices = validated.indices;
let min_index = validated
.min_index
.ok_or(ZerokitMerkleTreeError::InvalidIndices)?;
let max_index = validated.max_index.unwrap_or(start);
if min_index >= max_index {
return Err(ZerokitMerkleTreeError::InvalidIndices);
}
let mut set_values = vec![Self::Hasher::default_leaf(); max_index - min_index];
for i in min_index..start {
if !indices.contains(&i) {
let value = self.get(i)?;
set_values[i - min_index] = value;
}
}
for i in 0..leaves_vec.len() {
set_values[start - min_index + i] = leaves_vec[i];
}
for i in indices {
self.cached_leaves_indices[i] = 0;
}
self.set_range(start, set_values.into_iter())
}
fn update_next(&mut self, leaf: H::Fr) -> Result<(), ZerokitMerkleTreeError> {
self.set(self.next_index, leaf)?;
Ok(())
}
fn delete(&mut self, index: usize) -> Result<(), ZerokitMerkleTreeError> {
if index < self.next_index {
self.set(index, H::default_leaf())?;
self.cached_leaves_indices[index] = 0;
}
Ok(())
}
fn proof(&self, index: usize) -> Result<Self::Proof, ZerokitMerkleTreeError> {
if index >= self.capacity() {
return Err(ZerokitMerkleTreeError::InvalidLeaf);
}
let mut witness = Vec::<(H::Fr, u8)>::with_capacity(self.depth);
let mut i = index;
let mut depth = self.depth;
loop {
i ^= 1;
witness.push((self.get_node(depth, i), (1 - (i & 1)) as u8));
i >>= 1;
depth -= 1;
if depth == 0 {
break;
}
}
if i != 0 {
Err(ZerokitMerkleTreeError::ComputingProofError)
} else {
Ok(OptimalMerkleProof(witness))
}
}
fn verify(
&self,
leaf: &H::Fr,
merkle_proof: &Self::Proof,
) -> Result<bool, ZerokitMerkleTreeError> {
if merkle_proof.length() != self.depth {
return Err(ZerokitMerkleTreeError::InvalidMerkleProof);
}
let expected_root = merkle_proof.compute_root_from(leaf)?;
Ok(expected_root.eq(&self.root()))
}
fn set_metadata(&mut self, metadata: &[u8]) -> Result<(), ZerokitMerkleTreeError> {
self.metadata = metadata.to_vec();
Ok(())
}
fn metadata(&self) -> Result<Vec<u8>, ZerokitMerkleTreeError> {
Ok(self.metadata.to_vec())
}
}
impl<H: Hasher> OptimalMerkleTree<H>
where
H: Hasher,
{
fn get_node(&self, depth: usize, index: usize) -> H::Fr {
*self
.nodes
.get(&(depth, index))
.unwrap_or(&self.cached_nodes[depth])
}
fn hash_couple(&self, depth: usize, index: usize) -> Result<H::Fr, ZerokitMerkleTreeError> {
let b = index & !1;
H::hash(&[self.get_node(depth, b), self.get_node(depth, b + 1)]).map_err(Into::into)
}
fn update_hashes(&mut self, start: usize, length: usize) -> Result<(), ZerokitMerkleTreeError> {
let mut current_depth = self.depth;
let mut current_index = start & !1;
let mut current_index_max = (start + length + 1) & !1;
while current_depth > 0 {
let parent_depth = current_depth - 1;
if current_index_max - current_index >= MIN_PARALLEL_NODES {
#[allow(clippy::type_complexity)]
let updates: Result<
Vec<((usize, usize), H::Fr)>,
ZerokitMerkleTreeError,
> = (current_index..current_index_max)
.step_by(2)
.collect::<Vec<_>>()
.into_par_iter()
.map(|index| {
let hash = self.hash_couple(current_depth, index)?;
Ok(((parent_depth, index >> 1), hash))
})
.collect();
for (parent, hash) in updates? {
self.nodes.insert(parent, hash);
}
} else {
for index in (current_index..current_index_max).step_by(2) {
let hash = self.hash_couple(current_depth, index)?;
self.nodes.insert((parent_depth, index >> 1), hash);
}
}
current_index >>= 1;
current_index_max = (current_index_max + 1) >> 1;
current_depth -= 1;
}
Ok(())
}
}
impl<H: Hasher> ZerokitMerkleProof for OptimalMerkleProof<H>
where
H: Hasher,
{
type Index = u8;
type Hasher = H;
fn length(&self) -> usize {
self.0.len()
}
fn leaf_index(&self) -> usize {
let mut binary_repr = self.get_path_index();
binary_repr.reverse();
binary_repr
.into_iter()
.fold(0, |acc, digit| (acc << 1) + usize::from(digit))
}
fn get_path_elements(&self) -> Vec<H::Fr> {
self.0.iter().map(|x| x.0).collect()
}
fn get_path_index(&self) -> Vec<u8> {
self.0.iter().map(|x| x.1).collect()
}
fn compute_root_from(&self, leaf: &H::Fr) -> Result<H::Fr, ZerokitMerkleTreeError> {
let mut acc: H::Fr = *leaf;
for w in self.0.iter() {
if w.1 == 0 {
acc = H::hash(&[acc, w.0]).map_err(Into::into)?;
} else {
acc = H::hash(&[w.0, acc]).map_err(Into::into)?;
}
}
Ok(acc)
}
}
impl<H> Debug for OptimalMerkleProof<H>
where
H: Hasher,
H::Fr: Debug,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_tuple("Proof").field(&self.0).finish()
}
}