1use crate::{
4 common::DownloadContext, SnapAccountStore, SnapStorageStore, SnapSyncError, StorageChunk,
5 StorageProgress, VerifiedRange, MAX_HASH,
6};
7use alloy_primitives::{B256, U256};
8use futures::future::join_all;
9use reth_db_api::transaction::DbTxMut;
10use reth_downloaders::snap::{
11 StorageRangeDownloader, StorageRangeOutcome, VerifiedAccountBatch, VerifiedStorageRanges,
12};
13use reth_eth_wire_types::snap::GetStorageRangesMessage;
14use reth_network_p2p::snap::client::SnapClient;
15use reth_network_peers::PeerId;
16use reth_storage_api::{
17 DBProvider, DatabaseProviderFactory, MetadataProvider, MetadataWriter, StateWriter,
18};
19use reth_tasks::Runtime;
20use std::fmt;
21
22pub const DEFAULT_STORAGE_ACCOUNTS: usize = 128;
24
25pub const DEFAULT_REPAIR_SLOTS: usize = 128;
27
28pub struct StorageRangeDownload<C, F> {
33 context: DownloadContext<C, F>,
34 max_accounts: usize,
36 max_repair_slots: usize,
38}
39
40impl<C, F> StorageRangeDownload<C, F> {
41 pub const fn new(client: C, factory: F, runtime: Runtime) -> Self {
43 Self {
44 context: DownloadContext::new(client, factory, runtime),
45 max_accounts: DEFAULT_STORAGE_ACCOUNTS,
46 max_repair_slots: DEFAULT_REPAIR_SLOTS,
47 }
48 }
49
50 pub const fn with_response_bytes(mut self, response_bytes: u64) -> Self {
52 self.context.set_response_bytes(response_bytes);
53 self
54 }
55
56 pub const fn with_max_accounts(mut self, max_accounts: usize) -> Self {
58 self.max_accounts = if max_accounts == 0 { 1 } else { max_accounts };
59 self
60 }
61
62 pub const fn with_max_repair_slots(mut self, max_repair_slots: usize) -> Self {
65 self.max_repair_slots = if max_repair_slots == 0 { 1 } else { max_repair_slots };
66 self
67 }
68}
69
70impl<C, F> StorageRangeDownload<C, F>
71where
72 C: SnapClient + Clone + Unpin,
73 F: DatabaseProviderFactory + Clone + 'static,
74 F::Provider: MetadataProvider,
75 F::ProviderRW: MetadataProvider + MetadataWriter + StateWriter + DBProvider<Tx: DbTxMut>,
76{
77 pub async fn next(&mut self, range: &VerifiedRange) -> Result<StorageRangeStep, SnapSyncError> {
82 let (write, origin) = (range.write(), range.origin());
83 let progress =
84 self.context.factory().database_provider_ro()?.storage_progress(write, origin)?;
85 let contracts = range.range().storage_batch();
86 let Some(first) =
87 contracts.accounts().iter().position(|(account, _)| !progress.is_complete(*account))
88 else {
89 return Ok(StorageRangeStep::Complete)
90 };
91 let end = progress.request_end(contracts.accounts(), first, self.max_accounts);
92 let batch = contracts.range(first..end).expect("positions are inside the batch");
93 let from = progress.resume_at(batch.accounts()[0].0).expect("first contract is incomplete");
94
95 let request = GetStorageRangesMessage {
96 request_id: self.context.next_request_id(),
97 root_hash: batch.state_root(),
98 account_hashes: batch.accounts().iter().map(|(account, _)| *account).collect(),
99 starting_hash: from.into(),
100 limit_hash: MAX_HASH.into(),
101 response_bytes: self.context.response_bytes(),
102 };
103 let downloader = StorageRangeDownloader::new(
104 self.context.client().clone(),
105 request,
106 &batch,
107 self.context.runtime().clone(),
108 )?;
109 let ranges = match downloader.await? {
110 StorageRangeOutcome::Verified(ranges) => ranges,
111 StorageRangeOutcome::Unavailable { peer_id } => {
112 return Ok(StorageRangeStep::Unavailable { peer_id })
113 }
114 };
115 let chunks = chunks(ranges, batch, from)?;
116 let committed = self
118 .context
119 .commit(move |provider| {
120 let mut progress = StorageProgress::START;
121 for chunk in chunks {
122 progress = provider.commit_storage_chunk(write, origin, chunk)?;
123 }
124 Ok(progress)
125 })
126 .await?;
127 Ok(StorageRangeStep::Committed(committed))
128 }
129
130 pub async fn repair_slots(
134 &mut self,
135 range: &VerifiedRange,
136 ) -> Result<Option<Vec<(B256, U256)>>, SnapSyncError> {
137 let (write, account) = (range.write(), range.origin());
138 let repairs = self.context.factory().database_provider_ro()?.snap_repairs(write)?;
139 let contracts = range.range().storage_batch();
141 let Some(batch) = contracts.range(0..1).filter(|batch| batch.accounts()[0].0 == account)
142 else {
143 return Ok(Some(Vec::new()))
144 };
145
146 let mut requests = Vec::new();
147 for slot in repairs.slots(account).take(self.max_repair_slots) {
148 let request = GetStorageRangesMessage {
149 request_id: self.context.next_request_id(),
150 root_hash: batch.state_root(),
151 account_hashes: vec![account],
152 starting_hash: slot.into(),
153 limit_hash: slot.into(),
154 response_bytes: self.context.response_bytes(),
155 };
156 let downloader = StorageRangeDownloader::new(
157 self.context.client().clone(),
158 request,
159 &batch,
160 self.context.runtime().clone(),
161 )?;
162 requests.push(async move { downloader.await.map(|outcome| (slot, outcome)) });
163 }
164
165 let requested = requests.len();
166 let mut values = Vec::new();
167 for response in join_all(requests).await {
168 let (slot, StorageRangeOutcome::Verified(ranges)) = response? else { continue };
169 let value = ranges
171 .into_ranges()
172 .into_iter()
173 .next()
174 .and_then(|range| range.slots.first().copied())
175 .filter(|(key, _)| *key == slot)
176 .map_or(U256::ZERO, |(_, value)| value);
177 values.push((slot, value));
178 }
179 Ok((!values.is_empty() || requested == 0).then_some(values))
181 }
182}
183
184impl<C, F> fmt::Debug for StorageRangeDownload<C, F> {
185 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
186 f.debug_struct("StorageRangeDownload")
187 .field("context", &self.context)
188 .field("max_accounts", &self.max_accounts)
189 .field("max_repair_slots", &self.max_repair_slots)
190 .finish()
191 }
192}
193
194#[derive(Debug)]
196pub enum StorageRangeStep {
197 Committed(StorageProgress),
199 Unavailable {
201 peer_id: PeerId,
203 },
204 Complete,
206}
207
208fn chunks(
211 ranges: VerifiedStorageRanges,
212 batch: VerifiedAccountBatch<'_>,
213 from: B256,
214) -> Result<Vec<StorageChunk>, SnapSyncError> {
215 let roots: Vec<B256> =
216 batch.accounts().iter().map(|(_, account)| account.storage_root).collect();
217 let resume = ranges.follow_up(0, batch)?.map(|(request, _)| {
218 (request.account_hashes[0], request.starting_hash.unwrap_or(B256::ZERO))
219 });
220 Ok(ranges
221 .into_ranges()
222 .into_iter()
223 .zip(roots)
224 .enumerate()
225 .map(|(index, (range, storage_root))| {
226 let from = if index == 0 { from } else { B256::ZERO };
228 let next =
229 resume.filter(|(account, _)| *account == range.account_hash).map(|(_, slot)| slot);
230 StorageChunk::new(range.account_hash, storage_root, from, range.slots, next)
231 })
232 .collect())
233}
234
235#[cfg(test)]
236mod tests {
237 use super::*;
238 use crate::{
239 test_utils::{
240 account, generation, hashed_factory, insert_generation_headers, key, state_root,
241 storage_ranges, storage_root_of, stored_slots, verified_range, verified_repair,
242 ScriptedSnapClient,
243 },
244 SnapAccountStore, SnapAttemptStore, SnapCatchUpStore, StateRepairs,
245 };
246 use alloy_eips::BlockNumHash;
247 use alloy_primitives::U256;
248 use reth_eth_wire_types::snap::StorageRangesMessage;
249 use reth_network_p2p::{error::PeerRequestResult, snap::client::SnapResponse};
250 use reth_network_peers::WithPeerId;
251 use reth_provider::{test_utils::MockNodeTypesWithDB, ProviderFactory};
252 use reth_trie_common::TrieAccount;
253 use std::sync::Arc;
254
255 const FAR: B256 = B256::repeat_byte(0xaa);
256
257 type Factory = ProviderFactory<MockNodeTypesWithDB>;
258 type Download = StorageRangeDownload<Arc<ScriptedSnapClient>, Factory>;
259
260 fn large() -> Vec<(B256, U256)> {
262 vec![(key(1), U256::from(11)), (key(2), U256::from(12)), (key(3), U256::from(13))]
263 }
264
265 fn small() -> Vec<(B256, U256)> {
266 vec![(key(9), U256::from(19))]
267 }
268
269 fn contract(nonce: u64, slots: &[(B256, U256)]) -> TrieAccount {
270 let mut contract = account(nonce);
271 contract.storage_root = storage_root_of(slots);
272 contract
273 }
274
275 fn accounts() -> Vec<(B256, TrieAccount)> {
277 vec![
278 (key(1), account(1)),
279 (key(2), contract(2, &large())),
280 (key(3), contract(3, &small())),
281 (FAR, account(4)),
282 ]
283 }
284
285 fn started(accounts: &[(B256, TrieAccount)]) -> (Factory, VerifiedRange) {
287 let factory = hashed_factory();
288 insert_generation_headers(&factory);
289 let provider = factory.database_provider_rw().unwrap();
290 let write = provider.start_snap_attempt(generation(1, state_root(accounts))).unwrap();
291 provider.start_account_coverage(write).unwrap();
292 provider.commit().unwrap();
293 let range = verified_range(accounts, 0..accounts.len(), B256::ZERO, &[]);
294 (factory, VerifiedRange::new(write, range))
295 }
296
297 fn download(
298 responses: impl IntoIterator<Item = PeerRequestResult<SnapResponse>>,
299 factory: Factory,
300 ) -> (Arc<ScriptedSnapClient>, Download) {
301 let client = Arc::new(ScriptedSnapClient::new(responses));
302 (Arc::clone(&client), StorageRangeDownload::new(client, factory, Runtime::test()))
303 }
304
305 async fn committed(download: &mut Download, range: &VerifiedRange) -> StorageProgress {
306 match download.next(range).await.unwrap() {
307 StorageRangeStep::Committed(progress) => progress,
308 step => panic!("expected a commit, got {step:?}"),
309 }
310 }
311
312 fn slots_of(factory: &Factory, account: B256) -> Vec<(B256, U256)> {
313 stored_slots(&factory.database_provider_ro().unwrap(), account)
314 }
315
316 #[tokio::test]
317 async fn a_contract_spanning_responses_is_committed_response_by_response() {
318 let accounts = accounts();
319 let (factory, range) = started(&accounts);
320 let (large, small) = (large(), small());
321 let responses = [
322 storage_ranges(1, &[&large[..1]], &large, &[B256::ZERO, key(1)]),
324 storage_ranges(2, &[&large[1..]], &large, &[key(2), key(3)]),
326 storage_ranges(3, &[&small[..]], &small, &[]),
327 ];
328 let (client, mut download) = download(responses, factory.clone());
329
330 let progress = committed(&mut download, &range).await;
331 assert_eq!(progress.resume_at(key(2)), Some(key(2)));
332 assert_eq!(slots_of(&factory, key(2)), large[..1]);
333
334 assert!(committed(&mut download, &range).await.is_complete(key(2)));
335 assert!(committed(&mut download, &range).await.is_complete(key(3)));
336 assert!(matches!(download.next(&range).await.unwrap(), StorageRangeStep::Complete));
337 assert_eq!(
338 *client.storage_requests(),
339 [
340 (vec![key(2), key(3)], B256::ZERO),
341 (vec![key(2), key(3)], key(2)),
342 (vec![key(3)], B256::ZERO),
343 ]
344 );
345
346 let provider = factory.database_provider_rw().unwrap();
348 let coverage = provider
349 .commit_account_range(range.write(), range.range(), Default::default(), Vec::new())
350 .unwrap();
351 provider.commit().unwrap();
352 assert!(coverage.is_complete());
353 assert_eq!(slots_of(&factory, key(2)), large);
354 assert_eq!(slots_of(&factory, key(3)), small);
355 }
356
357 #[tokio::test]
358 async fn an_unavailable_response_leaves_the_progress_in_place() {
359 let accounts = accounts();
360 let (factory, range) = started(&accounts);
361 let large = large();
362 let peer = PeerId::random();
363 let empty = StorageRangesMessage { request_id: 2, slots: Vec::new(), proof: Vec::new() };
364 let responses = [
365 storage_ranges(1, &[&large[..1]], &large, &[B256::ZERO, key(1)]),
366 Ok(WithPeerId::new(peer, SnapResponse::StorageRanges(empty))),
367 storage_ranges(3, &[&large[1..]], &large, &[key(2), key(3)]),
368 ];
369 let (client, mut download) = download(responses, factory.clone());
370 let progress = committed(&mut download, &range).await;
371
372 let step = download.next(&range).await.unwrap();
373
374 assert!(matches!(step, StorageRangeStep::Unavailable { peer_id } if peer_id == peer));
375 let provider = factory.database_provider_ro().unwrap();
376 assert_eq!(provider.storage_progress(range.write(), range.origin()).unwrap(), progress);
377 drop(provider);
378 assert!(committed(&mut download, &range).await.is_complete(key(2)));
379 let origins: Vec<_> = client.storage_requests().iter().map(|(_, from)| *from).collect();
380 assert_eq!(origins, [B256::ZERO, key(2), key(2)]);
381 }
382
383 #[tokio::test]
384 async fn a_new_download_resumes_from_the_last_committed_response() {
385 let accounts = accounts();
386 let (factory, range) = started(&accounts);
387 let (large, small) = (large(), small());
388 let first = storage_ranges(1, &[&large[..1]], &large, &[B256::ZERO, key(1)]);
389 let (_, mut interrupted) = download([first], factory.clone());
390 committed(&mut interrupted, &range).await;
391 drop(interrupted);
392
393 let responses = [
394 storage_ranges(1, &[&large[1..]], &large, &[key(2), key(3)]),
395 storage_ranges(2, &[&small[..]], &small, &[]),
396 ];
397 let (client, mut resumed) = download(responses, factory.clone());
398 committed(&mut resumed, &range).await;
399 committed(&mut resumed, &range).await;
400
401 assert!(matches!(resumed.next(&range).await.unwrap(), StorageRangeStep::Complete));
402 assert_eq!(
403 *client.storage_requests(),
404 [(vec![key(2), key(3)], key(2)), (vec![key(3)], B256::ZERO)]
405 );
406 assert_eq!(slots_of(&factory, key(2)), large);
407 }
408
409 #[tokio::test]
410 async fn requests_ask_for_at_most_the_configured_contracts() {
411 let accounts = accounts();
412 let (factory, range) = started(&accounts);
413 let (large, small) = (large(), small());
414 let responses = [
415 storage_ranges(1, &[&large[..]], &large, &[]),
416 storage_ranges(2, &[&small[..]], &small, &[]),
417 ];
418 let (client, download) = download(responses, factory);
419 let mut download = download.with_max_accounts(1);
420
421 committed(&mut download, &range).await;
422 committed(&mut download, &range).await;
423
424 assert!(matches!(download.next(&range).await.unwrap(), StorageRangeStep::Complete));
425 assert_eq!(
426 *client.storage_requests(),
427 [(vec![key(2)], B256::ZERO), (vec![key(3)], B256::ZERO)]
428 );
429 }
430
431 #[tokio::test]
432 async fn a_range_without_contracts_requests_no_storage() {
433 let accounts = vec![(key(1), account(1)), (FAR, account(2))];
434 let (factory, range) = started(&accounts);
435 let (client, mut download) = download([], factory);
436
437 assert!(matches!(download.next(&range).await.unwrap(), StorageRangeStep::Complete));
438 assert!(client.storage_requests().is_empty());
439 }
440
441 #[tokio::test]
442 async fn storage_fetched_before_the_pivot_moved_is_not_resumed() {
443 let accounts = accounts();
444 let (factory, range) = started(&accounts);
445 let large = large();
446 let first = storage_ranges(1, &[&large[..1]], &large, &[B256::ZERO, key(1)]);
447 let (client, mut download) = download([first], factory.clone());
448 committed(&mut download, &range).await;
449 let provider = factory.database_provider_rw().unwrap();
450 provider.advance_snap_pivot(range.write(), generation(2, state_root(&accounts))).unwrap();
451 provider.commit().unwrap();
452
453 assert!(matches!(download.next(&range).await, Err(SnapSyncError::StaleWrite { .. })));
454 assert_eq!(client.storage_requests().len(), 1);
455 }
456
457 #[tokio::test]
458 async fn storage_part_way_through_resumes_at_the_new_root() {
459 let accounts = accounts();
460 let (factory, range) = started(&accounts);
461 let (large, small) = (large(), small());
462 let mut moved = accounts.clone();
464 moved[0].1 = contract(1, &small);
465 let responses = [
466 storage_ranges(1, &[&large[..1]], &large, &[B256::ZERO, key(1)]),
467 storage_ranges(2, &[&small[..]], &small, &[]),
468 storage_ranges(3, &[&large[1..]], &large, &[key(2), key(3)]),
469 storage_ranges(4, &[&small[..]], &small, &[]),
470 ];
471 let (client, mut download) = download(responses, factory.clone());
472 committed(&mut download, &range).await;
473
474 let provider = factory.database_provider_rw().unwrap();
475 let write =
476 provider.advance_snap_pivot(range.write(), generation(2, state_root(&moved))).unwrap();
477 provider.commit().unwrap();
478 let range =
479 VerifiedRange::new(write, verified_range(&moved, 0..moved.len(), B256::ZERO, &[]));
480 for _ in 0..3 {
481 committed(&mut download, &range).await;
482 }
483 assert!(matches!(download.next(&range).await.unwrap(), StorageRangeStep::Complete));
484 assert_eq!(
485 *client.storage_requests(),
486 [
487 (vec![key(2), key(3)], B256::ZERO),
488 (vec![key(1)], B256::ZERO),
489 (vec![key(2), key(3)], key(2)),
490 (vec![key(3)], B256::ZERO),
491 ]
492 );
493
494 let provider = factory.database_provider_rw().unwrap();
496 assert!(matches!(
497 provider.commit_account_range(write, range.range(), Default::default(), Vec::new()),
498 Err(SnapSyncError::CatchUpBehindPivot { applied: 1, pivot: 2 })
499 ));
500 let block = BlockNumHash::new(2, B256::repeat_byte(2));
501 provider.commit_block_access_list(write, block, B256::repeat_byte(1), &[]).unwrap();
502 let coverage = provider
503 .commit_account_range(write, range.range(), Default::default(), Vec::new())
504 .unwrap();
505 provider.commit().unwrap();
506 assert!(coverage.is_complete());
507 assert_eq!(slots_of(&factory, key(1)), small);
508 assert_eq!(slots_of(&factory, key(2)), large);
509 }
510
511 #[tokio::test]
512 async fn a_page_ending_before_a_carried_contract_hands_it_to_the_next_range() {
513 let accounts = accounts();
514 let (factory, range) = started(&accounts);
515 let large = large();
516 let first = storage_ranges(1, &[&large[..1]], &large, &[B256::ZERO, key(1)]);
517 let (_, mut download) = download([first], factory.clone());
518 committed(&mut download, &range).await;
519 let provider = factory.database_provider_rw().unwrap();
520 let write = provider
521 .advance_snap_pivot(range.write(), generation(2, state_root(&accounts)))
522 .unwrap();
523 let block = BlockNumHash::new(2, B256::repeat_byte(2));
524 provider.commit_block_access_list(write, block, B256::repeat_byte(1), &[]).unwrap();
525
526 let page = verified_range(&accounts, 0..1, B256::ZERO, &[B256::ZERO, key(1)]);
528 let next = provider
529 .commit_account_range(write, &page, Default::default(), Vec::new())
530 .unwrap()
531 .next()
532 .unwrap();
533
534 assert!(next <= key(2));
535 assert_eq!(provider.storage_progress(write, next).unwrap().resume_at(key(2)), Some(key(2)));
536 }
537
538 fn repairing(accounts: &[(B256, TrieAccount)], slots: &[B256]) -> (Factory, VerifiedRange) {
540 let (factory, _) = started(accounts);
541 let provider = factory.database_provider_rw().unwrap();
542 let write = provider.active_snap_write().unwrap().unwrap();
543 let mut repairs = StateRepairs::default();
544 for slot in slots {
545 repairs.insert_slot(key(2), *slot);
546 }
547 provider.schedule_snap_repairs(write, repairs).unwrap();
548 provider.commit().unwrap();
549 let range = verified_repair(accounts, 1..2, key(2), &[key(2)]);
550 (factory, VerifiedRange::new(write, range))
551 }
552
553 #[tokio::test]
554 async fn repair_slots_read_zero_where_the_pivot_holds_none() {
555 let accounts = accounts();
556 let (factory, range) = repairing(&accounts, &[key(1), key(5)]);
557 let large = large();
558 let responses = [
559 storage_ranges(1, &[&large[..1]], &large, &[key(1)]),
560 storage_ranges(2, &[&[]], &large, &[key(5)]),
562 ];
563 let (client, mut download) = download(responses, factory);
564
565 let values = download.repair_slots(&range).await.unwrap();
566
567 assert_eq!(values, Some(vec![(key(1), U256::from(11)), (key(5), U256::ZERO)]));
568 assert_eq!(*client.storage_requests(), [(vec![key(2)], key(1)), (vec![key(2)], key(5))]);
569 }
570
571 #[tokio::test]
572 async fn a_contract_scheduled_without_slots_fetches_none() {
573 let accounts = accounts();
574 let (factory, range) = repairing(&accounts, &[]);
575 let provider = factory.database_provider_rw().unwrap();
576 let mut repairs = StateRepairs::default();
577 repairs.insert_account(key(2));
578 provider.schedule_snap_repairs(range.write(), repairs).unwrap();
579 provider.commit().unwrap();
580 let (client, mut download) = download([], factory);
581
582 assert_eq!(download.repair_slots(&range).await.unwrap(), Some(Vec::new()));
583 assert!(client.storage_requests().is_empty());
584 }
585
586 #[tokio::test]
587 async fn an_unserved_repair_slot_fetches_nothing() {
588 let accounts = accounts();
589 let (factory, range) = repairing(&accounts, &[key(1)]);
590 let (_, mut download) = download([storage_ranges(1, &[], &[], &[])], factory);
591
592 assert_eq!(download.repair_slots(&range).await.unwrap(), None);
593 }
594
595 #[tokio::test]
596 async fn served_repair_slots_survive_an_unserved_one() {
597 let accounts = accounts();
598 let (factory, range) = repairing(&accounts, &[key(1), key(5)]);
599 let large = large();
600 let responses = [
601 storage_ranges(1, &[&large[..1]], &large, &[key(1)]),
602 storage_ranges(2, &[], &[], &[]),
603 ];
604 let (_, mut download) = download(responses, factory);
605
606 let values = download.repair_slots(&range).await.unwrap();
607
608 assert_eq!(values, Some(vec![(key(1), U256::from(11))]));
610 }
611
612 #[tokio::test]
613 async fn repair_slots_are_fetched_in_bounded_batches() {
614 let accounts = accounts();
615 let (factory, range) = repairing(&accounts, &[key(1), key(5)]);
616 let large = large();
617 let responses = [storage_ranges(1, &[&large[..1]], &large, &[key(1)])];
618 let (client, download) = download(responses, factory);
619 let mut download = download.with_max_repair_slots(1);
620
621 let values = download.repair_slots(&range).await.unwrap();
622
623 assert_eq!(values, Some(vec![(key(1), U256::from(11))]));
625 assert_eq!(*client.storage_requests(), [(vec![key(2)], key(1))]);
626 }
627}