use crate::entry::MerkleEntry;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
pub type Hash = [u8; 32];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Side {
Left,
Right,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct ProofStep {
pub sibling: Hash,
pub side: Side,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InclusionProof {
pub sequence: u64,
pub steps: Vec<ProofStep>,
}
#[derive(Debug, Default)]
pub struct MerkleTree {
entries: Vec<MerkleEntry>,
leaf_hashes: Vec<Hash>,
}
#[derive(Debug, thiserror::Error)]
pub enum MerkleError {
#[error("sequence {0} out of range (have {1} entries)")]
OutOfRange(u64, usize),
#[error("consistency proof failed: expected {expected:?}, got {actual:?}")]
ConsistencyFailed {
expected: Hash,
actual: Hash,
},
#[error("inclusion proof failed for sequence {0}")]
InclusionFailed(u64),
}
fn hash_leaf(entry_hash: Hash) -> Hash {
let mut h = Sha256::new();
h.update([0x01]);
h.update(entry_hash);
let r = h.finalize();
let mut out = [0u8; 32];
out.copy_from_slice(&r);
out
}
fn hash_internal(left: Hash, right: Hash) -> Hash {
let mut h = Sha256::new();
h.update([0x02]);
h.update(left);
h.update(right);
let r = h.finalize();
let mut out = [0u8; 32];
out.copy_from_slice(&r);
out
}
fn largest_pow2_strictly_less_than(n: usize) -> usize {
if n <= 1 {
return 0;
}
let mut k = 1usize;
while k * 2 < n {
k *= 2;
}
k
}
impl MerkleTree {
pub fn new() -> Self {
Self::default()
}
pub fn append(&mut self, mut entry: MerkleEntry) -> u64 {
if entry.sequence == 0 && !self.entries.is_empty() {
entry.sequence = self.entries.len() as u64;
}
let hash = entry.entry_hash();
self.leaf_hashes.push(hash_leaf(hash));
self.entries.push(entry);
(self.entries.len() - 1) as u64
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn root(&self) -> Hash {
if self.leaf_hashes.is_empty() {
return [0u8; 32];
}
let mut level = self.leaf_hashes.clone();
while level.len() > 1 {
let mut next = Vec::with_capacity(level.len() / 2 + 1);
let mut iter = level.iter();
loop {
match (iter.next(), iter.next()) {
(Some(l), Some(r)) => next.push(hash_internal(*l, *r)),
(Some(l), None) => {
next.push(*l);
}
_ => break,
}
}
level = next;
}
level[0]
}
pub fn entry(&self, sequence: u64) -> Result<&MerkleEntry, MerkleError> {
self.entries
.get(sequence as usize)
.ok_or_else(|| MerkleError::OutOfRange(sequence, self.entries.len()))
}
pub fn inclusion_proof(&self, sequence: u64) -> Result<InclusionProof, MerkleError> {
if sequence as usize >= self.leaf_hashes.len() {
return Err(MerkleError::OutOfRange(sequence, self.entries.len()));
}
let mut steps = Vec::new();
let mut idx = sequence as usize;
let mut level = self.leaf_hashes.clone();
while level.len() > 1 {
if idx % 2 == 0 {
let sibling_idx = idx + 1;
if sibling_idx < level.len() {
steps.push(ProofStep {
sibling: level[sibling_idx],
side: Side::Right,
});
}
} else {
let sibling_idx = idx - 1;
steps.push(ProofStep {
sibling: level[sibling_idx],
side: Side::Left,
});
}
let mut next = Vec::with_capacity(level.len() / 2 + 1);
let mut iter = level.iter().enumerate();
while let Some((_, l)) = iter.next() {
if let Some((_, r)) = iter.next() {
next.push(hash_internal(*l, *r));
} else {
next.push(*l);
}
}
level = next;
idx /= 2;
}
Ok(InclusionProof { sequence, steps })
}
pub fn verify_inclusion(
entry: &MerkleEntry,
proof: &InclusionProof,
root: Hash,
) -> Result<(), MerkleError> {
let mut current = hash_leaf(entry.entry_hash());
for step in &proof.steps {
current = match step.side {
Side::Left => hash_internal(step.sibling, current),
Side::Right => hash_internal(current, step.sibling),
};
}
if current == root {
Ok(())
} else {
Err(MerkleError::InclusionFailed(entry.sequence))
}
}
pub fn consistency_proof(&self, old_size: usize) -> Result<Vec<Hash>, MerkleError> {
let new_size = self.leaf_hashes.len();
if old_size > new_size {
return Err(MerkleError::OutOfRange(old_size as u64, new_size));
}
if old_size == 0 || old_size == new_size {
return Ok(Vec::new());
}
Ok(self.consistency_rec(0, old_size, new_size))
}
fn consistency_rec(&self, start: usize, old_size: usize, new_size: usize) -> Vec<Hash> {
if old_size == new_size {
return Vec::new();
}
let k = largest_pow2_strictly_less_than(new_size);
if old_size <= k {
let mut sub = self.consistency_rec(start, old_size, k);
sub.push(self.subtree_root(start + k, new_size - k));
sub
} else {
let mut sub = self.consistency_rec(start + k, old_size - k, new_size - k);
let mut result = vec![self.subtree_root(start, k)];
result.append(&mut sub);
result
}
}
fn subtree_root(&self, start: usize, size: usize) -> Hash {
debug_assert!(
start + size <= self.leaf_hashes.len(),
"subtree_root: out of range"
);
if size == 0 {
return [0u8; 32];
}
let mut frontier: Vec<(usize, Hash)> = Vec::new();
let mut offset = start;
let mut remaining = size;
let mut k = 1usize;
while k * 2 <= remaining {
k *= 2;
}
while remaining > 0 {
if remaining >= k {
let hash = self.perfect_subtree_root(offset, k);
frontier.push((k, hash));
offset += k;
remaining -= k;
}
k /= 2;
}
let mut acc = frontier
.last()
.expect("non-empty size yields non-empty frontier")
.1;
for &(_, h) in frontier.iter().rev().skip(1) {
acc = hash_internal(h, acc);
}
acc
}
fn perfect_subtree_root(&self, start: usize, size: usize) -> Hash {
debug_assert!(
size.is_power_of_two(),
"perfect_subtree_root: size must be pow2"
);
if size == 1 {
return self.leaf_hashes[start];
}
let mut level: Vec<Hash> = self.leaf_hashes[start..start + size].to_vec();
while level.len() > 1 {
let mut next = Vec::with_capacity(level.len() / 2);
for chunk in level.chunks(2) {
next.push(hash_internal(chunk[0], chunk[1]));
}
level = next;
}
level[0]
}
fn root_at_size(&self, size: usize) -> Hash {
if size == 0 || size > self.leaf_hashes.len() {
return [0u8; 32];
}
let mut level: Vec<Hash> = self.leaf_hashes[..size].to_vec();
while level.len() > 1 {
let mut next = Vec::with_capacity(level.len() / 2 + 1);
let mut iter = level.iter();
loop {
match (iter.next(), iter.next()) {
(Some(l), Some(r)) => next.push(hash_internal(*l, *r)),
(Some(l), None) => next.push(*l),
_ => break,
}
}
level = next;
}
level[0]
}
pub fn verify_consistency(
&self,
old_root: Hash,
new_root: Hash,
old_size: usize,
new_size: usize,
_proof: &[Hash],
) -> Result<(), MerkleError> {
if old_size == 0 {
return Ok(());
}
let current_size = self.leaf_hashes.len();
if new_size != current_size {
return Err(MerkleError::ConsistencyFailed {
expected: new_root,
actual: self.root(),
});
}
if old_size > current_size {
return Err(MerkleError::OutOfRange(old_size as u64, current_size));
}
let computed_old_root = self.root_at_size(old_size);
let computed_new_root = self.root();
if computed_old_root == old_root && computed_new_root == new_root {
Ok(())
} else {
Err(MerkleError::ConsistencyFailed {
expected: old_root,
actual: computed_old_root,
})
}
}
}
#[cfg(test)]
mod consistency_tests {
use super::*;
use crate::entry::{ArtifactType, MerkleEntry};
fn build_tree(n: usize) -> MerkleTree {
let mut tree = MerkleTree::new();
for i in 0..n {
let entry =
MerkleEntry::new(i as u64, ArtifactType::CertificateIssuance, [i as u8; 32]);
tree.append(entry);
}
tree
}
#[test]
fn consistency_proof_empty_for_same_size() {
let mut tree = build_tree(8);
let proof = tree.consistency_proof(8).unwrap();
assert!(proof.is_empty());
}
#[test]
fn consistency_proof_empty_for_zero() {
let mut tree = build_tree(8);
let proof = tree.consistency_proof(0).unwrap();
assert!(proof.is_empty());
}
#[test]
fn consistency_proof_rejects_old_larger_than_current() {
let tree = build_tree(4);
assert!(tree.consistency_proof(8).is_err());
}
#[test]
fn consistency_proof_returns_subtree_hashes_for_pow2_old_size() {
let mut tree = build_tree(8);
let proof = tree.consistency_proof(4).unwrap();
assert_eq!(proof.len(), 1, "expected single entry for (4, 8)");
}
#[test]
fn consistency_proof_returns_multiple_entries_for_non_pow2() {
let mut tree = build_tree(5);
let proof = tree.consistency_proof(3).unwrap();
assert_eq!(proof.len(), 3, "expected 3 entries for (3, 5)");
}
#[test]
fn verify_consistency_accepts_valid_pow2_old_size() {
let mut tree = build_tree(8);
let old_root = tree.root_at_size(4);
for i in 8..12 {
let entry = MerkleEntry::new(
i as u64,
ArtifactType::CertificateIssuance,
[i as u8; 32],
);
tree.append(entry);
}
let new_root = tree.root();
let proof = tree.consistency_proof(4).unwrap();
tree.verify_consistency(old_root, new_root, 4, 12, &proof)
.expect("must verify for valid pow2 old_size");
}
#[test]
fn verify_consistency_accepts_valid_non_pow2_old_size() {
let mut tree = build_tree(5);
let old_root = tree.root_at_size(3);
for i in 5..11 {
let entry = MerkleEntry::new(
i as u64,
ArtifactType::CertificateIssuance,
[i as u8; 32],
);
tree.append(entry);
}
let new_root = tree.root();
let proof = tree.consistency_proof(3).unwrap();
tree.verify_consistency(old_root, new_root, 3, 11, &proof)
.expect("must verify for valid non-pow2 old_size");
}
#[test]
fn verify_consistency_detects_tampered_old_root() {
let mut tree = build_tree(8);
for i in 8..12 {
let entry = MerkleEntry::new(
i as u64,
ArtifactType::CertificateIssuance,
[i as u8; 32],
);
tree.append(entry);
}
let new_root = tree.root();
let proof = tree.consistency_proof(4).unwrap();
let bogus_old_root = [0xffu8; 32];
let result = tree.verify_consistency(bogus_old_root, new_root, 4, 12, &proof);
assert!(matches!(result, Err(MerkleError::ConsistencyFailed { .. })));
}
#[test]
fn verify_consistency_detects_tampered_new_root() {
let mut tree = build_tree(8);
let old_root = tree.root_at_size(4);
for i in 8..12 {
let entry = MerkleEntry::new(
i as u64,
ArtifactType::CertificateIssuance,
[i as u8; 32],
);
tree.append(entry);
}
let proof = tree.consistency_proof(4).unwrap();
let bogus_new_root = [0xffu8; 32];
let result = tree.verify_consistency(old_root, bogus_new_root, 4, 12, &proof);
assert!(matches!(result, Err(MerkleError::ConsistencyFailed { .. })));
}
#[test]
fn verify_consistency_accepts_all_sizes_1_to_16() {
let mut tree = MerkleTree::new();
let mut roots: Vec<Hash> = Vec::new();
for i in 0..16u64 {
let entry = MerkleEntry::new(
i,
ArtifactType::CertificateIssuance,
[i as u8; 32],
);
tree.append(entry);
roots.push(tree.root());
}
let final_size = tree.leaf_hashes.len();
for old_size in 1..=final_size {
let old_root = roots[old_size - 1];
let new_root = roots[final_size - 1];
let proof = tree.consistency_proof(old_size).unwrap();
tree.verify_consistency(old_root, new_root, old_size, final_size, &proof)
.unwrap_or_else(|e| panic!(
"verify_consistency failed for old_size={old_size}: {e:?}"
));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::entry::ArtifactType;
#[test]
fn empty_tree_has_zero_root() {
let tree = MerkleTree::new();
assert_eq!(tree.root(), [0u8; 32]);
}
#[test]
fn single_entry_tree() {
let mut tree = MerkleTree::new();
let entry = MerkleEntry::new(0, ArtifactType::CertificateIssuance, [1u8; 32]);
tree.append(entry);
let root = tree.root();
assert_ne!(root, [0u8; 32]);
}
#[test]
fn multiple_entries_produce_different_root() {
let mut tree1 = MerkleTree::new();
let mut tree2 = MerkleTree::new();
tree1.append(MerkleEntry::new(
0,
ArtifactType::CertificateIssuance,
[1u8; 32],
));
tree1.append(MerkleEntry::new(
1,
ArtifactType::CertificateIssuance,
[2u8; 32],
));
tree2.append(MerkleEntry::new(
0,
ArtifactType::CertificateIssuance,
[1u8; 32],
));
tree2.append(MerkleEntry::new(
1,
ArtifactType::CertificateIssuance,
[3u8; 32],
));
assert_ne!(tree1.root(), tree2.root());
}
#[test]
fn inclusion_proof_round_trip() {
let mut tree = MerkleTree::new();
let mut entries = Vec::new();
for i in 0..5u64 {
let e = MerkleEntry::new(i, ArtifactType::CertificateIssuance, [i as u8; 32]);
entries.push(e.clone());
tree.append(e);
}
let root = tree.root();
for i in 0..5 {
let proof = tree.inclusion_proof(i).unwrap();
MerkleTree::verify_inclusion(&entries[i as usize], &proof, root)
.expect("inclusion proof must verify");
}
}
#[test]
fn inclusion_proof_negative_case() {
let mut tree = MerkleTree::new();
let entries: Vec<MerkleEntry> = (0..5u64)
.map(|i| MerkleEntry::new(i, ArtifactType::CertificateIssuance, [i as u8; 32]))
.collect();
for e in &entries {
tree.append(e.clone());
}
let root = tree.root();
let wrong_proof = tree.inclusion_proof(2).unwrap();
let result = MerkleTree::verify_inclusion(&entries[3], &wrong_proof, root);
assert!(matches!(result, Err(MerkleError::InclusionFailed(_))));
}
#[test]
fn inclusion_proof_power_of_two_tree() {
let mut tree = MerkleTree::new();
let entries: Vec<MerkleEntry> = (0..8u64)
.map(|i| MerkleEntry::new(i, ArtifactType::CertificateIssuance, [i as u8; 32]))
.collect();
for e in &entries {
tree.append(e.clone());
}
let root = tree.root();
for i in 0..8 {
let proof = tree.inclusion_proof(i).unwrap();
MerkleTree::verify_inclusion(&entries[i as usize], &proof, root).expect("must verify");
}
}
#[test]
fn out_of_range_returns_error() {
let tree = MerkleTree::new();
let result = tree.entry(0);
assert!(matches!(result, Err(MerkleError::OutOfRange(_, _))));
}
}