Skip to main content

miden_crypto/merkle/smt/full/
leaf.rs

1use alloc::{string::ToString, vec::Vec};
2
3use super::EMPTY_WORD;
4use crate::{
5    Felt, Word,
6    hash::poseidon2::Poseidon2,
7    merkle::smt::{LEAF_DOMAIN, LeafIndex, MAX_LEAF_ENTRIES, SMT_DEPTH, SmtLeafError},
8    utils::{ByteReader, ByteWriter, Deserializable, DeserializationError, Serializable},
9};
10
11/// The number of field elements in a key-value pair (two Words, 4 Felts each).
12const DOUBLE_WORD_LEN: usize = 8;
13
14/// Represents a leaf node in the Sparse Merkle Tree.
15///
16/// A leaf can be empty, hold a single key-value pair, or multiple key-value pairs.
17#[derive(Clone, Debug, PartialEq, Eq)]
18pub enum SmtLeaf {
19    /// An empty leaf at the specified index.
20    Empty(LeafIndex<SMT_DEPTH>),
21    /// A leaf containing a single key-value pair.
22    Single((Word, Word)),
23    /// A leaf containing multiple key-value pairs.
24    Multiple(Vec<(Word, Word)>),
25}
26
27impl SmtLeaf {
28    // CONSTRUCTORS
29    // ---------------------------------------------------------------------------------------------
30
31    /// Returns a new leaf with the specified entries
32    ///
33    /// # Errors
34    ///   - Returns an error if 2 keys in `entries` map to a different leaf index
35    ///   - Returns an error if 1 or more keys in `entries` map to a leaf index different from
36    ///     `leaf_index`
37    pub fn new(
38        entries: Vec<(Word, Word)>,
39        leaf_index: LeafIndex<SMT_DEPTH>,
40    ) -> Result<Self, SmtLeafError> {
41        match entries.len() {
42            0 => Ok(Self::new_empty(leaf_index)),
43            1 => {
44                let (key, value) = entries[0];
45
46                let computed_index = LeafIndex::<SMT_DEPTH>::from(key);
47                if computed_index != leaf_index {
48                    return Err(SmtLeafError::InconsistentSingleLeafIndices {
49                        key,
50                        expected_leaf_index: leaf_index,
51                        actual_leaf_index: computed_index,
52                    });
53                }
54
55                Ok(Self::new_single(key, value))
56            },
57            _ => {
58                let leaf = Self::new_multiple(entries)?;
59
60                // `new_multiple()` checked that all keys map to the same leaf index. We still need
61                // to ensure that leaf index is `leaf_index`.
62                if leaf.index() != leaf_index {
63                    Err(SmtLeafError::InconsistentMultipleLeafIndices {
64                        leaf_index_from_keys: leaf.index(),
65                        leaf_index_supplied: leaf_index,
66                    })
67                } else {
68                    Ok(leaf)
69                }
70            },
71        }
72    }
73
74    /// Returns a new empty leaf with the specified leaf index
75    pub fn new_empty(leaf_index: LeafIndex<SMT_DEPTH>) -> Self {
76        Self::Empty(leaf_index)
77    }
78
79    /// Returns a new single leaf with the specified entry. The leaf index is derived from the
80    /// entry's key.
81    pub fn new_single(key: Word, value: Word) -> Self {
82        Self::Single((key, value))
83    }
84
85    /// Returns a new multiple leaf with the specified entries. The leaf index is derived from the
86    /// entries' keys. Entries must be sorted by key in strictly increasing order, which is the
87    /// form that `SmtLeaf::insert` and `SmtLeaf::remove` maintain.
88    ///
89    /// # Errors
90    ///   - Returns an error if 2 keys in `entries` map to a different leaf index
91    ///   - Returns an error if the keys are not sorted in strictly increasing order (this also
92    ///     rejects repeated keys)
93    ///   - Returns an error if the number of entries exceeds [`MAX_LEAF_ENTRIES`]
94    pub fn new_multiple(entries: Vec<(Word, Word)>) -> Result<Self, SmtLeafError> {
95        if entries.len() < 2 {
96            return Err(SmtLeafError::MultipleLeafRequiresTwoEntries(entries.len()));
97        }
98
99        if entries.len() > MAX_LEAF_ENTRIES {
100            return Err(SmtLeafError::TooManyLeafEntries { actual: entries.len() });
101        }
102
103        // Check that all keys map to the same leaf index and are strictly increasing, since
104        // `insert()` and `remove()` binary-search the entries.
105        {
106            let mut keys = entries.iter().map(|(key, _)| key);
107
108            let first_key = *keys.next().expect("ensured at least 2 entries");
109            let first_leaf_index: LeafIndex<SMT_DEPTH> = first_key.into();
110            let mut previous_key = first_key;
111
112            for &next_key in keys {
113                let next_leaf_index: LeafIndex<SMT_DEPTH> = next_key.into();
114
115                if next_leaf_index != first_leaf_index {
116                    return Err(SmtLeafError::InconsistentMultipleLeafKeys {
117                        key_1: first_key,
118                        key_2: next_key,
119                    });
120                }
121
122                if next_key <= previous_key {
123                    return Err(SmtLeafError::UnsortedMultipleLeafKeys {
124                        previous: previous_key,
125                        next: next_key,
126                    });
127                }
128                previous_key = next_key;
129            }
130        }
131
132        Ok(Self::Multiple(entries))
133    }
134
135    // PUBLIC ACCESSORS
136    // ---------------------------------------------------------------------------------------------
137
138    /// Returns the value associated with `key` in the leaf, or `None` if `key` maps to another
139    /// leaf.
140    pub fn get_value(&self, key: &Word) -> Option<Word> {
141        // Ensure that `key` maps to this leaf
142        if self.index() != (*key).into() {
143            return None;
144        }
145
146        match self {
147            SmtLeaf::Empty(_) => Some(EMPTY_WORD),
148            SmtLeaf::Single((key_in_leaf, value_in_leaf)) => {
149                if key == key_in_leaf {
150                    Some(*value_in_leaf)
151                } else {
152                    Some(EMPTY_WORD)
153                }
154            },
155            SmtLeaf::Multiple(kv_pairs) => {
156                for (key_in_leaf, value_in_leaf) in kv_pairs {
157                    if key == key_in_leaf {
158                        return Some(*value_in_leaf);
159                    }
160                }
161
162                Some(EMPTY_WORD)
163            },
164        }
165    }
166
167    /// Returns true if the leaf is empty
168    pub fn is_empty(&self) -> bool {
169        matches!(self, Self::Empty(_))
170    }
171
172    /// Returns the leaf's index in the [`super::Smt`]
173    pub fn index(&self) -> LeafIndex<SMT_DEPTH> {
174        match self {
175            SmtLeaf::Empty(leaf_index) => *leaf_index,
176            SmtLeaf::Single((key, _)) => (*key).into(),
177            SmtLeaf::Multiple(entries) => {
178                // Note: All keys are guaranteed to have the same leaf index
179                let (first_key, _) = entries[0];
180                first_key.into()
181            },
182        }
183    }
184
185    /// Returns the number of entries stored in the leaf
186    pub fn num_entries(&self) -> usize {
187        match self {
188            SmtLeaf::Empty(_) => 0,
189            SmtLeaf::Single(_) => 1,
190            SmtLeaf::Multiple(entries) => entries.len(),
191        }
192    }
193
194    /// Computes the hash of the leaf
195    pub fn hash(&self) -> Word {
196        match self {
197            SmtLeaf::Empty(_) => EMPTY_WORD,
198            SmtLeaf::Single((key, value)) => {
199                Poseidon2::merge_in_domain(&[*key, *value], LEAF_DOMAIN)
200            },
201            SmtLeaf::Multiple(kvs) => {
202                let elements: Vec<Felt> = kvs.iter().copied().flat_map(kv_to_elements).collect();
203                Poseidon2::hash_elements_in_domain(&elements, LEAF_DOMAIN)
204            },
205        }
206    }
207
208    // ITERATORS
209    // ---------------------------------------------------------------------------------------------
210
211    /// Returns a slice with key-value pairs in the leaf.
212    pub fn entries(&self) -> &[(Word, Word)] {
213        match self {
214            SmtLeaf::Empty(_) => &[],
215            SmtLeaf::Single(kv_pair) => core::slice::from_ref(kv_pair),
216            SmtLeaf::Multiple(kv_pairs) => kv_pairs,
217        }
218    }
219
220    // CONVERSIONS
221    // ---------------------------------------------------------------------------------------------
222
223    /// Returns an iterator over the field elements representing this leaf.
224    pub fn to_elements(&self) -> impl Iterator<Item = Felt> + '_ {
225        self.entries().iter().copied().flat_map(kv_to_elements)
226    }
227
228    /// Returns an iterator over the key-value pairs in the leaf.
229    pub fn to_entries(&self) -> impl Iterator<Item = (&Word, &Word)> + '_ {
230        // Needed for type conversion from `&(T, T)` to `(&T, &T)`.
231        self.entries().iter().map(|(k, v)| (k, v))
232    }
233
234    /// Converts a leaf to a list of field elements.
235    pub fn into_elements(self) -> Vec<Felt> {
236        self.into_entries().into_iter().flat_map(kv_to_elements).collect()
237    }
238
239    /// Converts a leaf the key-value pairs in the leaf
240    pub fn into_entries(self) -> Vec<(Word, Word)> {
241        match self {
242            SmtLeaf::Empty(_) => Vec::new(),
243            SmtLeaf::Single(kv_pair) => vec![kv_pair],
244            SmtLeaf::Multiple(kv_pairs) => kv_pairs,
245        }
246    }
247
248    /// Converts a list of elements into a leaf
249    pub fn try_from_elements(
250        elements: &[Felt],
251        leaf_index: LeafIndex<SMT_DEPTH>,
252    ) -> Result<SmtLeaf, SmtLeafError> {
253        if elements.is_empty() {
254            return Ok(SmtLeaf::new_empty(leaf_index));
255        }
256
257        // Elements should be organized into a contiguous array of K/V Words (4 Felts each).
258        if !elements.len().is_multiple_of(DOUBLE_WORD_LEN) {
259            return Err(SmtLeafError::DecodingError(
260                "elements length is not a multiple of 8".into(),
261            ));
262        }
263
264        let mut entries = Vec::with_capacity(elements.len() / DOUBLE_WORD_LEN);
265        for entry in elements.as_chunks::<DOUBLE_WORD_LEN>().0 {
266            let key = Word::new([entry[0], entry[1], entry[2], entry[3]]);
267            let value = Word::new([entry[4], entry[5], entry[6], entry[7]]);
268            entries.push((key, value));
269        }
270
271        SmtLeaf::new(entries, leaf_index)
272    }
273
274    // HELPERS
275    // ---------------------------------------------------------------------------------------------
276
277    /// Inserts key-value pair into the leaf; returns the previous value associated with `key`, if
278    /// any.
279    ///
280    /// The caller needs to ensure that `key` has the same leaf index as all other keys in the leaf
281    ///
282    /// # Errors
283    /// Returns an error if inserting the key-value pair would exceed [`MAX_LEAF_ENTRIES`] (1024
284    /// entries) in the leaf.
285    pub(in crate::merkle::smt) fn insert(
286        &mut self,
287        key: Word,
288        value: Word,
289    ) -> Result<Option<Word>, SmtLeafError> {
290        match self {
291            SmtLeaf::Empty(_) => {
292                *self = SmtLeaf::new_single(key, value);
293                Ok(None)
294            },
295            SmtLeaf::Single(kv_pair) => {
296                if kv_pair.0 == key {
297                    // the key is already in this leaf. Update the value and return the previous
298                    // value
299                    let old_value = kv_pair.1;
300                    kv_pair.1 = value;
301                    Ok(Some(old_value))
302                } else {
303                    // Another entry is present in this leaf. Transform the entry into a list
304                    // entry, and make sure the key-value pairs are sorted by key
305                    // This stays within MAX_LEAF_ENTRIES limit. We're only adding one entry to a
306                    // single leaf
307                    let mut pairs = vec![*kv_pair, (key, value)];
308                    pairs.sort_by_key(|(key, _)| *key);
309                    *self = SmtLeaf::Multiple(pairs);
310                    Ok(None)
311                }
312            },
313            SmtLeaf::Multiple(kv_pairs) => {
314                match kv_pairs.binary_search_by(|kv_pair| kv_pair.0.cmp(&key)) {
315                    Ok(pos) => {
316                        let old_value = kv_pairs[pos].1;
317                        kv_pairs[pos].1 = value;
318                        Ok(Some(old_value))
319                    },
320                    Err(pos) => {
321                        if kv_pairs.len() >= MAX_LEAF_ENTRIES {
322                            return Err(SmtLeafError::TooManyLeafEntries {
323                                actual: kv_pairs.len() + 1,
324                            });
325                        }
326                        kv_pairs.insert(pos, (key, value));
327                        Ok(None)
328                    },
329                }
330            },
331        }
332    }
333
334    /// Removes key-value pair from the leaf stored at key; returns the previous value associated
335    /// with `key`, if any. Also returns an `is_empty` flag, indicating whether the leaf became
336    /// empty, and must be removed from the data structure it is contained in.
337    pub(in crate::merkle::smt) fn remove(&mut self, key: Word) -> (Option<Word>, bool) {
338        match self {
339            SmtLeaf::Empty(_) => (None, false),
340            SmtLeaf::Single((key_at_leaf, value_at_leaf)) => {
341                if *key_at_leaf == key {
342                    // our key was indeed stored in the leaf, so we return the value that was stored
343                    // in it, and indicate that the leaf should be removed
344                    let old_value = *value_at_leaf;
345
346                    // Note: this is not strictly needed, since the caller is expected to drop this
347                    // `SmtLeaf` object.
348                    *self = SmtLeaf::new_empty(key.into());
349
350                    (Some(old_value), true)
351                } else {
352                    // another key is stored at leaf; nothing to update
353                    (None, false)
354                }
355            },
356            SmtLeaf::Multiple(kv_pairs) => {
357                match kv_pairs.binary_search_by(|kv_pair| kv_pair.0.cmp(&key)) {
358                    Ok(pos) => {
359                        let old_value = kv_pairs[pos].1;
360
361                        let _ = kv_pairs.remove(pos);
362                        debug_assert!(!kv_pairs.is_empty());
363
364                        if kv_pairs.len() == 1 {
365                            // convert the leaf into `Single`
366                            *self = SmtLeaf::Single(kv_pairs[0]);
367                        }
368
369                        (Some(old_value), false)
370                    },
371                    Err(_) => {
372                        // other keys are stored at leaf; nothing to update
373                        (None, false)
374                    },
375                }
376            },
377        }
378    }
379}
380
381impl Serializable for SmtLeaf {
382    fn write_into<W: ByteWriter>(&self, target: &mut W) {
383        // Write: num entries
384        self.num_entries().write_into(target);
385
386        // Write: leaf index
387        let leaf_index: u64 = self.index().position();
388        leaf_index.write_into(target);
389
390        // Write: entries
391        for (key, value) in self.entries() {
392            key.write_into(target);
393            value.write_into(target);
394        }
395    }
396}
397
398impl Deserializable for SmtLeaf {
399    fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
400        // Read: num entries
401        let num_entries = source.read_usize()?;
402
403        // Read: leaf index
404        let leaf_index: LeafIndex<SMT_DEPTH> = {
405            let value = source.read_u64()?;
406            LeafIndex::new_max_depth(value)
407        };
408
409        // Read: entries using read_many_iter to avoid eager allocation
410        let entries: Vec<(Word, Word)> =
411            source.read_many_iter(num_entries)?.collect::<Result<_, _>>()?;
412
413        Self::new(entries, leaf_index)
414            .map_err(|err| DeserializationError::InvalidValue(err.to_string()))
415    }
416
417    /// Minimum serialized size: vint64 (num_entries) + u64 (leaf_index) with 0 entries.
418    fn min_serialized_size() -> usize {
419        1 + 8
420    }
421}
422
423// HELPER FUNCTIONS
424// ================================================================================================
425
426/// Converts a key-value tuple to an iterator of `Felt`s
427pub(crate) fn kv_to_elements((key, value): (Word, Word)) -> impl Iterator<Item = Felt> {
428    let key_elements = key.into_iter();
429    let value_elements = value.into_iter();
430
431    key_elements.chain(value_elements)
432}