use crate::crypto::{CryptoPolicyDefault, hexdigest};
use crate::error::{Error, Result};
#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash, PartialOrd, Ord)]
pub struct LeafHash(pub [u8; 32]);
#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash, PartialOrd, Ord)]
pub struct NodeHash(pub [u8; 32]);
#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash, PartialOrd, Ord)]
pub struct MerkleRoot(pub [u8; 32]);
impl LeafHash {
pub fn to_hex(&self) -> String {
hex::encode(self.0)
}
pub fn from_hex(s: &str) -> Result<Self> {
let bytes = hex::decode(s)?;
if bytes.len() != 32 {
return Err(Error::msg(format!(
"leaf hash must be 32 bytes, got {}",
bytes.len()
)));
}
let mut arr = [0u8; 32];
arr.copy_from_slice(&bytes);
Ok(LeafHash(arr))
}
}
impl NodeHash {
pub fn to_hex(&self) -> String {
hex::encode(self.0)
}
}
impl MerkleRoot {
pub fn to_hex(&self) -> String {
hex::encode(self.0)
}
}
impl std::fmt::Display for MerkleRoot {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.to_hex())
}
}
const LEAF_TAG: u8 = 0x00;
const INTERNAL_TAG: u8 = 0x01;
pub fn hash_leaf(data: &[u8]) -> Result<LeafHash> {
let policy = CryptoPolicyDefault {};
let mut input = Vec::with_capacity(1 + data.len());
input.push(LEAF_TAG);
input.extend_from_slice(data);
let hex = hexdigest("sha3-256", &input, &policy)?;
let mut arr = [0u8; 32];
arr.copy_from_slice(&hex::decode(hex)?);
Ok(LeafHash(arr))
}
pub fn hash_internal(left: &LeafHash, right: &LeafHash) -> Result<NodeHash> {
let policy = CryptoPolicyDefault {};
let mut input = Vec::with_capacity(1 + 32 + 32);
input.push(INTERNAL_TAG);
input.extend_from_slice(&left.0);
input.extend_from_slice(&right.0);
let hex = hexdigest("sha3-256", &input, &policy)?;
let mut arr = [0u8; 32];
arr.copy_from_slice(&hex::decode(hex)?);
Ok(NodeHash(arr))
}
#[derive(Clone, Debug)]
pub struct MerkleTree {
levels: Vec<Vec<LeafHash>>,
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub struct LeafIndex(pub usize);
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub struct ProofStep {
pub sibling: LeafHash,
pub sibling_is_left: bool,
}
#[derive(Clone, Debug)]
pub struct MerkleProof {
pub leaf: LeafHash,
pub steps: Vec<ProofStep>,
}
impl MerkleTree {
pub fn from_leaves(leaves: &[Vec<u8>]) -> Result<Self> {
if leaves.is_empty() {
return Ok(MerkleTree { levels: Vec::new() });
}
let mut levels: Vec<Vec<LeafHash>> = Vec::new();
let mut current: Vec<LeafHash> =
leaves.iter().map(|d| hash_leaf(d)).collect::<Result<_>>()?;
levels.push(current.clone());
while current.len() > 1 {
let mut next: Vec<LeafHash> = Vec::with_capacity(current.len().div_ceil(2));
let mut i = 0;
while i < current.len() {
if i + 1 < current.len() {
let node = hash_internal(¤t[i], ¤t[i + 1])?;
next.push(LeafHash(node.0));
i += 2;
} else {
next.push(current[i]);
i += 1;
}
}
levels.push(next.clone());
current = next;
}
Ok(MerkleTree { levels })
}
pub fn leaf_count(&self) -> usize {
self.levels.first().map(|l| l.len()).unwrap_or(0)
}
pub fn is_empty(&self) -> bool {
self.leaf_count() == 0
}
pub fn root(&self) -> Result<MerkleRoot> {
if self.is_empty() {
return Err(Error::msg("Merkle tree is empty; no root"));
}
let top = self.levels.last().unwrap();
debug_assert_eq!(top.len(), 1);
Ok(MerkleRoot(top[0].0))
}
pub fn leaves(&self) -> &[LeafHash] {
self.levels.first().map(|l| l.as_slice()).unwrap_or(&[])
}
pub fn proof(&self, idx: usize) -> Result<MerkleProof> {
if self.is_empty() {
return Err(Error::msg("Merkle tree is empty; no proofs"));
}
let leaves = self.leaves();
if idx >= leaves.len() {
return Err(Error::msg(format!(
"leaf index {} out of bounds (tree has {} leaves)",
idx,
leaves.len()
)));
}
let leaf = leaves[idx];
let mut steps = Vec::new();
let mut i = idx;
for level in 0..self.levels.len() - 1 {
let layer = &self.levels[level];
let sibling_idx = i ^ 1;
if sibling_idx < layer.len() && sibling_idx != i {
let sibling = layer[sibling_idx];
let sibling_is_left = sibling_idx < i;
steps.push(ProofStep {
sibling,
sibling_is_left,
});
} else {
}
i /= 2;
}
Ok(MerkleProof { leaf, steps })
}
pub fn debug_dump(&self) -> String {
let mut out = String::new();
for (level_num, level) in self.levels.iter().enumerate() {
out.push_str(&format!("level {}: ", level_num));
for (i, h) in level.iter().enumerate() {
if i > 0 {
out.push(' ');
}
out.push_str(&h.to_hex());
}
out.push('\n');
}
out
}
}
pub fn verify_proof(root: &MerkleRoot, proof: &MerkleProof) -> Result<bool> {
let mut current = proof.leaf;
for step in &proof.steps {
let node = if step.sibling_is_left {
hash_internal(&step.sibling, ¤t)?
} else {
hash_internal(¤t, &step.sibling)?
};
current = LeafHash(node.0);
}
Ok(current.0 == root.0)
}
#[cfg(test)]
mod tests {
use super::*;
fn leaf(s: &str) -> Vec<u8> {
s.as_bytes().to_vec()
}
#[test]
fn empty_tree_has_no_root() {
let t = MerkleTree::from_leaves(&[]).unwrap();
assert!(t.is_empty());
assert!(t.root().is_err());
}
#[test]
fn single_leaf_tree_root_equals_leaf_hash() {
let t = MerkleTree::from_leaves(&[leaf("hello")]).unwrap();
let root = t.root().unwrap();
let expected = hash_leaf(b"hello").unwrap();
assert_eq!(root.0, expected.0);
}
#[test]
fn root_is_stable_for_same_input() {
let leaves = vec![leaf("a"), leaf("b"), leaf("c")];
let t1 = MerkleTree::from_leaves(&leaves).unwrap();
let t2 = MerkleTree::from_leaves(&leaves).unwrap();
assert_eq!(t1.root().unwrap(), t2.root().unwrap());
}
#[test]
fn root_changes_when_leaf_changes() {
let t1 = MerkleTree::from_leaves(&[leaf("a"), leaf("b")]).unwrap();
let t2 = MerkleTree::from_leaves(&[leaf("a"), leaf("c")]).unwrap();
assert_ne!(t1.root().unwrap(), t2.root().unwrap());
}
#[test]
fn odd_width_promotes_last_node() {
let t = MerkleTree::from_leaves(&[leaf("a"), leaf("b"), leaf("c")]).unwrap();
assert_eq!(t.levels[0].len(), 3);
assert_eq!(t.levels[1].len(), 2);
assert_eq!(t.levels[2].len(), 1);
}
#[test]
fn proof_round_trips() {
let leaves = vec![leaf("a"), leaf("b"), leaf("c"), leaf("d")];
let t = MerkleTree::from_leaves(&leaves).unwrap();
let root = t.root().unwrap();
for i in 0..leaves.len() {
let proof = t.proof(i).unwrap();
assert!(verify_proof(&root, &proof).unwrap());
}
}
#[test]
fn proof_round_trips_odd_width() {
let leaves = vec![leaf("a"), leaf("b"), leaf("c")];
let t = MerkleTree::from_leaves(&leaves).unwrap();
let root = t.root().unwrap();
for i in 0..leaves.len() {
let proof = t.proof(i).unwrap();
assert!(verify_proof(&root, &proof).unwrap());
}
}
#[test]
fn proof_fails_with_wrong_root() {
let leaves = vec![leaf("a"), leaf("b"), leaf("c"), leaf("d")];
let t = MerkleTree::from_leaves(&leaves).unwrap();
let proof = t.proof(0).unwrap();
let other = MerkleTree::from_leaves(&[leaf("x"), leaf("y"), leaf("z"), leaf("w")]).unwrap();
let wrong_root = other.root().unwrap();
assert!(!verify_proof(&wrong_root, &proof).unwrap());
}
#[test]
fn proof_fails_with_tampered_leaf() {
let leaves = vec![leaf("a"), leaf("b"), leaf("c"), leaf("d")];
let t = MerkleTree::from_leaves(&leaves).unwrap();
let root = t.root().unwrap();
let mut proof = t.proof(0).unwrap();
proof.leaf = hash_leaf(b"X").unwrap();
assert!(!verify_proof(&root, &proof).unwrap());
}
#[test]
fn proof_fails_with_tampered_sibling() {
let leaves = vec![leaf("a"), leaf("b"), leaf("c"), leaf("d")];
let t = MerkleTree::from_leaves(&leaves).unwrap();
let root = t.root().unwrap();
let mut proof = t.proof(0).unwrap();
assert!(!proof.steps.is_empty());
proof.steps[0].sibling = hash_leaf(b"X").unwrap();
assert!(!verify_proof(&root, &proof).unwrap());
}
#[test]
fn out_of_bounds_index_rejected() {
let t = MerkleTree::from_leaves(&[leaf("a")]).unwrap();
assert!(t.proof(5).is_err());
}
#[test]
fn leaf_hash_changes_with_data() {
let h1 = hash_leaf(b"hello").unwrap();
let h2 = hash_leaf(b"world").unwrap();
assert_ne!(h1, h2);
}
#[test]
fn leaf_and_internal_tags_differ() {
let leaf_hash = hash_leaf(&[0u8; 64]).unwrap();
let policy = CryptoPolicyDefault {};
let untagged = hexdigest("sha3-256", &[0u8; 64], &policy).unwrap();
assert_ne!(leaf_hash.to_hex(), untagged);
}
#[test]
fn large_tree_root_computes() {
let leaves: Vec<Vec<u8>> = (0..100)
.map(|i| format!("leaf-{}", i).into_bytes())
.collect();
let t1 = MerkleTree::from_leaves(&leaves).unwrap();
let t2 = MerkleTree::from_leaves(&leaves).unwrap();
assert_eq!(t1.root().unwrap(), t2.root().unwrap());
assert_eq!(t1.leaf_count(), 100);
}
#[test]
fn debug_dump_is_nonempty_for_nonempty_tree() {
let t = MerkleTree::from_leaves(&[leaf("a"), leaf("b")]).unwrap();
let dump = t.debug_dump();
assert!(dump.contains("level 0"));
assert!(dump.contains("level 1"));
}
}