Skip to main content

miden_crypto/merkle/partial_mt/
mod.rs

1use alloc::{
2    collections::{BTreeMap, BTreeSet},
3    string::String,
4    vec::Vec,
5};
6use core::fmt;
7
8use super::{
9    EMPTY_WORD, InnerNodeInfo, MerkleError, MerklePath, MerkleProof, NodeIndex, Poseidon2, Word,
10};
11use crate::utils::{
12    ByteReader, ByteWriter, Deserializable, DeserializationError, Serializable, word_to_hex,
13};
14
15#[cfg(test)]
16mod tests;
17
18// CONSTANTS
19// ================================================================================================
20
21/// Index of the root node.
22const ROOT_INDEX: NodeIndex = NodeIndex::root();
23
24/// An Word consisting of 4 ZERO elements.
25const EMPTY_DIGEST: Word = EMPTY_WORD;
26
27// PARTIAL MERKLE TREE
28// ================================================================================================
29
30/// A partial Merkle tree with NodeIndex keys and 4-element [Word] leaf values. Partial Merkle
31/// Tree allows to create Merkle Tree by providing Merkle paths of different lengths.
32///
33/// The root of the tree is recomputed on each new leaf update.
34#[derive(Debug, Clone, PartialEq, Eq)]
35pub struct PartialMerkleTree {
36    max_depth: u8,
37    nodes: BTreeMap<NodeIndex, Word>,
38    leaves: BTreeSet<NodeIndex>,
39}
40
41impl Default for PartialMerkleTree {
42    fn default() -> Self {
43        Self::new()
44    }
45}
46
47impl PartialMerkleTree {
48    // CONSTANTS
49    // --------------------------------------------------------------------------------------------
50
51    /// Minimum supported depth.
52    pub const MIN_DEPTH: u8 = 1;
53
54    /// Maximum supported depth.
55    pub const MAX_DEPTH: u8 = 64;
56
57    // CONSTRUCTORS
58    // --------------------------------------------------------------------------------------------
59
60    /// Returns a new empty [PartialMerkleTree].
61    pub fn new() -> Self {
62        PartialMerkleTree {
63            max_depth: 0,
64            nodes: BTreeMap::new(),
65            leaves: BTreeSet::new(),
66        }
67    }
68
69    /// Appends the provided paths iterator into the set.
70    ///
71    /// Analogous to [Self::add_path].
72    pub fn with_paths<I>(paths: I) -> Result<Self, MerkleError>
73    where
74        I: IntoIterator<Item = (u64, Word, MerklePath)>,
75    {
76        // create an empty tree
77        let tree = PartialMerkleTree::new();
78
79        paths.into_iter().try_fold(tree, |mut tree, (index, value, path)| {
80            tree.add_path(index, value, path)?;
81            Ok(tree)
82        })
83    }
84
85    /// Returns a new [PartialMerkleTree] instantiated with leaves map as specified by the provided
86    /// entries.
87    ///
88    /// # Errors
89    /// Returns an error if:
90    /// - Any entry has depth 0 or is greater than 64.
91    /// - The number of entries exceeds the maximum tree capacity, that is 2^{depth}.
92    /// - The provided entries contain an insufficient set of nodes.
93    /// - Any entry is an ancestor of another entry (creates hash ambiguity).
94    ///
95    /// An empty input returns an empty tree.
96    pub fn with_leaves<R, I>(entries: R) -> Result<Self, MerkleError>
97    where
98        R: IntoIterator<IntoIter = I>,
99        I: Iterator<Item = (NodeIndex, Word)> + ExactSizeIterator,
100    {
101        let entries = entries.into_iter();
102        if entries.len() == 0 {
103            return Ok(PartialMerkleTree::new());
104        }
105
106        let mut layers: BTreeMap<u8, Vec<u64>> = BTreeMap::new();
107        let mut leaves = BTreeSet::new();
108        let mut nodes = BTreeMap::new();
109
110        // add data to the leaves and nodes maps and also fill layers map, where the key is the
111        // depth of the node and value is its index.
112        for (node_index, hash) in entries {
113            Self::check_depth(node_index.depth())?;
114            leaves.insert(node_index);
115            nodes.insert(node_index, hash);
116            layers
117                .entry(node_index.depth())
118                .and_modify(|layer_vec| layer_vec.push(node_index.position()))
119                .or_insert(vec![node_index.position()]);
120        }
121
122        // Get maximum depth
123        let max_depth = *layers.keys().next_back().unwrap_or(&0);
124
125        // fill layers without nodes with empty vector
126        for depth in 0..max_depth {
127            layers.entry(depth).or_default();
128        }
129
130        let mut layer_iter = layers.into_values().rev();
131        let mut parent_layer = layer_iter.next().unwrap();
132        let mut current_layer;
133
134        for depth in (1..max_depth + 1).rev() {
135            // set current_layer = parent_layer and parent_layer = layer_iter.next()
136            current_layer = layer_iter.next().unwrap();
137            core::mem::swap(&mut current_layer, &mut parent_layer);
138
139            // Siblings have the same parent position. Sort them together and retain one child per
140            // parent so each parent is computed once.
141            current_layer.sort_unstable();
142            current_layer.dedup_by_key(|position| *position / 2);
143
144            for index_value in current_layer {
145                // get the parent node index
146                let parent_node = NodeIndex::new(depth - 1, index_value / 2)?;
147
148                // A user-provided parent cannot also have a descendant in the input set.
149                if leaves.contains(&parent_node) {
150                    return Err(MerkleError::EntryIsNotLeaf { node: parent_node });
151                }
152
153                // create current node index
154                let index = NodeIndex::new(depth, index_value)?;
155
156                // get hash of the current node
157                let node = nodes.get(&index).ok_or(MerkleError::NodeIndexNotFoundInTree(index))?;
158                // get hash of the sibling node
159                let sibling = nodes
160                    .get(&index.sibling())
161                    .ok_or(MerkleError::NodeIndexNotFoundInTree(index.sibling()))?;
162                // get parent hash
163                let parent = Poseidon2::merge(&index.build_node(*node, *sibling));
164
165                // add index value of the calculated node to the parents layer
166                parent_layer.push(parent_node.position());
167                // add index and hash to the nodes map
168                nodes.insert(parent_node, parent);
169            }
170        }
171
172        Ok(PartialMerkleTree { max_depth, nodes, leaves })
173    }
174
175    // PUBLIC ACCESSORS
176    // --------------------------------------------------------------------------------------------
177
178    /// Returns the root of this Merkle tree.
179    pub fn root(&self) -> Word {
180        self.nodes.get(&ROOT_INDEX).cloned().unwrap_or(EMPTY_DIGEST)
181    }
182
183    /// Returns the depth of this Merkle tree.
184    pub fn max_depth(&self) -> u8 {
185        self.max_depth
186    }
187
188    /// Returns a node at the specified NodeIndex.
189    ///
190    /// # Errors
191    /// Returns an error if the specified NodeIndex is not contained in the nodes map.
192    pub fn get_node(&self, index: NodeIndex) -> Result<Word, MerkleError> {
193        self.nodes
194            .get(&index)
195            .ok_or(MerkleError::NodeIndexNotFoundInTree(index))
196            .copied()
197    }
198
199    /// Returns true if provided index contains in the leaves set, false otherwise.
200    pub fn is_leaf(&self, index: NodeIndex) -> bool {
201        self.leaves.contains(&index)
202    }
203
204    /// Returns a vector of paths from every leaf to the root.
205    pub fn to_paths(&self) -> Vec<(NodeIndex, MerkleProof)> {
206        let mut paths = Vec::new();
207        self.leaves.iter().for_each(|&leaf| {
208            paths.push((
209                leaf,
210                MerkleProof {
211                    value: self.get_node(leaf).expect("Failed to get leaf node"),
212                    path: self.get_path(leaf).expect("Failed to get path"),
213                },
214            ));
215        });
216        paths
217    }
218
219    /// Returns a Merkle path from the node at the specified index to the root.
220    ///
221    /// The node itself is not included in the path.
222    ///
223    /// # Errors
224    /// Returns an error if:
225    /// - the specified index has depth set to 0 or the depth is greater than the depth of this
226    ///   Merkle tree.
227    /// - the specified index is not contained in the nodes map.
228    pub fn get_path(&self, mut index: NodeIndex) -> Result<MerklePath, MerkleError> {
229        if index.is_root() {
230            return Err(MerkleError::DepthTooSmall(index.depth()));
231        } else if index.depth() > self.max_depth() {
232            return Err(MerkleError::DepthTooBig(index.depth() as u64));
233        }
234
235        if !self.nodes.contains_key(&index) {
236            return Err(MerkleError::NodeIndexNotFoundInTree(index));
237        }
238
239        let mut path = Vec::new();
240        for _ in 0..index.depth() {
241            let sibling_index = index.sibling();
242            index.move_up();
243            let sibling =
244                self.nodes.get(&sibling_index).cloned().expect("Sibling node not in the map");
245            path.push(sibling);
246        }
247        Ok(MerklePath::new(path))
248    }
249
250    // ITERATORS
251    // --------------------------------------------------------------------------------------------
252
253    /// Returns an iterator over the leaves of this [PartialMerkleTree].
254    pub fn leaves(&self) -> impl Iterator<Item = (NodeIndex, Word)> + '_ {
255        self.leaves.iter().map(|&leaf| {
256            (
257                leaf,
258                self.get_node(leaf)
259                    .unwrap_or_else(|_| panic!("Leaf with {leaf} is not in the nodes map")),
260            )
261        })
262    }
263
264    /// Returns an iterator over the inner nodes of this Merkle tree.
265    pub fn inner_nodes(&self) -> impl Iterator<Item = InnerNodeInfo> + '_ {
266        let inner_nodes = self.nodes.iter().filter(|(index, _)| !self.leaves.contains(index));
267        inner_nodes.map(|(index, digest)| {
268            let left_hash =
269                self.nodes.get(&index.left_child()).expect("Failed to get left child hash");
270            let right_hash =
271                self.nodes.get(&index.right_child()).expect("Failed to get right child hash");
272            InnerNodeInfo {
273                value: *digest,
274                left: *left_hash,
275                right: *right_hash,
276            }
277        })
278    }
279
280    // STATE MUTATORS
281    // --------------------------------------------------------------------------------------------
282
283    /// Adds the nodes of the specified Merkle path to this [PartialMerkleTree]. The `index_value`
284    /// and `value` parameters specify the leaf node at which the path starts.
285    ///
286    /// # Errors
287    /// Returns an error if:
288    /// - The depth of the specified node_index is greater than 64 or smaller than 1.
289    /// - The specified path is not consistent with other paths in the set (i.e., resolves to a
290    ///   different root).
291    pub fn add_path(
292        &mut self,
293        index_value: u64,
294        value: Word,
295        path: MerklePath,
296    ) -> Result<(), MerkleError> {
297        let index_value = NodeIndex::new(path.len() as u8, index_value)?;
298
299        Self::check_depth(index_value.depth())?;
300        self.update_depth(index_value.depth());
301
302        // add provided node and its sibling to the leaves set
303        self.leaves.insert(index_value);
304        let sibling_node_index = index_value.sibling();
305        self.leaves.insert(sibling_node_index);
306
307        // add provided node and its sibling to the nodes map
308        self.nodes.insert(index_value, value);
309        self.nodes.insert(sibling_node_index, path[0]);
310
311        // traverse to the root, updating the nodes
312        let mut index_value = index_value;
313        let node = Poseidon2::merge(&index_value.build_node(value, path[0]));
314        let root = path.iter().skip(1).copied().fold(node, |node, hash| {
315            index_value.move_up();
316            // insert calculated node to the nodes map
317            self.nodes.insert(index_value, node);
318
319            // if the calculated node was a leaf, remove it from leaves set.
320            self.leaves.remove(&index_value);
321
322            let sibling_node = index_value.sibling();
323
324            // Insert node from Merkle path to the nodes map. This sibling node becomes a leaf only
325            // if it is a new node (it wasn't in nodes map).
326            // Node can be in 3 states: internal node, leaf of the tree and not a tree node at all.
327            // - Internal node can only stay in this state -- addition of a new path can't make it
328            // a leaf or remove it from the tree.
329            // - Leaf node can stay in the same state (remain a leaf) or can become an internal
330            // node. In the first case we don't need to do anything, and the second case is handled
331            // by the call of `self.leaves.remove(&index_value);`
332            // - New node can be a calculated node or a "sibling" node from a Merkle Path:
333            // --- Calculated node, obviously, never can be a leaf.
334            // --- Sibling node can be only a leaf, because otherwise it is not a new node.
335            if self.nodes.insert(sibling_node, hash).is_none() {
336                self.leaves.insert(sibling_node);
337            }
338
339            Poseidon2::merge(&index_value.build_node(node, hash))
340        });
341
342        // if the path set is empty (the root is all ZEROs), set the root to the root of the added
343        // path; otherwise, the root of the added path must be identical to the current root
344        if self.root() == EMPTY_DIGEST {
345            self.nodes.insert(ROOT_INDEX, root);
346        } else if self.root() != root {
347            return Err(MerkleError::ConflictingRoots {
348                expected_root: self.root(),
349                actual_root: root,
350            });
351        }
352
353        Ok(())
354    }
355
356    /// Updates value of the leaf at the specified index returning the old leaf value.
357    ///
358    /// By default the specified index is assumed to belong to the deepest layer. If the considered
359    /// node does not belong to the tree, the first node on the way to the root will be changed.
360    ///
361    /// This also recomputes all hashes between the leaf and the root, updating the root itself.
362    ///
363    /// # Errors
364    /// Returns an error if:
365    /// - No entry exists at the specified index.
366    /// - The specified index is greater than the maximum number of nodes on the deepest layer.
367    pub fn update_leaf(&mut self, index: u64, value: Word) -> Result<Word, MerkleError> {
368        let mut node_index = NodeIndex::new(self.max_depth(), index)?;
369
370        // proceed to the leaf
371        for _ in 0..node_index.depth() {
372            if !self.leaves.contains(&node_index) {
373                node_index.move_up();
374            }
375        }
376
377        // add node value to the nodes Map
378        let old_value = self
379            .nodes
380            .insert(node_index, value)
381            .ok_or(MerkleError::NodeIndexNotFoundInTree(node_index))?;
382
383        // if the old value and new value are the same, there is nothing to update
384        if value == old_value {
385            return Ok(old_value);
386        }
387
388        let mut value = value;
389        for _ in 0..node_index.depth() {
390            let sibling = self.nodes.get(&node_index.sibling()).expect("sibling should exist");
391            value = Poseidon2::merge(&node_index.build_node(value, *sibling));
392            node_index.move_up();
393            self.nodes.insert(node_index, value);
394        }
395
396        Ok(old_value)
397    }
398
399    // UTILITY FUNCTIONS
400    // --------------------------------------------------------------------------------------------
401
402    /// Utility to visualize a [PartialMerkleTree] in text.
403    pub fn print(&self) -> Result<String, fmt::Error> {
404        let indent = "  ";
405        let mut s = String::new();
406        s.push_str("root: ");
407        s.push_str(&word_to_hex(&self.root())?);
408        s.push('\n');
409        for d in 1..=self.max_depth() {
410            let entries = 2u64.pow(d.into());
411            for i in 0..entries {
412                let index = NodeIndex::new(d, i).expect("The index must always be valid");
413                let node = self.get_node(index);
414                let node = match node {
415                    Err(_) => continue,
416                    Ok(node) => node,
417                };
418
419                for _ in 0..d {
420                    s.push_str(indent);
421                }
422                s.push_str(&format!("({}, {}): ", index.depth(), index.position()));
423                s.push_str(&word_to_hex(&node)?);
424                s.push('\n');
425            }
426        }
427
428        Ok(s)
429    }
430
431    // HELPER METHODS
432    // --------------------------------------------------------------------------------------------
433
434    /// Updates depth value with the maximum of current and provided depth.
435    fn update_depth(&mut self, new_depth: u8) {
436        self.max_depth = new_depth.max(self.max_depth);
437    }
438
439    /// Returns an error if the depth is 0 or is greater than 64.
440    fn check_depth(depth: u8) -> Result<(), MerkleError> {
441        // validate the range of the depth.
442        if depth < Self::MIN_DEPTH {
443            return Err(MerkleError::DepthTooSmall(depth));
444        } else if Self::MAX_DEPTH < depth {
445            return Err(MerkleError::DepthTooBig(depth as u64));
446        }
447        Ok(())
448    }
449}
450
451// SERIALIZATION
452// ================================================================================================
453
454impl Serializable for PartialMerkleTree {
455    fn write_into<W: ByteWriter>(&self, target: &mut W) {
456        // write leaf nodes
457        target.write_u64(self.leaves.len() as u64);
458        for leaf_index in self.leaves.iter() {
459            leaf_index.write_into(target);
460            self.get_node(*leaf_index).expect("Leaf hash not found").write_into(target);
461        }
462    }
463}
464
465impl Deserializable for PartialMerkleTree {
466    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
467        let leaves_len_u64 = source.read_u64()?;
468        let leaves_len = usize::try_from(leaves_len_u64).map_err(|_| {
469            DeserializationError::InvalidValue("PartialMerkleTree leaf count too large".into())
470        })?;
471
472        // Use read_many_iter to avoid eager allocation and respect BudgetedReader limits
473        let leaf_nodes: Vec<(NodeIndex, Word)> =
474            source.read_many_iter(leaves_len)?.collect::<Result<_, _>>()?;
475
476        let pmt = PartialMerkleTree::with_leaves(leaf_nodes).map_err(|_| {
477            DeserializationError::InvalidValue("Invalid data for PartialMerkleTree creation".into())
478        })?;
479
480        Ok(pmt)
481    }
482
483    /// Minimum serialized size: u64 length prefix (0 entries).
484    fn min_serialized_size() -> usize {
485        8
486    }
487}