Skip to main content

reth_trie/
walker.rs

1use crate::{
2    prefix_set::PrefixSet,
3    trie_cursor::{subnode::SubNodePosition, CursorSubNode, TrieCursor},
4    BranchNodeCompact, Nibbles,
5};
6use alloy_primitives::{map::HashSet, B256};
7use alloy_trie::proof::AddedRemovedKeys;
8use reth_storage_errors::db::DatabaseError;
9use tracing::{instrument, trace};
10
11#[cfg(test)]
12use crate::trie_cursor::{mock::MockTrieCursorFactory, TrieCursorFactory};
13
14#[cfg(test)]
15use alloy_primitives::map::B256Map;
16
17#[cfg(test)]
18use alloy_trie::TrieMask;
19
20#[cfg(test)]
21use std::collections::BTreeMap;
22
23#[cfg(feature = "metrics")]
24use crate::metrics::WalkerMetrics;
25
26/// Traverses the trie in lexicographic order.
27///
28/// This iterator depends on the ordering guarantees of [`TrieCursor`].
29#[derive(Debug)]
30pub struct TrieWalker<C, K = AddedRemovedKeys> {
31    /// A mutable reference to a trie cursor instance used for navigating the trie.
32    pub cursor: C,
33    /// A vector containing the trie nodes that have been visited.
34    pub stack: Vec<CursorSubNode>,
35    /// A flag indicating whether the current node can be skipped when traversing the trie. This
36    /// is determined by whether the current key's prefix is included in the prefix set and if the
37    /// hash flag is set.
38    pub can_skip_current_node: bool,
39    /// A `PrefixSet` representing the changes to be applied to the trie.
40    pub changes: PrefixSet,
41    /// When enabled, all children of a branch become unskippable if the branch path itself
42    /// matches the prefix set, even if a given child path does not.
43    walk_all_changed_branch_children: bool,
44    /// The retained trie node keys that need to be removed.
45    removed_keys: Option<HashSet<Nibbles>>,
46    /// Provided when it's necessary not to skip certain nodes during proof generation.
47    /// Specifically we don't skip certain branch nodes even when they are not in the `PrefixSet`,
48    /// when they might be required to support leaf removal.
49    added_removed_keys: Option<K>,
50    #[cfg(feature = "metrics")]
51    /// Walker metrics.
52    metrics: WalkerMetrics,
53}
54
55impl<C: TrieCursor, K: AsRef<AddedRemovedKeys>> TrieWalker<C, K> {
56    /// Constructs a new `TrieWalker` for the state trie from existing stack and a cursor.
57    pub fn state_trie_from_stack(cursor: C, stack: Vec<CursorSubNode>, changes: PrefixSet) -> Self {
58        Self::from_stack(
59            cursor,
60            stack,
61            changes,
62            #[cfg(feature = "metrics")]
63            crate::TrieType::State,
64        )
65    }
66
67    /// Constructs a new `TrieWalker` for the storage trie from existing stack and a cursor.
68    pub fn storage_trie_from_stack(
69        cursor: C,
70        stack: Vec<CursorSubNode>,
71        changes: PrefixSet,
72    ) -> Self {
73        Self::from_stack(
74            cursor,
75            stack,
76            changes,
77            #[cfg(feature = "metrics")]
78            crate::TrieType::Storage,
79        )
80    }
81
82    /// Constructs a new `TrieWalker` from existing stack and a cursor.
83    fn from_stack(
84        cursor: C,
85        stack: Vec<CursorSubNode>,
86        changes: PrefixSet,
87        #[cfg(feature = "metrics")] trie_type: crate::TrieType,
88    ) -> Self {
89        let mut this = Self {
90            cursor,
91            changes,
92            stack,
93            can_skip_current_node: false,
94            walk_all_changed_branch_children: false,
95            removed_keys: None,
96            added_removed_keys: None,
97            #[cfg(feature = "metrics")]
98            metrics: WalkerMetrics::new(trie_type),
99        };
100        this.update_skip_node();
101        this
102    }
103
104    /// Sets the flag whether the trie updates should be stored.
105    pub fn with_deletions_retained(mut self, retained: bool) -> Self {
106        if retained {
107            self.removed_keys = Some(HashSet::default());
108        }
109        self
110    }
111
112    /// Configures the walker to not skip certain branch nodes, even when they are not in the
113    /// `PrefixSet`, when they might be needed to support leaf removal.
114    pub fn with_added_removed_keys<K2>(self, added_removed_keys: Option<K2>) -> TrieWalker<C, K2> {
115        TrieWalker {
116            cursor: self.cursor,
117            stack: self.stack,
118            can_skip_current_node: self.can_skip_current_node,
119            changes: self.changes,
120            walk_all_changed_branch_children: self.walk_all_changed_branch_children,
121            removed_keys: self.removed_keys,
122            added_removed_keys,
123            #[cfg(feature = "metrics")]
124            metrics: self.metrics,
125        }
126    }
127
128    /// Configures the walker to treat every child of a matching branch path as unskippable.
129    pub const fn with_walk_all_changed_branch_children(mut self, enabled: bool) -> Self {
130        self.walk_all_changed_branch_children = enabled;
131        self
132    }
133
134    /// Split the walker into stack and trie updates.
135    pub fn split(mut self) -> (Vec<CursorSubNode>, HashSet<Nibbles>) {
136        let keys = self.take_removed_keys();
137        (self.stack, keys)
138    }
139
140    /// Take removed keys from the walker.
141    pub fn take_removed_keys(&mut self) -> HashSet<Nibbles> {
142        self.removed_keys.take().unwrap_or_default()
143    }
144
145    /// Prints the current stack of trie nodes.
146    pub fn print_stack(&self) {
147        println!("====================== STACK ======================");
148        for node in &self.stack {
149            println!("{node:?}");
150        }
151        println!("====================== END STACK ======================\n");
152    }
153
154    /// The current length of the removed keys.
155    pub fn removed_keys_len(&self) -> usize {
156        self.removed_keys.as_ref().map_or(0, |u| u.len())
157    }
158
159    /// Returns the current key in the trie.
160    pub fn key(&self) -> Option<&Nibbles> {
161        self.stack.last().map(|n| n.full_key())
162    }
163
164    /// Returns the current hash in the trie, if any.
165    pub fn hash(&self) -> Option<B256> {
166        self.stack.last().and_then(|n| n.hash())
167    }
168
169    /// Returns the current hash in the trie, if any.
170    ///
171    /// Differs from [`Self::hash`] in that it returns `None` if the subnode is positioned at the
172    /// child without a hash mask bit set. [`Self::hash`] panics in that case.
173    pub fn maybe_hash(&self) -> Option<B256> {
174        self.stack.last().and_then(|n| n.maybe_hash())
175    }
176
177    /// Indicates whether the children of the current node are present in the trie.
178    pub fn children_are_in_trie(&self) -> bool {
179        self.stack.last().is_some_and(|n| n.tree_flag())
180    }
181
182    /// Returns the next unprocessed key in the trie along with its raw [`Nibbles`] representation.
183    #[instrument(level = "trace", skip(self), ret)]
184    pub fn next_unprocessed_key(&self) -> Option<(B256, Nibbles)> {
185        self.key()
186            .and_then(|key| if self.can_skip_current_node { key.increment() } else { Some(*key) })
187            .map(|key| (B256::right_padding_from(&key.pack()), key))
188    }
189
190    /// Updates the skip node flag based on the walker's current state.
191    fn update_skip_node(&mut self) {
192        let old = self.can_skip_current_node;
193        self.can_skip_current_node = self.stack.last().is_some_and(|node| {
194            // If the current key is not removed according to the [`AddedRemovedKeys`], and all of
195            // its siblings are removed, then we don't want to skip it. This allows the
196            // `ProofRetainer` to include this node in the returned proofs. Required to support
197            // leaf removal.
198            let key_is_only_nonremoved_child =
199                self.added_removed_keys.as_ref().is_some_and(|added_removed_keys| {
200                    node.full_key_is_only_nonremoved_child(added_removed_keys.as_ref())
201                });
202
203            trace!(
204                target: "trie::walker",
205                ?key_is_only_nonremoved_child,
206                full_key=?node.full_key(),
207                "Checked for only non-removed child",
208            );
209
210            let branch_path_matches_prefix_set = self.walk_all_changed_branch_children &&
211                node.position().is_child() &&
212                self.changes.contains(&node.key);
213
214            !self.changes.contains(node.full_key()) &&
215                !branch_path_matches_prefix_set &&
216                node.hash_flag() &&
217                !key_is_only_nonremoved_child
218        });
219        trace!(
220            target: "trie::walker",
221            old,
222            new = self.can_skip_current_node,
223            last = ?self.stack.last(),
224            "updated skip node flag"
225        );
226    }
227
228    /// Constructs a new [`TrieWalker`] for the state trie.
229    pub fn state_trie(cursor: C, changes: PrefixSet) -> Self {
230        Self::new(
231            cursor,
232            changes,
233            #[cfg(feature = "metrics")]
234            crate::TrieType::State,
235        )
236    }
237
238    /// Constructs a new [`TrieWalker`] for the storage trie.
239    pub fn storage_trie(cursor: C, changes: PrefixSet) -> Self {
240        Self::new(
241            cursor,
242            changes,
243            #[cfg(feature = "metrics")]
244            crate::TrieType::Storage,
245        )
246    }
247
248    /// Constructs a new `TrieWalker`, setting up the initial state of the stack and cursor.
249    fn new(
250        cursor: C,
251        changes: PrefixSet,
252        #[cfg(feature = "metrics")] trie_type: crate::TrieType,
253    ) -> Self {
254        // Initialize the walker with a single empty stack element.
255        let mut this = Self {
256            cursor,
257            changes,
258            stack: vec![CursorSubNode::default()],
259            can_skip_current_node: false,
260            walk_all_changed_branch_children: false,
261            removed_keys: None,
262            added_removed_keys: Default::default(),
263            #[cfg(feature = "metrics")]
264            metrics: WalkerMetrics::new(trie_type),
265        };
266
267        // Set up the root node of the trie in the stack, if it exists.
268        if let Some((key, value)) = this.node(true).unwrap() {
269            this.stack[0] = CursorSubNode::new(key, Some(value));
270        }
271
272        // Update the skip state for the root node.
273        this.update_skip_node();
274        this
275    }
276
277    /// Advances the walker to the next trie node and updates the skip node flag.
278    /// The new key can then be obtained via `key()`.
279    ///
280    /// # Returns
281    ///
282    /// * `Result<(), Error>` - Unit on success or an error.
283    pub fn advance(&mut self) -> Result<(), DatabaseError> {
284        if let Some(last) = self.stack.last() {
285            if !self.can_skip_current_node && self.children_are_in_trie() {
286                trace!(
287                    target: "trie::walker",
288                    position = ?last.position(),
289                    "cannot skip current node and children are in the trie"
290                );
291                // If we can't skip the current node and the children are in the trie,
292                // either consume the next node or move to the next sibling.
293                match last.position() {
294                    SubNodePosition::ParentBranch => self.move_to_next_sibling(true)?,
295                    SubNodePosition::Child(_) => self.consume_node()?,
296                }
297            } else {
298                trace!(target: "trie::walker", "can skip current node");
299                // If we can skip the current node, move to the next sibling.
300                self.move_to_next_sibling(false)?;
301            }
302
303            // Update the skip node flag based on the new position in the trie.
304            self.update_skip_node();
305        }
306
307        Ok(())
308    }
309
310    /// Retrieves the current root node from the DB, seeking either the exact node or the next one.
311    fn node(&mut self, exact: bool) -> Result<Option<(Nibbles, BranchNodeCompact)>, DatabaseError> {
312        let key = self.key().expect("key must exist");
313        let entry = if exact { self.cursor.seek_exact(*key)? } else { self.cursor.seek(*key)? };
314        #[cfg(feature = "metrics")]
315        self.metrics.inc_branch_nodes_seeked();
316
317        if let Some((_, node)) = &entry {
318            assert!(!node.state_mask.is_empty());
319        }
320
321        Ok(entry)
322    }
323
324    /// Consumes the next node in the trie, updating the stack.
325    #[instrument(level = "trace", skip(self), ret)]
326    fn consume_node(&mut self) -> Result<(), DatabaseError> {
327        let Some((key, node)) = self.node(false)? else {
328            // If no next node is found, clear the stack.
329            self.stack.clear();
330            return Ok(())
331        };
332
333        // Overwrite the root node's first nibble
334        // We need to sync the stack with the trie structure when consuming a new node. This is
335        // necessary for proper traversal and accurately representing the trie in the stack.
336        if !key.is_empty() && !self.stack.is_empty() {
337            self.stack[0].set_nibble(key.get_unchecked(0));
338        }
339
340        // The current tree mask might have been set incorrectly.
341        // Sanity check that the newly retrieved trie node key is the child of the last item
342        // on the stack. If not, advance to the next sibling instead of adding the node to the
343        // stack.
344        if let Some(subnode) = self.stack.last() &&
345            !key.starts_with(subnode.full_key())
346        {
347            #[cfg(feature = "metrics")]
348            self.metrics.inc_out_of_order_subnode(1);
349            self.move_to_next_sibling(false)?;
350            return Ok(())
351        }
352
353        // Create a new CursorSubNode and push it to the stack.
354        let subnode = CursorSubNode::new(key, Some(node));
355        let position = subnode.position();
356        self.stack.push(subnode);
357        self.update_skip_node();
358
359        // Delete the current node if it's included in the prefix set or it doesn't contain the root
360        // hash.
361        if (!self.can_skip_current_node || position.is_child()) &&
362            let Some((keys, key)) = self.removed_keys.as_mut().zip(self.cursor.current()?)
363        {
364            keys.insert(key);
365        }
366
367        Ok(())
368    }
369
370    /// Moves to the next sibling node in the trie, updating the stack.
371    #[instrument(level = "trace", skip(self), ret)]
372    fn move_to_next_sibling(
373        &mut self,
374        allow_root_to_child_nibble: bool,
375    ) -> Result<(), DatabaseError> {
376        let Some(subnode) = self.stack.last_mut() else { return Ok(()) };
377
378        // Check if the walker needs to backtrack to the previous level in the trie during its
379        // traversal.
380        if subnode.position().is_last_child() ||
381            (subnode.position().is_parent() && !allow_root_to_child_nibble)
382        {
383            self.stack.pop();
384            self.move_to_next_sibling(false)?;
385            return Ok(())
386        }
387
388        subnode.inc_nibble();
389
390        if subnode.node.is_none() {
391            return self.consume_node()
392        }
393
394        // Find the next sibling with state.
395        loop {
396            let position = subnode.position();
397            if subnode.state_flag() {
398                trace!(target: "trie::walker", ?position, "found next sibling with state");
399                return Ok(())
400            }
401            if position.is_last_child() {
402                trace!(target: "trie::walker", ?position, "checked all siblings");
403                break
404            }
405            subnode.inc_nibble();
406        }
407
408        // Pop the current node and move to the next sibling.
409        self.stack.pop();
410        self.move_to_next_sibling(false)?;
411
412        Ok(())
413    }
414}
415
416#[cfg(test)]
417mod tests {
418    use super::*;
419    use crate::prefix_set::PrefixSetMut;
420    use alloy_primitives::B256;
421
422    fn branch_node(state_mask: u16, tree_mask: u16, hash_mask: u16) -> BranchNodeCompact {
423        let hash_count = hash_mask.count_ones() as usize;
424        BranchNodeCompact::new(
425            TrieMask::new(state_mask),
426            TrieMask::new(tree_mask),
427            TrieMask::new(hash_mask),
428            vec![B256::ZERO; hash_count],
429            None,
430        )
431    }
432
433    fn root_branch_node(state_mask: u16, tree_mask: u16, hash_mask: u16) -> BranchNodeCompact {
434        let hash_count = hash_mask.count_ones() as usize;
435        BranchNodeCompact::new(
436            TrieMask::new(state_mask),
437            TrieMask::new(tree_mask),
438            TrieMask::new(hash_mask),
439            vec![B256::ZERO; hash_count],
440            Some(B256::ZERO),
441        )
442    }
443
444    fn walker_for_matching_branch_children_test(
445        walk_all_changed_branch_children: bool,
446    ) -> TrieWalker<crate::trie_cursor::mock::MockTrieCursor> {
447        let trie_nodes = BTreeMap::from([
448            (Nibbles::default(), root_branch_node(1 << 2, 1 << 2, 1 << 2)),
449            (
450                Nibbles::from_nibbles([0x2]),
451                branch_node((1 << 3) | (1 << 4), 0, (1 << 3) | (1 << 4)),
452            ),
453        ]);
454        let factory = MockTrieCursorFactory::new(trie_nodes, B256Map::default());
455
456        let mut prefix_set = PrefixSetMut::default();
457        prefix_set.insert(Nibbles::from_nibbles([0x2, 0x3, 0x1]));
458
459        TrieWalker::state_trie(factory.account_trie_cursor().unwrap(), prefix_set.freeze())
460            .with_walk_all_changed_branch_children(walk_all_changed_branch_children)
461    }
462
463    #[test]
464    fn branch_siblings_remain_skippable_by_default() {
465        let mut walker = walker_for_matching_branch_children_test(false);
466
467        assert_eq!(walker.key().copied(), Some(Nibbles::default()));
468        assert!(!walker.can_skip_current_node);
469
470        walker.advance().unwrap();
471        assert_eq!(walker.key().copied(), Some(Nibbles::from_nibbles([0x2])));
472        assert!(!walker.can_skip_current_node);
473
474        walker.advance().unwrap();
475        assert_eq!(walker.key().copied(), Some(Nibbles::from_nibbles([0x2, 0x3])));
476        assert_eq!(walker.stack.last().unwrap().position(), SubNodePosition::Child(0x3));
477        assert!(!walker.can_skip_current_node);
478
479        walker.advance().unwrap();
480        assert_eq!(walker.key().copied(), Some(Nibbles::from_nibbles([0x2, 0x4])));
481        assert!(walker.can_skip_current_node);
482    }
483
484    #[test]
485    fn matching_branch_path_can_make_all_children_unskippable() {
486        let mut walker = walker_for_matching_branch_children_test(true);
487
488        walker.advance().unwrap();
489        walker.advance().unwrap();
490        walker.advance().unwrap();
491        assert_eq!(walker.key().copied(), Some(Nibbles::from_nibbles([0x2, 0x4])));
492        assert!(!walker.can_skip_current_node);
493    }
494}