1use crate::{common::DownloadContext, SnapBytecodeStore, SnapSyncError, VerifiedRange};
4use reth_db_api::transaction::DbTxMut;
5use reth_downloaders::snap::{BytecodeDownloader, BytecodeOutcome};
6use reth_eth_wire_types::snap::GetByteCodesMessage;
7use reth_network_p2p::snap::client::SnapClient;
8use reth_network_peers::PeerId;
9use reth_storage_api::{
10 DBProvider, DatabaseProviderFactory, MetadataProvider, MetadataWriter, StateWriter,
11};
12use reth_tasks::Runtime;
13use std::fmt;
14
15pub const DEFAULT_CODE_HASHES: usize = 128;
17
18pub struct BytecodeDownload<C, F> {
23 context: DownloadContext<C, F>,
24 max_hashes: usize,
26}
27
28impl<C, F> BytecodeDownload<C, F> {
29 pub const fn new(client: C, factory: F, runtime: Runtime) -> Self {
31 Self {
32 context: DownloadContext::new(client, factory, runtime),
33 max_hashes: DEFAULT_CODE_HASHES,
34 }
35 }
36
37 pub const fn with_response_bytes(mut self, response_bytes: u64) -> Self {
39 self.context.set_response_bytes(response_bytes);
40 self
41 }
42
43 pub const fn with_max_hashes(mut self, max_hashes: usize) -> Self {
45 self.max_hashes = if max_hashes == 0 { 1 } else { max_hashes };
46 self
47 }
48}
49
50impl<C, F> BytecodeDownload<C, F>
51where
52 C: SnapClient + Clone + Unpin,
53 F: DatabaseProviderFactory + Clone + 'static,
54 F::Provider: MetadataProvider,
55 F::ProviderRW: MetadataProvider + MetadataWriter + StateWriter + DBProvider<Tx: DbTxMut>,
56{
57 pub async fn next(&mut self, range: &VerifiedRange) -> Result<BytecodeStep, SnapSyncError> {
62 let write = range.write();
63 let referenced = range.range().code_hashes();
64 if referenced.is_empty() {
65 return Ok(BytecodeStep::Complete)
66 }
67 let limit = self.max_hashes;
68 let missing = self
69 .context
70 .read(move |provider| provider.missing_code(write, &referenced, limit))
71 .await?;
72 if missing.is_empty() {
73 return Ok(BytecodeStep::Complete)
74 }
75
76 let request = GetByteCodesMessage {
77 request_id: self.context.next_request_id(),
78 hashes: missing,
79 response_bytes: self.context.response_bytes(),
80 };
81 let downloader = BytecodeDownloader::new(
82 self.context.client().clone(),
83 request,
84 self.context.runtime().clone(),
85 )
86 .expect("missing code is never empty");
87 let codes = match downloader.await? {
88 BytecodeOutcome::Verified(verified) => verified.into_codes(),
89 BytecodeOutcome::Unavailable { peer_id } => {
90 return Ok(BytecodeStep::Unavailable { peer_id })
91 }
92 };
93
94 let persisted =
95 self.context.commit(move |provider| provider.commit_bytecodes(write, codes)).await?;
96 Ok(BytecodeStep::Committed { persisted })
97 }
98}
99
100impl<C, F> fmt::Debug for BytecodeDownload<C, F> {
101 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
102 f.debug_struct("BytecodeDownload")
103 .field("context", &self.context)
104 .field("max_hashes", &self.max_hashes)
105 .finish()
106 }
107}
108
109#[derive(Debug)]
111pub enum BytecodeStep {
112 Committed {
114 persisted: usize,
116 },
117 Unavailable {
119 peer_id: PeerId,
121 },
122 Complete,
124}
125
126#[cfg(test)]
127mod tests {
128 use super::*;
129 use crate::{
130 test_utils::{
131 account, byte_codes, generation, hashed_factory, insert_generation_headers, key,
132 state_root, verified_range, verified_repair, ScriptedSnapClient,
133 },
134 SnapAccountStore, SnapAttemptStore,
135 };
136 use alloy_primitives::{keccak256, Bytes, B256};
137 use reth_db_api::{tables, transaction::DbTx};
138 use reth_eth_wire_types::snap::ByteCodesMessage;
139 use reth_network_p2p::{error::PeerRequestResult, snap::client::SnapResponse};
140 use reth_network_peers::WithPeerId;
141 use reth_provider::{test_utils::MockNodeTypesWithDB, ProviderFactory};
142 use reth_trie_common::TrieAccount;
143 use std::sync::Arc;
144
145 type Factory = ProviderFactory<MockNodeTypesWithDB>;
146 type Download = BytecodeDownload<Arc<ScriptedSnapClient>, Factory>;
147
148 fn code(byte: u8) -> Bytes {
149 Bytes::from(vec![byte; 4])
150 }
151
152 fn contract(nonce: u64, code: &Bytes) -> TrieAccount {
153 let mut contract = account(nonce);
154 contract.code_hash = keccak256(code);
155 contract
156 }
157
158 fn accounts() -> Vec<(B256, TrieAccount)> {
160 vec![
161 (key(1), contract(1, &code(1))),
162 (key(2), account(2)),
163 (key(3), contract(3, &code(1))),
164 (key(4), contract(4, &code(2))),
165 ]
166 }
167
168 fn started(accounts: &[(B256, TrieAccount)]) -> (Factory, VerifiedRange) {
170 let factory = hashed_factory();
171 insert_generation_headers(&factory);
172 let provider = factory.database_provider_rw().unwrap();
173 let write = provider.start_snap_attempt(generation(1, state_root(accounts))).unwrap();
174 provider.start_account_coverage(write).unwrap();
175 provider.commit().unwrap();
176 let range = verified_range(accounts, 0..accounts.len(), B256::ZERO, &[]);
177 (factory, VerifiedRange::new(write, range))
178 }
179
180 fn download(
181 responses: impl IntoIterator<Item = PeerRequestResult<SnapResponse>>,
182 factory: Factory,
183 ) -> (Arc<ScriptedSnapClient>, Download) {
184 let client = Arc::new(ScriptedSnapClient::new(responses));
185 (Arc::clone(&client), BytecodeDownload::new(client, factory, Runtime::test()))
186 }
187
188 async fn committed(download: &mut Download, range: &VerifiedRange) -> usize {
189 match download.next(range).await.unwrap() {
190 BytecodeStep::Committed { persisted } => persisted,
191 step => panic!("expected a commit, got {step:?}"),
192 }
193 }
194
195 fn stored(factory: &Factory, code: &Bytes) -> bool {
196 let provider = factory.database_provider_ro().unwrap();
197 provider.tx_ref().get::<tables::Bytecodes>(keccak256(code)).unwrap().is_some()
198 }
199
200 #[tokio::test]
201 async fn shared_code_is_requested_once_and_persisted_for_every_account() {
202 let accounts = accounts();
203 let (factory, range) = started(&accounts);
204 let (client, mut download) =
205 download([byte_codes(1, &[code(1), code(2)])], factory.clone());
206
207 assert_eq!(committed(&mut download, &range).await, 2);
208
209 assert!(matches!(download.next(&range).await.unwrap(), BytecodeStep::Complete));
210 assert_eq!(*client.code_requests(), [vec![keccak256(code(1)), keccak256(code(2))]]);
211 assert!(stored(&factory, &code(1)) && stored(&factory, &code(2)));
212
213 let provider = factory.database_provider_rw().unwrap();
215 let coverage = provider
216 .commit_account_range(range.write(), range.range(), Default::default(), Vec::new())
217 .unwrap();
218 provider.commit().unwrap();
219 assert!(coverage.is_complete());
220 }
221
222 #[tokio::test]
223 async fn a_repair_requests_only_the_code_of_its_account() {
224 let accounts = accounts();
225 let (factory, range) = started(&accounts);
226 let write = range.write();
227 let (client, mut download) = download([byte_codes(1, &[code(2)])], factory);
228
229 let plain =
231 VerifiedRange::new(write, verified_repair(&accounts, 1..3, key(2), &[key(2), key(3)]));
232 assert!(matches!(download.next(&plain).await.unwrap(), BytecodeStep::Complete));
233 let contract =
234 VerifiedRange::new(write, verified_repair(&accounts, 3..4, key(4), &[key(4)]));
235 assert!(matches!(
236 download.next(&contract).await.unwrap(),
237 BytecodeStep::Committed { persisted: 1 }
238 ));
239
240 assert_eq!(*client.code_requests(), [vec![keccak256(code(2))]]);
241 }
242
243 #[tokio::test]
244 async fn code_already_stored_is_never_requested() {
245 let accounts = accounts();
246 let (factory, range) = started(&accounts);
247 let (_, mut first) = download([byte_codes(1, &[code(1), code(2)])], factory.clone());
248 committed(&mut first, &range).await;
249 drop(first);
250
251 let (client, mut resumed) = download([], factory);
252
253 assert!(matches!(resumed.next(&range).await.unwrap(), BytecodeStep::Complete));
254 assert!(client.code_requests().is_empty());
255 }
256
257 #[tokio::test]
258 async fn code_a_peer_does_not_have_stays_missing() {
259 let accounts = accounts();
260 let (factory, range) = started(&accounts);
261 let peer = PeerId::random();
262 let empty = ByteCodesMessage { request_id: 1, codes: Vec::new() };
263 let responses = [
264 Ok(WithPeerId::new(peer, SnapResponse::ByteCodes(empty))),
265 byte_codes(2, &[code(1)]),
267 byte_codes(3, &[code(2)]),
268 ];
269 let (client, mut download) = download(responses, factory.clone());
270
271 let step = download.next(&range).await.unwrap();
272
273 assert!(matches!(step, BytecodeStep::Unavailable { peer_id } if peer_id == peer));
274 assert!(!stored(&factory, &code(1)));
275 assert_eq!(committed(&mut download, &range).await, 1);
276 assert_eq!(committed(&mut download, &range).await, 1);
278 assert!(matches!(download.next(&range).await.unwrap(), BytecodeStep::Complete));
279 assert_eq!(
280 *client.code_requests(),
281 [
282 vec![keccak256(code(1)), keccak256(code(2))],
283 vec![keccak256(code(1)), keccak256(code(2))],
284 vec![keccak256(code(2))],
285 ]
286 );
287 }
288
289 #[tokio::test]
290 async fn unresolved_code_prevents_the_range_from_committing() {
291 let accounts = accounts();
292 let (factory, range) = started(&accounts);
293 let (_, mut download) = download([byte_codes(1, &[code(1)])], factory.clone());
294 committed(&mut download, &range).await;
295
296 let provider = factory.database_provider_rw().unwrap();
297 let refused = provider.commit_account_range(
298 range.write(),
299 range.range(),
300 Default::default(),
301 Vec::new(),
302 );
303
304 assert!(matches!(
305 refused,
306 Err(SnapSyncError::MissingCode { hash }) if hash == keccak256(code(2))
307 ));
308 }
309
310 #[tokio::test]
311 async fn requests_ask_for_at_most_the_configured_hashes() {
312 let accounts = accounts();
313 let (factory, range) = started(&accounts);
314 let responses = [byte_codes(1, &[code(1)]), byte_codes(2, &[code(2)])];
315 let (client, download) = download(responses, factory);
316 let mut download = download.with_max_hashes(1);
317
318 committed(&mut download, &range).await;
319 committed(&mut download, &range).await;
320
321 assert!(matches!(download.next(&range).await.unwrap(), BytecodeStep::Complete));
322 assert_eq!(*client.code_requests(), [vec![keccak256(code(1))], vec![keccak256(code(2))]]);
323 }
324
325 #[tokio::test]
326 async fn a_range_without_contracts_requests_no_code() {
327 let accounts = vec![(key(1), account(1)), (key(2), account(2))];
328 let (factory, range) = started(&accounts);
329 let (client, mut download) = download([], factory);
330
331 assert!(matches!(download.next(&range).await.unwrap(), BytecodeStep::Complete));
332 assert!(client.code_requests().is_empty());
333 }
334
335 #[tokio::test]
336 async fn code_fetched_before_the_pivot_moved_is_not_committed() {
337 let accounts = accounts();
338 let (factory, range) = started(&accounts);
339 let (client, mut download) =
340 download([byte_codes(1, &[code(1), code(2)])], factory.clone());
341 let provider = factory.database_provider_rw().unwrap();
342 provider.advance_snap_pivot(range.write(), generation(2, state_root(&accounts))).unwrap();
343 provider.commit().unwrap();
344
345 assert!(matches!(download.next(&range).await, Err(SnapSyncError::StaleWrite { .. })));
346 assert!(client.code_requests().is_empty());
347 }
348}