1use crate::{
7 common::SnapRecord, SnapAccountStore, SnapAttemptStore, SnapCatchUpStore, SnapSyncError,
8 SnapWrite,
9};
10use alloy_eips::BlockNumHash;
11use alloy_primitives::{B256, KECCAK256_EMPTY};
12use reth_db_api::{cursor::DbCursorRO, tables, transaction::DbTx, RawKey, RawTable};
13use reth_primitives_traits::{AlloyBlockHeader, GotExpected};
14use reth_stages_types::{StageCheckpoint, StageId};
15use reth_storage_api::{
16 BlockHashReader, DBProvider, HeaderProvider, MetadataProvider, MetadataWriter, SnapAttemptId,
17 StageCheckpointReader, StageCheckpointWriter,
18};
19use reth_storage_errors::provider::{ProviderError, RootMismatch};
20use serde::{Deserialize, Serialize};
21use tokio_util::sync::CancellationToken;
22
23pub const DEFAULT_SCAN_CHUNK: u64 = 100_000;
26
27pub trait SnapStateVerifier {
31 fn start_trie_rebuild(
36 &self,
37 write: SnapWrite,
38 chunk: u64,
39 cancel: &CancellationToken,
40 ) -> Result<(), SnapSyncError>
41 where
42 Self: MetadataWriter + StageCheckpointWriter + DBProvider;
43
44 fn verify_completeness(
47 &self,
48 write: SnapWrite,
49 chunk: u64,
50 cancel: &CancellationToken,
51 ) -> Result<(), SnapSyncError>
52 where
53 Self: DBProvider;
54
55 fn is_trie_rebuild_started(&self, write: SnapWrite) -> Result<bool, SnapSyncError>;
58
59 fn verify_state_root(&self, write: SnapWrite) -> Result<VerifiedSnapState, SnapSyncError>
64 where
65 Self: BlockHashReader + HeaderProvider + MetadataWriter + StageCheckpointReader;
66}
67
68#[derive(Clone, Copy, Debug, Eq, PartialEq)]
70pub struct VerifiedSnapState {
71 attempt: SnapAttemptId,
73 target: BlockNumHash,
75 state_root: B256,
77}
78
79impl VerifiedSnapState {
80 pub const fn attempt(&self) -> SnapAttemptId {
82 self.attempt
83 }
84
85 pub const fn target(&self) -> BlockNumHash {
87 self.target
88 }
89
90 pub const fn state_root(&self) -> B256 {
92 self.state_root
93 }
94}
95
96impl<T: MetadataProvider> SnapStateVerifier for T {
97 fn start_trie_rebuild(
98 &self,
99 write: SnapWrite,
100 chunk: u64,
101 cancel: &CancellationToken,
102 ) -> Result<(), SnapSyncError>
103 where
104 Self: MetadataWriter + StageCheckpointWriter + DBProvider,
105 {
106 if self.authorize_snap_write(write)?.pivot().number == 0 {
109 return Err(SnapSyncError::GenesisPivot)
110 }
111 self.verify_completeness(write, chunk, cancel)?;
112 self.save_stage_checkpoint(StageId::MerkleExecute, StageCheckpoint::default())?;
115 self.save_stage_checkpoint_progress(StageId::MerkleExecute, Vec::new())?;
116 StoredRebuild::new(write).write(self)?;
117 Ok(())
118 }
119
120 fn verify_completeness(
121 &self,
122 write: SnapWrite,
123 chunk: u64,
124 cancel: &CancellationToken,
125 ) -> Result<(), SnapSyncError>
126 where
127 Self: DBProvider,
128 {
129 let attempt = self.authorize_snap_write(write)?;
130 let coverage = self.account_coverage(write)?.ok_or(SnapSyncError::NoCoverage)?;
131 if let Some(next) = coverage.next() {
132 return Err(SnapSyncError::IncompleteAccounts { next })
133 }
134 let applied =
135 self.catch_up_progress(write)?.ok_or(SnapSyncError::NoCatchUpProgress)?.applied();
136 if applied != attempt.pivot() {
137 return Err(SnapSyncError::CatchUpBehindPivot {
138 applied: applied.number,
139 pivot: attempt.pivot().number,
140 })
141 }
142 let repairs = self.snap_repairs(write)?;
143 if !repairs.is_empty() {
144 return Err(SnapSyncError::PendingRepairs { accounts: repairs.len() })
145 }
146 ensure_code_present(self.tx_ref(), chunk, cancel)
147 }
148
149 fn is_trie_rebuild_started(&self, write: SnapWrite) -> Result<bool, SnapSyncError> {
150 Ok(StoredRebuild::read(self)?.is_some_and(|stored| stored.write == write))
151 }
152
153 fn verify_state_root(&self, write: SnapWrite) -> Result<VerifiedSnapState, SnapSyncError>
154 where
155 Self: BlockHashReader + HeaderProvider + MetadataWriter + StageCheckpointReader,
156 {
157 let attempt = self.authorize_canonical_snap_write(write)?;
158 let target = attempt.pivot();
159 let header = self
160 .sealed_header(target.number)?
161 .filter(|header| header.hash() == target.hash)
162 .ok_or(SnapSyncError::MissingHeader { block: target.number })?;
163 if header.state_root() != attempt.state_root() {
165 return Err(ProviderError::StateRootMismatch(Box::new(RootMismatch {
166 root: GotExpected { got: attempt.state_root(), expected: header.state_root() },
167 block_number: target.number,
168 block_hash: target.hash,
169 }))
170 .into())
171 }
172 let handed_off = self.is_trie_rebuild_started(write)?;
175 let rebuilt = self.get_stage_checkpoint(StageId::MerkleExecute)?;
176 if !handed_off || rebuilt.map(|checkpoint| checkpoint.block_number) != Some(target.number) {
177 return Err(ProviderError::StateForNumberNotFound(target.number).into())
178 }
179 self.verify_snap_attempt(write)?;
180 Ok(VerifiedSnapState { attempt: attempt.id(), target, state_root: header.state_root() })
181 }
182}
183
184#[derive(Serialize, Deserialize)]
186pub(crate) struct StoredRebuild {
187 version: u32,
189 write: SnapWrite,
191}
192
193impl SnapRecord for StoredRebuild {
194 const KEY: &'static str = "snap_trie_rebuild";
195 const VERSION: u32 = 1;
196}
197
198impl StoredRebuild {
199 const fn new(write: SnapWrite) -> Self {
201 Self { version: Self::VERSION, write }
202 }
203}
204
205fn ensure_code_present(
208 tx: &impl DbTx,
209 chunk: u64,
210 cancel: &CancellationToken,
211) -> Result<(), SnapSyncError> {
212 let chunk = chunk.max(1);
213 let mut cursor = tx.cursor_read::<tables::HashedAccounts>()?;
214 for (scanned, entry) in cursor.walk(None)?.enumerate() {
215 if (scanned as u64).is_multiple_of(chunk) && cancel.is_cancelled() {
216 return Err(SnapSyncError::Cancelled)
217 }
218 let (_, account) = entry?;
219 if let Some(hash) = account.bytecode_hash.filter(|hash| *hash != KECCAK256_EMPTY) &&
221 tx.get::<RawTable<tables::Bytecodes>>(RawKey::new(hash))?.is_none()
222 {
223 return Err(SnapSyncError::MissingCode { hash })
224 }
225 }
226 Ok(())
227}
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232 use crate::{
233 test_utils::{account, hashed_factory, header, key, state_root},
234 SnapGeneration, SnapStorageStore, StorageChunk,
235 };
236 use alloy_primitives::{map::B256Map, Bytes, U256};
237 use reth_db_api::transaction::DbTxMut;
238 use reth_primitives_traits::SealedHeader;
239 use reth_provider::{
240 test_utils::{insert_headers, MockNodeTypesWithDB},
241 DatabaseProviderFactory, ProviderFactory,
242 };
243 use reth_stages::stages::MerkleStage;
244 use reth_stages_api::{ExecInput, Stage, StageError};
245 use reth_trie_common::{root::storage_root_unsorted, HashedStorage, TrieAccount};
246 use revm::bytecode::Bytecode;
247
248 type Factory = ProviderFactory<MockNodeTypesWithDB>;
249 type Provider = <Factory as DatabaseProviderFactory>::ProviderRW;
250
251 const CONTRACT: B256 = B256::repeat_byte(0xaa);
252 const SLOT: B256 = B256::repeat_byte(0x55);
253
254 fn code() -> Bytecode {
255 Bytecode::new_raw(Bytes::from_static(&[0x60, 0x00]))
256 }
257
258 fn storage() -> HashedStorage {
259 HashedStorage::from_iter([(SLOT, U256::from(7))])
260 }
261
262 fn accounts() -> Vec<(B256, TrieAccount)> {
264 let mut contract = account(3);
265 contract.storage_root = storage_root_unsorted([(SLOT, U256::from(7))]);
266 contract.code_hash = code().hash_slow();
267 vec![(key(1), account(1)), (key(2), account(2)), (CONTRACT, contract)]
268 }
269
270 fn insert_chain(factory: &Factory, header_root: B256) -> Vec<BlockNumHash> {
272 let mut parent = B256::ZERO;
273 let headers: Vec<_> = (0..=2)
274 .map(|number| {
275 let mut header = header(number, parent, None);
276 header.state_root = header_root;
277 let sealed = SealedHeader::seal_slow(header);
278 parent = sealed.hash();
279 sealed
280 })
281 .collect();
282 insert_headers(factory, &headers);
283 headers.iter().map(|header| BlockNumHash::new(header.number, header.hash())).collect()
284 }
285
286 fn downloaded(header_root: B256, served: usize) -> (Factory, SnapWrite, Vec<BlockNumHash>) {
289 let accounts = accounts();
290 let factory = hashed_factory();
291 let blocks = insert_chain(&factory, header_root);
292 let provider = factory.database_provider_rw().unwrap();
293 let write = provider
294 .start_snap_attempt(SnapGeneration::new(blocks[1], state_root(&accounts)))
295 .unwrap();
296 provider.start_account_coverage(write).unwrap();
297 let targets = (served < accounts.len()).then(|| accounts[served - 1].0);
298 let range =
299 crate::test_utils::verified_range(&accounts, 0..served, B256::ZERO, targets.as_slice());
300 let complete = served == accounts.len();
301 let (storages, bytecodes) = if complete {
302 (B256Map::from_iter([(CONTRACT, storage())]), vec![(code().hash_slow(), code())])
303 } else {
304 Default::default()
305 };
306 provider.commit_account_range(write, &range, storages, bytecodes).unwrap();
307 provider.commit().unwrap();
308 (factory, write, blocks)
309 }
310
311 fn start(provider: &Provider, write: SnapWrite) -> Result<(), SnapSyncError> {
312 provider.start_trie_rebuild(write, 1, &CancellationToken::new())
313 }
314
315 fn run_merkle(provider: &Provider, target: u64) -> Result<(), StageError> {
317 let mut stage = MerkleStage::default_execution();
318 loop {
319 let checkpoint = provider.get_stage_checkpoint(StageId::MerkleExecute).unwrap();
320 let output = stage.execute(provider, ExecInput { target: Some(target), checkpoint })?;
321 provider.save_stage_checkpoint(StageId::MerkleExecute, output.checkpoint).unwrap();
322 if output.done {
323 return Ok(())
324 }
325 }
326 }
327
328 fn merkle_checkpoint(provider: &Provider) -> Option<u64> {
329 provider.get_stage_checkpoint(StageId::MerkleExecute).unwrap().map(|c| c.block_number)
330 }
331
332 #[test]
333 fn complete_state_is_verified_once_the_merkle_stage_reaches_the_pivot() {
334 let root = state_root(&accounts());
335 let (factory, write, blocks) = downloaded(root, accounts().len());
336 let provider = factory.database_provider_rw().unwrap();
337
338 start(&provider, write).unwrap();
339 run_merkle(&provider, 1).unwrap();
340 let verified = provider.verify_state_root(write).unwrap();
341 provider.commit().unwrap();
342
343 assert_eq!((verified.target(), verified.state_root()), (blocks[1], root));
344 let provider = factory.database_provider_rw().unwrap();
345 assert!(provider.snap_attempt().unwrap().unwrap().is_verified());
346 assert!(matches!(provider.verify_state_root(write), Err(SnapSyncError::StaleWrite { .. })));
348 }
349
350 #[test]
351 fn unfinished_accounts_prevent_the_hand_off() {
352 let (factory, write, _) = downloaded(state_root(&accounts()), 1);
353 let provider = factory.database_provider_rw().unwrap();
354
355 assert!(matches!(
356 start(&provider, write),
357 Err(SnapSyncError::IncompleteAccounts { next }) if next == key(2)
358 ));
359 assert_eq!(merkle_checkpoint(&provider), None);
360 }
361
362 #[test]
363 fn storage_persisted_ahead_of_its_range_prevents_the_hand_off() {
364 let (factory, write, _) = downloaded(state_root(&accounts()), 2);
365 let provider = factory.database_provider_rw().unwrap();
366 let origin = provider.account_coverage(write).unwrap().unwrap().next().unwrap();
367 let root = accounts()[2].1.storage_root;
368 let chunk =
369 StorageChunk::new(CONTRACT, root, B256::ZERO, vec![(SLOT, U256::from(7))], None);
370 provider.commit_storage_chunk(write, origin, chunk).unwrap();
371
372 assert!(matches!(start(&provider, write), Err(SnapSyncError::IncompleteAccounts { .. })));
373 }
374
375 #[test]
376 fn missing_code_prevents_the_hand_off() {
377 let (factory, write, _) = downloaded(state_root(&accounts()), accounts().len());
378 let provider = factory.database_provider_rw().unwrap();
379 provider.tx_ref().delete::<tables::Bytecodes>(code().hash_slow(), None).unwrap();
380
381 assert!(matches!(
382 start(&provider, write),
383 Err(SnapSyncError::MissingCode { hash }) if hash == code().hash_slow()
384 ));
385 }
386
387 #[test]
388 fn catch_up_short_of_the_pivot_prevents_the_hand_off() {
389 let root = state_root(&accounts());
390 let (factory, write, blocks) = downloaded(root, accounts().len());
391 let provider = factory.database_provider_rw().unwrap();
392 let write =
393 provider.advance_snap_pivot(write, SnapGeneration::new(blocks[2], root)).unwrap();
394
395 assert!(matches!(
396 start(&provider, write),
397 Err(SnapSyncError::CatchUpBehindPivot { applied: 1, pivot: 2 })
398 ));
399 }
400
401 #[test]
402 fn a_cancelled_session_stops_the_completeness_scan() {
403 let (factory, write, _) = downloaded(state_root(&accounts()), accounts().len());
404 let provider = factory.database_provider_rw().unwrap();
405 let cancel = CancellationToken::new();
406 cancel.cancel();
407
408 assert!(matches!(
409 provider.start_trie_rebuild(write, 1, &cancel),
410 Err(SnapSyncError::Cancelled)
411 ));
412 assert_eq!(merkle_checkpoint(&provider), None);
413 }
414
415 #[test]
416 fn corrupted_state_never_reaches_the_pivot() {
417 let (factory, write, _) = downloaded(state_root(&accounts()), accounts().len());
418 let provider = factory.database_provider_rw().unwrap();
419 start(&provider, write).unwrap();
420 provider.tx_ref().delete::<tables::HashedStorages>(CONTRACT, None).unwrap();
421
422 assert!(matches!(run_merkle(&provider, 1), Err(StageError::Block { .. })));
424 assert!(matches!(
425 provider.verify_state_root(write),
426 Err(SnapSyncError::Provider(ProviderError::StateForNumberNotFound(1)))
427 ));
428 assert!(provider.snap_attempt().unwrap().unwrap().is_unfinished());
429 }
430
431 #[test]
432 fn an_interrupted_rebuild_is_not_verified() {
433 let (factory, write, _) = downloaded(state_root(&accounts()), accounts().len());
434 let provider = factory.database_provider_rw().unwrap();
435 provider.save_stage_checkpoint(StageId::MerkleExecute, StageCheckpoint::new(1)).unwrap();
437
438 start(&provider, write).unwrap();
439
440 assert!(matches!(
441 provider.verify_state_root(write),
442 Err(SnapSyncError::Provider(ProviderError::StateForNumberNotFound(1)))
443 ));
444 run_merkle(&provider, 1).unwrap();
445 provider.verify_state_root(write).unwrap();
446 }
447
448 #[test]
449 fn a_header_committing_to_another_root_is_refused() {
450 let header_root = B256::repeat_byte(0xcc);
452 let (factory, write, _) = downloaded(header_root, accounts().len());
453 let provider = factory.database_provider_rw().unwrap();
454 start(&provider, write).unwrap();
455
456 assert!(matches!(run_merkle(&provider, 1), Err(StageError::Block { .. })));
457 match provider.verify_state_root(write) {
458 Err(SnapSyncError::Provider(ProviderError::StateRootMismatch(mismatch))) => {
459 assert_eq!(
460 mismatch.root,
461 GotExpected { got: state_root(&accounts()), expected: header_root }
462 );
463 }
464 other => panic!("expected a state root mismatch, got {other:?}"),
465 }
466 }
467
468 #[test]
469 fn a_rebuild_for_an_earlier_attempt_is_not_trusted() {
470 let root = state_root(&accounts());
471 let (factory, write, blocks) = downloaded(root, accounts().len());
472 let provider = factory.database_provider_rw().unwrap();
473 start(&provider, write).unwrap();
474 run_merkle(&provider, 1).unwrap();
475
476 let restarted = provider.start_snap_attempt(SnapGeneration::new(blocks[1], root)).unwrap();
478
479 assert_eq!(merkle_checkpoint(&provider), Some(1));
480 assert!(matches!(
481 provider.verify_state_root(restarted),
482 Err(SnapSyncError::Provider(ProviderError::StateForNumberNotFound(1)))
483 ));
484 }
485
486 #[test]
487 fn moving_the_pivot_after_the_hand_off_requires_a_new_one() {
488 let root = state_root(&accounts());
489 let (factory, write, blocks) = downloaded(root, accounts().len());
490 let provider = factory.database_provider_rw().unwrap();
491 start(&provider, write).unwrap();
492 let advanced =
493 provider.advance_snap_pivot(write, SnapGeneration::new(blocks[2], root)).unwrap();
494
495 run_merkle(&provider, 2).unwrap();
496
497 assert!(matches!(
498 provider.verify_state_root(advanced),
499 Err(SnapSyncError::Provider(ProviderError::StateForNumberNotFound(2)))
500 ));
501 }
502
503 #[test]
504 fn a_genesis_pivot_is_not_handed_off() {
505 let root = state_root(&accounts());
506 let (factory, _, blocks) = downloaded(root, accounts().len());
507 let provider = factory.database_provider_rw().unwrap();
508 let write = provider.start_snap_attempt(SnapGeneration::new(blocks[0], root)).unwrap();
509
510 assert!(matches!(start(&provider, write), Err(SnapSyncError::GenesisPivot)));
511 assert_eq!(merkle_checkpoint(&provider), None);
512 }
513}