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
21impl 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
53impl 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
82impl 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
114impl 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
162impl 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
186impl 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
211impl 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}