1use thiserror::Error;
8
9#[derive(Debug, Error, PartialEq)]
15pub enum MerkleError {
16 #[error("Merkle tree is empty")]
18 EmptyTree,
19
20 #[error("leaf index {index} is out of bounds for tree of size {tree_size}")]
22 LeafIndexOutOfBounds { index: usize, tree_size: usize },
23
24 #[error("invalid proof: {reason}")]
26 InvalidProof { reason: String },
27
28 #[error("duplicate leaf index {index}")]
30 DuplicateLeaf { index: usize },
31}
32
33#[derive(Clone, Debug)]
39pub enum MerkleNode {
40 Leaf {
42 index: usize,
44 hash: u64,
46 },
47 Internal {
49 left: u64,
51 right: u64,
53 hash: u64,
55 },
56}
57
58impl MerkleNode {
59 pub fn hash(&self) -> u64 {
61 match self {
62 MerkleNode::Leaf { hash, .. } => *hash,
63 MerkleNode::Internal { hash, .. } => *hash,
64 }
65 }
66}
67
68#[derive(Debug, Clone)]
74pub struct BatchProof {
75 pub leaf_indices: Vec<usize>,
77 pub sibling_hashes: Vec<(usize, u64)>,
80 pub root: u64,
82}
83
84impl BatchProof {
85 pub fn proof_size(&self) -> usize {
87 self.sibling_hashes.len()
88 }
89}
90
91const FNV_OFFSET: u64 = 14_695_981_039_346_656_037;
97const FNV_PRIME: u64 = 1_099_511_628_211;
98
99#[inline]
101fn fnv1a(data: &[u8]) -> u64 {
102 let mut hash = FNV_OFFSET;
103 for &byte in data {
104 hash ^= u64::from(byte);
105 hash = hash.wrapping_mul(FNV_PRIME);
106 }
107 hash
108}
109
110#[derive(Debug)]
131pub struct MerkleBatchProver {
132 pub leaves: Vec<u64>,
134}
135
136impl MerkleBatchProver {
137 pub fn new(data: &[&[u8]]) -> Result<Self, MerkleError> {
147 if data.is_empty() {
148 return Err(MerkleError::EmptyTree);
149 }
150 let leaves = data.iter().map(|d| Self::leaf_hash(d)).collect();
151 Ok(Self { leaves })
152 }
153
154 pub fn tree_size(&self) -> usize {
160 next_power_of_two(self.leaves.len())
161 }
162
163 pub fn root(&self) -> u64 {
167 let size = self.tree_size();
168 let mut level: Vec<u64> = (0..size)
169 .map(|i| {
170 if i < self.leaves.len() {
171 self.leaves[i]
172 } else {
173 0u64
174 }
175 })
176 .collect();
177
178 while level.len() > 1 {
179 level = level
180 .chunks(2)
181 .map(|pair| Self::combine(pair[0], pair[1]))
182 .collect();
183 }
184
185 level[0]
186 }
187
188 pub fn prove_batch(&self, indices: &[usize]) -> Result<BatchProof, MerkleError> {
195 let n = self.leaves.len();
197 let mut sorted = indices.to_vec();
198 sorted.sort_unstable();
199
200 for &idx in &sorted {
201 if idx >= n {
202 return Err(MerkleError::LeafIndexOutOfBounds {
203 index: idx,
204 tree_size: n,
205 });
206 }
207 }
208 for window in sorted.windows(2) {
209 if window[0] == window[1] {
210 return Err(MerkleError::DuplicateLeaf { index: window[0] });
211 }
212 }
213
214 let size = self.tree_size();
216 let tree = build_tree(&self.leaves, size);
217 let height = tree.len() - 1; let mut sibling_hashes: Vec<(usize, u64)> = Vec::new();
224
225 let mut covered: std::collections::BTreeSet<usize> = sorted.iter().copied().collect();
229
230 for (level, level_nodes) in tree.iter().enumerate().take(height) {
231 let mut next_covered: std::collections::BTreeSet<usize> =
232 std::collections::BTreeSet::new();
233 for &pos in &covered {
235 let sibling = if pos % 2 == 0 { pos + 1 } else { pos - 1 };
236 let parent = pos / 2;
237 next_covered.insert(parent);
238
239 if !covered.contains(&sibling) {
241 let sib_hash = level_nodes.get(sibling).copied().unwrap_or(0);
243 let entry = (level, sib_hash);
245 if !sibling_hashes.contains(&entry) {
246 sibling_hashes.push(entry);
247 }
248 }
249 }
250 covered = next_covered;
251 }
252
253 Ok(BatchProof {
254 leaf_indices: sorted,
255 sibling_hashes,
256 root: self.root(),
257 })
258 }
259
260 pub fn verify_batch(
270 &self,
271 proof: &BatchProof,
272 data_items: &[&[u8]],
273 ) -> Result<bool, MerkleError> {
274 if data_items.len() != proof.leaf_indices.len() {
275 return Err(MerkleError::InvalidProof {
276 reason: format!(
277 "data_items length {} does not match leaf_indices length {}",
278 data_items.len(),
279 proof.leaf_indices.len()
280 ),
281 });
282 }
283
284 let size = self.tree_size();
285 let height = size.trailing_zeros() as usize; let mut known: std::collections::HashMap<(usize, usize), u64> =
290 std::collections::HashMap::new();
291
292 for (i, &leaf_idx) in proof.leaf_indices.iter().enumerate() {
294 let hash = Self::leaf_hash(data_items[i]);
295 known.insert((0, leaf_idx), hash);
296 }
297
298 let mut covered: std::collections::BTreeSet<usize> =
302 proof.leaf_indices.iter().copied().collect();
303 let mut sibling_iter = proof.sibling_hashes.iter();
304
305 for level in 0..height {
306 let mut next_covered: std::collections::BTreeSet<usize> =
307 std::collections::BTreeSet::new();
308 for &pos in &covered {
309 let sibling = if pos % 2 == 0 { pos + 1 } else { pos - 1 };
310 let parent = pos / 2;
311 next_covered.insert(parent);
312
313 if !covered.contains(&sibling) {
314 match sibling_iter.next() {
316 Some(&(_lv, sib_hash)) => {
317 known.insert((level, sibling), sib_hash);
318 }
319 None => {
320 return Err(MerkleError::InvalidProof {
321 reason: "not enough sibling hashes in proof".to_string(),
322 });
323 }
324 }
325 }
326 }
327 covered = next_covered;
328 }
329
330 let mut current_level_known: std::collections::HashMap<usize, u64> = known
333 .iter()
334 .filter(|((lvl, _), _)| *lvl == 0)
335 .map(|((_, pos), &hash)| (*pos, hash))
336 .collect();
337
338 for level in 0..height {
339 for ((lvl, pos), &hash) in &known {
341 if *lvl == level {
342 current_level_known.insert(*pos, hash);
343 }
344 }
345
346 let mut next_level: std::collections::HashMap<usize, u64> =
347 std::collections::HashMap::new();
348
349 let mut positions: Vec<usize> = current_level_known.keys().copied().collect();
351 positions.sort_unstable();
352 positions.dedup();
353
354 let mut processed_parents: std::collections::HashSet<usize> =
356 std::collections::HashSet::new();
357 for pos in positions {
358 let parent = pos / 2;
359 if processed_parents.contains(&parent) {
360 continue;
361 }
362 let left_pos = parent * 2;
363 let right_pos = parent * 2 + 1;
364 if let (Some(&left_h), Some(&right_h)) = (
365 current_level_known.get(&left_pos),
366 current_level_known.get(&right_pos),
367 ) {
368 let parent_hash = Self::combine(left_h, right_h);
369 next_level.insert(parent, parent_hash);
370 processed_parents.insert(parent);
371 }
372 }
373
374 current_level_known = next_level;
375 }
376
377 match current_level_known.get(&0) {
379 Some(&reconstructed_root) => Ok(reconstructed_root == proof.root),
380 None => Err(MerkleError::InvalidProof {
381 reason: "could not reconstruct root from proof".to_string(),
382 }),
383 }
384 }
385
386 pub fn leaf_hash(data: &[u8]) -> u64 {
392 fnv1a(data)
393 }
394
395 pub fn combine(left: u64, right: u64) -> u64 {
400 let mut buf = [0u8; 16];
401 buf[..8].copy_from_slice(&left.to_le_bytes());
402 buf[8..].copy_from_slice(&right.to_le_bytes());
403 fnv1a(&buf)
404 }
405}
406
407fn next_power_of_two(n: usize) -> usize {
414 if n <= 1 {
415 return 1;
416 }
417 let mut p = 1usize;
418 while p < n {
419 p <<= 1;
420 }
421 p
422}
423
424fn build_tree(leaves: &[u64], size: usize) -> Vec<Vec<u64>> {
429 let mut level: Vec<u64> = (0..size)
430 .map(|i| if i < leaves.len() { leaves[i] } else { 0u64 })
431 .collect();
432
433 let mut tree = vec![level.clone()];
434 while level.len() > 1 {
435 level = level
436 .chunks(2)
437 .map(|pair| MerkleBatchProver::combine(pair[0], pair[1]))
438 .collect();
439 tree.push(level.clone());
440 }
441 tree
442}
443
444#[cfg(test)]
449mod tests {
450 use super::*;
451
452 #[test]
457 fn test_new_single_leaf() {
458 let prover = MerkleBatchProver::new(&[b"hello"]).unwrap();
459 assert_eq!(prover.leaves.len(), 1);
460 assert_eq!(prover.leaves[0], MerkleBatchProver::leaf_hash(b"hello"));
461 }
462
463 #[test]
464 fn test_new_empty_returns_empty_tree() {
465 let result = MerkleBatchProver::new(&[]);
466 assert_eq!(result.unwrap_err(), MerkleError::EmptyTree);
467 }
468
469 #[test]
474 fn test_root_single_leaf_equals_leaf_hash() {
475 let prover = MerkleBatchProver::new(&[b"abc"]).unwrap();
476 assert_eq!(prover.root(), MerkleBatchProver::leaf_hash(b"abc"));
477 }
478
479 #[test]
480 fn test_root_two_leaves_equals_combine() {
481 let prover = MerkleBatchProver::new(&[b"left", b"right"]).unwrap();
482 let h0 = MerkleBatchProver::leaf_hash(b"left");
483 let h1 = MerkleBatchProver::leaf_hash(b"right");
484 assert_eq!(prover.root(), MerkleBatchProver::combine(h0, h1));
485 }
486
487 #[test]
488 fn test_root_power_of_two_tree_correct() {
489 let data: &[&[u8]] = &[b"a", b"b", b"c", b"d"];
491 let prover = MerkleBatchProver::new(data).unwrap();
492 let h0 = MerkleBatchProver::leaf_hash(b"a");
493 let h1 = MerkleBatchProver::leaf_hash(b"b");
494 let h2 = MerkleBatchProver::leaf_hash(b"c");
495 let h3 = MerkleBatchProver::leaf_hash(b"d");
496 let expected = MerkleBatchProver::combine(
497 MerkleBatchProver::combine(h0, h1),
498 MerkleBatchProver::combine(h2, h3),
499 );
500 assert_eq!(prover.root(), expected);
501 }
502
503 #[test]
504 fn test_root_non_power_of_two_pads_with_zeros() {
505 let data: &[&[u8]] = &[b"x", b"y", b"z"];
507 let prover = MerkleBatchProver::new(data).unwrap();
508 let h0 = MerkleBatchProver::leaf_hash(b"x");
509 let h1 = MerkleBatchProver::leaf_hash(b"y");
510 let h2 = MerkleBatchProver::leaf_hash(b"z");
511 let h3 = 0u64; let expected = MerkleBatchProver::combine(
513 MerkleBatchProver::combine(h0, h1),
514 MerkleBatchProver::combine(h2, h3),
515 );
516 assert_eq!(prover.root(), expected);
517 }
518
519 #[test]
524 fn test_prove_batch_out_of_bounds_returns_error() {
525 let prover = MerkleBatchProver::new(&[b"a", b"b"]).unwrap();
526 let result = prover.prove_batch(&[5]);
527 assert!(matches!(
528 result.unwrap_err(),
529 MerkleError::LeafIndexOutOfBounds {
530 index: 5,
531 tree_size: 2
532 }
533 ));
534 }
535
536 #[test]
537 fn test_prove_batch_duplicate_index_returns_error() {
538 let prover = MerkleBatchProver::new(&[b"a", b"b", b"c"]).unwrap();
539 let result = prover.prove_batch(&[1, 1]);
540 assert!(matches!(
541 result.unwrap_err(),
542 MerkleError::DuplicateLeaf { index: 1 }
543 ));
544 }
545
546 #[test]
551 fn test_prove_batch_single_leaf_produces_proof() {
552 let prover = MerkleBatchProver::new(&[b"a", b"b", b"c", b"d"]).unwrap();
553 let proof = prover.prove_batch(&[0]).unwrap();
554 assert_eq!(proof.leaf_indices, vec![0]);
555 assert_eq!(proof.proof_size(), 2);
557 }
558
559 #[test]
560 fn test_prove_batch_two_leaves_smaller_than_two_individual() {
561 let data: &[&[u8]] = &[b"a", b"b", b"c", b"d"];
562 let prover = MerkleBatchProver::new(data).unwrap();
563 let batch_proof = prover.prove_batch(&[0, 1]).unwrap();
565 let proof_0 = prover.prove_batch(&[0]).unwrap();
566 let proof_1 = prover.prove_batch(&[1]).unwrap();
567 assert!(batch_proof.proof_size() < proof_0.proof_size() + proof_1.proof_size());
568 }
569
570 #[test]
575 fn test_verify_batch_valid_proof_returns_true() {
576 let data: &[&[u8]] = &[b"hello", b"world", b"foo", b"bar"];
577 let prover = MerkleBatchProver::new(data).unwrap();
578 let proof = prover.prove_batch(&[1, 3]).unwrap();
579 let result = prover.verify_batch(&proof, &[b"world", b"bar"]).unwrap();
580 assert!(result);
581 }
582
583 #[test]
584 fn test_verify_batch_tampered_data_returns_false() {
585 let data: &[&[u8]] = &[b"hello", b"world", b"foo", b"bar"];
586 let prover = MerkleBatchProver::new(data).unwrap();
587 let proof = prover.prove_batch(&[0, 2]).unwrap();
588 let result = prover.verify_batch(&proof, &[b"tampered", b"foo"]).unwrap();
590 assert!(!result);
591 }
592
593 #[test]
594 fn test_verify_batch_indices_0_and_1_in_4_leaf_tree() {
595 let data: &[&[u8]] = &[b"a", b"b", b"c", b"d"];
596 let prover = MerkleBatchProver::new(data).unwrap();
597 let proof = prover.prove_batch(&[0, 1]).unwrap();
598 let result = prover.verify_batch(&proof, &[b"a", b"b"]).unwrap();
599 assert!(result);
600 }
601
602 #[test]
607 fn test_leaf_hash_deterministic() {
608 let h1 = MerkleBatchProver::leaf_hash(b"deterministic");
609 let h2 = MerkleBatchProver::leaf_hash(b"deterministic");
610 assert_eq!(h1, h2);
611 }
612
613 #[test]
614 fn test_combine_is_non_commutative() {
615 let a = MerkleBatchProver::leaf_hash(b"left");
616 let b = MerkleBatchProver::leaf_hash(b"right");
617 assert_ne!(a, b); assert_ne!(
619 MerkleBatchProver::combine(a, b),
620 MerkleBatchProver::combine(b, a)
621 );
622 }
623
624 #[test]
625 fn test_tree_size_is_next_power_of_two() {
626 let cases: &[(&[&[u8]], usize)] = &[
627 (&[b"a"], 1),
628 (&[b"a", b"b"], 2),
629 (&[b"a", b"b", b"c"], 4),
630 (&[b"a", b"b", b"c", b"d"], 4),
631 (&[b"a", b"b", b"c", b"d", b"e"], 8),
632 ];
633 for (data, expected) in cases {
634 let prover = MerkleBatchProver::new(data).unwrap();
635 assert_eq!(prover.tree_size(), *expected, "data.len()={}", data.len());
636 }
637 }
638
639 #[test]
644 fn test_verify_all_leaves_single_leaf_tree() {
645 let prover = MerkleBatchProver::new(&[b"solo"]).unwrap();
646 let proof = prover.prove_batch(&[0]).unwrap();
647 assert!(prover.verify_batch(&proof, &[b"solo"]).unwrap());
648 }
649
650 #[test]
651 fn test_root_eight_leaf_tree() {
652 let data: &[&[u8]] = &[b"1", b"2", b"3", b"4", b"5", b"6", b"7", b"8"];
653 let prover = MerkleBatchProver::new(data).unwrap();
654 let proof = prover.prove_batch(&[0, 1, 2, 3, 4, 5, 6, 7]).unwrap();
656 assert!(prover
657 .verify_batch(&proof, &[b"1", b"2", b"3", b"4", b"5", b"6", b"7", b"8"])
658 .unwrap());
659 }
660}