use crate::poseidon::FieldHasher;
use anyhow::{Error, Result};
use sp_std::{
borrow::ToOwned,
collections::{btree_map::BTreeMap, btree_set::BTreeSet},
marker::PhantomData,
};
use zkstd::common::{vec, Decode, Encode, FftField, Vec};
#[derive(Debug)]
pub enum MerkleError {
InvalidLeaf,
InvalidPathNodes,
}
impl core::fmt::Display for MerkleError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let msg = match self {
MerkleError::InvalidLeaf => "Invalid leaf".to_owned(),
MerkleError::InvalidPathNodes => "Path nodes are not consistent".to_owned(),
};
write!(f, "{}", msg)
}
}
impl ark_std::error::Error for MerkleError {}
#[derive(Clone, Debug, PartialEq, Eq, Encode, Decode)]
pub struct MerkleProof<F: FftField, H: FieldHasher<F, 2>, const N: usize> {
pub path: Vec<(F, F)>,
pub path_pos: Vec<u64>,
pub marker: PhantomData<H>,
}
impl<F: FftField, H: FieldHasher<F, 2>, const N: usize> Default for MerkleProof<F, H, N> {
fn default() -> Self {
let empty: [F; N] =
gen_empty_hashes(&H::default(), &[0; 64]).expect("Failed to generate empty hashes");
Self {
path: empty.into_iter().take(N - 1).map(|x| (x, x)).collect(),
path_pos: vec![0; N - 1],
marker: Default::default(),
}
}
}
impl<F: FftField, H: FieldHasher<F, 2>, const N: usize> MerkleProof<F, H, N> {
pub fn check_membership(&self, root_hash: &F, leaf: &F, hasher: &H) -> Result<bool, Error> {
let root = self.calculate_root(leaf, hasher)?;
Ok(root == *root_hash)
}
pub fn calculate_root(&self, leaf: &F, hasher: &H) -> Result<F, Error> {
if *leaf != self.path[0].0 && *leaf != self.path[0].1 {
return Err(MerkleError::InvalidLeaf.into());
}
let mut prev = *leaf;
for &(ref left_hash, ref right_hash) in &self.path {
if &prev != left_hash && &prev != right_hash {
return Err(MerkleError::InvalidPathNodes.into());
}
prev = hasher.hash([*left_hash, *right_hash])?;
}
Ok(prev)
}
pub fn get_index(&self, root_hash: &F, leaf: &F, hasher: &H) -> Result<F, Error> {
if !self.check_membership(root_hash, leaf, hasher)? {
return Err(MerkleError::InvalidLeaf.into());
}
let mut prev = *leaf;
let mut index = F::zero();
let mut twopower = F::one();
for &(ref left_hash, ref right_hash) in &self.path {
if &prev != left_hash {
index += twopower;
}
twopower = twopower + twopower;
prev = hasher.hash([*left_hash, *right_hash])?;
}
Ok(index)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Encode, Decode)]
pub struct SparseMerkleTree<F: FftField, H: FieldHasher<F, 2>, const N: usize> {
pub tree: BTreeMap<u64, F>,
empty_hashes: [F; N],
marker: PhantomData<H>,
}
impl<F: FftField, H: FieldHasher<F, 2>, const N: usize> Default for SparseMerkleTree<F, H, N> {
fn default() -> Self {
Self {
tree: Default::default(),
empty_hashes: [F::zero(); N],
marker: Default::default(),
}
}
}
impl<F: FftField, H: FieldHasher<F, 2>, const N: usize> SparseMerkleTree<F, H, N> {
pub fn insert_batch(&mut self, leaves: &BTreeMap<u32, F>, hasher: &H) -> Result<(), Error> {
let last_level_index: u64 = (1u64 << (N - 1)) - 1;
let mut level_idxs: BTreeSet<u64> = BTreeSet::new();
for (i, leaf) in leaves {
let true_index = last_level_index + (*i as u64);
self.tree.insert(true_index, *leaf);
level_idxs.insert(parent(true_index).unwrap());
}
for level in 0..N {
let mut new_idxs: BTreeSet<u64> = BTreeSet::new();
let empty_hash = self.empty_hashes[level];
for i in level_idxs {
let left_index = left_child(i);
let right_index = right_child(i);
let left = self.tree.get(&left_index).unwrap_or(&empty_hash);
let right = self.tree.get(&right_index).unwrap_or(&empty_hash);
self.tree.insert(i, hasher.hash([*left, *right])?);
let parent = match parent(i) {
Some(i) => i,
None => break,
};
new_idxs.insert(parent);
}
level_idxs = new_idxs;
}
Ok(())
}
pub fn update(&mut self, index: u64, val: F, hasher: &H) -> Result<(), Error> {
self.insert_batch(&BTreeMap::from([(index as u32, val)]), hasher)
}
pub fn delete(&mut self, index: u64, hasher: &H) -> Result<(), Error> {
self.update(index, self.empty_hashes[0], hasher)
}
pub fn new_empty(hasher: &H, empty_leaf: &[u8; 64]) -> Result<Self, Error> {
Self::new_sequential(&[], hasher, empty_leaf)
}
pub fn new(
leaves: &BTreeMap<u32, F>,
hasher: &H,
empty_leaf: &[u8; 64],
) -> Result<Self, Error> {
let last_level_size = leaves.len().next_power_of_two();
let tree_size = 2 * last_level_size - 1;
let tree_height = tree_height(tree_size as u64);
assert!(tree_height <= N as u32);
let tree: BTreeMap<u64, F> = BTreeMap::new();
let empty_hashes = gen_empty_hashes(hasher, empty_leaf)?;
let mut smt = SparseMerkleTree::<F, H, N> {
tree,
empty_hashes,
marker: PhantomData,
};
smt.insert_batch(leaves, hasher)?;
Ok(smt)
}
pub fn new_sequential(leaves: &[F], hasher: &H, empty_leaf: &[u8; 64]) -> Result<Self, Error> {
let pairs: BTreeMap<u32, F> = leaves
.iter()
.enumerate()
.map(|(i, l)| (i as u32, *l))
.collect();
let smt = Self::new(&pairs, hasher, empty_leaf)?;
Ok(smt)
}
pub fn root(&self) -> F {
self.tree
.get(&0)
.cloned()
.unwrap_or(*self.empty_hashes.last().unwrap())
}
pub fn generate_membership_proof(&self, index: u64) -> MerkleProof<F, H, N> {
let mut path = vec![(F::zero(), F::zero()); N - 1];
let mut path_pos = vec![0; N - 1];
let tree_index = convert_index_to_last_level(index, N);
let mut current_node = tree_index;
let mut level = 0;
while !is_root(current_node) {
let sibling_node = sibling(current_node).unwrap();
let empty_hash = &self.empty_hashes[level];
let current = self.tree.get(¤t_node).cloned().unwrap_or(*empty_hash);
let sibling = self.tree.get(&sibling_node).cloned().unwrap_or(*empty_hash);
if is_left_child(current_node) {
path[level] = (current, sibling);
} else {
path[level] = (sibling, current);
path_pos[level] = 1;
}
current_node = parent(current_node).unwrap();
level += 1;
}
MerkleProof {
path,
path_pos,
marker: PhantomData,
}
}
}
pub fn gen_empty_hashes<F: FftField, H: FieldHasher<F, 2>, const N: usize>(
hasher: &H,
default_leaf: &[u8; 64],
) -> Result<[F; N], Error> {
let mut empty_hashes = [F::zero(); N];
let mut empty_hash = F::from_bytes_wide(default_leaf);
for item in empty_hashes.iter_mut().take(N) {
*item = empty_hash;
empty_hash = hasher.hash([empty_hash, empty_hash])?;
}
Ok(empty_hashes)
}
fn convert_index_to_last_level(index: u64, height: usize) -> u64 {
index + (1u64 << (height - 1)) - 1
}
#[inline]
fn log2(number: u64) -> u32 {
ark_std::log2(number as usize)
}
#[inline]
fn tree_height(tree_size: u64) -> u32 {
log2(tree_size)
}
#[inline]
fn is_root(index: u64) -> bool {
index == 0
}
#[inline]
fn left_child(index: u64) -> u64 {
2 * index + 1
}
#[inline]
fn right_child(index: u64) -> u64 {
2 * index + 2
}
#[inline]
fn sibling(index: u64) -> Option<u64> {
if index == 0 {
None
} else if is_left_child(index) {
Some(index + 1)
} else {
Some(index - 1)
}
}
#[inline]
fn is_left_child(index: u64) -> bool {
index % 2 == 1
}
#[inline]
fn parent(index: u64) -> Option<u64> {
if index > 0 {
Some((index - 1) >> 1)
} else {
None
}
}
#[cfg(test)]
mod test {
use super::{gen_empty_hashes, SparseMerkleTree};
use crate::poseidon::{FieldHasher, Poseidon};
use jub_jub::Fp;
use rand::rngs::OsRng;
use zkstd::common::{FftField, Group};
fn create_merkle_tree<F: FftField, H: FieldHasher<F, 2>, const N: usize>(
hasher: H,
leaves: &[F],
default_leaf: &[u8; 64],
) -> SparseMerkleTree<F, H, N> {
SparseMerkleTree::<F, H, N>::new_sequential(leaves, &hasher, default_leaf).unwrap()
}
#[test]
fn should_create_tree_poseidon() {
let poseidon = Poseidon::<Fp, 2>::new();
let default_leaf = [0u8; 64];
let rng = OsRng;
let leaves = [Fp::random(rng), Fp::random(rng), Fp::random(rng)];
const HEIGHT: usize = 3;
let smt =
create_merkle_tree::<Fp, Poseidon<Fp, 2>, HEIGHT>(poseidon, &leaves, &default_leaf);
let root = smt.root();
let empty_hashes =
gen_empty_hashes::<Fp, Poseidon<Fp, 2>, HEIGHT>(&poseidon, &default_leaf).unwrap();
let hash1 = leaves[0];
let hash2 = leaves[1];
let hash3 = leaves[2];
let hash12 = poseidon.hash([hash1, hash2]).unwrap();
let hash34 = poseidon.hash([hash3, empty_hashes[0]]).unwrap();
let calc_root = poseidon.hash([hash12, hash34]).unwrap();
assert_eq!(root, calc_root);
}
#[test]
fn should_generate_and_validate_proof_poseidon() {
let poseidon = Poseidon::<Fp, 2>::new();
let default_leaf = [0u8; 64];
let rng = OsRng;
let leaves = [Fp::random(rng), Fp::random(rng), Fp::random(rng)];
const HEIGHT: usize = 3;
let smt =
create_merkle_tree::<Fp, Poseidon<Fp, 2>, HEIGHT>(poseidon, &leaves, &default_leaf);
let proof = smt.generate_membership_proof(0);
let res = proof
.check_membership(&smt.root(), &leaves[0], &poseidon)
.unwrap();
assert!(res);
}
#[test]
fn should_find_the_index_poseidon() {
let poseidon = Poseidon::<Fp, 2>::new();
let default_leaf = [0u8; 64];
let rng = OsRng;
let leaves = [Fp::random(rng), Fp::random(rng), Fp::random(rng)];
const HEIGHT: usize = 3;
let smt =
create_merkle_tree::<Fp, Poseidon<Fp, 2>, HEIGHT>(poseidon, &leaves, &default_leaf);
let index = 2;
let proof = smt.generate_membership_proof(index);
let res = proof
.get_index(&smt.root(), &leaves[index as usize], &poseidon)
.unwrap();
let desired_res = Fp::from(index);
assert_eq!(res, desired_res);
}
#[test]
fn should_update_leaf_poseidon() {
let poseidon = Poseidon::<Fp, 2>::new();
let default_leaf = [0u8; 64];
let rng = OsRng;
let leaves = [Fp::random(rng), Fp::random(rng), Fp::random(rng)];
const HEIGHT: usize = 3;
let mut smt =
create_merkle_tree::<Fp, Poseidon<Fp, 2>, HEIGHT>(poseidon, &leaves, &default_leaf);
let empty_hashes =
gen_empty_hashes::<Fp, Poseidon<Fp, 2>, HEIGHT>(&poseidon, &default_leaf).unwrap();
let root = smt.root();
let index = 2;
let new_leaf = Fp::from(2_u64);
smt.update(index, new_leaf, &poseidon).unwrap();
let hash1 = leaves[0];
let hash2 = leaves[1];
let hash3 = new_leaf;
let hash12 = poseidon.hash([hash1, hash2]).unwrap();
let hash34 = poseidon.hash([hash3, empty_hashes[0]]).unwrap();
let calc_root = poseidon.hash([hash12, hash34]).unwrap();
let new_root = smt.root();
assert_ne!(root, new_root);
assert_eq!(calc_root, new_root);
}
#[test]
fn should_delete_leaf_poseidon() {
let poseidon = Poseidon::<Fp, 2>::new();
let default_leaf = [0u8; 64];
let rng = OsRng;
let leaves = [Fp::random(rng), Fp::random(rng), Fp::random(rng)];
const HEIGHT: usize = 3;
let mut smt =
create_merkle_tree::<Fp, Poseidon<Fp, 2>, HEIGHT>(poseidon, &leaves, &default_leaf);
let empty_hashes =
gen_empty_hashes::<Fp, Poseidon<Fp, 2>, HEIGHT>(&poseidon, &default_leaf).unwrap();
let root = smt.root();
let index = 2;
smt.delete(index, &poseidon).unwrap();
let hash1 = leaves[0];
let hash2 = leaves[1];
let hash3 = Fp::zero();
let hash12 = poseidon.hash([hash1, hash2]).unwrap();
let hash34 = poseidon.hash([hash3, empty_hashes[0]]).unwrap();
let calc_root = poseidon.hash([hash12, hash34]).unwrap();
let new_root = smt.root();
assert_ne!(root, new_root);
assert_eq!(calc_root, new_root);
}
}