use serde::Serialize;
use crate::Digest;
use crate::hash::Hasher;
pub const DIGEST_LEN: usize = 32;
pub const LEAF_PREFIX: [u8; 1] = [0x00];
pub const INNER_PREFIX: [u8; 1] = [0x01];
pub const EMPTY_NODE: [u8; DIGEST_LEN] = [0u8; DIGEST_LEN];
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
pub enum Node {
Empty,
Digest([u8; DIGEST_LEN]),
}
impl Node {
pub fn bytes(&self) -> [u8; DIGEST_LEN] {
match self {
Self::Empty => EMPTY_NODE,
Self::Digest(value) => *value,
}
}
}
impl AsRef<[u8]> for Node {
fn as_ref(&self) -> &[u8] {
match self {
Self::Empty => &EMPTY_NODE,
Self::Digest(value) => value.as_ref(),
}
}
}
impl From<[u8; DIGEST_LEN]> for Node {
fn from(value: [u8; DIGEST_LEN]) -> Self {
Self::Digest(value)
}
}
impl From<Digest> for Node {
fn from(value: Digest) -> Self {
Self::Digest(value.into_inner())
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum MerkleError {
InvalidProof,
InvalidInput,
}
impl std::fmt::Display for MerkleError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InvalidProof => f.write_str("invalid merkle proof"),
Self::InvalidInput => f.write_str("invalid merkle input"),
}
}
}
impl std::error::Error for MerkleError {}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MerkleProof {
path: Vec<Node>,
}
impl MerkleProof {
pub fn new(path: Vec<Node>) -> Self {
Self { path }
}
pub fn path(&self) -> &[Node] {
&self.path
}
pub fn verify_proof_with_leaf_bytes(
&self,
root: &Node,
leaf_bytes: &[u8],
leaf_index: usize,
) -> Result<(), MerkleError> {
match self.compute_root(leaf_bytes, leaf_index) {
Some(computed) if &computed == root => Ok(()),
_ => Err(MerkleError::InvalidProof),
}
}
pub fn verify_proof<L: Serialize>(
&self,
root: &Node,
leaf: &L,
leaf_index: usize,
) -> Result<(), MerkleError> {
let bytes = bcs::to_bytes(leaf).map_err(|_| MerkleError::InvalidInput)?;
self.verify_proof_with_leaf_bytes(root, &bytes, leaf_index)
}
pub fn compute_root(&self, leaf: &[u8], leaf_index: usize) -> Option<Node> {
if leaf_index >> self.path.len() != 0 {
return None;
}
let mut current = leaf_hash(leaf);
let mut level_index = leaf_index;
for sibling in &self.path {
current = if level_index.is_multiple_of(2) {
inner_hash(¤t, sibling)
} else {
inner_hash(sibling, ¤t)
};
level_index /= 2;
}
Some(current)
}
pub fn is_right_most(&self, leaf_index: usize) -> bool {
let mut level_index = leaf_index;
for sibling in &self.path {
if level_index.is_multiple_of(2) && sibling.as_ref() != EMPTY_NODE.as_ref() {
return false;
}
level_index /= 2;
}
true
}
}
#[derive(Debug)]
pub struct MerkleTree {
nodes: Vec<Node>,
n_leaves: usize,
}
impl MerkleTree {
pub fn build_from_serialized<I>(iter: I) -> Self
where
I: IntoIterator,
I::IntoIter: ExactSizeIterator,
I::Item: AsRef<[u8]>,
{
let leaf_hashes = iter
.into_iter()
.map(|leaf| leaf_hash(leaf.as_ref()))
.collect::<Vec<_>>();
Self::build_from_leaf_hashes(leaf_hashes)
}
pub fn build_from_unserialized<I>(iter: I) -> Result<Self, MerkleError>
where
I: IntoIterator,
I::IntoIter: ExactSizeIterator,
I::Item: Serialize,
{
let leaf_hashes = iter
.into_iter()
.map(|leaf| {
bcs::to_bytes(&leaf)
.map_err(|_| MerkleError::InvalidInput)
.map(|bytes| leaf_hash(&bytes))
})
.collect::<Result<Vec<_>, _>>()?;
Ok(Self::build_from_leaf_hashes(leaf_hashes))
}
pub fn build_from_leaf_hashes<I>(iter: I) -> Self
where
I: IntoIterator,
I::IntoIter: ExactSizeIterator<Item = Node>,
{
let iter = iter.into_iter();
let mut nodes = Vec::with_capacity(n_nodes(iter.len()));
nodes.extend(iter);
let n_leaves = nodes.len();
let mut level_nodes = n_leaves;
let mut prev_level_index = 0;
while level_nodes > 1 {
if level_nodes.is_multiple_of(2) {
} else {
nodes.push(Node::Empty);
level_nodes += 1;
}
let new_level_index = prev_level_index + level_nodes;
(prev_level_index..new_level_index)
.step_by(2)
.for_each(|index| nodes.push(inner_hash(&nodes[index], &nodes[index + 1])));
prev_level_index = new_level_index;
level_nodes /= 2;
}
Self { nodes, n_leaves }
}
pub fn root(&self) -> Node {
self.nodes.last().copied().unwrap_or(Node::Empty)
}
pub fn n_leaves(&self) -> usize {
self.n_leaves
}
pub fn get_proof(&self, leaf_index: usize) -> Result<MerkleProof, MerkleError> {
if leaf_index >= self.n_leaves {
return Err(MerkleError::InvalidInput);
}
let path_capacity = self
.n_leaves
.checked_ilog2()
.map(|log| log as usize + 1)
.unwrap_or(0);
let mut path = Vec::with_capacity(path_capacity);
let mut level_index = leaf_index;
let mut n_level = self.n_leaves;
let mut level_base_index = 0;
while n_level > 1 {
n_level = n_level.next_multiple_of(2);
let sibling_index = if level_index.is_multiple_of(2) {
level_base_index + level_index + 1
} else {
level_base_index + level_index - 1
};
path.push(self.nodes[sibling_index]);
level_index /= 2;
level_base_index += n_level;
n_level /= 2;
}
Ok(MerkleProof { path })
}
pub fn compute_non_inclusion_proof<L>(
&self,
sorted_leaves: &[L],
target: &L,
) -> Result<MerkleNonInclusionProof<L>, MerkleError>
where
L: Ord + Clone,
{
let position = sorted_leaves.partition_point(|leaf| leaf <= target);
if position > 0 && &sorted_leaves[position - 1] == target {
return Err(MerkleError::InvalidInput);
}
let left_leaf = if position > 0 {
Some((
sorted_leaves[position - 1].clone(),
self.get_proof(position - 1)?,
))
} else {
None
};
let right_leaf = if position < sorted_leaves.len() {
Some((sorted_leaves[position].clone(), self.get_proof(position)?))
} else {
None
};
Ok(MerkleNonInclusionProof {
index: position,
left_leaf,
right_leaf,
})
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MerkleNonInclusionProof<L> {
index: usize,
left_leaf: Option<(L, MerkleProof)>,
right_leaf: Option<(L, MerkleProof)>,
}
impl<L> MerkleNonInclusionProof<L> {
pub fn new(
index: usize,
left_leaf: Option<(L, MerkleProof)>,
right_leaf: Option<(L, MerkleProof)>,
) -> Self {
Self {
index,
left_leaf,
right_leaf,
}
}
pub fn index(&self) -> usize {
self.index
}
pub fn left_leaf(&self) -> Option<&(L, MerkleProof)> {
self.left_leaf.as_ref()
}
pub fn right_leaf(&self) -> Option<&(L, MerkleProof)> {
self.right_leaf.as_ref()
}
}
impl<L> MerkleNonInclusionProof<L>
where
L: Serialize,
{
pub fn verify_proof_by_key<K, F>(
&self,
root: &Node,
target_key: &K,
key_of: F,
) -> Result<(), MerkleError>
where
K: Ord + ?Sized,
F: Fn(&L) -> &K,
{
if root.as_ref() == EMPTY_NODE.as_ref() {
return Ok(());
}
let right_leaf_index = self.index;
let left_leaf_with_index = self.left_leaf.as_ref().zip(self.index.checked_sub(1));
if let Some(((left_leaf, left_proof), left_leaf_index)) = left_leaf_with_index {
left_proof.verify_proof(root, left_leaf, left_leaf_index)?;
if key_of(left_leaf) >= target_key {
return Err(MerkleError::InvalidProof);
}
} else if right_leaf_index != 0 || self.right_leaf.is_none() {
return Err(MerkleError::InvalidProof);
}
if let Some((right_leaf, right_proof)) = &self.right_leaf {
right_proof.verify_proof(root, right_leaf, right_leaf_index)?;
if key_of(right_leaf) <= target_key {
return Err(MerkleError::InvalidProof);
}
} else if let Some(((_, left_proof), left_leaf_index)) = left_leaf_with_index {
if !left_proof.is_right_most(left_leaf_index) {
return Err(MerkleError::InvalidProof);
}
} else {
return Err(MerkleError::InvalidProof);
}
Ok(())
}
}
impl<L> MerkleNonInclusionProof<L>
where
L: Ord + Serialize,
{
pub fn verify_proof(&self, root: &Node, target: &L) -> Result<(), MerkleError> {
self.verify_proof_by_key(root, target, |leaf| leaf)
}
}
pub(crate) fn leaf_hash(input: &[u8]) -> Node {
let mut hasher = Hasher::new();
hasher.update(LEAF_PREFIX);
hasher.update(input);
Node::Digest(hasher.finalize().into_inner())
}
fn inner_hash(left: &Node, right: &Node) -> Node {
let mut hasher = Hasher::new();
hasher.update(INNER_PREFIX);
hasher.update(left.bytes());
hasher.update(right.bytes());
Node::Digest(hasher.finalize().into_inner())
}
pub(crate) fn n_nodes(n_leaves: usize) -> usize {
let mut lvl_nodes = n_leaves;
let mut tot_nodes = 0;
while lvl_nodes > 1 {
lvl_nodes += lvl_nodes % 2;
tot_nodes += lvl_nodes;
lvl_nodes /= 2;
}
tot_nodes + lvl_nodes
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(target_arch = "wasm32")]
use wasm_bindgen_test::wasm_bindgen_test as test;
const TEST_INPUT: [&[u8]; 9] = [
b"foo", b"bar", b"fizz", b"baz", b"buzz", b"fizz", b"foobar", b"walrus", b"fizz",
];
#[test]
fn n_nodes_formula() {
assert_eq!(n_nodes(0), 0);
assert_eq!(n_nodes(1), 1);
assert_eq!(n_nodes(2), 3);
assert_eq!(n_nodes(3), 7);
assert_eq!(n_nodes(4), 7);
assert_eq!(n_nodes(5), 13);
assert_eq!(n_nodes(6), 13);
assert_eq!(n_nodes(7), 15);
assert_eq!(n_nodes(8), 15);
assert_eq!(n_nodes(9), 23);
}
#[test]
fn empty_tree_root_is_empty_node() {
let tree = MerkleTree::build_from_serialized::<[&[u8]; 0]>([]);
assert_eq!(tree.root().bytes(), EMPTY_NODE);
}
#[test]
fn single_element_tree_root_is_leaf_hash() {
let leaf = b"Test";
let tree = MerkleTree::build_from_serialized([leaf.as_ref()]);
let mut hasher = Hasher::new();
hasher.update(LEAF_PREFIX);
hasher.update(leaf);
assert_eq!(tree.root().bytes(), hasher.finalize().into_inner());
}
#[test]
fn empty_element_tree_root_is_hash_of_empty_leaf() {
let tree = MerkleTree::build_from_serialized([&[][..]]);
let mut hasher = Hasher::new();
hasher.update(LEAF_PREFIX);
hasher.update::<&[u8]>(&[]);
assert_eq!(tree.root().bytes(), hasher.finalize().into_inner());
}
#[test]
fn get_proof_out_of_bounds() {
for i in 0..TEST_INPUT.len() {
let tree = MerkleTree::build_from_serialized(&TEST_INPUT[..i]);
assert_eq!(
tree.get_proof(i.next_power_of_two()),
Err(MerkleError::InvalidInput),
);
}
}
#[test]
fn every_proof_round_trips() {
for i in 0..TEST_INPUT.len() {
let tree = MerkleTree::build_from_serialized(&TEST_INPUT[..i]);
for (index, leaf) in TEST_INPUT[..i].iter().enumerate() {
let proof = tree.get_proof(index).unwrap();
proof
.verify_proof_with_leaf_bytes(&tree.root(), leaf, index)
.unwrap();
}
}
}
#[test]
fn proof_fails_for_wrong_index() {
for i in 1..TEST_INPUT.len() {
let tree = MerkleTree::build_from_serialized(&TEST_INPUT[..i]);
for (index, leaf) in TEST_INPUT[..i].iter().enumerate() {
let proof = tree.get_proof(index).unwrap();
assert_eq!(
proof.verify_proof_with_leaf_bytes(&tree.root(), leaf, index + 1),
Err(MerkleError::InvalidProof),
);
}
}
}
#[test]
fn proof_fails_for_tampered_leaf() {
let tree = MerkleTree::build_from_serialized(TEST_INPUT);
let proof = tree.get_proof(3).unwrap();
let tampered = b"not-the-real-leaf";
assert_eq!(
proof.verify_proof_with_leaf_bytes(&tree.root(), tampered, 3),
Err(MerkleError::InvalidProof),
);
}
#[test]
fn proof_fails_against_wrong_root() {
let tree = MerkleTree::build_from_serialized(TEST_INPUT);
let proof = tree.get_proof(2).unwrap();
let wrong_root = Node::Digest([0xab; DIGEST_LEN]);
assert_eq!(
proof.verify_proof_with_leaf_bytes(&wrong_root, TEST_INPUT[2], 2),
Err(MerkleError::InvalidProof),
);
}
#[test]
fn is_right_most_detects_last_leaf() {
for i in 1..TEST_INPUT.len() {
let tree = MerkleTree::build_from_serialized(&TEST_INPUT[..i]);
for j in 0..i {
let proof = tree.get_proof(j).unwrap();
let expected = j == i - 1;
assert_eq!(proof.is_right_most(j), expected);
}
}
}
#[test]
fn non_inclusion_empty_tree() {
let tree = MerkleTree::build_from_unserialized::<[&[u8]; 0]>([]).unwrap();
let leaves: [&[u8]; 0] = [];
let proof = tree
.compute_non_inclusion_proof(&leaves, &b"foo".as_ref())
.unwrap();
assert!(proof.left_leaf().is_none());
assert!(proof.right_leaf().is_none());
assert_eq!(proof.index(), 0);
proof.verify_proof(&tree.root(), &b"foo".as_ref()).unwrap();
proof.verify_proof(&tree.root(), &b"bar".as_ref()).unwrap();
}
#[test]
fn non_inclusion_single_leaf() {
let leaves: [&[u8]; 1] = [b"foo"];
let tree = MerkleTree::build_from_unserialized(&leaves).unwrap();
let proof = tree
.compute_non_inclusion_proof(&leaves, &b"bar".as_ref())
.unwrap();
proof.verify_proof(&tree.root(), &b"bar".as_ref()).unwrap();
assert_eq!(
proof.verify_proof(&tree.root(), &b"foo".as_ref()),
Err(MerkleError::InvalidProof),
);
assert_eq!(
tree.compute_non_inclusion_proof(&leaves, &b"foo".as_ref()),
Err(MerkleError::InvalidInput),
);
}
#[test]
fn non_inclusion_multiple_leaves() {
const RAW: [&str; 9] = [
"foo", "bar", "fizz", "baz", "buzz", "fizz", "foobar", "walrus", "fizz",
];
let mut sorted: Vec<&str> = RAW.to_vec();
sorted.sort();
sorted.dedup();
let tree = MerkleTree::build_from_unserialized(&sorted).unwrap();
let probes = ["fuzz", "yankee", "aloha", "foo", "bar", "fizz", "walrus"];
for probe in probes {
let result = tree.compute_non_inclusion_proof(&sorted, &probe);
if sorted.contains(&probe) {
assert_eq!(result, Err(MerkleError::InvalidInput));
} else {
let proof = result.unwrap();
proof.verify_proof(&tree.root(), &probe).unwrap();
}
}
}
#[test]
fn non_inclusion_at_extremes() {
let sorted: [&str; 3] = ["bar", "foo", "qux"];
let tree = MerkleTree::build_from_unserialized(&sorted).unwrap();
let before = tree.compute_non_inclusion_proof(&sorted, &"aaa").unwrap();
assert!(before.left_leaf().is_none());
assert!(before.right_leaf().is_some());
assert_eq!(before.index(), 0);
before.verify_proof(&tree.root(), &"aaa").unwrap();
let after = tree.compute_non_inclusion_proof(&sorted, &"zzz").unwrap();
assert!(after.left_leaf().is_some());
assert!(after.right_leaf().is_none());
assert_eq!(after.index(), sorted.len());
after.verify_proof(&tree.root(), &"zzz").unwrap();
}
#[test]
fn non_inclusion_forged_zero_index_with_left_leaf() {
let leaves: [&[u8]; 1] = [b"foo"];
let tree = MerkleTree::build_from_unserialized(&leaves).unwrap();
let fake: &[u8] = b"fake";
let forged = MerkleNonInclusionProof::new(
0,
Some((fake, tree.get_proof(0).unwrap())),
Some((fake, tree.get_proof(0).unwrap())),
);
assert_eq!(
forged.verify_proof(&tree.root(), &fake),
Err(MerkleError::InvalidProof),
);
}
#[test]
fn root_matches_upstream_for_known_input() {
const EXPECTED_ROOT: [u8; DIGEST_LEN] = [
0x8d, 0x01, 0x06, 0x76, 0xde, 0x3d, 0x66, 0x08, 0x77, 0xcc, 0x8c, 0x27, 0xa4, 0x2d,
0xcf, 0xf9, 0xc1, 0x15, 0x97, 0x20, 0x36, 0x1a, 0x82, 0x36, 0xd2, 0xd2, 0x07, 0xb6,
0x8b, 0x72, 0x9b, 0x0c,
];
let tree = MerkleTree::build_from_serialized(TEST_INPUT);
assert_eq!(tree.root().bytes(), EXPECTED_ROOT);
}
#[cfg(feature = "proptest")]
mod proptests {
use super::*;
use proptest::collection::vec;
use proptest::prelude::*;
use test_strategy::proptest;
#[cfg(target_arch = "wasm32")]
use wasm_bindgen_test::wasm_bindgen_test as test;
fn small_leaves() -> impl Strategy<Value = Vec<u32>> {
vec(any::<u32>(), 1..=32)
}
fn sorted_unique_leaves() -> impl Strategy<Value = Vec<u32>> {
vec(any::<u32>(), 0..=32).prop_map(|mut v| {
v.sort();
v.dedup();
v
})
}
#[proptest]
fn inclusion_proof_round_trips(#[strategy(small_leaves())] leaves: Vec<u32>) {
let tree = MerkleTree::build_from_unserialized(leaves.iter()).unwrap();
let root = tree.root();
for (index, leaf) in leaves.iter().enumerate() {
let proof = tree.get_proof(index).unwrap();
proof
.verify_proof(&root, leaf, index)
.expect("freshly-built proof must verify");
}
}
#[proptest]
fn inclusion_proof_rejects_wrong_leaf(#[strategy(small_leaves())] leaves: Vec<u32>) {
prop_assume!(leaves.len() >= 2);
let tree = MerkleTree::build_from_unserialized(leaves.iter()).unwrap();
let root = tree.root();
for (index, leaf) in leaves.iter().enumerate() {
let other_index = (index + 1) % leaves.len();
if leaves[other_index] == *leaf {
continue;
}
let proof = tree.get_proof(index).unwrap();
prop_assert!(
proof
.verify_proof(&root, &leaves[other_index], index)
.is_err(),
"proof for index {index} must not verify a different leaf at the same index",
);
}
}
#[proptest]
fn non_inclusion_proof_round_trips(
#[strategy(sorted_unique_leaves())] leaves: Vec<u32>,
target: u32,
) {
prop_assume!(leaves.binary_search(&target).is_err());
let tree = MerkleTree::build_from_unserialized(leaves.iter()).unwrap();
let root = tree.root();
let proof = tree.compute_non_inclusion_proof(&leaves, &target).unwrap();
proof
.verify_proof(&root, &target)
.expect("freshly-built non-inclusion proof must verify");
}
#[proptest]
fn non_inclusion_rejects_present_target(
#[strategy(sorted_unique_leaves())] leaves: Vec<u32>,
) {
prop_assume!(!leaves.is_empty());
let tree = MerkleTree::build_from_unserialized(leaves.iter()).unwrap();
for present in &leaves {
prop_assert_eq!(
tree.compute_non_inclusion_proof(&leaves, present),
Err(MerkleError::InvalidInput),
);
}
}
#[proptest]
fn is_right_most_only_for_last_leaf(#[strategy(small_leaves())] leaves: Vec<u32>) {
let tree = MerkleTree::build_from_unserialized(leaves.iter()).unwrap();
let last = leaves.len() - 1;
for j in 0..leaves.len() {
let proof = tree.get_proof(j).unwrap();
prop_assert_eq!(proof.is_right_most(j), j == last);
}
}
}
}