use std::collections::VecDeque;
use std::fmt::Write as _;
use serde::{Deserialize, Serialize};
use crate::bytes::{Bytes32, decode_hex, encode_hex};
use crate::error::{Error, Result};
use crate::hashes::NodeHashFn;
#[inline]
const fn left_child(i: usize) -> usize {
2 * i + 1
}
#[inline]
const fn right_child(i: usize) -> usize {
2 * i + 2
}
#[inline]
const fn parent(i: usize) -> Result<usize> {
if i == 0 {
Err(Error::RootHasNoParent)
} else {
Ok((i - 1) / 2)
}
}
#[inline]
const fn sibling(i: usize) -> Result<usize> {
if i == 0 {
Err(Error::RootHasNoSibling)
} else if i % 2 == 1 {
Ok(i + 1)
} else {
Ok(i - 1)
}
}
#[inline]
const fn is_internal(tree_len: usize, i: usize) -> bool {
left_child(i) < tree_len
}
#[inline]
const fn is_leaf(tree_len: usize, i: usize) -> bool {
i < tree_len && !is_internal(tree_len, i)
}
const fn check_leaf(tree_len: usize, i: usize) -> Result<()> {
if is_leaf(tree_len, i) {
Ok(())
} else {
Err(Error::NotALeaf(i))
}
}
pub(crate) fn build(leaves: &[Bytes32], node_hash: NodeHashFn) -> Result<Vec<Bytes32>> {
if leaves.is_empty() {
return Err(Error::EmptyLeaves);
}
let n = leaves.len();
let tree_len = 2 * n - 1;
let mut tree = vec![[0u8; 32]; tree_len];
let leaf_start = tree_len - n;
for (slot, leaf) in tree
.get_mut(leaf_start..)
.ok_or(Error::EmptyLeaves)?
.iter_mut()
.rev()
.zip(leaves)
{
*slot = *leaf;
}
for i in (0..leaf_start).rev() {
let l = left_child(i);
let r = right_child(i);
let hash = node_hash(
tree.get(l).ok_or(Error::IndexOutOfBounds {
index: l,
len: tree_len,
})?,
tree.get(r).ok_or(Error::IndexOutOfBounds {
index: r,
len: tree_len,
})?,
);
*tree.get_mut(i).ok_or(Error::IndexOutOfBounds {
index: i,
len: tree_len,
})? = hash;
}
Ok(tree)
}
pub(crate) fn proof(tree: &[Bytes32], index: usize) -> Result<Vec<Bytes32>> {
check_leaf(tree.len(), index)?;
let mut result = Vec::new();
let mut idx = index;
while idx > 0 {
let sib = sibling(idx)?;
result.push(*tree.get(sib).ok_or(Error::IndexOutOfBounds {
index: sib,
len: tree.len(),
})?);
idx = parent(idx)?;
}
Ok(result)
}
pub(crate) fn process_proof(leaf: &Bytes32, proof: &[Bytes32], node_hash: NodeHashFn) -> Bytes32 {
let mut current = *leaf;
for sib in proof {
current = node_hash(¤t, sib);
}
current
}
pub(crate) fn is_valid(tree: &[Bytes32], node_hash: NodeHashFn) -> bool {
if tree.is_empty() {
return false;
}
for i in 0..tree.len() {
let l = left_child(i);
let r = right_child(i);
match (tree.get(l), tree.get(r)) {
(Some(lv), Some(rv)) => {
let Some(node) = tree.get(i) else {
return false;
};
if *node != node_hash(lv, rv) {
return false;
}
}
(Some(_), None) => {
return false;
}
_ => {}
}
}
true
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MultiProof {
pub leaves: Vec<Bytes32>,
pub proof: Vec<Bytes32>,
pub proof_flags: Vec<bool>,
}
pub(crate) fn multi_proof(tree: &[Bytes32], indices: &[usize]) -> Result<MultiProof> {
for &i in indices {
check_leaf(tree.len(), i)?;
}
let mut sorted: Vec<usize> = indices.to_vec();
sorted.sort_unstable_by(|a, b| b.cmp(a));
for pair in sorted.windows(2) {
if let [a, b] = *pair
&& a == b
{
return Err(Error::DuplicateIndex(a));
}
}
let mut queue: VecDeque<usize> = sorted.iter().copied().collect();
let mut proof_nodes = Vec::new();
let mut flags = Vec::new();
while let Some(j) = queue.pop_front() {
if j == 0 {
break;
}
let s = sibling(j)?;
let p = parent(j)?;
if queue.front() == Some(&s) {
flags.push(true);
queue.pop_front();
} else {
flags.push(false);
proof_nodes.push(*tree.get(s).ok_or(Error::IndexOutOfBounds {
index: s,
len: tree.len(),
})?);
}
queue.push_back(p);
}
if indices.is_empty() {
proof_nodes.push(*tree.first().ok_or(Error::EmptyLeaves)?);
}
let leaves: Vec<Bytes32> = sorted
.iter()
.map(|&i| {
tree.get(i).copied().ok_or(Error::IndexOutOfBounds {
index: i,
len: tree.len(),
})
})
.collect::<Result<_>>()?;
Ok(MultiProof {
leaves,
proof: proof_nodes,
proof_flags: flags,
})
}
pub(crate) fn process_multi_proof(mp: &MultiProof, node_hash: NodeHashFn) -> Result<Bytes32> {
let proof_needed = mp.proof_flags.iter().filter(|&&f| !f).count();
if mp.proof.len() < proof_needed {
return Err(Error::InvalidMultiproof {
expected: proof_needed,
got: mp.proof.len(),
});
}
if mp.leaves.len() + mp.proof.len() != mp.proof_flags.len() + 1 {
return Err(Error::IncompatibleMultiproof {
leaves: mp.leaves.len(),
proof: mp.proof.len(),
flags: mp.proof_flags.len(),
});
}
let mut stack: VecDeque<Bytes32> = mp.leaves.iter().copied().collect();
let mut proof_iter = mp.proof.iter();
for &flag in &mp.proof_flags {
let a = stack.pop_front().ok_or(Error::MultiproofStackEmpty)?;
let b = if flag {
stack.pop_front().ok_or(Error::MultiproofStackEmpty)?
} else {
*proof_iter.next().ok_or(Error::MultiproofProofExhausted)?
};
stack.push_back(node_hash(&a, &b));
}
let remaining: usize = stack.len() + proof_iter.count();
if remaining != 1 {
return Err(Error::MultiproofNotConverged);
}
stack.pop_front().ok_or(Error::MultiproofStackEmpty)
}
pub(crate) fn render(tree: &[Bytes32]) -> Result<String> {
if tree.is_empty() {
return Err(Error::EmptyLeaves);
}
let mut output = String::new();
let mut stack: Vec<(usize, Vec<bool>)> = vec![(0, vec![])];
while let Some((i, path)) = stack.pop() {
for &is_continuation in path.iter().take(path.len().saturating_sub(1)) {
output.push_str(if is_continuation { "│ " } else { " " });
}
if let Some(&is_left) = path.last() {
output.push_str(if is_left { "├─ " } else { "└─ " });
}
let node = tree.get(i).ok_or(Error::IndexOutOfBounds {
index: i,
len: tree.len(),
})?;
_ = writeln!(output, "{i}) {}", encode_hex(node));
let r = right_child(i);
if r < tree.len() {
let mut right_path = path.clone();
right_path.push(false);
stack.push((r, right_path));
let mut left_path = path;
left_path.push(true);
stack.push((left_child(i), left_path));
}
}
if output.ends_with('\n') {
output.pop();
}
Ok(output)
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct MultiProofJson {
pub leaves: Vec<String>,
pub proof: Vec<String>,
pub proof_flags: Vec<bool>,
}
impl TryFrom<MultiProofJson> for MultiProof {
type Error = Error;
fn try_from(json: MultiProofJson) -> Result<Self> {
let leaves = json
.leaves
.iter()
.map(|s| decode_hex(s))
.collect::<Result<Vec<_>>>()?;
let proof = json
.proof
.iter()
.map(|s| decode_hex(s))
.collect::<Result<Vec<_>>>()?;
Ok(Self {
leaves,
proof,
proof_flags: json.proof_flags,
})
}
}
impl From<&MultiProof> for MultiProofJson {
fn from(mp: &MultiProof) -> Self {
Self {
leaves: mp.leaves.iter().map(encode_hex).collect(),
proof: mp.proof.iter().map(encode_hex).collect(),
proof_flags: mp.proof_flags.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hashes::{keccak256, standard_node_hash};
fn test_leaves(count: usize) -> Vec<Bytes32> {
(0..count)
.map(|i| {
#[expect(clippy::cast_possible_truncation, reason = "test helper, i < 256")]
let b = i as u8;
keccak256(&[b])
})
.collect()
}
#[test]
fn build_and_validate() {
let leaves = test_leaves(4);
let tree = build(&leaves, standard_node_hash).unwrap();
assert_eq!(tree.len(), 7);
assert!(is_valid(&tree, standard_node_hash));
}
#[test]
fn single_leaf() {
let leaves = test_leaves(1);
let tree = build(&leaves, standard_node_hash).unwrap();
assert_eq!(tree.len(), 1);
assert!(is_valid(&tree, standard_node_hash));
}
#[test]
fn empty_leaves_rejected() {
let result = build(&[], standard_node_hash);
assert!(matches!(result, Err(Error::EmptyLeaves)));
}
#[test]
fn proof_roundtrip() {
let leaves = test_leaves(8);
let tree = build(&leaves, standard_node_hash).unwrap();
let first_leaf = tree.len() - leaves.len();
for i in first_leaf..tree.len() {
let p = proof(&tree, i).unwrap();
let root = process_proof(tree.get(i).unwrap(), &p, standard_node_hash);
assert_eq!(root, *tree.first().unwrap(), "proof failed for index {i}");
}
}
#[test]
fn multi_proof_roundtrip() {
let leaves = test_leaves(4);
let tree = build(&leaves, standard_node_hash).unwrap();
let mp = multi_proof(&tree, &[4, 5]).unwrap();
let root = process_multi_proof(&mp, standard_node_hash).unwrap();
assert_eq!(root, *tree.first().unwrap());
}
#[test]
fn multi_proof_all_leaves() {
let leaves = test_leaves(4);
let tree = build(&leaves, standard_node_hash).unwrap();
let indices: Vec<usize> = (tree.len() - leaves.len()..tree.len()).collect();
let mp = multi_proof(&tree, &indices).unwrap();
let root = process_multi_proof(&mp, standard_node_hash).unwrap();
assert_eq!(root, *tree.first().unwrap());
}
#[test]
fn multi_proof_empty_indices() {
let leaves = test_leaves(4);
let tree = build(&leaves, standard_node_hash).unwrap();
let mp = multi_proof(&tree, &[]).unwrap();
assert!(mp.leaves.is_empty());
assert_eq!(mp.proof.len(), 1);
assert_eq!(*mp.proof.first().unwrap(), *tree.first().unwrap());
}
#[test]
fn duplicate_index_rejected() {
let leaves = vec![[0u8; 32]; 2];
let tree = build(&leaves, standard_node_hash).unwrap();
let result = multi_proof(&tree, &[1, 1]);
assert!(matches!(result, Err(Error::DuplicateIndex(1))));
}
#[test]
fn proof_for_internal_node_rejected() {
let leaves = vec![[0u8; 32]; 2];
let tree = build(&leaves, standard_node_hash).unwrap();
assert!(matches!(proof(&tree, 0), Err(Error::NotALeaf(0))));
}
#[test]
fn invalid_trees() {
assert!(!is_valid(&[], standard_node_hash));
assert!(!is_valid(&[[0u8; 32]; 2], standard_node_hash));
assert!(!is_valid(&[[0u8; 32]; 3], standard_node_hash));
}
#[test]
fn render_tree() {
let leaves = test_leaves(2);
let tree = build(&leaves, standard_node_hash).unwrap();
let text = render(&tree).unwrap();
assert!(text.contains("0)"), "should contain root index");
assert!(text.contains("0x"), "should contain hex hashes");
}
#[test]
fn multi_proof_json_roundtrip() {
let mp = MultiProof {
leaves: vec![[0u8; 32]],
proof: vec![[1u8; 32]],
proof_flags: vec![true, false],
};
let json = MultiProofJson::from(&mp);
let json_str = serde_json::to_string(&json).unwrap();
assert!(json_str.contains("proofFlags"));
let parsed: MultiProofJson = serde_json::from_str(&json_str).unwrap();
let recovered = MultiProof::try_from(parsed).unwrap();
assert_eq!(mp, recovered);
}
#[test]
fn power_of_two_and_non_power() {
for count in [2, 3, 4, 5, 7, 8, 9, 15, 16] {
let leaves = test_leaves(count);
let tree = build(&leaves, standard_node_hash).unwrap();
assert_eq!(tree.len(), 2 * count - 1);
assert!(is_valid(&tree, standard_node_hash));
}
}
}