Skip to main content

reth_stages/stages/
merkle.rs

1use alloy_consensus::BlockHeader;
2use alloy_primitives::{BlockNumber, Sealable, B256};
3use reth_codecs::Compact;
4use reth_consensus::ConsensusError;
5use reth_db_api::{
6    tables,
7    transaction::{DbTx, DbTxMut},
8};
9use reth_primitives_traits::{GotExpected, SealedHeader};
10use reth_provider::{
11    ChangeSetReader, DBProvider, HeaderProvider, ProviderError, StageCheckpointReader,
12    StageCheckpointWriter, StatsReader, StorageChangeSetReader, StorageSettingsCache, TrieWriter,
13};
14use reth_stages_api::{
15    BlockErrorKind, EntitiesCheckpoint, ExecInput, ExecOutput, MerkleCheckpoint, Stage,
16    StageCheckpoint, StageError, StageId, StorageRootMerkleCheckpoint, UnwindInput, UnwindOutput,
17};
18use reth_trie::{IntermediateStateRootState, StateRoot, StateRootProgress, StoredSubNode};
19use reth_trie_db::DatabaseStateRoot;
20
21use std::fmt::Debug;
22
23type DbStateRoot<'a, TX, A> = StateRoot<
24    reth_trie_db::DatabaseTrieCursorFactory<&'a TX, A>,
25    reth_trie_db::DatabaseHashedCursorFactory<&'a TX>,
26>;
27use tracing::*;
28
29// TODO: automate the process outlined below so the user can just send in a debugging package
30/// The error message that we include in invalid state root errors to tell users what information
31/// they should include in a bug report, since true state root errors can be impossible to debug
32/// with just basic logs.
33pub const INVALID_STATE_ROOT_ERROR_MESSAGE: &str = r#"
34Invalid state root error on stage verification!
35This is an error that likely requires a report to the reth team with additional information.
36Please include the following information in your report:
37 * This error message
38 * The state root of the block that was rejected
39 * The output of `reth db stats --checksum` from the database that was being used. This will take a long time to run!
40 * 50-100 lines of logs before and after the first occurrence of the log message with the state root of the block that was rejected.
41 * The debug logs from __the same time period__. To find the default location for these logs, run:
42   `reth --help | grep -A 4 'log.file.directory'`
43
44Once you have this information, please submit a github issue at https://github.com/paradigmxyz/reth/issues/new
45"#;
46
47/// The default threshold (in number of blocks) for switching from incremental trie building
48/// of changes to whole rebuild.
49pub const MERKLE_STAGE_DEFAULT_REBUILD_THRESHOLD: u64 = 100_000;
50
51/// The default threshold (in number of blocks) to run the stage in incremental mode. The
52/// incremental mode will calculate the state root for a large range of blocks by calculating the
53/// new state root for this many blocks, in batches, repeating until we reach the desired block
54/// number.
55pub const MERKLE_STAGE_DEFAULT_INCREMENTAL_THRESHOLD: u64 = 7_000;
56
57/// The merkle hashing stage uses input from
58/// [`AccountHashingStage`][crate::stages::AccountHashingStage] and
59/// [`StorageHashingStage`][crate::stages::StorageHashingStage] to calculate intermediate hashes
60/// and state roots.
61///
62/// This stage should be run with the above two stages, otherwise it is a no-op.
63///
64/// This stage is split in two: one for calculating hashes and one for unwinding.
65///
66/// When run in execution, it's going to be executed AFTER the hashing stages, to generate
67/// the state root. When run in unwind mode, it's going to be executed BEFORE the hashing stages,
68/// so that it unwinds the intermediate hashes based on the unwound hashed state from the hashing
69/// stages. The order of these two variants is important. The unwind variant should be added to the
70/// pipeline before the execution variant.
71///
72/// An example pipeline to only hash state would be:
73///
74/// - [`MerkleStage::Unwind`]
75/// - [`AccountHashingStage`][crate::stages::AccountHashingStage]
76/// - [`StorageHashingStage`][crate::stages::StorageHashingStage]
77/// - [`MerkleStage::Execution`]
78#[derive(Debug, Clone)]
79pub enum MerkleStage {
80    /// The execution portion of the merkle stage.
81    Execution {
82        // TODO: make struct for holding incremental settings, for code reuse between `Execution`
83        // variant and `Both`
84        /// The threshold (in number of blocks) for switching from incremental trie building
85        /// of changes to whole rebuild.
86        rebuild_threshold: u64,
87        /// The threshold (in number of blocks) to run the stage in incremental mode. The
88        /// incremental mode will calculate the state root by calculating the new state root for
89        /// some number of blocks, repeating until we reach the desired block number.
90        incremental_threshold: u64,
91    },
92    /// The unwind portion of the merkle stage.
93    Unwind {
94        /// Whether every child of a changed branch path should be walked.
95        walk_all_changed_branch_children: bool,
96    },
97    /// Able to execute and unwind. Used for tests
98    #[cfg(any(test, feature = "test-utils"))]
99    Both {
100        /// The threshold (in number of blocks) for switching from incremental trie building
101        /// of changes to whole rebuild.
102        rebuild_threshold: u64,
103        /// The threshold (in number of blocks) to run the stage in incremental mode. The
104        /// incremental mode will calculate the state root by calculating the new state root for
105        /// some number of blocks, repeating until we reach the desired block number.
106        incremental_threshold: u64,
107    },
108}
109
110impl MerkleStage {
111    /// Stage default for the [`MerkleStage::Execution`].
112    pub const fn default_execution() -> Self {
113        Self::Execution {
114            rebuild_threshold: MERKLE_STAGE_DEFAULT_REBUILD_THRESHOLD,
115            incremental_threshold: MERKLE_STAGE_DEFAULT_INCREMENTAL_THRESHOLD,
116        }
117    }
118
119    /// Stage default for the [`MerkleStage::Unwind`].
120    pub const fn default_unwind() -> Self {
121        Self::new_unwind(false)
122    }
123
124    /// Create a new instance of [`MerkleStage::Unwind`].
125    pub const fn new_unwind(walk_all_changed_branch_children: bool) -> Self {
126        Self::Unwind { walk_all_changed_branch_children }
127    }
128
129    /// Create new instance of [`MerkleStage::Execution`].
130    pub const fn new_execution(rebuild_threshold: u64, incremental_threshold: u64) -> Self {
131        Self::Execution { rebuild_threshold, incremental_threshold }
132    }
133
134    /// Gets the hashing progress
135    pub fn get_execution_checkpoint(
136        &self,
137        provider: &impl StageCheckpointReader,
138    ) -> Result<Option<MerkleCheckpoint>, StageError> {
139        let buf =
140            provider.get_stage_checkpoint_progress(StageId::MerkleExecute)?.unwrap_or_default();
141
142        if buf.is_empty() {
143            return Ok(None)
144        }
145
146        let (checkpoint, _) = MerkleCheckpoint::from_compact(&buf, buf.len());
147        Ok(Some(checkpoint))
148    }
149
150    /// Saves the hashing progress
151    pub fn save_execution_checkpoint(
152        &self,
153        provider: &impl StageCheckpointWriter,
154        checkpoint: Option<MerkleCheckpoint>,
155    ) -> Result<(), StageError> {
156        let mut buf = vec![];
157        if let Some(checkpoint) = checkpoint {
158            debug!(
159                target: "sync::stages::merkle::exec",
160                last_account_key = ?checkpoint.last_account_key,
161                "Saving inner merkle checkpoint"
162            );
163            checkpoint.to_compact(&mut buf);
164        }
165        Ok(provider.save_stage_checkpoint_progress(StageId::MerkleExecute, buf)?)
166    }
167}
168
169impl<Provider> Stage<Provider> for MerkleStage
170where
171    Provider: DBProvider<Tx: DbTxMut>
172        + TrieWriter
173        + StatsReader
174        + HeaderProvider
175        + ChangeSetReader
176        + StorageChangeSetReader
177        + StorageSettingsCache
178        + StageCheckpointReader
179        + StageCheckpointWriter,
180{
181    /// Return the id of the stage
182    fn id(&self) -> StageId {
183        match self {
184            Self::Execution { .. } => StageId::MerkleExecute,
185            Self::Unwind { .. } => StageId::MerkleUnwind,
186            #[cfg(any(test, feature = "test-utils"))]
187            Self::Both { .. } => StageId::Other("MerkleBoth"),
188        }
189    }
190
191    /// Execute the stage.
192    fn execute(&mut self, provider: &Provider, input: ExecInput) -> Result<ExecOutput, StageError> {
193        let (threshold, incremental_threshold) = match self {
194            Self::Unwind { .. } => {
195                info!(target: "sync::stages::merkle::unwind", "Stage is always skipped");
196                return Ok(ExecOutput::done(StageCheckpoint::new(input.target())))
197            }
198            Self::Execution { rebuild_threshold, incremental_threshold } => {
199                (*rebuild_threshold, *incremental_threshold)
200            }
201            #[cfg(any(test, feature = "test-utils"))]
202            Self::Both { rebuild_threshold, incremental_threshold } => {
203                (*rebuild_threshold, *incremental_threshold)
204            }
205        };
206
207        let range = input.next_block_range();
208        let (from_block, to_block) = range.clone().into_inner();
209        let current_block_number = input.checkpoint().block_number;
210
211        let target_block = provider
212            .header_by_number(to_block)?
213            .ok_or_else(|| ProviderError::HeaderNotFound(to_block.into()))?;
214        let target_block_root = target_block.state_root();
215
216        let (trie_root, entities_checkpoint) = if range.is_empty() {
217            (target_block_root, input.checkpoint().entities_stage_checkpoint().unwrap_or_default())
218        } else if to_block - from_block > threshold || from_block == 1 {
219            let mut checkpoint = self.get_execution_checkpoint(provider)?;
220
221            // if there are more blocks than threshold it is faster to rebuild the trie
222            let mut entities_checkpoint = if let Some(checkpoint) =
223                checkpoint.as_ref().filter(|c| c.target_block == to_block)
224            {
225                debug!(
226                    target: "sync::stages::merkle::exec",
227                    current = ?current_block_number,
228                    target = ?to_block,
229                    last_account_key = ?checkpoint.last_account_key,
230                    "Continuing inner merkle checkpoint"
231                );
232
233                input.checkpoint().entities_stage_checkpoint()
234            } else {
235                debug!(
236                    target: "sync::stages::merkle::exec",
237                    current = ?current_block_number,
238                    target = ?to_block,
239                    previous_checkpoint = ?checkpoint,
240                    "Rebuilding trie"
241                );
242                // Reset the checkpoint and clear trie tables
243                checkpoint = None;
244                self.save_execution_checkpoint(provider, None)?;
245                provider.tx_ref().clear::<tables::AccountsTrie>()?;
246                provider.tx_ref().clear::<tables::StoragesTrie>()?;
247
248                None
249            }
250            .unwrap_or(EntitiesCheckpoint {
251                processed: 0,
252                total: (provider.count_entries::<tables::HashedAccounts>()? +
253                    provider.count_entries::<tables::HashedStorages>()?)
254                    as u64,
255            });
256
257            let tx = provider.tx_ref();
258            let progress = reth_trie_db::with_adapter!(provider, |A| {
259                DbStateRoot::<_, A>::from_tx(tx)
260                    .with_intermediate_state(checkpoint.map(IntermediateStateRootState::from))
261                    .root_with_progress()
262            })
263            .map_err(|e| {
264                error!(target: "sync::stages::merkle", %e, ?current_block_number, ?to_block, "State root with progress failed! {INVALID_STATE_ROOT_ERROR_MESSAGE}");
265                StageError::Fatal(Box::new(e))
266            })?;
267            match progress {
268                StateRootProgress::Progress(state, hashed_entries_walked, updates) => {
269                    provider.write_trie_updates(updates)?;
270
271                    let mut checkpoint = MerkleCheckpoint::new(
272                        to_block,
273                        state.account_root_state.last_hashed_key,
274                        state
275                            .account_root_state
276                            .walker_stack
277                            .into_iter()
278                            .map(StoredSubNode::from)
279                            .collect(),
280                        state.account_root_state.hash_builder.into(),
281                    );
282
283                    // Save storage root state if present
284                    if let Some(storage_state) = state.storage_root_state {
285                        checkpoint.storage_root_checkpoint =
286                            Some(StorageRootMerkleCheckpoint::new(
287                                storage_state.state.last_hashed_key,
288                                storage_state
289                                    .state
290                                    .walker_stack
291                                    .into_iter()
292                                    .map(StoredSubNode::from)
293                                    .collect(),
294                                storage_state.state.hash_builder.into(),
295                                storage_state.account,
296                            ));
297                    }
298                    self.save_execution_checkpoint(provider, Some(checkpoint))?;
299
300                    entities_checkpoint.processed += hashed_entries_walked as u64;
301
302                    return Ok(ExecOutput {
303                        checkpoint: input
304                            .checkpoint()
305                            .with_entities_stage_checkpoint(entities_checkpoint),
306                        done: false,
307                    })
308                }
309                StateRootProgress::Complete(root, hashed_entries_walked, updates) => {
310                    provider.write_trie_updates(updates)?;
311
312                    entities_checkpoint.processed += hashed_entries_walked as u64;
313
314                    (root, entities_checkpoint)
315                }
316            }
317        } else {
318            debug!(target: "sync::stages::merkle::exec", current = ?current_block_number, target = ?to_block, "Updating trie in chunks");
319            let mut final_root = None;
320            for start_block in range.step_by(incremental_threshold as usize) {
321                let chunk_to = std::cmp::min(start_block + incremental_threshold - 1, to_block);
322                let chunk_range = start_block..=chunk_to;
323                debug!(
324                    target: "sync::stages::merkle::exec",
325                    current = ?current_block_number,
326                    target = ?to_block,
327                    incremental_threshold,
328                    chunk_range = ?chunk_range,
329                    "Processing chunk"
330                );
331                let (root, updates) = reth_trie_db::with_adapter!(provider, |A| {
332                    DbStateRoot::<_, A>::incremental_root_with_updates(provider, chunk_range)
333                })
334                .map_err(|e| {
335                    error!(target: "sync::stages::merkle", %e, ?current_block_number, ?to_block, "Incremental state root failed! {INVALID_STATE_ROOT_ERROR_MESSAGE}");
336                    StageError::Fatal(Box::new(e))
337                })?;
338                provider.write_trie_updates(updates)?;
339                final_root = Some(root);
340            }
341
342            // if we had no final root, we must have not looped above, which should not be possible
343            let final_root = final_root.ok_or(StageError::Fatal(
344                "Incremental merkle hashing did not produce a final root".into(),
345            ))?;
346
347            let total_hashed_entries = (provider.count_entries::<tables::HashedAccounts>()? +
348                provider.count_entries::<tables::HashedStorages>()?)
349                as u64;
350
351            let entities_checkpoint = EntitiesCheckpoint {
352                // This is fine because `range` doesn't have an upper bound, so in this `else`
353                // branch we're just hashing all remaining accounts and storage slots we have in the
354                // database.
355                processed: total_hashed_entries,
356                total: total_hashed_entries,
357            };
358            // Save the checkpoint
359            (final_root, entities_checkpoint)
360        };
361
362        // Reset the checkpoint
363        self.save_execution_checkpoint(provider, None)?;
364
365        validate_state_root(trie_root, SealedHeader::seal_slow(target_block), to_block)?;
366
367        Ok(ExecOutput {
368            checkpoint: StageCheckpoint::new(to_block)
369                .with_entities_stage_checkpoint(entities_checkpoint),
370            done: true,
371        })
372    }
373
374    /// Unwind the stage.
375    fn unwind(
376        &mut self,
377        provider: &Provider,
378        input: UnwindInput,
379    ) -> Result<UnwindOutput, StageError> {
380        let tx = provider.tx_ref();
381        let range = input.unwind_block_range();
382        if matches!(self, Self::Execution { .. }) {
383            info!(target: "sync::stages::merkle::unwind", "Stage is always skipped");
384            return Ok(UnwindOutput { checkpoint: StageCheckpoint::new(input.unwind_to) })
385        }
386        let walk_all_changed_branch_children = match self {
387            Self::Unwind { walk_all_changed_branch_children } => *walk_all_changed_branch_children,
388            #[cfg(any(test, feature = "test-utils"))]
389            Self::Both { .. } => false,
390            Self::Execution { .. } => unreachable!(),
391        };
392
393        let mut entities_checkpoint =
394            input.checkpoint.entities_stage_checkpoint().unwrap_or(EntitiesCheckpoint {
395                processed: 0,
396                total: (tx.entries::<tables::HashedAccounts>()? +
397                    tx.entries::<tables::HashedStorages>()?) as u64,
398            });
399
400        if input.unwind_to == 0 {
401            tx.clear::<tables::AccountsTrie>()?;
402            tx.clear::<tables::StoragesTrie>()?;
403
404            entities_checkpoint.processed = 0;
405
406            return Ok(UnwindOutput {
407                checkpoint: StageCheckpoint::new(input.unwind_to)
408                    .with_entities_stage_checkpoint(entities_checkpoint),
409            })
410        }
411
412        // Unwind trie only if there are transitions
413        if range.is_empty() {
414            info!(target: "sync::stages::merkle::unwind", "Nothing to unwind");
415        } else {
416            let (block_root, updates) = reth_trie_db::with_adapter!(provider, |A| {
417                DbStateRoot::<_, A>::incremental_root_calculator(provider, range).and_then(
418                    |calculator| {
419                        calculator
420                            .with_walk_all_changed_branch_children(walk_all_changed_branch_children)
421                            .root_with_updates()
422                    },
423                )
424            })
425            .map_err(|e| StageError::Fatal(Box::new(e)))?;
426
427            // Validate the calculated state root
428            let target = provider
429                .header_by_number(input.unwind_to)?
430                .ok_or_else(|| ProviderError::HeaderNotFound(input.unwind_to.into()))?;
431
432            validate_state_root(block_root, SealedHeader::seal_slow(target), input.unwind_to)?;
433
434            // Validation passed, apply unwind changes to the database.
435            provider.write_trie_updates(updates)?;
436
437            // Update entities checkpoint to reflect the unwind operation
438            // Since we're unwinding, we need to recalculate the total entities at the target block
439            let accounts = tx.entries::<tables::HashedAccounts>()?;
440            let storages = tx.entries::<tables::HashedStorages>()?;
441            let total = (accounts + storages) as u64;
442            entities_checkpoint.total = total;
443            entities_checkpoint.processed = total;
444        }
445
446        Ok(UnwindOutput {
447            checkpoint: StageCheckpoint::new(input.unwind_to)
448                .with_entities_stage_checkpoint(entities_checkpoint),
449        })
450    }
451}
452
453/// Check that the computed state root matches the root in the expected header.
454#[inline]
455fn validate_state_root<H: BlockHeader + Sealable + Debug>(
456    got: B256,
457    expected: SealedHeader<H>,
458    target_block: BlockNumber,
459) -> Result<(), StageError> {
460    if got == expected.state_root() {
461        Ok(())
462    } else {
463        error!(target: "sync::stages::merkle", ?target_block, ?got, ?expected, "Failed to verify block state root! {INVALID_STATE_ROOT_ERROR_MESSAGE}");
464        Err(StageError::Block {
465            error: BlockErrorKind::Validation(ConsensusError::BodyStateRootDiff(
466                GotExpected { got, expected: expected.state_root() }.into(),
467            )),
468            block: Box::new(expected.block_with_parent()),
469        })
470    }
471}
472
473#[cfg(test)]
474mod tests {
475    use super::*;
476    use crate::test_utils::{
477        stage_test_suite_ext, ExecuteStageTestRunner, StageTestRunner, StorageKind,
478        TestRunnerError, TestStageDB, UnwindStageTestRunner,
479    };
480    use alloy_primitives::{keccak256, U256};
481    use assert_matches::assert_matches;
482    use reth_db_api::cursor::{DbCursorRO, DbCursorRW, DbDupCursorRO};
483    use reth_primitives_traits::{SealedBlock, StorageEntry};
484    use reth_provider::{providers::StaticFileWriter, StaticFileProviderFactory};
485    use reth_stages_api::StageUnitCheckpoint;
486    use reth_static_file_types::StaticFileSegment;
487    use reth_testing_utils::generators::{
488        self, random_block, random_block_range, random_changeset_range,
489        random_contract_account_range, BlockParams, BlockRangeParams,
490    };
491    use reth_trie::test_utils::{state_root, state_root_prehashed};
492    use std::collections::BTreeMap;
493
494    stage_test_suite_ext!(MerkleTestRunner, merkle);
495
496    /// Execute from genesis so as to merkelize whole state
497    #[tokio::test]
498    async fn execute_clean_merkle() {
499        let (previous_stage, stage_progress) = (500, 0);
500
501        // Set up the runner
502        let mut runner = MerkleTestRunner::default();
503        // set low threshold so we hash the whole storage
504        let input = ExecInput {
505            target: Some(previous_stage),
506            checkpoint: Some(StageCheckpoint::new(stage_progress)),
507        };
508
509        runner.seed_execution(input).expect("failed to seed execution");
510
511        let rx = runner.execute(input);
512
513        // Assert the successful result
514        let result = rx.await.unwrap();
515        assert_matches!(
516            result,
517            Ok(ExecOutput {
518                checkpoint: StageCheckpoint {
519                    block_number,
520                    stage_checkpoint: Some(StageUnitCheckpoint::Entities(EntitiesCheckpoint {
521                        processed,
522                        total
523                    }))
524                },
525                done: true
526            }) if block_number == previous_stage && processed == total &&
527                total == (
528                    runner.db.count_entries::<tables::HashedAccounts>().unwrap() +
529                    runner.db.count_entries::<tables::HashedStorages>().unwrap()
530                ) as u64
531        );
532
533        // Validate the stage execution
534        assert!(runner.validate_execution(input, result.ok()).is_ok(), "execution validation");
535    }
536
537    /// Update small trie
538    #[tokio::test]
539    async fn execute_small_merkle() {
540        let (previous_stage, stage_progress) = (2, 1);
541
542        // Set up the runner
543        let mut runner = MerkleTestRunner::default();
544        let input = ExecInput {
545            target: Some(previous_stage),
546            checkpoint: Some(StageCheckpoint::new(stage_progress)),
547        };
548
549        runner.seed_execution(input).expect("failed to seed execution");
550
551        let rx = runner.execute(input);
552
553        // Assert the successful result
554        let result = rx.await.unwrap();
555        assert_matches!(
556            result,
557            Ok(ExecOutput {
558                checkpoint: StageCheckpoint {
559                    block_number,
560                    stage_checkpoint: Some(StageUnitCheckpoint::Entities(EntitiesCheckpoint {
561                        processed,
562                        total
563                    }))
564                },
565                done: true
566            }) if block_number == previous_stage && processed == total &&
567                total == (
568                    runner.db.count_entries::<tables::HashedAccounts>().unwrap() +
569                    runner.db.count_entries::<tables::HashedStorages>().unwrap()
570                ) as u64
571        );
572
573        // Validate the stage execution
574        assert!(runner.validate_execution(input, result.ok()).is_ok(), "execution validation");
575    }
576
577    #[tokio::test]
578    async fn execute_chunked_merkle() {
579        let (previous_stage, stage_progress) = (200, 100);
580        let clean_threshold = 100;
581        let incremental_threshold = 10;
582
583        // Set up the runner
584        let mut runner =
585            MerkleTestRunner { db: TestStageDB::default(), clean_threshold, incremental_threshold };
586
587        let input = ExecInput {
588            target: Some(previous_stage),
589            checkpoint: Some(StageCheckpoint::new(stage_progress)),
590        };
591
592        runner.seed_execution(input).expect("failed to seed execution");
593        let rx = runner.execute(input);
594
595        // Assert the successful result
596        let result = rx.await.unwrap();
597        assert_matches!(
598            result,
599            Ok(ExecOutput {
600                checkpoint: StageCheckpoint {
601                    block_number,
602                    stage_checkpoint: Some(StageUnitCheckpoint::Entities(EntitiesCheckpoint {
603                        processed,
604                        total
605                    }))
606                },
607                done: true
608            }) if block_number == previous_stage && processed == total &&
609                total == (
610                    runner.db.count_entries::<tables::HashedAccounts>().unwrap() +
611                    runner.db.count_entries::<tables::HashedStorages>().unwrap()
612                ) as u64
613        );
614
615        // Validate the stage execution
616        let provider = runner.db.factory.provider().unwrap();
617        let header = provider.header_by_number(previous_stage).unwrap().unwrap();
618        let expected_root = header.state_root;
619
620        let actual_root = runner
621            .db
622            .query_with_provider(|provider| {
623                Ok(reth_trie_db::with_adapter!(provider, |A| {
624                    DbStateRoot::<_, A>::incremental_root_with_updates(
625                        &provider,
626                        stage_progress + 1..=previous_stage,
627                    )
628                }))
629            })
630            .unwrap();
631
632        assert_eq!(
633            actual_root.unwrap().0,
634            expected_root,
635            "State root mismatch after chunked processing"
636        );
637    }
638
639    struct MerkleTestRunner {
640        db: TestStageDB,
641        clean_threshold: u64,
642        incremental_threshold: u64,
643    }
644
645    impl Default for MerkleTestRunner {
646        fn default() -> Self {
647            Self {
648                db: TestStageDB::default(),
649                clean_threshold: 10000,
650                incremental_threshold: 10000,
651            }
652        }
653    }
654
655    impl StageTestRunner for MerkleTestRunner {
656        type S = MerkleStage;
657
658        fn db(&self) -> &TestStageDB {
659            &self.db
660        }
661
662        fn stage(&self) -> Self::S {
663            Self::S::Both {
664                rebuild_threshold: self.clean_threshold,
665                incremental_threshold: self.incremental_threshold,
666            }
667        }
668    }
669
670    impl ExecuteStageTestRunner for MerkleTestRunner {
671        type Seed = Vec<SealedBlock<reth_ethereum_primitives::Block>>;
672
673        #[allow(clippy::clone_on_copy)]
674        fn seed_execution(&mut self, input: ExecInput) -> Result<Self::Seed, TestRunnerError> {
675            let stage_progress = input.checkpoint().block_number;
676            let start = stage_progress + 1;
677            let end = input.target();
678            let mut rng = generators::rng();
679
680            let mut preblocks = vec![];
681            if stage_progress > 0 {
682                preblocks.append(&mut random_block_range(
683                    &mut rng,
684                    0..=stage_progress - 1,
685                    BlockRangeParams {
686                        parent: Some(B256::ZERO),
687                        tx_count: 0..1,
688                        ..Default::default()
689                    },
690                ));
691                self.db.insert_blocks(preblocks.iter(), StorageKind::Static)?;
692            }
693
694            let num_of_accounts = 31;
695            let accounts = random_contract_account_range(&mut rng, &mut (0..num_of_accounts))
696                .into_iter()
697                .collect::<BTreeMap<_, _>>();
698
699            self.db.insert_accounts_and_storages(
700                accounts.iter().map(|(addr, acc)| (*addr, (acc.clone(), std::iter::empty()))),
701            )?;
702
703            let (header, body) = random_block(
704                &mut rng,
705                stage_progress,
706                BlockParams { parent: preblocks.last().map(|b| b.hash()), ..Default::default() },
707            )
708            .split_sealed_header_body();
709            let mut header = header.unseal();
710
711            header.state_root = state_root(
712                accounts
713                    .clone()
714                    .into_iter()
715                    .map(|(address, account)| (address, (account, std::iter::empty()))),
716            );
717            let sealed_head = SealedBlock::<reth_ethereum_primitives::Block>::from_sealed_parts(
718                SealedHeader::seal_slow(header),
719                body,
720            );
721
722            let head_hash = sealed_head.hash();
723            let mut blocks = vec![sealed_head];
724            blocks.extend(random_block_range(
725                &mut rng,
726                start..=end,
727                BlockRangeParams { parent: Some(head_hash), tx_count: 0..3, ..Default::default() },
728            ));
729            let last_block = blocks.last().cloned().unwrap();
730            self.db.insert_blocks(blocks.iter(), StorageKind::Static)?;
731
732            let (transitions, final_state) = random_changeset_range(
733                &mut rng,
734                blocks.iter(),
735                accounts.into_iter().map(|(addr, acc)| (addr, (acc, Vec::new()))),
736                0..3,
737                0..256,
738            );
739            // add block changeset from block 1.
740            self.db.insert_changesets(transitions, Some(start))?;
741            self.db.insert_accounts_and_storages(final_state)?;
742
743            // Calculate state root
744            let root = self.db.query(|tx| {
745                let mut accounts = BTreeMap::default();
746                let mut accounts_cursor = tx.cursor_read::<tables::HashedAccounts>()?;
747                let mut storage_cursor = tx.cursor_dup_read::<tables::HashedStorages>()?;
748                for entry in accounts_cursor.walk_range(..)? {
749                    let (key, account) = entry?;
750                    let mut storage_entries = Vec::new();
751                    let mut entry = storage_cursor.seek_exact(key)?;
752                    while let Some((_, storage)) = entry {
753                        storage_entries.push(storage);
754                        entry = storage_cursor.next_dup()?;
755                    }
756                    let storage = storage_entries
757                        .into_iter()
758                        .filter(|v| !v.value.is_zero())
759                        .map(|v| (v.key, v.value))
760                        .collect::<Vec<_>>();
761                    accounts.insert(key, (account, storage));
762                }
763
764                Ok(state_root_prehashed(accounts))
765            })?;
766
767            let static_file_provider = self.db.factory.static_file_provider();
768            let mut writer =
769                static_file_provider.latest_writer(StaticFileSegment::Headers).unwrap();
770            let mut last_header = last_block.clone_sealed_header();
771            last_header.set_state_root(root);
772
773            let hash = last_header.hash_slow();
774            writer.prune_headers(1).unwrap();
775            writer.commit().unwrap();
776            writer.append_header(&last_header, &hash).unwrap();
777            writer.commit().unwrap();
778
779            Ok(blocks)
780        }
781
782        fn validate_execution(
783            &self,
784            _input: ExecInput,
785            _output: Option<ExecOutput>,
786        ) -> Result<(), TestRunnerError> {
787            // The execution is validated within the stage
788            Ok(())
789        }
790    }
791
792    impl UnwindStageTestRunner for MerkleTestRunner {
793        fn validate_unwind(&self, _input: UnwindInput) -> Result<(), TestRunnerError> {
794            // The unwind is validated within the stage
795            Ok(())
796        }
797
798        fn before_unwind(&self, input: UnwindInput) -> Result<(), TestRunnerError> {
799            let target_block = input.unwind_to + 1;
800
801            self.db
802                .commit(|tx| {
803                    let mut storage_changesets_cursor =
804                        tx.cursor_dup_read::<tables::StorageChangeSets>().unwrap();
805                    let mut storage_cursor =
806                        tx.cursor_dup_write::<tables::HashedStorages>().unwrap();
807
808                    let mut tree: BTreeMap<B256, BTreeMap<B256, U256>> = BTreeMap::new();
809
810                    let mut rev_changeset_walker =
811                        storage_changesets_cursor.walk_back(None).unwrap();
812                    while let Some((bn_address, entry)) =
813                        rev_changeset_walker.next().transpose().unwrap()
814                    {
815                        if bn_address.block_number() < target_block {
816                            break
817                        }
818
819                        tree.entry(keccak256(bn_address.address()))
820                            .or_default()
821                            .insert(keccak256(entry.key), entry.value);
822                    }
823                    for (hashed_address, storage) in tree {
824                        for (hashed_slot, value) in storage {
825                            let storage_entry = storage_cursor
826                                .seek_by_key_subkey(hashed_address, hashed_slot)
827                                .unwrap();
828                            if storage_entry.is_some_and(|v| v.key == hashed_slot) {
829                                storage_cursor.delete_current().unwrap();
830                            }
831
832                            if !value.is_zero() {
833                                let storage_entry = StorageEntry { key: hashed_slot, value };
834                                storage_cursor.upsert(hashed_address, &storage_entry).unwrap();
835                            }
836                        }
837                    }
838
839                    let mut changeset_cursor =
840                        tx.cursor_dup_write::<tables::AccountChangeSets>().unwrap();
841                    let mut rev_changeset_walker = changeset_cursor.walk_back(None).unwrap();
842
843                    while let Some((block_number, account_before_tx)) =
844                        rev_changeset_walker.next().transpose().unwrap()
845                    {
846                        if block_number < target_block {
847                            break
848                        }
849
850                        if let Some(acc) = account_before_tx.info {
851                            tx.put::<tables::HashedAccounts>(
852                                keccak256(account_before_tx.address),
853                                acc,
854                            )
855                            .unwrap();
856                        } else {
857                            tx.delete::<tables::HashedAccounts>(
858                                keccak256(account_before_tx.address),
859                                None,
860                            )
861                            .unwrap();
862                        }
863                    }
864                    Ok(())
865                })
866                .unwrap();
867            Ok(())
868        }
869    }
870}