1use crate::{
4 common::DownloadContext, AccountCoverage, SnapAccountStore, SnapAttemptStore, SnapSyncError,
5 SnapWrite, MAX_HASH,
6};
7use alloy_primitives::{map::B256Map, B256, U256};
8use reth_db_api::transaction::DbTxMut;
9use reth_downloaders::snap::{AccountRangeDownloader, AccountRangeOutcome, VerifiedAccountRange};
10use reth_eth_wire_types::snap::GetAccountRangeMessage;
11use reth_network_p2p::snap::client::SnapClient;
12use reth_network_peers::PeerId;
13use reth_storage_api::{
14 BlockHashReader, DBProvider, DatabaseProviderFactory, MetadataProvider, MetadataWriter,
15 StateWriter,
16};
17use reth_tasks::Runtime;
18use reth_trie_common::HashedStorage;
19use revm::bytecode::Bytecode;
20use std::fmt;
21
22pub struct AccountRangeDownload<C, F> {
27 context: DownloadContext<C, F>,
28 coverage: Option<AccountCoverage>,
30}
31
32impl<C, F> AccountRangeDownload<C, F> {
33 pub const fn new(client: C, factory: F, runtime: Runtime) -> Self {
35 Self { context: DownloadContext::new(client, factory, runtime), coverage: None }
36 }
37
38 pub const fn with_response_bytes(mut self, response_bytes: u64) -> Self {
40 self.context.set_response_bytes(response_bytes);
41 self
42 }
43
44 pub const fn coverage(&self) -> Option<AccountCoverage> {
46 self.coverage
47 }
48}
49
50impl<C, F> AccountRangeDownload<C, F>
51where
52 C: SnapClient + Clone + Unpin,
53 F: DatabaseProviderFactory + Clone + 'static,
54 F::Provider: MetadataProvider,
55 F::ProviderRW:
56 BlockHashReader + MetadataProvider + MetadataWriter + StateWriter + DBProvider<Tx: DbTxMut>,
57{
58 pub async fn next(&mut self) -> Result<Option<AccountRangeStep>, SnapSyncError> {
63 let (write, root_hash, coverage) = self.active_write()?;
64 self.coverage = Some(coverage);
65 let Some(origin) = coverage.next() else { return Ok(None) };
66 self.request(write, root_hash, origin, MAX_HASH).await.map(Some)
67 }
68
69 pub async fn commit(
74 &mut self,
75 verified: VerifiedRange,
76 storages: B256Map<HashedStorage>,
77 bytecodes: Vec<(B256, Bytecode)>,
78 ) -> Result<AccountCoverage, SnapSyncError> {
79 let coverage = self
80 .context
81 .commit(move |provider| {
82 let VerifiedRange { write, range } = verified;
83 provider.commit_account_range(write, &range, storages, bytecodes)
84 })
85 .await?;
86 self.coverage = Some(coverage);
87 Ok(coverage)
88 }
89
90 pub async fn next_repair(&mut self) -> Result<Option<AccountRangeStep>, SnapSyncError> {
94 let (write, root_hash, _) = self.active_write()?;
95 let repairs = self.context.factory().database_provider_ro()?.snap_repairs(write)?;
96 let Some(hashed_address) = repairs.first() else { return Ok(None) };
97 self.request(write, root_hash, hashed_address, hashed_address).await.map(Some)
98 }
99
100 pub async fn commit_repair(
103 &mut self,
104 verified: VerifiedRange,
105 slots: Vec<(B256, U256)>,
106 ) -> Result<usize, SnapSyncError> {
107 self.context
108 .commit(move |provider| {
109 let VerifiedRange { write, range } = verified;
110 provider.commit_account_repair(write, &range, slots)
111 })
112 .await
113 }
114
115 async fn request(
117 &mut self,
118 write: SnapWrite,
119 root_hash: B256,
120 origin: B256,
121 limit: B256,
122 ) -> Result<AccountRangeStep, SnapSyncError> {
123 let request = GetAccountRangeMessage {
124 request_id: self.context.next_request_id(),
125 root_hash,
126 starting_hash: origin,
127 limit_hash: limit,
128 response_bytes: self.context.response_bytes(),
129 };
130 let downloader = AccountRangeDownloader::new(
131 self.context.client().clone(),
132 request,
133 self.context.runtime().clone(),
134 )?;
135
136 Ok(match downloader.await? {
137 AccountRangeOutcome::Verified(range) => {
138 AccountRangeStep::Verified(VerifiedRange { write, range })
139 }
140 AccountRangeOutcome::Unavailable { peer_id } => {
141 AccountRangeStep::Unavailable { origin, peer_id }
142 }
143 })
144 }
145
146 fn active_write(&self) -> Result<(SnapWrite, B256, AccountCoverage), SnapSyncError> {
149 let provider = self.context.factory().database_provider_ro()?;
150 let write = provider.active_snap_write()?.ok_or(SnapSyncError::NoAttempt)?;
151 let root = provider.authorize_snap_write(write)?.state_root();
152 let coverage = provider.account_coverage(write)?.ok_or(SnapSyncError::NoCoverage)?;
153 Ok((write, root, coverage))
154 }
155}
156
157impl<C, F> fmt::Debug for AccountRangeDownload<C, F> {
158 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
159 f.debug_struct("AccountRangeDownload")
160 .field("context", &self.context)
161 .field("coverage", &self.coverage)
162 .finish()
163 }
164}
165
166#[derive(Debug)]
168pub enum AccountRangeStep {
169 Verified(VerifiedRange),
171 Unavailable {
173 origin: B256,
175 peer_id: PeerId,
177 },
178}
179
180#[derive(Debug)]
185pub struct VerifiedRange {
186 write: SnapWrite,
188 range: VerifiedAccountRange,
190}
191
192impl VerifiedRange {
193 #[cfg(test)]
194 pub(crate) const fn new(write: SnapWrite, range: VerifiedAccountRange) -> Self {
195 Self { write, range }
196 }
197
198 pub(crate) const fn write(&self) -> SnapWrite {
200 self.write
201 }
202
203 pub const fn origin(&self) -> B256 {
205 self.range.origin()
206 }
207
208 pub const fn range(&self) -> &VerifiedAccountRange {
210 &self.range
211 }
212}
213
214#[cfg(test)]
215mod tests {
216 use super::*;
217 use crate::{
218 test_utils::{
219 account, account_range, generation, hashed_factory, key, state_root, verified_range,
220 ScriptedSnapClient,
221 },
222 StateRepairs,
223 };
224 use reth_db_api::{cursor::DbCursorRO, tables, transaction::DbTx};
225 use reth_eth_wire_types::snap::AccountRangeMessage;
226 use reth_network_p2p::{
227 error::{PeerRequestResult, RequestError},
228 snap::client::SnapResponse,
229 };
230 use reth_network_peers::WithPeerId;
231 use reth_provider::{test_utils::MockNodeTypesWithDB, ProviderFactory};
232 use reth_trie_common::TrieAccount;
233 use std::sync::Arc;
234
235 const FAR: B256 = B256::repeat_byte(0xaa);
236
237 fn accounts() -> Vec<(B256, TrieAccount)> {
238 vec![(key(1), account(1)), (key(2), account(2)), (FAR, account(3))]
239 }
240
241 fn started(accounts: &[(B256, TrieAccount)]) -> ProviderFactory<MockNodeTypesWithDB> {
242 let factory = hashed_factory();
243 let provider = factory.database_provider_rw().unwrap();
244 let write = provider.start_snap_attempt(generation(1, state_root(accounts))).unwrap();
245 provider.start_account_coverage(write).unwrap();
246 provider.commit().unwrap();
247 factory
248 }
249
250 type Download =
251 AccountRangeDownload<Arc<ScriptedSnapClient>, ProviderFactory<MockNodeTypesWithDB>>;
252
253 fn download(
254 responses: impl IntoIterator<Item = PeerRequestResult<SnapResponse>>,
255 factory: ProviderFactory<MockNodeTypesWithDB>,
256 ) -> (Arc<ScriptedSnapClient>, Download) {
257 let client = Arc::new(ScriptedSnapClient::new(responses));
258 let download = AccountRangeDownload::new(Arc::clone(&client), factory, Runtime::test());
259 (client, download)
260 }
261
262 fn stored_accounts(factory: &ProviderFactory<MockNodeTypesWithDB>) -> Vec<B256> {
263 let provider = factory.database_provider_ro().unwrap();
264 let mut cursor = provider.tx_ref().cursor_read::<tables::HashedAccounts>().unwrap();
265 cursor.walk(None).unwrap().map(|entry| entry.unwrap().0).collect()
266 }
267
268 async fn verified(download: &mut Download) -> VerifiedRange {
269 match download.next().await.unwrap().unwrap() {
270 AccountRangeStep::Verified(verified) => verified,
271 AccountRangeStep::Unavailable { .. } => panic!("fixture serves the range"),
272 }
273 }
274
275 #[tokio::test]
276 async fn downloads_the_trie_in_key_order_and_stops() {
277 let accounts = accounts();
278 let factory = started(&accounts);
279 let responses = [
280 account_range(1, &accounts, 0..1, &[key(1)]),
282 account_range(2, &accounts, 1..3, &[key(2), FAR]),
284 ];
285 let (client, mut download) = download(responses, factory.clone());
286
287 let first = verified(&mut download).await;
288 assert_eq!(first.origin(), B256::ZERO);
289 assert_eq!(first.range().accounts().len(), 1);
290 let coverage = download.commit(first, Default::default(), Vec::new()).await.unwrap();
291 assert_eq!(coverage.next(), Some(key(2)));
292
293 let second = verified(&mut download).await;
294 assert_eq!(second.origin(), key(2));
295 let coverage = download.commit(second, Default::default(), Vec::new()).await.unwrap();
296 assert!(coverage.is_complete());
297
298 assert!(download.next().await.unwrap().is_none());
299 assert_eq!(*client.origins(), [B256::ZERO, key(2)]);
300 assert_eq!(stored_accounts(&factory), [key(1), key(2), FAR]);
301 }
302
303 #[tokio::test]
304 async fn an_unavailable_response_names_the_peer_and_leaves_the_cursor_in_place() {
305 let accounts = accounts();
306 let factory = started(&accounts);
307 let peer = PeerId::random();
308 let empty = AccountRangeMessage { request_id: 1, accounts: Vec::new(), proof: Vec::new() };
309 let responses = [
310 Ok(WithPeerId::new(peer, SnapResponse::AccountRange(empty))),
311 account_range(2, &accounts, 0..3, &[]),
312 ];
313 let (client, mut download) = download(responses, factory.clone());
314
315 let step = download.next().await.unwrap().unwrap();
316
317 assert!(matches!(
318 step,
319 AccountRangeStep::Unavailable { origin, peer_id }
320 if origin == B256::ZERO && peer_id == peer
321 ));
322 assert_eq!(download.coverage(), Some(AccountCoverage::START));
323
324 let range = verified(&mut download).await;
325 download.commit(range, Default::default(), Vec::new()).await.unwrap();
326 assert!(download.coverage().unwrap().is_complete());
327 assert_eq!(*client.origins(), [B256::ZERO, B256::ZERO]);
328 }
329
330 #[tokio::test]
331 async fn a_failed_request_leaves_the_cursor_in_place() {
332 let factory = started(&accounts());
333 let (_, mut download) =
334 download([Err(RequestError::UnsupportedCapability)], factory.clone());
335
336 let error = download.next().await.unwrap_err();
337
338 assert!(matches!(error, SnapSyncError::Request(RequestError::UnsupportedCapability)));
339 assert_eq!(download.coverage(), Some(AccountCoverage::START));
340 assert!(stored_accounts(&factory).is_empty());
341 }
342
343 #[tokio::test]
344 async fn the_download_continues_from_the_coverage_the_store_records() {
345 let accounts = accounts();
346 let factory = started(&accounts);
347 let provider = factory.database_provider_rw().unwrap();
349 let write = provider.active_snap_write().unwrap().unwrap();
350 let head = verified_range(&accounts, 0..1, B256::ZERO, &[key(1)]);
351 provider.commit_account_range(write, &head, Default::default(), Vec::new()).unwrap();
352 provider.commit().unwrap();
353 let (client, mut download) =
354 download([account_range(1, &accounts, 1..3, &[key(2), FAR])], factory.clone());
355
356 let rest = verified(&mut download).await;
357
358 assert_eq!(rest.origin(), key(2));
359 assert_eq!(*client.origins(), [key(2)]);
360 download.commit(rest, Default::default(), Vec::new()).await.unwrap();
361 assert!(download.coverage().unwrap().is_complete());
362 assert_eq!(stored_accounts(&factory), [key(1), key(2), FAR]);
363 }
364
365 #[tokio::test]
366 async fn a_range_fetched_under_a_replaced_attempt_is_refused_at_commit() {
367 let accounts = accounts();
368 let factory = started(&accounts);
369 let (_, mut download) = download([account_range(1, &accounts, 0..3, &[])], factory.clone());
370 let range = verified(&mut download).await;
371 let provider = factory.database_provider_rw().unwrap();
373 provider.start_snap_attempt(generation(1, state_root(&accounts))).unwrap();
374 provider.commit().unwrap();
375
376 let error = download.commit(range, Default::default(), Vec::new()).await.unwrap_err();
377
378 assert!(matches!(error, SnapSyncError::StaleWrite { .. }));
379 assert!(stored_accounts(&factory).is_empty());
380 assert_eq!(download.coverage(), Some(AccountCoverage::START));
381 }
382
383 #[tokio::test]
384 async fn a_range_missing_its_dependencies_is_refused_at_commit() {
385 let mut accounts = accounts();
386 accounts[1].1.code_hash = B256::repeat_byte(0x33);
387 let factory = started(&accounts);
388 let (_, mut download) = download([account_range(1, &accounts, 0..3, &[])], factory.clone());
389 let range = verified(&mut download).await;
390
391 let error = download.commit(range, Default::default(), Vec::new()).await.unwrap_err();
392
393 assert!(matches!(error, SnapSyncError::MissingCode { .. }));
394 assert!(stored_accounts(&factory).is_empty());
395 assert_eq!(download.coverage(), Some(AccountCoverage::START));
396 }
397
398 #[tokio::test]
399 async fn nothing_is_requested_without_an_active_attempt() {
400 let (client, mut download) = download([], hashed_factory());
401
402 assert!(matches!(download.next().await, Err(SnapSyncError::NoAttempt)));
403 assert!(client.origins().is_empty());
404 assert_eq!(download.coverage(), None);
405 }
406
407 #[tokio::test]
408 async fn nothing_is_requested_without_recorded_coverage() {
409 let factory = hashed_factory();
410 let provider = factory.database_provider_rw().unwrap();
411 provider.start_snap_attempt(generation(1, state_root(&accounts()))).unwrap();
412 provider.commit().unwrap();
413 let (client, mut download) = download([], factory);
414
415 assert!(matches!(download.next().await, Err(SnapSyncError::NoCoverage)));
416 assert!(client.origins().is_empty());
417 }
418
419 #[tokio::test]
420 async fn a_complete_coverage_requests_nothing() {
421 let accounts = accounts();
422 let factory = started(&accounts);
423 let provider = factory.database_provider_rw().unwrap();
424 let write = provider.active_snap_write().unwrap().unwrap();
425 let whole = verified_range(&accounts, 0..3, B256::ZERO, &[]);
426 provider.commit_account_range(write, &whole, Default::default(), Vec::new()).unwrap();
427 provider.commit().unwrap();
428 let (client, mut download) = download([], factory);
429
430 assert!(download.next().await.unwrap().is_none());
431 assert_eq!(download.coverage(), Some(AccountCoverage::COMPLETE));
432 assert!(client.origins().is_empty());
433 }
434
435 #[tokio::test]
436 async fn repairs_are_fetched_one_account_at_a_time() {
437 let accounts = accounts();
438 let factory = started(&accounts);
439 let responses = [account_range(1, &accounts, 1..2, &[key(2)])];
440 let (client, mut download) = download(responses, factory.clone());
441 assert!(download.next_repair().await.unwrap().is_none());
442
443 let provider = factory.database_provider_rw().unwrap();
444 let write = provider.active_snap_write().unwrap().unwrap();
445 let mut repairs = StateRepairs::default();
446 repairs.insert_account(key(2));
447 provider.schedule_snap_repairs(write, repairs).unwrap();
448 provider.commit().unwrap();
449
450 let Some(AccountRangeStep::Verified(range)) = download.next_repair().await.unwrap() else {
451 panic!("fixture serves the account")
452 };
453 assert_eq!(*client.origins(), [key(2)]);
454 assert_eq!(range.range().accounts(), &accounts[1..2]);
455 }
456}