Skip to main content

ipfrs_core/
merkle_batch.rs

1//! Merkle batch inclusion proof generation and verification.
2//!
3//! Provides [`MerkleBatchProver`] which generates and verifies batch inclusion
4//! proofs for multiple leaves in a single Merkle tree traversal, achieving
5//! O(n + log n) cost instead of O(n * log n) for individual proofs.
6
7use thiserror::Error;
8
9// ---------------------------------------------------------------------------
10// Error types
11// ---------------------------------------------------------------------------
12
13/// Errors that can occur during Merkle batch operations.
14#[derive(Debug, Error, PartialEq)]
15pub enum MerkleError {
16    /// The tree has no leaves.
17    #[error("Merkle tree is empty")]
18    EmptyTree,
19
20    /// A requested leaf index exceeds the number of leaves.
21    #[error("leaf index {index} is out of bounds for tree of size {tree_size}")]
22    LeafIndexOutOfBounds { index: usize, tree_size: usize },
23
24    /// Proof verification failed.
25    #[error("invalid proof: {reason}")]
26    InvalidProof { reason: String },
27
28    /// A leaf index was supplied more than once in a batch request.
29    #[error("duplicate leaf index {index}")]
30    DuplicateLeaf { index: usize },
31}
32
33// ---------------------------------------------------------------------------
34// MerkleNode
35// ---------------------------------------------------------------------------
36
37/// A node in a Merkle tree (either a leaf or an internal node).
38#[derive(Clone, Debug)]
39pub enum MerkleNode {
40    /// A leaf node storing the FNV-1a hash of the original data.
41    Leaf {
42        /// Position of the leaf in the original leaf array.
43        index: usize,
44        /// FNV-1a hash of the leaf's raw data.
45        hash: u64,
46    },
47    /// An internal node storing the FNV-1a hash of its two children combined.
48    Internal {
49        /// Hash of the left child.
50        left: u64,
51        /// Hash of the right child.
52        right: u64,
53        /// FNV-1a hash over `left.to_le_bytes() || right.to_le_bytes()`.
54        hash: u64,
55    },
56}
57
58impl MerkleNode {
59    /// Returns the hash stored in this node.
60    pub fn hash(&self) -> u64 {
61        match self {
62            MerkleNode::Leaf { hash, .. } => *hash,
63            MerkleNode::Internal { hash, .. } => *hash,
64        }
65    }
66}
67
68// ---------------------------------------------------------------------------
69// BatchProof
70// ---------------------------------------------------------------------------
71
72/// A batch inclusion proof for multiple leaves of a Merkle tree.
73#[derive(Debug, Clone)]
74pub struct BatchProof {
75    /// Sorted list of leaf indices covered by this proof.
76    pub leaf_indices: Vec<usize>,
77    /// Sibling hashes needed for verification, keyed by `(level, hash)`.
78    /// `level` 0 is the leaf level; the root is at `level = tree_height`.
79    pub sibling_hashes: Vec<(usize, u64)>,
80    /// The expected Merkle root.
81    pub root: u64,
82}
83
84impl BatchProof {
85    /// Returns the number of sibling hashes in this proof.
86    pub fn proof_size(&self) -> usize {
87        self.sibling_hashes.len()
88    }
89}
90
91// ---------------------------------------------------------------------------
92// FNV-1a helpers
93// ---------------------------------------------------------------------------
94
95/// FNV-1a 64-bit offset basis and prime.
96const FNV_OFFSET: u64 = 14_695_981_039_346_656_037;
97const FNV_PRIME: u64 = 1_099_511_628_211;
98
99/// Compute the FNV-1a 64-bit hash of a byte slice.
100#[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// ---------------------------------------------------------------------------
111// MerkleBatchProver
112// ---------------------------------------------------------------------------
113
114/// Generates and verifies batch Merkle inclusion proofs.
115///
116/// The tree is built over a padded-to-next-power-of-2 leaf array (padding with
117/// zero hashes). Leaf hashes use FNV-1a; internal nodes combine children with
118/// FNV-1a over their concatenated little-endian representations.
119///
120/// # Example
121///
122/// ```rust
123/// use ipfrs_core::merkle_batch::MerkleBatchProver;
124///
125/// let data: &[&[u8]] = &[b"hello", b"world", b"foo", b"bar"];
126/// let prover = MerkleBatchProver::new(data).unwrap();
127/// let proof = prover.prove_batch(&[0, 2]).unwrap();
128/// assert!(prover.verify_batch(&proof, &[b"hello", b"foo"]).unwrap());
129/// ```
130#[derive(Debug)]
131pub struct MerkleBatchProver {
132    /// FNV-1a hashes of the original leaf data.
133    pub leaves: Vec<u64>,
134}
135
136impl MerkleBatchProver {
137    // -----------------------------------------------------------------------
138    // Construction
139    // -----------------------------------------------------------------------
140
141    /// Create a new prover from raw leaf data.
142    ///
143    /// # Errors
144    ///
145    /// Returns [`MerkleError::EmptyTree`] if `data` is empty.
146    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    // -----------------------------------------------------------------------
155    // Public API
156    // -----------------------------------------------------------------------
157
158    /// Returns the number of leaves padded to the next power of two.
159    pub fn tree_size(&self) -> usize {
160        next_power_of_two(self.leaves.len())
161    }
162
163    /// Compute the Merkle root.
164    ///
165    /// The leaf layer is padded to `tree_size()` with zero hashes.
166    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    /// Generate a batch inclusion proof for the given leaf indices.
189    ///
190    /// # Errors
191    ///
192    /// - [`MerkleError::LeafIndexOutOfBounds`] if any index >= `leaves.len()`.
193    /// - [`MerkleError::DuplicateLeaf`] if the same index appears more than once.
194    pub fn prove_batch(&self, indices: &[usize]) -> Result<BatchProof, MerkleError> {
195        // Validate: bounds and duplicates.
196        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        // Build full padded tree.
215        let size = self.tree_size();
216        let tree = build_tree(&self.leaves, size);
217        let height = tree.len() - 1; // number of levels above leaves
218
219        // Collect sibling hashes using a set to dedup.
220        // We track which positions at each level are "already covered" by the
221        // batch (i.e., their hash will be recomputed from children), so we only
222        // emit siblings for uncovered positions.
223        let mut sibling_hashes: Vec<(usize, u64)> = Vec::new();
224
225        // At each level we track the set of node positions that will be
226        // recomputed (either they are a target leaf or the parent of two
227        // recomputed children).  Siblings of recomputed nodes are needed.
228        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 every covered node, its sibling may be needed unless also covered.
234            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                // The sibling is needed iff it is NOT itself in `covered`.
240                if !covered.contains(&sibling) {
241                    // Fetch sibling hash from the tree at the current level.
242                    let sib_hash = level_nodes.get(sibling).copied().unwrap_or(0);
243                    // Avoid emitting the same (level, hash) pair twice.
244                    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    /// Verify a [`BatchProof`] against the provided data items.
261    ///
262    /// `data_items` must correspond 1-to-1 with `proof.leaf_indices` (same
263    /// order and count).
264    ///
265    /// # Errors
266    ///
267    /// Returns [`MerkleError::InvalidProof`] if the lengths do not match or if
268    /// internal reconstruction fails.
269    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; // log2(size)
286
287        // Rebuild the partial tree bottom-up using the proof's sibling hashes.
288        // We store known hashes as a map: (level, position) -> hash.
289        let mut known: std::collections::HashMap<(usize, usize), u64> =
290            std::collections::HashMap::new();
291
292        // Insert leaf hashes.
293        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        // Insert sibling hashes level by level.
299        // We need the positions of siblings to insert them correctly.
300        // Re-derive sibling positions from the proof's leaf_indices.
301        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                    // Pull next sibling hash from the iterator.
315                    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        // Propagate upwards to reconstruct the root.
331        // Collect all positions at level 0 that are known.
332        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            // Add sibling-provided hashes at this level.
340            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            // Find pairs where both or at least one is known (with sibling).
350            let mut positions: Vec<usize> = current_level_known.keys().copied().collect();
351            positions.sort_unstable();
352            positions.dedup();
353
354            // Attempt to compute parents for all known positions.
355            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        // The root should be at position 0 of the topmost level.
378        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    // -----------------------------------------------------------------------
387    // Hash primitives
388    // -----------------------------------------------------------------------
389
390    /// Compute FNV-1a hash of leaf data.
391    pub fn leaf_hash(data: &[u8]) -> u64 {
392        fnv1a(data)
393    }
394
395    /// Combine two child hashes into a parent hash.
396    ///
397    /// Uses FNV-1a over the concatenation of the left and right hashes
398    /// in little-endian byte order, making it non-commutative.
399    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
407// ---------------------------------------------------------------------------
408// Internal tree-building helper
409// ---------------------------------------------------------------------------
410
411/// Returns the smallest power of two that is >= `n`.  Panics only for `n == 0`
412/// (which the public API guards against).
413fn 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
424/// Build the complete padded Merkle tree as a vector of levels.
425///
426/// `tree[0]` is the leaf level (padded to `size`), `tree[height]` is a
427/// single-element vec containing the root.
428fn 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// ---------------------------------------------------------------------------
445// Tests
446// ---------------------------------------------------------------------------
447
448#[cfg(test)]
449mod tests {
450    use super::*;
451
452    // ------------------------------------------------------------------
453    // Construction
454    // ------------------------------------------------------------------
455
456    #[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    // ------------------------------------------------------------------
470    // root()
471    // ------------------------------------------------------------------
472
473    #[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        // 4-leaf tree: root = combine(combine(h0,h1), combine(h2,h3))
490        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        // 3-leaf tree pads to 4; the 4th leaf hash is 0.
506        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; // padding
512        let expected = MerkleBatchProver::combine(
513            MerkleBatchProver::combine(h0, h1),
514            MerkleBatchProver::combine(h2, h3),
515        );
516        assert_eq!(prover.root(), expected);
517    }
518
519    // ------------------------------------------------------------------
520    // prove_batch() error cases
521    // ------------------------------------------------------------------
522
523    #[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    // ------------------------------------------------------------------
547    // prove_batch() success cases
548    // ------------------------------------------------------------------
549
550    #[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        // For a single leaf in a 4-leaf tree we need log2(4) = 2 siblings.
556        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        // Adjacent leaves 0 and 1 share a parent; only 1 sibling needed at level 1.
564        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    // ------------------------------------------------------------------
571    // verify_batch()
572    // ------------------------------------------------------------------
573
574    #[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        // Pass tampered data for index 0.
589        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    // ------------------------------------------------------------------
603    // Primitives
604    // ------------------------------------------------------------------
605
606    #[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); // sanity: different hashes
618        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    // ------------------------------------------------------------------
640    // Additional edge-case / regression tests
641    // ------------------------------------------------------------------
642
643    #[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        // Verify round-trip: prove all leaves, verify.
655        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}