1use alloc::collections::{BTreeMap, BTreeSet};
2use alloc::format;
3use alloc::vec::Vec;
4
5use miden_protocol::Word;
6use miden_protocol::crypto::merkle::mmr::{Forest, MmrDelta};
7use miden_protocol::crypto::merkle::smt::{
8 LeafIndex,
9 PartialSmt,
10 SMT_DEPTH,
11 SmtLeaf,
12 SmtProof,
13 UniqueNodes,
14};
15use miden_protocol::crypto::merkle::{MerklePath, NodeIndex, SparseMerklePath};
16
17use super::{MessageDecodeExt, required};
18use crate::{ConversionError, ConversionResultExt, proto};
19
20impl From<&MerklePath> for proto::primitives::MerklePath {
24 fn from(value: &MerklePath) -> Self {
25 let siblings = value.nodes().iter().map(Into::into).collect();
26 proto::primitives::MerklePath { siblings }
27 }
28}
29
30impl From<MerklePath> for proto::primitives::MerklePath {
31 fn from(value: MerklePath) -> Self {
32 (&value).into()
33 }
34}
35
36impl TryFrom<&proto::primitives::MerklePath> for MerklePath {
37 type Error = ConversionError;
38
39 fn try_from(merkle_path: &proto::primitives::MerklePath) -> Result<Self, Self::Error> {
40 merkle_path.siblings.iter().map(Word::try_from).collect()
41 }
42}
43
44impl TryFrom<proto::primitives::MerklePath> for MerklePath {
45 type Error = ConversionError;
46
47 fn try_from(merkle_path: proto::primitives::MerklePath) -> Result<Self, Self::Error> {
48 (&merkle_path).try_into()
49 }
50}
51
52impl From<SparseMerklePath> for proto::primitives::SparseMerklePath {
56 fn from(value: SparseMerklePath) -> Self {
57 let (empty_nodes_mask, siblings) = value.into_parts();
58 proto::primitives::SparseMerklePath {
59 empty_nodes_mask,
60 siblings: siblings.into_iter().map(Into::into).collect(),
61 }
62 }
63}
64
65impl TryFrom<proto::primitives::SparseMerklePath> for SparseMerklePath {
66 type Error = ConversionError;
67
68 fn try_from(merkle_path: proto::primitives::SparseMerklePath) -> Result<Self, Self::Error> {
69 Ok(SparseMerklePath::from_parts(
70 merkle_path.empty_nodes_mask,
71 merkle_path
72 .siblings
73 .into_iter()
74 .map(Word::try_from)
75 .collect::<Result<Vec<_>, _>>()
76 .context("siblings")?,
77 )?)
78 }
79}
80
81impl From<MmrDelta> for proto::primitives::MmrDelta {
85 fn from(value: MmrDelta) -> Self {
86 let update_data = value.data.into_iter().map(Into::into).collect();
87 proto::primitives::MmrDelta {
88 forest: value.forest.num_leaves() as u64,
89 update_data,
90 }
91 }
92}
93
94impl TryFrom<proto::primitives::MmrDelta> for MmrDelta {
95 type Error = ConversionError;
96
97 fn try_from(value: proto::primitives::MmrDelta) -> Result<Self, Self::Error> {
98 let data: Vec<_> = value
99 .update_data
100 .into_iter()
101 .map(Word::try_from)
102 .collect::<Result<_, _>>()
103 .context("update_data")?;
104
105 let forest_size = value.forest.try_into().context("forest size does not fit in usize")?;
106 let forest = Forest::new(forest_size).context("forest size out of range")?;
107
108 Ok(MmrDelta { forest, data })
109 }
110}
111
112impl TryFrom<proto::primitives::SmtLeaf> for SmtLeaf {
119 type Error = ConversionError;
120
121 fn try_from(value: proto::primitives::SmtLeaf) -> Result<Self, Self::Error> {
122 let decoder = value.decoder();
123 let leaf = required!(decoder, value.leaf)?;
124
125 match leaf {
126 proto::primitives::smt_leaf::Leaf::EmptyLeafIndex(leaf_index) => {
127 Ok(Self::new_empty(LeafIndex::new_max_depth(leaf_index)))
128 },
129 proto::primitives::smt_leaf::Leaf::Single(entry) => {
130 let (key, value) = entry.try_into().context("entry")?;
131
132 Ok(SmtLeaf::new_single(key, value))
133 },
134 proto::primitives::smt_leaf::Leaf::Multiple(entries) => {
135 let domain_entries = entries
136 .entries
137 .into_iter()
138 .map(TryInto::try_into)
139 .collect::<Result<_, _>>()
140 .context("entries")?;
141
142 Ok(SmtLeaf::new_multiple(domain_entries)?)
143 },
144 }
145 }
146}
147
148impl From<SmtLeaf> for proto::primitives::SmtLeaf {
149 fn from(smt_leaf: SmtLeaf) -> Self {
150 use proto::primitives::smt_leaf::Leaf;
151
152 let leaf = match smt_leaf {
153 SmtLeaf::Empty(leaf_index) => Leaf::EmptyLeafIndex(leaf_index.position()),
154 SmtLeaf::Single(entry) => Leaf::Single(entry.into()),
155 SmtLeaf::Multiple(entries) => Leaf::Multiple(proto::primitives::SmtLeafEntryList {
156 entries: entries.into_iter().map(Into::into).collect(),
157 }),
158 };
159
160 Self { leaf: Some(leaf) }
161 }
162}
163
164impl TryFrom<proto::primitives::SmtLeafEntry> for (Word, Word) {
168 type Error = ConversionError;
169
170 fn try_from(entry: proto::primitives::SmtLeafEntry) -> Result<Self, Self::Error> {
171 let decoder = entry.decoder();
172 let key = required!(decoder, entry.key)?;
173 let value = required!(decoder, entry.value)?;
174
175 Ok((key, value))
176 }
177}
178
179impl From<(Word, Word)> for proto::primitives::SmtLeafEntry {
180 fn from((key, value): (Word, Word)) -> Self {
181 Self {
182 key: Some(key.into()),
183 value: Some(value.into()),
184 }
185 }
186}
187
188impl TryFrom<proto::primitives::SmtOpening> for SmtProof {
192 type Error = ConversionError;
193
194 fn try_from(opening: proto::primitives::SmtOpening) -> Result<Self, Self::Error> {
195 let decoder = opening.decoder();
196 let path = required!(decoder, opening.path)?;
197 let leaf = required!(decoder, opening.leaf)?;
198
199 Ok(SmtProof::new(path, leaf)?)
200 }
201}
202
203impl From<SmtProof> for proto::primitives::SmtOpening {
204 fn from(proof: SmtProof) -> Self {
205 let (path, leaf) = proof.into_parts();
206 Self {
207 path: Some(path.into()),
208 leaf: Some(leaf.into()),
209 }
210 }
211}
212
213impl From<UniqueNodes> for proto::primitives::PartialSmt {
217 fn from(unique_nodes: UniqueNodes) -> Self {
218 let UniqueNodes { root, nodes, leaves, value_only_leaves } = unique_nodes;
219
220 let mut node_levels = Vec::new();
221 let mut nodes = nodes.into_iter().peekable();
222 while let Some((index, _)) = nodes.peek() {
223 let depth = index.depth();
224 let mut level_nodes = Vec::new();
225 while let Some((index, digest)) = nodes.next_if(|(index, _)| index.depth() == depth) {
226 level_nodes.push(proto::primitives::PartialSmtNode {
227 index: index.position(),
228 digest: Some(digest.into()),
229 });
230 }
231 node_levels.push(proto::primitives::PartialSmtNodeLevel {
232 depth: u32::from(depth),
233 nodes: level_nodes,
234 });
235 }
236 let leaves = leaves
237 .into_iter()
238 .map(|(index, leaf)| proto::primitives::IndexedSmtLeaf {
239 index,
240 leaf: Some(leaf.into()),
241 })
242 .collect();
243
244 let value_only_leaves = value_only_leaves
245 .into_iter()
246 .map(|(index, value)| proto::primitives::IndexedDigest {
247 index,
248 value: Some(value.into()),
249 })
250 .collect();
251
252 Self {
253 root: Some(root.into()),
254 node_levels,
255 leaves,
256 value_only_leaves,
257 }
258 }
259}
260
261impl TryFrom<proto::primitives::PartialSmt> for UniqueNodes {
262 type Error = ConversionError;
263
264 fn try_from(value: proto::primitives::PartialSmt) -> Result<Self, Self::Error> {
265 let decoder = value.decoder();
266 let proto::primitives::PartialSmt {
267 root,
268 node_levels,
269 leaves,
270 value_only_leaves,
271 } = value;
272
273 let root = required!(decoder, root)?;
274
275 let mut seen_depths = BTreeSet::new();
276 let mut decoded_nodes = BTreeMap::new();
277 for level in node_levels {
278 let depth = u8::try_from(level.depth).context("node_levels.depth")?;
279 if depth == 0 || depth >= SMT_DEPTH {
280 return Err(ConversionError::message(format!(
281 "partial SMT node depth {depth} must be in the range 1..{SMT_DEPTH}"
282 )));
283 }
284 if !seen_depths.insert(depth) {
285 return Err(ConversionError::message(format!(
286 "partial SMT contains duplicate node depth {depth}"
287 )));
288 }
289
290 for node in level.nodes {
291 let index = NodeIndex::new(depth, node.index).context("node_levels.nodes.index")?;
292 if decoded_nodes.contains_key(&index) {
293 return Err(ConversionError::message(format!(
294 "partial SMT contains duplicate node index {} at depth {depth}",
295 node.index
296 )));
297 }
298 let digest = node.digest.ok_or_else(|| {
299 ConversionError::missing_field::<proto::primitives::PartialSmtNode>("digest")
300 })?;
301 decoded_nodes.insert(index, digest.try_into().context("digest")?);
302 }
303 }
304
305 let mut seen_leaf_indices = BTreeSet::new();
306 let mut decoded_leaves = BTreeMap::new();
307 for indexed_leaf in leaves {
308 if !seen_leaf_indices.insert(indexed_leaf.index) {
309 return Err(ConversionError::message(format!(
310 "partial SMT contains duplicate leaf index {}",
311 indexed_leaf.index
312 )));
313 }
314 let decoder = indexed_leaf.decoder();
315 let leaf = required!(decoder, indexed_leaf.leaf)?;
316 decoded_leaves.insert(indexed_leaf.index, leaf);
317 }
318
319 let mut seen_value_only_indices = BTreeSet::new();
320 let mut decoded_value_only_leaves = BTreeMap::new();
321 for indexed_digest in value_only_leaves {
322 if !seen_value_only_indices.insert(indexed_digest.index) {
323 return Err(ConversionError::message(format!(
324 "partial SMT contains duplicate value-only leaf index {}",
325 indexed_digest.index
326 )));
327 }
328 if seen_leaf_indices.contains(&indexed_digest.index) {
329 return Err(ConversionError::message(format!(
330 "partial SMT leaf index {} has both a leaf and a value-only leaf",
331 indexed_digest.index
332 )));
333 }
334 let decoder = indexed_digest.decoder();
335 let digest = required!(decoder, indexed_digest.value)?;
336 decoded_value_only_leaves.insert(indexed_digest.index, digest);
337 }
338
339 Ok(UniqueNodes {
340 root,
341 nodes: decoded_nodes,
342 leaves: decoded_leaves,
343 value_only_leaves: decoded_value_only_leaves,
344 })
345 }
346}
347
348impl From<PartialSmt> for proto::primitives::PartialSmt {
349 fn from(partial_smt: PartialSmt) -> Self {
350 partial_smt.to_unique_nodes().into()
351 }
352}
353
354impl TryFrom<proto::primitives::PartialSmt> for PartialSmt {
355 type Error = ConversionError;
356
357 fn try_from(value: proto::primitives::PartialSmt) -> Result<Self, Self::Error> {
358 let unique_nodes = UniqueNodes::try_from(value)?;
359 PartialSmt::from_unique_nodes(unique_nodes)
360 .map_err(|err| ConversionError::deserialization("PartialSmt", err))
361 }
362}
363
364#[cfg(test)]
365mod tests {
366 use alloc::collections::BTreeMap;
367 use alloc::string::ToString;
368 use alloc::vec;
369
370 use miden_protocol::crypto::merkle::smt::{PartialSmt, Smt, UniqueNodes};
371 use prost::Message;
372
373 use super::*;
374
375 #[test]
376 fn partial_smt_round_trip() {
377 let key0 = Word::from([1, 2, 3, 4u32]);
378 let key1 = Word::from([5, 6, 7, 8u32]);
379 let missing_key = Word::from([9, 10, 11, 12u32]);
380 let value0 = Word::from([13, 14, 15, 16u32]);
381 let value1 = Word::from([17, 18, 19, 20u32]);
382 let smt = Smt::with_entries([(key0, value0), (key1, value1)]).unwrap();
383 let partial_smt =
384 PartialSmt::from_proofs([smt.open(&key0), smt.open(&missing_key)]).unwrap();
385
386 let encoded: proto::primitives::PartialSmt = partial_smt.clone().into();
387 assert!(encoded.node_levels.is_sorted_by_key(|level| level.depth));
388
389 let decoded = PartialSmt::try_from(encoded).unwrap();
390
391 assert_eq!(decoded, partial_smt);
392 assert_eq!(decoded.get_value(&key0).unwrap(), value0);
393 assert_eq!(decoded.get_value(&missing_key).unwrap(), Word::empty());
394 }
395
396 #[test]
397 fn partial_smt_encoding_is_canonical_for_equivalent_unique_nodes() {
398 let mut first = UniqueNodes::empty();
399 first.nodes.insert(NodeIndex::new(1, 1).unwrap(), Word::from([1, 2, 3, 4u32]));
400 first
401 .nodes
402 .insert(NodeIndex::new(1, 0).unwrap(), Word::from([9, 10, 11, 12u32]));
403 first.leaves = BTreeMap::from([
404 (2, SmtLeaf::new_empty(LeafIndex::new_max_depth(2))),
405 (1, SmtLeaf::new_empty(LeafIndex::new_max_depth(1))),
406 ]);
407 first.value_only_leaves =
408 BTreeMap::from([(2, Word::from([5, 6, 7, 8u32])), (1, Word::from([9, 10, 11, 12u32]))]);
409
410 let mut second = first.clone();
411 second.nodes = BTreeMap::from([
412 (NodeIndex::new(1, 0).unwrap(), Word::from([9, 10, 11, 12u32])),
413 (NodeIndex::new(1, 1).unwrap(), Word::from([1, 2, 3, 4u32])),
414 ]);
415 second.leaves = BTreeMap::from([
416 (1, SmtLeaf::new_empty(LeafIndex::new_max_depth(1))),
417 (2, SmtLeaf::new_empty(LeafIndex::new_max_depth(2))),
418 ]);
419 second.value_only_leaves =
420 BTreeMap::from([(1, Word::from([9, 10, 11, 12u32])), (2, Word::from([5, 6, 7, 8u32]))]);
421
422 let first: proto::primitives::PartialSmt = first.into();
423 let second: proto::primitives::PartialSmt = second.into();
424
425 assert_eq!(first, second);
426 assert_eq!(first.encode_to_vec(), second.encode_to_vec());
427 }
428
429 #[test]
430 fn partial_smt_encoding_preserves_nodes_at_every_depth() {
431 let expected_nodes = BTreeMap::from([
432 (NodeIndex::new(1, 0).unwrap(), Word::from([1, 2, 3, 4u32])),
433 (NodeIndex::new(1, 1).unwrap(), Word::from([5, 6, 7, 8u32])),
434 (NodeIndex::new(2, 0).unwrap(), Word::from([9, 10, 11, 12u32])),
435 (NodeIndex::new(2, 3).unwrap(), Word::from([13, 14, 15, 16u32])),
436 (NodeIndex::new(3, 5).unwrap(), Word::from([17, 18, 19, 20u32])),
437 ]);
438 let mut unique_nodes = UniqueNodes::empty();
439 unique_nodes.nodes = expected_nodes.clone();
440
441 let encoded: proto::primitives::PartialSmt = unique_nodes.into();
442
443 assert_eq!(
444 encoded.node_levels.iter().map(|level| level.depth).collect::<Vec<_>>(),
445 vec![1, 2, 3]
446 );
447 let decoded = UniqueNodes::try_from(encoded).unwrap();
448 assert_eq!(decoded.nodes, expected_nodes);
449 }
450
451 fn empty_partial_smt_message() -> proto::primitives::PartialSmt {
452 proto::primitives::PartialSmt {
453 root: Some(PartialSmt::EMPTY_ROOT.into()),
454 node_levels: vec![],
455 leaves: vec![],
456 value_only_leaves: vec![],
457 }
458 }
459
460 fn assert_partial_smt_decode_error(
461 encoded: proto::primitives::PartialSmt,
462 expected_error: &str,
463 ) {
464 let error = PartialSmt::try_from(encoded).unwrap_err();
465 assert_eq!(error.to_string(), expected_error);
466 }
467
468 #[test]
469 fn partial_smt_rejects_missing_root() {
470 let mut encoded = empty_partial_smt_message();
471 encoded.root = None;
472 assert_partial_smt_decode_error(
473 encoded,
474 "field miden_objects::proto::primitives::PartialSmt::root is missing",
475 );
476 }
477
478 #[test]
479 fn partial_smt_rejects_duplicate_depth() {
480 let mut encoded = empty_partial_smt_message();
481 encoded.node_levels = vec![
482 proto::primitives::PartialSmtNodeLevel { depth: 1, nodes: vec![] },
483 proto::primitives::PartialSmtNodeLevel { depth: 1, nodes: vec![] },
484 ];
485 assert_partial_smt_decode_error(encoded, "partial SMT contains duplicate node depth 1");
486 }
487
488 #[test]
489 fn partial_smt_rejects_invalid_node_index() {
490 let mut encoded = empty_partial_smt_message();
491 encoded.node_levels = vec![proto::primitives::PartialSmtNodeLevel {
492 depth: 1,
493 nodes: vec![proto::primitives::PartialSmtNode {
494 index: 2,
495 digest: Some(Word::empty().into()),
496 }],
497 }];
498 assert_partial_smt_decode_error(
499 encoded,
500 "node_levels.nodes.index: node index position 2 is not valid for depth 1",
501 );
502 }
503
504 #[test]
505 fn partial_smt_rejects_missing_node_digest() {
506 let mut encoded = empty_partial_smt_message();
507 encoded.node_levels = vec![proto::primitives::PartialSmtNodeLevel {
508 depth: 1,
509 nodes: vec![proto::primitives::PartialSmtNode { index: 0, digest: None }],
510 }];
511 assert_partial_smt_decode_error(
512 encoded,
513 "field miden_objects::proto::primitives::PartialSmtNode::digest is missing",
514 );
515 }
516
517 #[test]
518 fn partial_smt_rejects_missing_leaf() {
519 let mut encoded = empty_partial_smt_message();
520 encoded.leaves = vec![proto::primitives::IndexedSmtLeaf { index: 0, leaf: None }];
521 assert_partial_smt_decode_error(
522 encoded,
523 "field miden_objects::proto::primitives::IndexedSmtLeaf::leaf is missing",
524 );
525 }
526
527 #[test]
528 fn partial_smt_rejects_missing_value_only_leaf() {
529 let mut encoded = empty_partial_smt_message();
530 encoded.value_only_leaves =
531 vec![proto::primitives::IndexedDigest { index: 0, value: None }];
532 assert_partial_smt_decode_error(
533 encoded,
534 "field miden_objects::proto::primitives::IndexedDigest::value is missing",
535 );
536 }
537
538 #[test]
539 fn partial_smt_rejects_depth_overflow() {
540 let mut encoded = empty_partial_smt_message();
541 encoded.node_levels =
542 vec![proto::primitives::PartialSmtNodeLevel { depth: 256, nodes: vec![] }];
543 assert_partial_smt_decode_error(
544 encoded,
545 "node_levels.depth: out of range integral type conversion attempted",
546 );
547 }
548
549 #[test]
550 fn partial_smt_rejects_zero_depth() {
551 let mut encoded = empty_partial_smt_message();
552 encoded.node_levels =
553 vec![proto::primitives::PartialSmtNodeLevel { depth: 0, nodes: vec![] }];
554 assert_partial_smt_decode_error(
555 encoded,
556 "partial SMT node depth 0 must be in the range 1..64",
557 );
558 }
559
560 #[test]
561 fn partial_smt_rejects_smt_depth() {
562 let mut encoded = empty_partial_smt_message();
563 encoded.node_levels = vec![proto::primitives::PartialSmtNodeLevel {
564 depth: u32::from(SMT_DEPTH),
565 nodes: vec![],
566 }];
567 assert_partial_smt_decode_error(
568 encoded,
569 "partial SMT node depth 64 must be in the range 1..64",
570 );
571 }
572
573 #[test]
574 fn partial_smt_rejects_duplicate_node_index() {
575 let mut encoded = empty_partial_smt_message();
576 encoded.node_levels = vec![proto::primitives::PartialSmtNodeLevel {
577 depth: 1,
578 nodes: vec![
579 proto::primitives::PartialSmtNode {
580 index: 0,
581 digest: Some(Word::empty().into()),
582 },
583 proto::primitives::PartialSmtNode {
584 index: 0,
585 digest: Some(Word::empty().into()),
586 },
587 ],
588 }];
589 assert_partial_smt_decode_error(
590 encoded,
591 "partial SMT contains duplicate node index 0 at depth 1",
592 );
593 }
594
595 #[test]
596 fn partial_smt_rejects_duplicate_leaf_index() {
597 let mut encoded = empty_partial_smt_message();
598 encoded.leaves = vec![
599 proto::primitives::IndexedSmtLeaf {
600 index: 0,
601 leaf: Some(SmtLeaf::new_empty(LeafIndex::new_max_depth(0)).into()),
602 },
603 proto::primitives::IndexedSmtLeaf {
604 index: 0,
605 leaf: Some(SmtLeaf::new_empty(LeafIndex::new_max_depth(0)).into()),
606 },
607 ];
608 assert_partial_smt_decode_error(encoded, "partial SMT contains duplicate leaf index 0");
609 }
610
611 #[test]
612 fn partial_smt_rejects_duplicate_value_only_leaf_index() {
613 let mut encoded = empty_partial_smt_message();
614 encoded.value_only_leaves = vec![
615 proto::primitives::IndexedDigest {
616 index: 0,
617 value: Some(Word::empty().into()),
618 },
619 proto::primitives::IndexedDigest {
620 index: 0,
621 value: Some(Word::empty().into()),
622 },
623 ];
624 assert_partial_smt_decode_error(
625 encoded,
626 "partial SMT contains duplicate value-only leaf index 0",
627 );
628 }
629
630 #[test]
631 fn partial_smt_rejects_overlapping_leaf_index() {
632 let mut encoded = empty_partial_smt_message();
633 encoded.leaves = vec![proto::primitives::IndexedSmtLeaf {
634 index: 0,
635 leaf: Some(SmtLeaf::new_empty(LeafIndex::new_max_depth(0)).into()),
636 }];
637 encoded.value_only_leaves = vec![proto::primitives::IndexedDigest {
638 index: 0,
639 value: Some(Word::empty().into()),
640 }];
641 assert_partial_smt_decode_error(
642 encoded,
643 "partial SMT leaf index 0 has both a leaf and a value-only leaf",
644 );
645 }
646
647 #[test]
648 fn partial_smt_rejects_embedded_leaf_index_mismatch() {
649 let mut encoded = empty_partial_smt_message();
650 encoded.leaves = vec![proto::primitives::IndexedSmtLeaf {
651 index: 0,
652 leaf: Some(SmtLeaf::new_empty(LeafIndex::new_max_depth(1)).into()),
653 }];
654 assert_partial_smt_decode_error(
655 encoded,
656 "failed to deserialize PartialSmt: invalid value: Node index 0 did not match the embedded leaf index 1",
657 );
658 }
659
660 #[test]
661 fn partial_smt_rejects_reconstruction_missing_node() {
662 let mut encoded = empty_partial_smt_message();
663 encoded.node_levels = vec![proto::primitives::PartialSmtNodeLevel {
664 depth: 1,
665 nodes: vec![proto::primitives::PartialSmtNode {
666 index: 0,
667 digest: Some(Word::empty().into()),
668 }],
669 }];
670 assert_partial_smt_decode_error(
671 encoded,
672 "failed to deserialize PartialSmt: invalid value: inner node hash is inconsistent with parent",
673 );
674 }
675}