Skip to main content

reth_trie/trie_cursor/
in_memory.rs

1use super::{TrieCursor, TrieCursorFactory, TrieStorageCursor};
2use crate::{forward_cursor::ForwardInMemoryCursor, updates::TrieUpdatesSorted};
3use alloy_primitives::B256;
4use reth_storage_errors::db::DatabaseError;
5use reth_trie_common::{BranchNodeCompact, Nibbles};
6
7/// The trie cursor factory for the trie updates.
8#[derive(Debug, Clone)]
9pub struct InMemoryTrieCursorFactory<CF, T> {
10    /// Underlying trie cursor factory.
11    cursor_factory: CF,
12    /// Reference to sorted trie updates.
13    trie_updates: T,
14}
15
16impl<CF, T> InMemoryTrieCursorFactory<CF, T> {
17    /// Create a new trie cursor factory.
18    pub const fn new(cursor_factory: CF, trie_updates: T) -> Self {
19        Self { cursor_factory, trie_updates }
20    }
21}
22
23impl<'overlay, CF, T> TrieCursorFactory for InMemoryTrieCursorFactory<CF, &'overlay T>
24where
25    CF: TrieCursorFactory + 'overlay,
26    T: AsRef<TrieUpdatesSorted>,
27{
28    type AccountTrieCursor<'cursor>
29        = InMemoryTrieCursor<'overlay, CF::AccountTrieCursor<'cursor>>
30    where
31        Self: 'cursor;
32
33    type StorageTrieCursor<'cursor>
34        = InMemoryTrieCursor<'overlay, CF::StorageTrieCursor<'cursor>>
35    where
36        Self: 'cursor;
37
38    fn account_trie_cursor(&self) -> Result<Self::AccountTrieCursor<'_>, DatabaseError> {
39        let cursor = self.cursor_factory.account_trie_cursor()?;
40        Ok(InMemoryTrieCursor::new_account(cursor, self.trie_updates.as_ref()))
41    }
42
43    fn storage_trie_cursor(
44        &self,
45        hashed_address: B256,
46    ) -> Result<Self::StorageTrieCursor<'_>, DatabaseError> {
47        let trie_updates = self.trie_updates.as_ref();
48        let cursor = self.cursor_factory.storage_trie_cursor(hashed_address)?;
49        Ok(InMemoryTrieCursor::new_storage(cursor, trie_updates, hashed_address))
50    }
51}
52
53/// A cursor to iterate over trie updates and corresponding database entries.
54/// It will always give precedence to the data from the trie updates.
55#[derive(Debug)]
56pub struct InMemoryTrieCursor<'a, C> {
57    /// The underlying cursor.
58    cursor: C,
59    /// Tracks whether the DB cursor is available, positioned, or exhausted.
60    db_cursor_state: DbCursorState,
61    /// Forward-only in-memory cursor over storage trie nodes.
62    in_memory_cursor: ForwardInMemoryCursor<'a, Nibbles, Option<BranchNodeCompact>>,
63    /// The key most recently returned from the Cursor.
64    last_key: Option<Nibbles>,
65    #[cfg(debug_assertions)]
66    /// Whether an initial seek was called.
67    seeked: bool,
68    /// Reference to the full trie updates.
69    trie_updates: &'a TrieUpdatesSorted,
70}
71
72#[derive(Debug)]
73enum DbCursorState {
74    NeedsPosition,
75    Positioned((Nibbles, BranchNodeCompact)),
76    Exhausted,
77}
78
79impl DbCursorState {
80    const fn entry(&self) -> Option<&(Nibbles, BranchNodeCompact)> {
81        match self {
82            Self::Positioned(entry) => Some(entry),
83            Self::NeedsPosition | Self::Exhausted => None,
84        }
85    }
86
87    fn set_entry(&mut self, entry: Option<(Nibbles, BranchNodeCompact)>) {
88        *self = match entry {
89            Some(entry) => Self::Positioned(entry),
90            None => Self::Exhausted,
91        };
92    }
93}
94
95impl<'a, C: TrieCursor> InMemoryTrieCursor<'a, C> {
96    /// Create new account trie cursor which combines a DB cursor and the trie updates.
97    pub fn new_account(cursor: C, trie_updates: &'a TrieUpdatesSorted) -> Self {
98        let in_memory_cursor = ForwardInMemoryCursor::new(trie_updates.account_nodes_ref());
99        Self {
100            cursor,
101            db_cursor_state: DbCursorState::NeedsPosition,
102            in_memory_cursor,
103            last_key: None,
104            #[cfg(debug_assertions)]
105            seeked: false,
106            trie_updates,
107        }
108    }
109
110    /// Create new storage trie cursor with full trie updates reference.
111    /// This allows the cursor to switch between storage tries when `set_hashed_address` is called.
112    pub fn new_storage(
113        cursor: C,
114        trie_updates: &'a TrieUpdatesSorted,
115        hashed_address: B256,
116    ) -> Self {
117        let in_memory_cursor = Self::get_storage_overlay(trie_updates, hashed_address);
118        Self {
119            cursor,
120            db_cursor_state: DbCursorState::NeedsPosition,
121            in_memory_cursor,
122            last_key: None,
123            #[cfg(debug_assertions)]
124            seeked: false,
125            trie_updates,
126        }
127    }
128
129    /// Returns the storage overlay for `hashed_address`.
130    fn get_storage_overlay(
131        trie_updates: &'a TrieUpdatesSorted,
132        hashed_address: B256,
133    ) -> ForwardInMemoryCursor<'a, Nibbles, Option<BranchNodeCompact>> {
134        let storage_trie_updates = trie_updates.storage_tries_ref().get(&hashed_address);
135        let storage_nodes = storage_trie_updates.map(|u| u.storage_nodes_ref()).unwrap_or(&[]);
136
137        ForwardInMemoryCursor::new(storage_nodes)
138    }
139
140    const fn get_cursor_mut(&mut self) -> &mut C {
141        &mut self.cursor
142    }
143
144    /// Asserts that the next entry to be returned from the cursor is not previous to the last entry
145    /// returned.
146    fn set_last_key(&mut self, next_entry: &Option<(Nibbles, BranchNodeCompact)>) {
147        let next_key = next_entry.as_ref().map(|e| e.0);
148        debug_assert!(
149            self.last_key.is_none_or(|last| next_key.is_none_or(|next| next >= last)),
150            "Cannot return entry {:?} previous to the last returned entry at {:?}",
151            next_key,
152            self.last_key,
153        );
154        self.last_key = next_key;
155    }
156
157    /// Positions the DB cursor state using the underlying cursor when needed.
158    fn cursor_seek(&mut self, key: Nibbles) -> Result<(), DatabaseError> {
159        // Only seek if:
160        // 1. We have a cursor entry and need to seek forward (entry.0 < key), OR
161        // 2. The DB cursor needs to be positioned.
162        let should_seek = match &self.db_cursor_state {
163            DbCursorState::NeedsPosition => true,
164            DbCursorState::Positioned((entry_key, _)) => entry_key < &key,
165            DbCursorState::Exhausted => false,
166        };
167
168        if should_seek {
169            let entry = self.get_cursor_mut().seek(key)?;
170            self.db_cursor_state.set_entry(entry);
171        }
172
173        Ok(())
174    }
175
176    /// Advances the DB cursor state to the subsequent entry using the underlying cursor.
177    fn cursor_next(&mut self) -> Result<(), DatabaseError> {
178        #[cfg(debug_assertions)]
179        {
180            debug_assert!(self.seeked);
181            debug_assert!(!matches!(self.db_cursor_state, DbCursorState::NeedsPosition));
182        }
183
184        // The exhausted state is stable; only advance if the DB cursor currently points to an
185        // entry.
186        if matches!(self.db_cursor_state, DbCursorState::Positioned(_)) {
187            let entry = self.get_cursor_mut().next()?;
188            self.db_cursor_state.set_entry(entry);
189        }
190
191        Ok(())
192    }
193
194    /// Compares the current in-memory entry with the current entry of the cursor, and applies the
195    /// in-memory entry to the cursor entry as an overlay.
196    //
197    /// This may consume and move forward the current entries when the overlay indicates a removed
198    /// node.
199    fn choose_next_entry(&mut self) -> Result<Option<(Nibbles, BranchNodeCompact)>, DatabaseError> {
200        loop {
201            let mem_entry = self.in_memory_cursor.current().cloned();
202            let db_entry = self.db_cursor_state.entry();
203
204            match (mem_entry, db_entry) {
205                (Some((mem_key, None)), _)
206                    if db_entry.is_none_or(|(db_key, _)| &mem_key < db_key) =>
207                {
208                    // If overlay has a removed node but DB cursor is exhausted or ahead of the
209                    // in-memory cursor then move ahead in-memory, as there might be further
210                    // non-removed overlay nodes.
211                    self.in_memory_cursor.first_after(&mem_key);
212                }
213                (Some((mem_key, None)), Some((db_key, _))) if &mem_key == db_key => {
214                    // If overlay has a removed node which is returned from DB then move both
215                    // cursors ahead to the next key.
216                    self.in_memory_cursor.first_after(&mem_key);
217                    self.cursor_next()?;
218                }
219                (Some((mem_key, Some(node))), _)
220                    if db_entry.is_none_or(|(db_key, _)| &mem_key <= db_key) =>
221                {
222                    // If overlay returns a node prior to the DB's node, or the DB is exhausted,
223                    // then we return the overlay's node.
224                    return Ok(Some((mem_key, node)))
225                }
226                // All other cases:
227                // - mem_key > db_key
228                // - overlay is exhausted
229                // Return the db_entry. If DB is also exhausted then this returns None.
230                _ => return Ok(db_entry.cloned()),
231            }
232        }
233    }
234}
235
236impl<C: TrieCursor> TrieCursor for InMemoryTrieCursor<'_, C> {
237    fn seek_exact(
238        &mut self,
239        key: Nibbles,
240    ) -> Result<Option<(Nibbles, BranchNodeCompact)>, DatabaseError> {
241        let mem_entry = self.in_memory_cursor.seek(&key);
242
243        if let Some((mem_key, entry_inner)) = mem_entry &&
244            *mem_key == key
245        {
246            #[cfg(debug_assertions)]
247            {
248                self.seeked = true;
249            }
250
251            // An exact overlay hit can move the logical cursor ahead without touching the DB. If
252            // the DB cursor was still behind this key, force a re-seek before the next DB-backed
253            // operation so `next()` cannot return a stale earlier entry.
254            if matches!(&self.db_cursor_state, DbCursorState::Positioned((db_key, _)) if db_key < &key)
255            {
256                self.db_cursor_state = DbCursorState::NeedsPosition;
257            }
258
259            let entry = entry_inner.clone().map(|node| (key, node));
260            self.set_last_key(&entry);
261            return Ok(entry)
262        }
263
264        self.cursor_seek(key)?;
265
266        #[cfg(debug_assertions)]
267        {
268            self.seeked = true;
269        }
270
271        let entry = match self.db_cursor_state.entry() {
272            Some((db_key, node)) if db_key == &key => Some((key, node.clone())),
273            _ => None,
274        };
275
276        self.set_last_key(&entry);
277        Ok(entry)
278    }
279
280    fn seek(
281        &mut self,
282        key: Nibbles,
283    ) -> Result<Option<(Nibbles, BranchNodeCompact)>, DatabaseError> {
284        let mem_entry = self.in_memory_cursor.seek(&key);
285
286        if let Some((mem_key, Some(node))) = mem_entry &&
287            *mem_key == key
288        {
289            #[cfg(debug_assertions)]
290            {
291                self.seeked = true;
292            }
293
294            // An exact overlay hit is the first logical entry at or after `key`, so the DB cursor
295            // can stay lazy until a later operation needs it.
296            if matches!(&self.db_cursor_state, DbCursorState::Positioned((db_key, _)) if db_key < &key)
297            {
298                self.db_cursor_state = DbCursorState::NeedsPosition;
299            }
300
301            let entry = Some((key, node.clone()));
302            self.set_last_key(&entry);
303            return Ok(entry)
304        }
305
306        self.cursor_seek(key)?;
307
308        #[cfg(debug_assertions)]
309        {
310            self.seeked = true;
311        }
312
313        let entry = self.choose_next_entry()?;
314        self.set_last_key(&entry);
315        Ok(entry)
316    }
317
318    fn next(&mut self) -> Result<Option<(Nibbles, BranchNodeCompact)>, DatabaseError> {
319        #[cfg(debug_assertions)]
320        {
321            debug_assert!(self.seeked, "Cursor must be seek'd before next is called");
322        }
323
324        // A `last_key` of `None` indicates that the cursor is exhausted.
325        let Some(last_key) = self.last_key else {
326            return Ok(None);
327        };
328
329        // If either cursor is currently pointing to the last entry which was returned then consume
330        // that entry so that `choose_next_entry` is looking at the subsequent one.
331        if let Some((key, _)) = self.in_memory_cursor.current() &&
332            key == &last_key
333        {
334            self.in_memory_cursor.first_after(&last_key);
335        }
336
337        if matches!(self.db_cursor_state, DbCursorState::NeedsPosition) {
338            self.cursor_seek(last_key)?;
339        }
340
341        if let Some((key, _)) = self.db_cursor_state.entry() &&
342            key == &last_key
343        {
344            self.cursor_next()?;
345        }
346
347        let entry = self.choose_next_entry()?;
348        self.set_last_key(&entry);
349        Ok(entry)
350    }
351
352    fn current(&mut self) -> Result<Option<Nibbles>, DatabaseError> {
353        match &self.last_key {
354            Some(key) => Ok(Some(*key)),
355            None => self.get_cursor_mut().current(),
356        }
357    }
358
359    fn reset(&mut self) {
360        self.cursor.reset();
361        self.in_memory_cursor.reset();
362
363        self.db_cursor_state = DbCursorState::NeedsPosition;
364        self.last_key = None;
365        #[cfg(debug_assertions)]
366        {
367            self.seeked = false;
368        }
369    }
370}
371
372impl<C: TrieStorageCursor> TrieStorageCursor for InMemoryTrieCursor<'_, C> {
373    fn set_hashed_address(&mut self, hashed_address: B256) {
374        self.reset();
375        self.cursor.set_hashed_address(hashed_address);
376        let in_memory_cursor = Self::get_storage_overlay(self.trie_updates, hashed_address);
377        self.in_memory_cursor = in_memory_cursor;
378        self.db_cursor_state = DbCursorState::NeedsPosition;
379    }
380}
381
382#[cfg(test)]
383mod tests {
384    use super::*;
385    use crate::trie_cursor::mock::MockTrieCursor;
386    use parking_lot::Mutex;
387    use std::{collections::BTreeMap, sync::Arc};
388
389    #[derive(Debug)]
390    struct InMemoryTrieCursorTestCase {
391        db_nodes: Vec<(Nibbles, BranchNodeCompact)>,
392        in_memory_nodes: Vec<(Nibbles, Option<BranchNodeCompact>)>,
393        expected_results: Vec<(Nibbles, BranchNodeCompact)>,
394    }
395
396    fn execute_test(test_case: InMemoryTrieCursorTestCase) {
397        let db_nodes_map: BTreeMap<Nibbles, BranchNodeCompact> =
398            test_case.db_nodes.into_iter().collect();
399        let db_nodes_arc = Arc::new(db_nodes_map);
400        let visited_keys = Arc::new(Mutex::new(Vec::new()));
401        let mock_cursor = MockTrieCursor::new(db_nodes_arc, visited_keys);
402
403        let trie_updates = TrieUpdatesSorted::new(test_case.in_memory_nodes, Default::default());
404        let mut cursor = InMemoryTrieCursor::new_account(mock_cursor, &trie_updates);
405
406        let mut results = Vec::new();
407
408        if let Some(first_expected) = test_case.expected_results.first() &&
409            let Ok(Some(entry)) = cursor.seek(first_expected.0)
410        {
411            results.push(entry);
412        }
413
414        if !test_case.expected_results.is_empty() {
415            while let Ok(Some(entry)) = cursor.next() {
416                results.push(entry);
417            }
418        }
419
420        assert_eq!(
421            results, test_case.expected_results,
422            "Results mismatch.\nGot: {:?}\nExpected: {:?}",
423            results, test_case.expected_results
424        );
425    }
426
427    #[test]
428    fn test_empty_db_and_memory() {
429        let test_case = InMemoryTrieCursorTestCase {
430            db_nodes: vec![],
431            in_memory_nodes: vec![],
432            expected_results: vec![],
433        };
434        execute_test(test_case);
435    }
436
437    #[test]
438    fn test_only_db_nodes() {
439        let db_nodes = vec![
440            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0011, 0b0001, 0, vec![], None)),
441            (Nibbles::from_nibbles([0x2]), BranchNodeCompact::new(0b0011, 0b0010, 0, vec![], None)),
442            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
443        ];
444
445        let test_case = InMemoryTrieCursorTestCase {
446            db_nodes: db_nodes.clone(),
447            in_memory_nodes: vec![],
448            expected_results: db_nodes,
449        };
450        execute_test(test_case);
451    }
452
453    #[test]
454    fn test_only_in_memory_nodes() {
455        let in_memory_nodes = vec![
456            (
457                Nibbles::from_nibbles([0x1]),
458                Some(BranchNodeCompact::new(0b0011, 0b0001, 0, vec![], None)),
459            ),
460            (
461                Nibbles::from_nibbles([0x2]),
462                Some(BranchNodeCompact::new(0b0011, 0b0010, 0, vec![], None)),
463            ),
464            (
465                Nibbles::from_nibbles([0x3]),
466                Some(BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
467            ),
468        ];
469
470        let expected_results: Vec<(Nibbles, BranchNodeCompact)> = in_memory_nodes
471            .iter()
472            .filter_map(|(k, v)| v.as_ref().map(|node| (*k, node.clone())))
473            .collect();
474
475        let test_case =
476            InMemoryTrieCursorTestCase { db_nodes: vec![], in_memory_nodes, expected_results };
477        execute_test(test_case);
478    }
479
480    #[test]
481    fn test_in_memory_overwrites_db() {
482        let db_nodes = vec![
483            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0011, 0b0001, 0, vec![], None)),
484            (Nibbles::from_nibbles([0x2]), BranchNodeCompact::new(0b0011, 0b0010, 0, vec![], None)),
485        ];
486
487        let in_memory_nodes = vec![
488            (
489                Nibbles::from_nibbles([0x1]),
490                Some(BranchNodeCompact::new(0b1111, 0b1111, 0, vec![], None)),
491            ),
492            (
493                Nibbles::from_nibbles([0x3]),
494                Some(BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
495            ),
496        ];
497
498        let expected_results = vec![
499            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b1111, 0b1111, 0, vec![], None)),
500            (Nibbles::from_nibbles([0x2]), BranchNodeCompact::new(0b0011, 0b0010, 0, vec![], None)),
501            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
502        ];
503
504        let test_case = InMemoryTrieCursorTestCase { db_nodes, in_memory_nodes, expected_results };
505        execute_test(test_case);
506    }
507
508    #[test]
509    fn test_in_memory_deletes_db_nodes() {
510        let db_nodes = vec![
511            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0011, 0b0001, 0, vec![], None)),
512            (Nibbles::from_nibbles([0x2]), BranchNodeCompact::new(0b0011, 0b0010, 0, vec![], None)),
513            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
514        ];
515
516        let in_memory_nodes = vec![(Nibbles::from_nibbles([0x2]), None)];
517
518        let expected_results = vec![
519            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0011, 0b0001, 0, vec![], None)),
520            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
521        ];
522
523        let test_case = InMemoryTrieCursorTestCase { db_nodes, in_memory_nodes, expected_results };
524        execute_test(test_case);
525    }
526
527    #[test]
528    fn test_complex_interleaving() {
529        let db_nodes = vec![
530            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0001, 0b0001, 0, vec![], None)),
531            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
532            (Nibbles::from_nibbles([0x5]), BranchNodeCompact::new(0b0101, 0b0101, 0, vec![], None)),
533            (Nibbles::from_nibbles([0x7]), BranchNodeCompact::new(0b0111, 0b0111, 0, vec![], None)),
534        ];
535
536        let in_memory_nodes = vec![
537            (
538                Nibbles::from_nibbles([0x2]),
539                Some(BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)),
540            ),
541            (Nibbles::from_nibbles([0x3]), None),
542            (
543                Nibbles::from_nibbles([0x4]),
544                Some(BranchNodeCompact::new(0b0100, 0b0100, 0, vec![], None)),
545            ),
546            (
547                Nibbles::from_nibbles([0x6]),
548                Some(BranchNodeCompact::new(0b0110, 0b0110, 0, vec![], None)),
549            ),
550            (Nibbles::from_nibbles([0x7]), None),
551            (
552                Nibbles::from_nibbles([0x8]),
553                Some(BranchNodeCompact::new(0b1000, 0b1000, 0, vec![], None)),
554            ),
555        ];
556
557        let expected_results = vec![
558            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0001, 0b0001, 0, vec![], None)),
559            (Nibbles::from_nibbles([0x2]), BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)),
560            (Nibbles::from_nibbles([0x4]), BranchNodeCompact::new(0b0100, 0b0100, 0, vec![], None)),
561            (Nibbles::from_nibbles([0x5]), BranchNodeCompact::new(0b0101, 0b0101, 0, vec![], None)),
562            (Nibbles::from_nibbles([0x6]), BranchNodeCompact::new(0b0110, 0b0110, 0, vec![], None)),
563            (Nibbles::from_nibbles([0x8]), BranchNodeCompact::new(0b1000, 0b1000, 0, vec![], None)),
564        ];
565
566        let test_case = InMemoryTrieCursorTestCase { db_nodes, in_memory_nodes, expected_results };
567        execute_test(test_case);
568    }
569
570    #[test]
571    fn test_seek_exact() {
572        let db_nodes = vec![
573            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0001, 0b0001, 0, vec![], None)),
574            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
575        ];
576
577        let in_memory_nodes = vec![(
578            Nibbles::from_nibbles([0x2]),
579            Some(BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)),
580        )];
581
582        let db_nodes_map: BTreeMap<Nibbles, BranchNodeCompact> = db_nodes.into_iter().collect();
583        let db_nodes_arc = Arc::new(db_nodes_map);
584        let visited_keys = Arc::new(Mutex::new(Vec::new()));
585        let mock_cursor = MockTrieCursor::new(db_nodes_arc, visited_keys.clone());
586
587        let trie_updates = TrieUpdatesSorted::new(in_memory_nodes, Default::default());
588        let mut cursor = InMemoryTrieCursor::new_account(mock_cursor, &trie_updates);
589
590        let result = cursor.seek_exact(Nibbles::from_nibbles([0x2])).unwrap();
591        assert_eq!(
592            result,
593            Some((
594                Nibbles::from_nibbles([0x2]),
595                BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)
596            ))
597        );
598        assert!(visited_keys.lock().is_empty(), "exact overlay hit should not touch the DB cursor");
599
600        let result = cursor.seek_exact(Nibbles::from_nibbles([0x3])).unwrap();
601        assert_eq!(
602            result,
603            Some((
604                Nibbles::from_nibbles([0x3]),
605                BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)
606            ))
607        );
608
609        let result = cursor.seek_exact(Nibbles::from_nibbles([0x4])).unwrap();
610        assert_eq!(result, None);
611    }
612
613    #[test]
614    fn test_seek_overlay_exact_hit_does_not_touch_db_until_next() {
615        let db_nodes = vec![
616            (Nibbles::from_nibbles([0x2]), BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)),
617            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
618        ];
619
620        let in_memory_nodes = vec![(
621            Nibbles::from_nibbles([0x2]),
622            Some(BranchNodeCompact::new(0b1111, 0b1111, 0, vec![], None)),
623        )];
624
625        let db_nodes_map: BTreeMap<Nibbles, BranchNodeCompact> = db_nodes.into_iter().collect();
626        let db_nodes_arc = Arc::new(db_nodes_map);
627        let visited_keys = Arc::new(Mutex::new(Vec::new()));
628        let mock_cursor = MockTrieCursor::new(db_nodes_arc, visited_keys.clone());
629
630        let trie_updates = TrieUpdatesSorted::new(in_memory_nodes, Default::default());
631        let mut cursor = InMemoryTrieCursor::new_account(mock_cursor, &trie_updates);
632
633        let result = cursor.seek(Nibbles::from_nibbles([0x2])).unwrap();
634        assert_eq!(
635            result,
636            Some((
637                Nibbles::from_nibbles([0x2]),
638                BranchNodeCompact::new(0b1111, 0b1111, 0, vec![], None)
639            ))
640        );
641        assert!(visited_keys.lock().is_empty(), "exact overlay hit should not touch the DB cursor");
642
643        let result = cursor.next().unwrap();
644        assert_eq!(
645            result,
646            Some((
647                Nibbles::from_nibbles([0x3]),
648                BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)
649            ))
650        );
651        assert!(!visited_keys.lock().is_empty(), "next should lazily position the DB cursor");
652    }
653
654    #[test]
655    fn test_seek_overlay_exact_hit_repositions_stale_db_on_next() {
656        let db_nodes = vec![
657            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(0b0001, 0b0001, 0, vec![], None)),
658            (Nibbles::from_nibbles([0x3]), BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
659        ];
660
661        let in_memory_nodes = vec![(
662            Nibbles::from_nibbles([0x2]),
663            Some(BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)),
664        )];
665
666        let db_nodes_map: BTreeMap<Nibbles, BranchNodeCompact> = db_nodes.into_iter().collect();
667        let db_nodes_arc = Arc::new(db_nodes_map);
668        let visited_keys = Arc::new(Mutex::new(Vec::new()));
669        let mock_cursor = MockTrieCursor::new(db_nodes_arc, visited_keys.clone());
670
671        let trie_updates = TrieUpdatesSorted::new(in_memory_nodes, Default::default());
672        let mut cursor = InMemoryTrieCursor::new_account(mock_cursor, &trie_updates);
673
674        let result = cursor.seek(Nibbles::from_nibbles([0x1])).unwrap();
675        assert_eq!(
676            result,
677            Some((
678                Nibbles::from_nibbles([0x1]),
679                BranchNodeCompact::new(0b0001, 0b0001, 0, vec![], None)
680            ))
681        );
682        assert_eq!(visited_keys.lock().len(), 1);
683
684        let result = cursor.seek(Nibbles::from_nibbles([0x2])).unwrap();
685        assert_eq!(
686            result,
687            Some((
688                Nibbles::from_nibbles([0x2]),
689                BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)
690            ))
691        );
692        assert_eq!(visited_keys.lock().len(), 1, "exact overlay hit should not seek the DB");
693
694        let result = cursor.next().unwrap();
695        assert_eq!(
696            result,
697            Some((
698                Nibbles::from_nibbles([0x3]),
699                BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)
700            ))
701        );
702    }
703
704    #[test]
705    fn test_multiple_consecutive_deletes() {
706        let db_nodes: Vec<(Nibbles, BranchNodeCompact)> = (1..=10)
707            .map(|i| {
708                (
709                    Nibbles::from_nibbles([i]),
710                    BranchNodeCompact::new(i as u16, i as u16, 0, vec![], None),
711                )
712            })
713            .collect();
714
715        let in_memory_nodes = vec![
716            (Nibbles::from_nibbles([0x3]), None),
717            (Nibbles::from_nibbles([0x4]), None),
718            (Nibbles::from_nibbles([0x5]), None),
719            (Nibbles::from_nibbles([0x6]), None),
720        ];
721
722        let expected_results = vec![
723            (Nibbles::from_nibbles([0x1]), BranchNodeCompact::new(1, 1, 0, vec![], None)),
724            (Nibbles::from_nibbles([0x2]), BranchNodeCompact::new(2, 2, 0, vec![], None)),
725            (Nibbles::from_nibbles([0x7]), BranchNodeCompact::new(7, 7, 0, vec![], None)),
726            (Nibbles::from_nibbles([0x8]), BranchNodeCompact::new(8, 8, 0, vec![], None)),
727            (Nibbles::from_nibbles([0x9]), BranchNodeCompact::new(9, 9, 0, vec![], None)),
728            (Nibbles::from_nibbles([0xa]), BranchNodeCompact::new(10, 10, 0, vec![], None)),
729        ];
730
731        let test_case = InMemoryTrieCursorTestCase { db_nodes, in_memory_nodes, expected_results };
732        execute_test(test_case);
733    }
734
735    #[test]
736    fn test_empty_db_with_in_memory_deletes() {
737        let in_memory_nodes = vec![
738            (Nibbles::from_nibbles([0x1]), None),
739            (
740                Nibbles::from_nibbles([0x2]),
741                Some(BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None)),
742            ),
743            (Nibbles::from_nibbles([0x3]), None),
744        ];
745
746        let expected_results = vec![(
747            Nibbles::from_nibbles([0x2]),
748            BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None),
749        )];
750
751        let test_case =
752            InMemoryTrieCursorTestCase { db_nodes: vec![], in_memory_nodes, expected_results };
753        execute_test(test_case);
754    }
755
756    #[test]
757    fn test_current_key_tracking() {
758        let db_nodes = vec![(
759            Nibbles::from_nibbles([0x2]),
760            BranchNodeCompact::new(0b0010, 0b0010, 0, vec![], None),
761        )];
762
763        let in_memory_nodes = vec![
764            (
765                Nibbles::from_nibbles([0x1]),
766                Some(BranchNodeCompact::new(0b0001, 0b0001, 0, vec![], None)),
767            ),
768            (
769                Nibbles::from_nibbles([0x3]),
770                Some(BranchNodeCompact::new(0b0011, 0b0011, 0, vec![], None)),
771            ),
772        ];
773
774        let db_nodes_map: BTreeMap<Nibbles, BranchNodeCompact> = db_nodes.into_iter().collect();
775        let db_nodes_arc = Arc::new(db_nodes_map);
776        let visited_keys = Arc::new(Mutex::new(Vec::new()));
777        let mock_cursor = MockTrieCursor::new(db_nodes_arc, visited_keys);
778
779        let trie_updates = TrieUpdatesSorted::new(in_memory_nodes, Default::default());
780        let mut cursor = InMemoryTrieCursor::new_account(mock_cursor, &trie_updates);
781
782        assert_eq!(cursor.current().unwrap(), None);
783
784        cursor.seek(Nibbles::from_nibbles([0x1])).unwrap();
785        assert_eq!(cursor.current().unwrap(), Some(Nibbles::from_nibbles([0x1])));
786
787        cursor.next().unwrap();
788        assert_eq!(cursor.current().unwrap(), Some(Nibbles::from_nibbles([0x2])));
789
790        cursor.next().unwrap();
791        assert_eq!(cursor.current().unwrap(), Some(Nibbles::from_nibbles([0x3])));
792    }
793
794    #[test]
795    fn test_all_storage_nodes_deleted_exact_keys() {
796        use tracing::debug;
797        reth_tracing::init_test_tracing();
798
799        // This test reproduces an edge case where:
800        // - cursor is available
801        // - All in-memory entries are deletions (None values)
802        // - Database has corresponding entries
803        // - Expected: NO leaves should be returned (all deleted)
804
805        // Generate 42 trie node entries with keys distributed across the keyspace
806        let mut db_nodes: Vec<(Nibbles, BranchNodeCompact)> = (0..10)
807            .map(|i| {
808                let key_bytes = vec![(i * 6) as u8, i as u8]; // Spread keys across keyspace
809                let nibbles = Nibbles::unpack(key_bytes);
810                (nibbles, BranchNodeCompact::new(i as u16, i as u16, 0, vec![], None))
811            })
812            .collect();
813
814        db_nodes.sort_by_key(|(key, _)| *key);
815        db_nodes.dedup_by_key(|(key, _)| *key);
816
817        for (key, _) in &db_nodes {
818            debug!("node at {key:?}");
819        }
820
821        // Create in-memory entries with same keys but all None values (deletions)
822        let in_memory_nodes: Vec<(Nibbles, Option<BranchNodeCompact>)> =
823            db_nodes.iter().map(|(key, _)| (*key, None)).collect();
824
825        let db_nodes_map: BTreeMap<Nibbles, BranchNodeCompact> = db_nodes.into_iter().collect();
826        let db_nodes_arc = Arc::new(db_nodes_map);
827        let visited_keys = Arc::new(Mutex::new(Vec::new()));
828        let mock_cursor = MockTrieCursor::new(db_nodes_arc, visited_keys);
829
830        let trie_updates = TrieUpdatesSorted::new(in_memory_nodes, Default::default());
831        let mut cursor = InMemoryTrieCursor::new_account(mock_cursor, &trie_updates);
832
833        // Seek to beginning should return None (all nodes are deleted)
834        tracing::debug!("seeking to 0x");
835        let result = cursor.seek(Nibbles::default()).unwrap();
836        assert_eq!(
837            result, None,
838            "Expected no entries when all nodes are deleted, but got {:?}",
839            result
840        );
841
842        // Test seek operations at various positions - all should return None
843        let seek_keys = vec![
844            Nibbles::unpack([0x00]),
845            Nibbles::unpack([0x5d]),
846            Nibbles::unpack([0x5e]),
847            Nibbles::unpack([0x5f]),
848            Nibbles::unpack([0xc2]),
849            Nibbles::unpack([0xc5]),
850            Nibbles::unpack([0xc9]),
851            Nibbles::unpack([0xf0]),
852        ];
853
854        for seek_key in seek_keys {
855            tracing::debug!("seeking to {seek_key:?}");
856            let result = cursor.seek(seek_key).unwrap();
857            assert_eq!(
858                result, None,
859                "Expected None when seeking to {:?} but got {:?}",
860                seek_key, result
861            );
862        }
863
864        // next() should also always return None
865        let result = cursor.next().unwrap();
866        assert_eq!(result, None, "Expected None from next() but got {:?}", result);
867    }
868
869    mod proptest_tests {
870        use super::*;
871        use itertools::Itertools;
872        use proptest::prelude::*;
873
874        /// Merge `db_nodes` with `in_memory_nodes`, applying the in-memory overlay.
875        /// This properly handles deletions (None values in `in_memory_nodes`).
876        fn merge_with_overlay(
877            db_nodes: Vec<(Nibbles, BranchNodeCompact)>,
878            in_memory_nodes: Vec<(Nibbles, Option<BranchNodeCompact>)>,
879        ) -> Vec<(Nibbles, BranchNodeCompact)> {
880            db_nodes
881                .into_iter()
882                .merge_join_by(in_memory_nodes, |db_entry, mem_entry| db_entry.0.cmp(&mem_entry.0))
883                .filter_map(|entry| match entry {
884                    // Only in db: keep it
885                    itertools::EitherOrBoth::Left((key, node)) => Some((key, node)),
886                    // Only in memory: keep if not a deletion
887                    itertools::EitherOrBoth::Right((key, node_opt)) => {
888                        node_opt.map(|node| (key, node))
889                    }
890                    // In both: memory takes precedence (keep if not a deletion)
891                    itertools::EitherOrBoth::Both(_, (key, node_opt)) => {
892                        node_opt.map(|node| (key, node))
893                    }
894                })
895                .collect()
896        }
897
898        /// Generate a strategy for a `BranchNodeCompact` with simplified parameters.
899        /// The constraints are:
900        /// - `tree_mask` must be a subset of `state_mask`
901        /// - `hash_mask` must be a subset of `state_mask`
902        /// - `hash_mask.count_ones()` must equal `hashes.len()`
903        ///
904        /// To keep it simple, we use an empty hashes vec and `hash_mask` of 0.
905        fn branch_node_strategy() -> impl Strategy<Value = BranchNodeCompact> {
906            any::<u16>()
907                .prop_flat_map(|state_mask| {
908                    let tree_mask_strategy = any::<u16>().prop_map(move |tree| tree & state_mask);
909                    (Just(state_mask), tree_mask_strategy)
910                })
911                .prop_map(|(state_mask, tree_mask)| {
912                    BranchNodeCompact::new(state_mask, tree_mask, 0, vec![], None)
913                })
914        }
915
916        /// Generate a sorted vector of (Nibbles, `BranchNodeCompact`) entries
917        fn sorted_db_nodes_strategy() -> impl Strategy<Value = Vec<(Nibbles, BranchNodeCompact)>> {
918            prop::collection::vec(
919                (prop::collection::vec(any::<u8>(), 0..2), branch_node_strategy()),
920                0..20,
921            )
922            .prop_map(|entries| {
923                // Convert Vec<u8> to Nibbles and sort
924                let mut result: Vec<(Nibbles, BranchNodeCompact)> = entries
925                    .into_iter()
926                    .map(|(bytes, node)| (Nibbles::from_nibbles_unchecked(bytes), node))
927                    .collect();
928                result.sort_by_key(|a| a.0);
929                result.dedup_by(|a, b| a.0 == b.0);
930                result
931            })
932        }
933
934        /// Generate a sorted vector of (Nibbles, Option<BranchNodeCompact>) entries
935        fn sorted_in_memory_nodes_strategy(
936        ) -> impl Strategy<Value = Vec<(Nibbles, Option<BranchNodeCompact>)>> {
937            prop::collection::vec(
938                (
939                    prop::collection::vec(any::<u8>(), 0..2),
940                    prop::option::of(branch_node_strategy()),
941                ),
942                0..20,
943            )
944            .prop_map(|entries| {
945                // Convert Vec<u8> to Nibbles and sort
946                let mut result: Vec<(Nibbles, Option<BranchNodeCompact>)> = entries
947                    .into_iter()
948                    .map(|(bytes, node)| (Nibbles::from_nibbles_unchecked(bytes), node))
949                    .collect();
950                result.sort_by_key(|a| a.0);
951                result.dedup_by(|a, b| a.0 == b.0);
952                result
953            })
954        }
955
956        proptest! {
957            #![proptest_config(ProptestConfig::with_cases(10000))]
958
959            #[test]
960            fn proptest_in_memory_trie_cursor(
961                db_nodes in sorted_db_nodes_strategy(),
962                in_memory_nodes in sorted_in_memory_nodes_strategy(),
963                op_choices in prop::collection::vec(any::<u8>(), 10..500),
964            ) {
965                reth_tracing::init_test_tracing();
966                use tracing::debug;
967
968                debug!(
969                    db_paths=?db_nodes.iter().map(|(k, _)| k).collect::<Vec<_>>(),
970                    in_mem_nodes=?in_memory_nodes.iter().map(|(k, v)| (k, v.is_some())).collect::<Vec<_>>(),
971                    num_op_choices=?op_choices.len(),
972                    "Starting proptest!",
973                );
974
975                // Create the expected results by merging the two sorted vectors,
976                // properly handling deletions (None values in in_memory_nodes)
977                let expected_combined = merge_with_overlay(db_nodes.clone(), in_memory_nodes.clone());
978
979                // Collect all keys for operation generation
980                let all_keys: Vec<Nibbles> = expected_combined.iter().map(|(k, _)| *k).collect();
981
982                // Create a control cursor using the combined result with a mock cursor
983                let control_db_map: BTreeMap<Nibbles, BranchNodeCompact> =
984                    expected_combined.into_iter().collect();
985                let control_db_arc = Arc::new(control_db_map);
986                let control_visited_keys = Arc::new(Mutex::new(Vec::new()));
987                let mut control_cursor = MockTrieCursor::new(control_db_arc, control_visited_keys);
988
989                // Create the InMemoryTrieCursor being tested
990                let db_nodes_map: BTreeMap<Nibbles, BranchNodeCompact> =
991                    db_nodes.into_iter().collect();
992                let db_nodes_arc = Arc::new(db_nodes_map);
993                let visited_keys = Arc::new(Mutex::new(Vec::new()));
994                let mock_cursor = MockTrieCursor::new(db_nodes_arc, visited_keys);
995                let trie_updates = TrieUpdatesSorted::new(in_memory_nodes, Default::default());
996                let mut test_cursor = InMemoryTrieCursor::new_account(mock_cursor, &trie_updates);
997
998                // Test: seek to the beginning first
999                let control_first = control_cursor.seek(Nibbles::default()).unwrap();
1000                let test_first = test_cursor.seek(Nibbles::default()).unwrap();
1001                debug!(
1002                    control=?control_first.as_ref().map(|(k, _)| k),
1003                    test=?test_first.as_ref().map(|(k, _)| k),
1004                    "Initial seek returned",
1005                );
1006                assert_eq!(control_first, test_first, "Initial seek mismatch");
1007
1008                // If both cursors returned None, nothing to test
1009                if control_first.is_none() && test_first.is_none() {
1010                    return Ok(());
1011                }
1012
1013                // Track the last key returned from the cursor
1014                let mut last_returned_key = control_first.as_ref().map(|(k, _)| *k);
1015
1016                // Execute a sequence of random operations
1017                for choice in op_choices {
1018                    let op_type = choice % 3;
1019
1020                    match op_type {
1021                        0 => {
1022                            // Next operation
1023                            let control_result = control_cursor.next().unwrap();
1024                            let test_result = test_cursor.next().unwrap();
1025                            debug!(
1026                                control=?control_result.as_ref().map(|(k, _)| k),
1027                                test=?test_result.as_ref().map(|(k, _)| k),
1028                                "Next returned",
1029                            );
1030                            assert_eq!(control_result, test_result, "Next operation mismatch");
1031
1032                            last_returned_key = control_result.as_ref().map(|(k, _)| *k);
1033
1034                            // Stop if both cursors are exhausted
1035                            if control_result.is_none() && test_result.is_none() {
1036                                break;
1037                            }
1038                        }
1039                        1 => {
1040                            // Seek operation - choose a key >= last_returned_key
1041                            if all_keys.is_empty() {
1042                                continue;
1043                            }
1044
1045                            let valid_keys: Vec<_> = all_keys
1046                                .iter()
1047                                .filter(|k| last_returned_key.is_none_or(|last| **k >= last))
1048                                .collect();
1049
1050                            if valid_keys.is_empty() {
1051                                continue;
1052                            }
1053
1054                            let key = *valid_keys[choice as usize % valid_keys.len()];
1055
1056                            let control_result = control_cursor.seek(key).unwrap();
1057                            let test_result = test_cursor.seek(key).unwrap();
1058                            debug!(
1059                                control=?control_result.as_ref().map(|(k, _)| k),
1060                                test=?test_result.as_ref().map(|(k, _)| k),
1061                                ?key,
1062                                "Seek returned",
1063                            );
1064                            assert_eq!(control_result, test_result, "Seek operation mismatch for key {:?}", key);
1065
1066                            last_returned_key = control_result.as_ref().map(|(k, _)| *k);
1067
1068                            // Stop if both cursors are exhausted
1069                            if control_result.is_none() && test_result.is_none() {
1070                                break;
1071                            }
1072                        }
1073                        _ => {
1074                            // SeekExact operation - choose a key >= last_returned_key
1075                            if all_keys.is_empty() {
1076                                continue;
1077                            }
1078
1079                            let valid_keys: Vec<_> = all_keys
1080                                .iter()
1081                                .filter(|k| last_returned_key.is_none_or(|last| **k >= last))
1082                                .collect();
1083
1084                            if valid_keys.is_empty() {
1085                                continue;
1086                            }
1087
1088                            let key = *valid_keys[choice as usize  % valid_keys.len()];
1089
1090                            let control_result = control_cursor.seek_exact(key).unwrap();
1091                            let test_result = test_cursor.seek_exact(key).unwrap();
1092                            debug!(
1093                                control=?control_result.as_ref().map(|(k, _)| k),
1094                                test=?test_result.as_ref().map(|(k, _)| k),
1095                                ?key,
1096                                "SeekExact returned",
1097                            );
1098                            assert_eq!(control_result, test_result, "SeekExact operation mismatch for key {:?}", key);
1099
1100                            // seek_exact updates the last_key internally but only if it found something
1101                            last_returned_key = control_result.as_ref().map(|(k, _)| *k);
1102                        }
1103                    }
1104                }
1105            }
1106        }
1107    }
1108}