Skip to main content

miden_node_proto/domain/
merkle.rs

1use std::collections::BTreeSet;
2
3use miden_protocol::Word;
4use miden_protocol::crypto::merkle::mmr::{Forest, MmrDelta};
5use miden_protocol::crypto::merkle::smt::{
6    LeafIndex,
7    NodeValue,
8    PartialSmt,
9    SMT_DEPTH,
10    SmtLeaf,
11    SmtProof,
12    UniqueNodes,
13};
14use miden_protocol::crypto::merkle::{MerklePath, NodeIndex, SparseMerklePath};
15
16use crate::decode::{ConversionResultExt, GrpcDecodeExt};
17use crate::domain::{convert, try_convert};
18use crate::errors::ConversionError;
19use crate::{decode, generated as proto};
20
21// MERKLE PATH
22// ================================================================================================
23
24impl From<&MerklePath> for proto::primitives::MerklePath {
25    fn from(value: &MerklePath) -> Self {
26        let siblings = value.nodes().iter().map(proto::primitives::Digest::from).collect();
27        proto::primitives::MerklePath { siblings }
28    }
29}
30
31impl From<MerklePath> for proto::primitives::MerklePath {
32    fn from(value: MerklePath) -> Self {
33        (&value).into()
34    }
35}
36
37impl TryFrom<&proto::primitives::MerklePath> for MerklePath {
38    type Error = ConversionError;
39
40    fn try_from(merkle_path: &proto::primitives::MerklePath) -> Result<Self, Self::Error> {
41        merkle_path.siblings.iter().map(Word::try_from).collect()
42    }
43}
44
45impl TryFrom<proto::primitives::MerklePath> for MerklePath {
46    type Error = ConversionError;
47
48    fn try_from(merkle_path: proto::primitives::MerklePath) -> Result<Self, Self::Error> {
49        (&merkle_path).try_into()
50    }
51}
52
53// SPARSE MERKLE PATH
54// ================================================================================================
55
56impl From<SparseMerklePath> for proto::primitives::SparseMerklePath {
57    fn from(value: SparseMerklePath) -> Self {
58        let (empty_nodes_mask, siblings) = value.into_parts();
59        proto::primitives::SparseMerklePath {
60            empty_nodes_mask,
61            siblings: siblings.into_iter().map(proto::primitives::Digest::from).collect(),
62        }
63    }
64}
65
66impl TryFrom<proto::primitives::SparseMerklePath> for SparseMerklePath {
67    type Error = ConversionError;
68
69    fn try_from(merkle_path: proto::primitives::SparseMerklePath) -> Result<Self, Self::Error> {
70        Ok(SparseMerklePath::from_parts(
71            merkle_path.empty_nodes_mask,
72            merkle_path
73                .siblings
74                .into_iter()
75                .map(Word::try_from)
76                .collect::<Result<Vec<_>, _>>()
77                .context("siblings")?,
78        )?)
79    }
80}
81
82// MMR DELTA
83// ================================================================================================
84
85impl From<MmrDelta> for proto::primitives::MmrDelta {
86    fn from(value: MmrDelta) -> Self {
87        let data = value.data.into_iter().map(proto::primitives::Digest::from).collect();
88        proto::primitives::MmrDelta {
89            forest: value.forest.num_leaves() as u64,
90            data,
91        }
92    }
93}
94
95impl TryFrom<proto::primitives::MmrDelta> for MmrDelta {
96    type Error = ConversionError;
97
98    fn try_from(value: proto::primitives::MmrDelta) -> Result<Self, Self::Error> {
99        let data: Vec<_> = value
100            .data
101            .into_iter()
102            .map(Word::try_from)
103            .collect::<Result<_, _>>()
104            .context("data")?;
105
106        let forest_size: usize =
107            value.forest.try_into().context("forest size does not fit in usize")?;
108        let forest = Forest::new(forest_size).context("forest size out of range")?;
109
110        Ok(MmrDelta { forest, data })
111    }
112}
113
114// SPARSE MERKLE TREE
115// ================================================================================================
116
117// SMT LEAF
118// ------------------------------------------------------------------------------------------------
119
120impl TryFrom<proto::primitives::SmtLeaf> for SmtLeaf {
121    type Error = ConversionError;
122
123    fn try_from(value: proto::primitives::SmtLeaf) -> Result<Self, Self::Error> {
124        let decoder = value.decoder();
125        let leaf = decode!(decoder, value.leaf)?;
126
127        match leaf {
128            proto::primitives::smt_leaf::Leaf::EmptyLeafIndex(leaf_index) => {
129                Ok(Self::new_empty(LeafIndex::new_max_depth(leaf_index)))
130            },
131            proto::primitives::smt_leaf::Leaf::Single(entry) => {
132                let (key, value): (Word, Word) = entry.try_into().context("entry")?;
133
134                Ok(SmtLeaf::new_single(key, value))
135            },
136            proto::primitives::smt_leaf::Leaf::Multiple(entries) => {
137                let domain_entries: Vec<(Word, Word)> =
138                    try_convert(entries.entries).collect::<Result<_, _>>().context("entries")?;
139
140                Ok(SmtLeaf::new_multiple(domain_entries)?)
141            },
142        }
143    }
144}
145
146impl From<SmtLeaf> for proto::primitives::SmtLeaf {
147    fn from(smt_leaf: SmtLeaf) -> Self {
148        use proto::primitives::smt_leaf::Leaf;
149
150        let leaf = match smt_leaf {
151            SmtLeaf::Empty(leaf_index) => Leaf::EmptyLeafIndex(leaf_index.position()),
152            SmtLeaf::Single(entry) => Leaf::Single(entry.into()),
153            SmtLeaf::Multiple(entries) => Leaf::Multiple(proto::primitives::SmtLeafEntryList {
154                entries: convert(entries).collect(),
155            }),
156        };
157
158        Self { leaf: Some(leaf) }
159    }
160}
161
162// SMT LEAF ENTRY
163// ------------------------------------------------------------------------------------------------
164
165impl TryFrom<proto::primitives::SmtLeafEntry> for (Word, Word) {
166    type Error = ConversionError;
167
168    fn try_from(entry: proto::primitives::SmtLeafEntry) -> Result<Self, Self::Error> {
169        let decoder = entry.decoder();
170        let key: Word = decode!(decoder, entry.key)?;
171        let value: Word = decode!(decoder, entry.value)?;
172
173        Ok((key, value))
174    }
175}
176
177impl From<(Word, Word)> for proto::primitives::SmtLeafEntry {
178    fn from((key, value): (Word, Word)) -> Self {
179        Self {
180            key: Some(key.into()),
181            value: Some(value.into()),
182        }
183    }
184}
185
186// SMT PROOF
187// ------------------------------------------------------------------------------------------------
188
189impl TryFrom<proto::primitives::SmtOpening> for SmtProof {
190    type Error = ConversionError;
191
192    fn try_from(opening: proto::primitives::SmtOpening) -> Result<Self, Self::Error> {
193        let decoder = opening.decoder();
194        let path: SparseMerklePath = decode!(decoder, opening.path)?;
195        let leaf: SmtLeaf = decode!(decoder, opening.leaf)?;
196
197        Ok(SmtProof::new(path, leaf)?)
198    }
199}
200
201impl From<SmtProof> for proto::primitives::SmtOpening {
202    fn from(proof: SmtProof) -> Self {
203        let (path, leaf) = proof.into_parts();
204        Self {
205            path: Some(path.into()),
206            leaf: Some(leaf.into()),
207        }
208    }
209}
210
211// PARTIAL SMT
212// ------------------------------------------------------------------------------------------------
213
214impl From<UniqueNodes> for proto::primitives::PartialSmt {
215    fn from(unique_nodes: UniqueNodes) -> Self {
216        use proto::primitives::partial_smt_node::Value;
217
218        let UniqueNodes { root, nodes, leaves, value_only_leaves } = unique_nodes;
219
220        let mut node_levels = nodes.into_iter().collect::<Vec<_>>();
221        node_levels.sort_by_key(|(depth, _)| *depth);
222        let node_levels = node_levels
223            .into_iter()
224            .map(|(depth, nodes)| {
225                let nodes = nodes
226                    .into_iter()
227                    .map(|(index, value)| {
228                        let value = match value {
229                            NodeValue::EmptySubtreeRoot => Value::EmptySubtreeRoot(true),
230                            NodeValue::Present(value) => Value::Digest(value.into()),
231                        };
232                        proto::primitives::PartialSmtNode { index, value: Some(value) }
233                    })
234                    .collect();
235
236                proto::primitives::PartialSmtNodeLevel { depth: u32::from(depth), nodes }
237            })
238            .collect();
239
240        let leaves = leaves
241            .into_iter()
242            .map(|(index, leaf)| proto::primitives::IndexedSmtLeaf {
243                index,
244                leaf: Some(leaf.into()),
245            })
246            .collect();
247
248        let value_only_leaves = value_only_leaves
249            .into_iter()
250            .map(|(index, value)| proto::primitives::IndexedDigest {
251                index,
252                value: Some(value.into()),
253            })
254            .collect();
255
256        Self {
257            root: Some(root.into()),
258            node_levels,
259            leaves,
260            value_only_leaves,
261        }
262    }
263}
264
265impl TryFrom<proto::primitives::PartialSmt> for UniqueNodes {
266    type Error = ConversionError;
267
268    fn try_from(value: proto::primitives::PartialSmt) -> Result<Self, Self::Error> {
269        use proto::primitives::partial_smt_node::Value;
270
271        let decoder = value.decoder();
272        let proto::primitives::PartialSmt {
273            root,
274            node_levels,
275            leaves,
276            value_only_leaves,
277        } = value;
278
279        let root = decode!(decoder, root)?;
280
281        let mut seen_depths = BTreeSet::new();
282        let mut decoded_levels = Vec::with_capacity(node_levels.len());
283        for level in node_levels {
284            let depth = u8::try_from(level.depth).context("node_levels.depth")?;
285            if depth == 0 || depth >= SMT_DEPTH {
286                return Err(ConversionError::message(format!(
287                    "partial SMT node depth {depth} must be in the range 1..{SMT_DEPTH}"
288                )));
289            }
290            if !seen_depths.insert(depth) {
291                return Err(ConversionError::message(format!(
292                    "partial SMT contains duplicate node depth {depth}"
293                )));
294            }
295
296            let mut seen_indices = BTreeSet::new();
297            let mut decoded_nodes = Vec::with_capacity(level.nodes.len());
298            for node in level.nodes {
299                NodeIndex::new(depth, node.index).context("node_levels.nodes.index")?;
300                if !seen_indices.insert(node.index) {
301                    return Err(ConversionError::message(format!(
302                        "partial SMT contains duplicate node index {} at depth {depth}",
303                        node.index
304                    )));
305                }
306
307                let node_value = match node.value.ok_or_else(|| {
308                    ConversionError::missing_field::<proto::primitives::PartialSmtNode>("value")
309                })? {
310                    Value::Digest(value) => NodeValue::Present(value.try_into().context("digest")?),
311                    Value::EmptySubtreeRoot(true) => NodeValue::EmptySubtreeRoot,
312                    Value::EmptySubtreeRoot(false) => {
313                        return Err(ConversionError::message(
314                            "partial SMT empty_subtree_root marker must be true",
315                        ));
316                    },
317                };
318                decoded_nodes.push((node.index, node_value));
319            }
320            decoded_levels.push((depth, decoded_nodes));
321        }
322
323        let mut seen_leaf_indices = BTreeSet::new();
324        let mut decoded_leaves = Vec::with_capacity(leaves.len());
325        for indexed_leaf in leaves {
326            if !seen_leaf_indices.insert(indexed_leaf.index) {
327                return Err(ConversionError::message(format!(
328                    "partial SMT contains duplicate leaf index {}",
329                    indexed_leaf.index
330                )));
331            }
332            let decoder = indexed_leaf.decoder();
333            let leaf = decode!(decoder, indexed_leaf.leaf)?;
334            decoded_leaves.push((indexed_leaf.index, leaf));
335        }
336
337        let mut seen_value_only_indices = BTreeSet::new();
338        let mut decoded_value_only_leaves = Vec::with_capacity(value_only_leaves.len());
339        for indexed_digest in value_only_leaves {
340            if !seen_value_only_indices.insert(indexed_digest.index) {
341                return Err(ConversionError::message(format!(
342                    "partial SMT contains duplicate value-only leaf index {}",
343                    indexed_digest.index
344                )));
345            }
346            if seen_leaf_indices.contains(&indexed_digest.index) {
347                return Err(ConversionError::message(format!(
348                    "partial SMT leaf index {} has both a leaf and a value-only leaf",
349                    indexed_digest.index
350                )));
351            }
352            let decoder = indexed_digest.decoder();
353            let digest = decode!(decoder, indexed_digest.value)?;
354            decoded_value_only_leaves.push((indexed_digest.index, digest));
355        }
356
357        Ok(UniqueNodes {
358            root,
359            nodes: decoded_levels.into_iter().collect(),
360            leaves: decoded_leaves,
361            value_only_leaves: decoded_value_only_leaves,
362        })
363    }
364}
365
366impl From<PartialSmt> for proto::primitives::PartialSmt {
367    fn from(partial_smt: PartialSmt) -> Self {
368        partial_smt.to_unique_nodes().into()
369    }
370}
371
372impl TryFrom<proto::primitives::PartialSmt> for PartialSmt {
373    type Error = ConversionError;
374
375    fn try_from(value: proto::primitives::PartialSmt) -> Result<Self, Self::Error> {
376        let unique_nodes = UniqueNodes::try_from(value)?;
377        PartialSmt::from_unique_nodes(unique_nodes)
378            .map_err(|err| ConversionError::deserialization("PartialSmt", err))
379    }
380}
381
382#[cfg(test)]
383mod tests {
384    use miden_protocol::crypto::merkle::smt::Smt;
385
386    use super::*;
387
388    #[test]
389    fn partial_smt_round_trip() {
390        let key0 = Word::from([1, 2, 3, 4u32]);
391        let key1 = Word::from([5, 6, 7, 8u32]);
392        let missing_key = Word::from([9, 10, 11, 12u32]);
393        let value0 = Word::from([13, 14, 15, 16u32]);
394        let value1 = Word::from([17, 18, 19, 20u32]);
395        let smt = Smt::with_entries([(key0, value0), (key1, value1)]).unwrap();
396        let partial_smt =
397            PartialSmt::from_proofs([smt.open(&key0), smt.open(&missing_key)]).unwrap();
398
399        let encoded: proto::primitives::PartialSmt = partial_smt.clone().into();
400        assert!(encoded.node_levels.is_sorted_by_key(|level| level.depth));
401
402        let decoded_unique_nodes = UniqueNodes::try_from(encoded).unwrap();
403        let decoded = PartialSmt::from_unique_nodes(decoded_unique_nodes).unwrap();
404
405        assert_eq!(decoded, partial_smt);
406        assert_eq!(decoded.get_value(&key0).unwrap(), value0);
407        assert_eq!(decoded.get_value(&missing_key).unwrap(), Word::empty());
408    }
409
410    #[test]
411    fn partial_smt_rejects_false_empty_subtree_marker() {
412        use proto::primitives::partial_smt_node::Value;
413
414        let encoded = proto::primitives::PartialSmt {
415            root: Some(PartialSmt::EMPTY_ROOT.into()),
416            node_levels: vec![proto::primitives::PartialSmtNodeLevel {
417                depth: 1,
418                nodes: vec![proto::primitives::PartialSmtNode {
419                    index: 0,
420                    value: Some(Value::EmptySubtreeRoot(false)),
421                }],
422            }],
423            leaves: vec![],
424            value_only_leaves: vec![],
425        };
426
427        let err = UniqueNodes::try_from(encoded).unwrap_err();
428        assert!(err.to_string().contains("must be true"));
429    }
430
431    #[test]
432    fn partial_smt_rejects_duplicate_depths() {
433        let encoded = proto::primitives::PartialSmt {
434            root: Some(PartialSmt::EMPTY_ROOT.into()),
435            node_levels: vec![
436                proto::primitives::PartialSmtNodeLevel { depth: 1, nodes: vec![] },
437                proto::primitives::PartialSmtNodeLevel { depth: 1, nodes: vec![] },
438            ],
439            leaves: vec![],
440            value_only_leaves: vec![],
441        };
442
443        let err = UniqueNodes::try_from(encoded).unwrap_err();
444        assert!(err.to_string().contains("duplicate node depth"));
445    }
446
447    #[test]
448    fn partial_smt_rejects_invalid_node_index() {
449        use proto::primitives::partial_smt_node::Value;
450
451        let encoded = proto::primitives::PartialSmt {
452            root: Some(PartialSmt::EMPTY_ROOT.into()),
453            node_levels: vec![proto::primitives::PartialSmtNodeLevel {
454                depth: 1,
455                nodes: vec![proto::primitives::PartialSmtNode {
456                    index: 2,
457                    value: Some(Value::EmptySubtreeRoot(true)),
458                }],
459            }],
460            leaves: vec![],
461            value_only_leaves: vec![],
462        };
463
464        let err = UniqueNodes::try_from(encoded).unwrap_err();
465        assert!(err.to_string().contains("not valid for depth"));
466    }
467}