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#[derive(Debug)]
30pub struct TrieWalker<C, K = AddedRemovedKeys> {
31 pub cursor: C,
33 pub stack: Vec<CursorSubNode>,
35 pub can_skip_current_node: bool,
39 pub changes: PrefixSet,
41 walk_all_changed_branch_children: bool,
44 removed_keys: Option<HashSet<Nibbles>>,
46 added_removed_keys: Option<K>,
50 #[cfg(feature = "metrics")]
51 metrics: WalkerMetrics,
53}
54
55impl<C: TrieCursor, K: AsRef<AddedRemovedKeys>> TrieWalker<C, K> {
56 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 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 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 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 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 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 pub fn split(mut self) -> (Vec<CursorSubNode>, HashSet<Nibbles>) {
136 let keys = self.take_removed_keys();
137 (self.stack, keys)
138 }
139
140 pub fn take_removed_keys(&mut self) -> HashSet<Nibbles> {
142 self.removed_keys.take().unwrap_or_default()
143 }
144
145 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 pub fn removed_keys_len(&self) -> usize {
156 self.removed_keys.as_ref().map_or(0, |u| u.len())
157 }
158
159 pub fn key(&self) -> Option<&Nibbles> {
161 self.stack.last().map(|n| n.full_key())
162 }
163
164 pub fn hash(&self) -> Option<B256> {
166 self.stack.last().and_then(|n| n.hash())
167 }
168
169 pub fn maybe_hash(&self) -> Option<B256> {
174 self.stack.last().and_then(|n| n.maybe_hash())
175 }
176
177 pub fn children_are_in_trie(&self) -> bool {
179 self.stack.last().is_some_and(|n| n.tree_flag())
180 }
181
182 #[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 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 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 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 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 fn new(
250 cursor: C,
251 changes: PrefixSet,
252 #[cfg(feature = "metrics")] trie_type: crate::TrieType,
253 ) -> Self {
254 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 if let Some((key, value)) = this.node(true).unwrap() {
269 this.stack[0] = CursorSubNode::new(key, Some(value));
270 }
271
272 this.update_skip_node();
274 this
275 }
276
277 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 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 self.move_to_next_sibling(false)?;
301 }
302
303 self.update_skip_node();
305 }
306
307 Ok(())
308 }
309
310 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 #[instrument(level = "trace", skip(self), ret)]
326 fn consume_node(&mut self) -> Result<(), DatabaseError> {
327 let Some((key, node)) = self.node(false)? else {
328 self.stack.clear();
330 return Ok(())
331 };
332
333 if !key.is_empty() && !self.stack.is_empty() {
337 self.stack[0].set_nibble(key.get_unchecked(0));
338 }
339
340 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 let subnode = CursorSubNode::new(key, Some(node));
355 let position = subnode.position();
356 self.stack.push(subnode);
357 self.update_skip_node();
358
359 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 #[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 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 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 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}