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
}
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))
}
}
}
#[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(_, _))));
}
}