1use crate::{common::SnapRecord, SnapAccountStore, SnapAttemptStore, SnapSyncError, SnapWrite};
7use alloy_primitives::{B256, U256};
8use reth_db_api::{
9 cursor::DbDupCursorRO,
10 tables,
11 transaction::{DbTx, DbTxMut},
12};
13use reth_storage_api::{DBProvider, MetadataProvider, MetadataWriter, SnapAttemptId, StateWriter};
14use reth_trie_common::{root::storage_root, HashedPostState, HashedStorage, EMPTY_ROOT_HASH};
15use serde::{Deserialize, Serialize};
16
17pub trait SnapStorageStore {
22 fn storage_progress(
26 &self,
27 write: SnapWrite,
28 origin: B256,
29 ) -> Result<StorageProgress, SnapSyncError>;
30
31 fn commit_storage_chunk(
36 &self,
37 write: SnapWrite,
38 origin: B256,
39 chunk: StorageChunk,
40 ) -> Result<StorageProgress, SnapSyncError>
41 where
42 Self: MetadataWriter + StateWriter + DBProvider<Tx: DbTxMut>;
43}
44
45#[derive(Clone, Debug, PartialEq, Eq)]
47pub struct StorageChunk {
48 account: B256,
50 storage_root: B256,
52 from: B256,
54 slots: Vec<(B256, U256)>,
56 next: Option<B256>,
58}
59
60impl StorageChunk {
61 pub const fn new(
63 account: B256,
64 storage_root: B256,
65 from: B256,
66 slots: Vec<(B256, U256)>,
67 next: Option<B256>,
68 ) -> Self {
69 Self { account, storage_root, from, slots, next }
70 }
71}
72
73#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
78pub struct StorageProgress {
79 complete: Option<B256>,
81 partial: Vec<PartialStorage>,
83}
84
85impl StorageProgress {
86 pub const START: Self = Self { complete: None, partial: Vec::new() };
88
89 pub fn is_complete(&self, account: B256) -> bool {
91 self.complete.is_some_and(|complete| account <= complete)
92 }
93
94 pub fn has_slots(&self, account: B256) -> bool {
96 self.is_complete(account) || self.partial_of(account).is_some()
97 }
98
99 fn partial_of(&self, account: B256) -> Option<&PartialStorage> {
101 self.partial.iter().find(|partial| partial.account == account)
102 }
103
104 pub(crate) fn request_end<T>(
107 &self,
108 contracts: &[(B256, T)],
109 first: usize,
110 max: usize,
111 ) -> usize {
112 let end = contracts.len().min(first.saturating_add(max));
113 contracts[first + 1..end]
114 .iter()
115 .position(|(account, _)| self.partial_of(*account).is_some())
116 .map_or(end, |offset| first + 1 + offset)
117 }
118
119 pub(crate) fn carry_to(
121 mut self,
122 provider: &impl MetadataWriter,
123 write: SnapWrite,
124 next: B256,
125 ) -> Result<(), SnapSyncError> {
126 self.complete = self.complete.filter(|complete| *complete >= next);
127 self.partial.retain(|partial| partial.account >= next);
128 if self != Self::START {
129 StoredProgress::new(write, next, self).write(provider)?;
130 }
131 Ok(())
132 }
133
134 pub fn resume_at(&self, account: B256) -> Option<B256> {
136 if self.is_complete(account) {
137 return None
138 }
139 Some(self.partial_of(account).map_or(B256::ZERO, |partial| partial.next))
140 }
141
142 fn carried(mut self) -> Self {
144 for partial in &mut self.partial {
145 partial.storage_root = None;
146 }
147 self
148 }
149
150 fn advance(mut self, chunk: &StorageChunk) -> Result<Self, SnapSyncError> {
155 let blocked = self
156 .partial
157 .iter()
158 .any(|partial| partial.account < chunk.account && partial.storage_root.is_some());
159 if blocked || self.resume_at(chunk.account) != Some(chunk.from) {
160 return Err(SnapSyncError::OutOfOrderStorage {
161 account: chunk.account,
162 from: chunk.from,
163 })
164 }
165 if let Some(partial) = self.partial_of(chunk.account) &&
166 let Some(expected) = partial.storage_root &&
167 expected != chunk.storage_root
168 {
169 return Err(SnapSyncError::StorageRootMismatch {
170 account: chunk.account,
171 expected,
172 got: chunk.storage_root,
173 })
174 }
175 if chunk.next.is_some_and(|next| next <= chunk.from) {
176 return Err(SnapSyncError::NoProgress { origin: chunk.from })
177 }
178 self.partial.retain(|partial| partial.account > chunk.account);
181 match chunk.next {
182 Some(next) => {
183 let storage_root = Some(chunk.storage_root);
184 self.partial
185 .insert(0, PartialStorage { account: chunk.account, storage_root, next });
186 }
187 None => self.complete = Some(chunk.account),
188 }
189 Ok(self)
190 }
191}
192
193#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
195struct PartialStorage {
196 account: B256,
198 storage_root: Option<B256>,
200 next: B256,
202}
203
204#[derive(Serialize, Deserialize)]
206pub(crate) struct StoredProgress {
207 version: u32,
209 attempt: SnapAttemptId,
211 state_version: u64,
213 origin: B256,
215 progress: StorageProgress,
217}
218
219impl SnapRecord for StoredProgress {
220 const KEY: &'static str = "snap_storage_progress";
221 const VERSION: u32 = 2;
222}
223
224impl StoredProgress {
225 const fn new(write: SnapWrite, origin: B256, progress: StorageProgress) -> Self {
227 Self {
228 version: Self::VERSION,
229 attempt: write.attempt(),
230 state_version: write.state_version(),
231 origin,
232 progress,
233 }
234 }
235
236 fn into_progress(self, write: SnapWrite, origin: B256) -> Option<StorageProgress> {
238 if self.attempt != write.attempt() || self.origin != origin {
239 return None
240 }
241 Some(if self.state_version == write.state_version() {
242 self.progress
243 } else {
244 self.progress.carried()
245 })
246 }
247}
248
249impl<T: MetadataProvider> SnapStorageStore for T {
250 fn storage_progress(
252 &self,
253 write: SnapWrite,
254 origin: B256,
255 ) -> Result<StorageProgress, SnapSyncError> {
256 self.authorize_snap_write(write)?;
257 let Some(stored) = StoredProgress::read(self)? else { return Ok(StorageProgress::START) };
258 Ok(stored.into_progress(write, origin).unwrap_or(StorageProgress::START))
259 }
260
261 fn commit_storage_chunk(
263 &self,
264 write: SnapWrite,
265 origin: B256,
266 chunk: StorageChunk,
267 ) -> Result<StorageProgress, SnapSyncError>
268 where
269 Self: MetadataWriter + StateWriter + DBProvider<Tx: DbTxMut>,
270 {
271 let coverage = self.account_coverage(write)?.ok_or(SnapSyncError::NoCoverage)?;
272 if coverage.next() != Some(origin) {
273 return Err(SnapSyncError::OutOfOrderRange { expected: coverage.next(), got: origin })
274 }
275 if chunk.account < origin || chunk.storage_root == EMPTY_ROOT_HASH {
276 return Err(SnapSyncError::UnexpectedStorage { account: chunk.account })
277 }
278 let progress = self.storage_progress(write, origin)?;
279 let complete = progress.complete;
280 let progress = progress.advance(&chunk)?;
281
282 if chunk.next.is_none() {
283 let start = complete
286 .map_or(origin, |account| (U256::from_be_bytes(account.0) + U256::from(1)).into());
287 self.remove::<tables::HashedStorages>(start..chunk.account)?;
288 }
289 if chunk.from == B256::ZERO {
290 self.remove::<tables::HashedStorages>(chunk.account..=chunk.account)?;
291 }
292 let state = HashedPostState::default()
293 .with_storages([(chunk.account, HashedStorage::from_iter(chunk.slots))])
294 .into_sorted();
295 self.write_hashed_state(&state)?;
296 let stored = StoredProgress::new(write, origin, progress);
297 stored.write(self)?;
298 Ok(stored.progress)
299 }
300}
301
302pub(crate) fn persisted_storage_root(tx: &impl DbTx, account: B256) -> Result<B256, SnapSyncError> {
305 let mut cursor = tx.cursor_dup_read::<tables::HashedStorages>()?;
306 let mut failed = None;
307 let slots = cursor.walk_dup(Some(account), None)?.map_while(|entry| {
308 entry.map(|(_, slot)| (slot.key, slot.value)).map_err(|error| failed = Some(error)).ok()
309 });
310 let root = storage_root(slots);
311 failed.map_or(Ok(root), |error| Err(error.into()))
312}
313
314#[cfg(test)]
315mod tests {
316 use super::*;
317 use crate::test_utils::{
318 account, generation, hashed_factory, insert_generation_headers, state_root,
319 storage_root_of, stored_slots,
320 };
321 use reth_provider::{
322 test_utils::MockNodeTypesWithDB, DatabaseProviderFactory, ProviderFactory,
323 };
324 use std::ops::Range;
325
326 const CONTRACT: B256 = B256::repeat_byte(0x22);
327
328 fn slot(value: u8) -> B256 {
329 B256::with_last_byte(value)
330 }
331
332 fn slots() -> Vec<(B256, U256)> {
333 vec![(slot(1), U256::from(11)), (slot(2), U256::from(12)), (slot(3), U256::from(13))]
334 }
335
336 fn root() -> B256 {
337 storage_root_of(&slots())
338 }
339
340 fn chunk(served: Range<usize>, from: B256, next: Option<B256>) -> StorageChunk {
342 StorageChunk::new(CONTRACT, root(), from, slots()[served].to_vec(), next)
343 }
344
345 fn started() -> (ProviderFactory<MockNodeTypesWithDB>, SnapWrite) {
347 let factory = hashed_factory();
348 insert_generation_headers(&factory);
349 let provider = factory.database_provider_rw().unwrap();
350 let mut contract = account(1);
351 contract.storage_root = root();
352 let generation = generation(1, state_root(&[(CONTRACT, contract)]));
353 let write = provider.start_snap_attempt(generation).unwrap();
354 provider.start_account_coverage(write).unwrap();
355 provider.commit().unwrap();
356 (factory, write)
357 }
358
359 #[test]
360 fn an_interrupted_chunk_commit_keeps_the_last_committed_progress() {
361 let (factory, write) = started();
362 let provider = factory.database_provider_rw().unwrap();
363 let first = chunk(0..1, B256::ZERO, Some(slot(2)));
364 let progress = provider.commit_storage_chunk(write, B256::ZERO, first).unwrap();
365 provider.commit().unwrap();
366 assert_eq!(progress.resume_at(CONTRACT), Some(slot(2)));
367
368 let provider = factory.database_provider_rw().unwrap();
369 let rest = chunk(1..3, slot(2), None);
370 assert!(provider
371 .commit_storage_chunk(write, B256::ZERO, rest)
372 .unwrap()
373 .is_complete(CONTRACT));
374 drop(provider);
375
376 let provider = factory.database_provider_rw().unwrap();
377 assert_eq!(provider.storage_progress(write, B256::ZERO).unwrap(), progress);
378 assert_eq!(stored_slots(&provider, CONTRACT), slots()[..1]);
379 }
380
381 #[test]
382 fn a_chunk_must_continue_its_contract_before_another_starts() {
383 let (factory, write) = started();
384 let provider = factory.database_provider_rw().unwrap();
385 let first = chunk(0..1, B256::ZERO, Some(slot(2)));
386 let progress = provider.commit_storage_chunk(write, B256::ZERO, first).unwrap();
387
388 let refused = [
389 chunk(0..1, B256::ZERO, Some(slot(2))),
391 chunk(2..3, slot(3), None),
392 StorageChunk::new(B256::repeat_byte(0x44), root(), B256::ZERO, slots(), None),
393 ];
394 for chunk in refused {
395 assert!(matches!(
396 provider.commit_storage_chunk(write, B256::ZERO, chunk),
397 Err(SnapSyncError::OutOfOrderStorage { .. })
398 ));
399 }
400 let stalled = chunk(1..1, slot(2), Some(slot(2)));
401 assert!(matches!(
402 provider.commit_storage_chunk(write, B256::ZERO, stalled),
403 Err(SnapSyncError::NoProgress { .. })
404 ));
405 let other_root = StorageChunk::new(
406 CONTRACT,
407 B256::repeat_byte(0x33),
408 slot(2),
409 slots()[1..].to_vec(),
410 None,
411 );
412 assert!(matches!(
413 provider.commit_storage_chunk(write, B256::ZERO, other_root),
414 Err(SnapSyncError::StorageRootMismatch { .. })
415 ));
416
417 assert_eq!(provider.storage_progress(write, B256::ZERO).unwrap(), progress);
418 assert_eq!(stored_slots(&provider, CONTRACT), slots()[..1]);
419 }
420
421 #[test]
422 fn starting_a_contract_replaces_what_an_earlier_attempt_left() {
423 let (factory, write) = started();
424 let provider = factory.database_provider_rw().unwrap();
425 let leftover = HashedStorage::from_iter([(slot(7), U256::from(77))]);
426 let leftover = HashedPostState::default().with_storages([(CONTRACT, leftover)]);
427 provider.write_hashed_state(&leftover.into_sorted()).unwrap();
428
429 provider.commit_storage_chunk(write, B256::ZERO, chunk(0..3, B256::ZERO, None)).unwrap();
430
431 assert_eq!(stored_slots(&provider, CONTRACT), slots());
432 assert_eq!(persisted_storage_root(provider.tx_ref(), CONTRACT).unwrap(), root());
433 }
434
435 fn carried() -> (ProviderFactory<MockNodeTypesWithDB>, SnapWrite, SnapWrite) {
437 let (factory, write) = started();
438 let provider = factory.database_provider_rw().unwrap();
439 let first = chunk(0..1, B256::ZERO, Some(slot(2)));
440 provider.commit_storage_chunk(write, B256::ZERO, first).unwrap();
441 let advanced =
442 provider.advance_snap_pivot(write, generation(2, B256::repeat_byte(0xcc))).unwrap();
443 provider.commit().unwrap();
444 (factory, write, advanced)
445 }
446
447 fn whole(account: B256) -> StorageChunk {
449 StorageChunk::new(account, root(), B256::ZERO, slots(), None)
450 }
451
452 #[test]
453 fn progress_survives_the_pivot_moving_with_its_root_unproved() {
454 let (factory, write, advanced) = carried();
455 let provider = factory.database_provider_rw().unwrap();
456
457 assert_eq!(provider.storage_progress(advanced, slot(9)).unwrap(), StorageProgress::START);
458 let progress = provider.storage_progress(advanced, B256::ZERO).unwrap();
459 assert_eq!(progress.resume_at(CONTRACT), Some(slot(2)));
460 assert!(matches!(
461 provider.commit_storage_chunk(write, B256::ZERO, chunk(1..3, slot(2), None)),
462 Err(SnapSyncError::StaleWrite { .. })
463 ));
464
465 let rest = StorageChunk::new(
467 CONTRACT,
468 B256::repeat_byte(0x33),
469 slot(2),
470 slots()[1..].to_vec(),
471 None,
472 );
473 let progress = provider.commit_storage_chunk(advanced, B256::ZERO, rest).unwrap();
474 assert!(progress.is_complete(CONTRACT));
475 assert_eq!(stored_slots(&provider, CONTRACT), slots());
476 }
477
478 #[test]
479 fn a_carried_contract_does_not_block_the_contracts_the_new_root_holds() {
480 let (factory, _, advanced) = carried();
481 let provider = factory.database_provider_rw().unwrap();
482
483 let before = B256::repeat_byte(0x11);
485 let progress = provider.commit_storage_chunk(advanced, B256::ZERO, whole(before)).unwrap();
486 assert!(progress.is_complete(before));
487 assert_eq!(progress.resume_at(CONTRACT), Some(slot(2)));
488
489 let after = B256::repeat_byte(0x44);
491 let progress = provider.commit_storage_chunk(advanced, B256::ZERO, whole(after)).unwrap();
492 assert!(progress.is_complete(after));
493 assert!(progress.partial.is_empty());
494 }
495
496 #[test]
497 fn a_new_contract_part_way_keeps_the_carried_one_resumable() {
498 let (factory, _, advanced) = carried();
499 let provider = factory.database_provider_rw().unwrap();
500 let before = B256::repeat_byte(0x11);
501 let first =
502 StorageChunk::new(before, root(), B256::ZERO, slots()[..1].to_vec(), Some(slot(2)));
503 provider.commit_storage_chunk(advanced, B256::ZERO, first).unwrap();
504
505 let advanced =
507 provider.advance_snap_pivot(advanced, generation(3, B256::repeat_byte(0xdd))).unwrap();
508 let progress = provider.storage_progress(advanced, B256::ZERO).unwrap();
509 assert_eq!(progress.resume_at(before), Some(slot(2)));
510 assert_eq!(progress.resume_at(CONTRACT), Some(slot(2)));
511
512 let rest = StorageChunk::new(before, root(), slot(2), slots()[1..].to_vec(), None);
513 let progress = provider.commit_storage_chunk(advanced, B256::ZERO, rest).unwrap();
514 assert!(progress.is_complete(before));
515 assert_eq!(progress.resume_at(CONTRACT), Some(slot(2)));
516 }
517
518 #[test]
519 fn a_contract_resuming_part_way_leads_its_request() {
520 let progress = StorageProgress {
521 complete: None,
522 partial: vec![PartialStorage { account: CONTRACT, storage_root: None, next: slot(2) }],
523 };
524 let contracts = [(slot(1), ()), (CONTRACT, ()), (B256::repeat_byte(0x33), ())];
525
526 assert_eq!(progress.request_end(&contracts, 0, 10), 1);
527 assert_eq!(progress.request_end(&contracts, 1, 10), 3);
528 assert_eq!(progress.request_end(&contracts, 1, 1), 2);
529 assert_eq!(StorageProgress::START.request_end(&contracts, 0, 10), 3);
530 }
531
532 #[test]
533 fn storage_outside_the_range_being_downloaded_is_refused() {
534 let (factory, write) = started();
535 let provider = factory.database_provider_rw().unwrap();
536
537 let ahead = provider.commit_storage_chunk(write, slot(9), chunk(0..3, B256::ZERO, None));
538 let empty = StorageChunk::new(CONTRACT, EMPTY_ROOT_HASH, B256::ZERO, Vec::new(), None);
539 let empty = provider.commit_storage_chunk(write, B256::ZERO, empty);
540
541 assert!(matches!(ahead, Err(SnapSyncError::OutOfOrderRange { .. })));
542 assert!(matches!(empty, Err(SnapSyncError::UnexpectedStorage { .. })));
543 assert!(stored_slots(&provider, CONTRACT).is_empty());
544 }
545
546 #[test]
547 fn a_record_this_build_cannot_read_is_reported() {
548 let (factory, write) = started();
549
550 for record in [br#"{"version":999}"#.to_vec(), b"{}".to_vec()] {
551 let provider = factory.database_provider_rw().unwrap();
552 provider.write_metadata(StoredProgress::KEY, record).unwrap();
553
554 assert!(matches!(
555 provider.storage_progress(write, B256::ZERO),
556 Err(SnapSyncError::UnsupportedRecord { .. })
557 ));
558 }
559 }
560}