1use alloy_primitives::{Address, B256, U256};
2use core::ops::{Deref, DerefMut};
3use reth_primitives_traits::Account;
4use reth_storage_api::{AccountReader, BytecodeReader, EvmStateProvider};
5use reth_storage_errors::provider::{ProviderError, ProviderResult};
6use revm::{bytecode::Bytecode, state::AccountInfo, Database, DatabaseRef};
7
8#[derive(Clone)]
11pub struct StateProviderDatabase<DB>(pub DB);
12
13impl<DB> StateProviderDatabase<DB> {
14 pub const fn new(db: DB) -> Self {
16 Self(db)
17 }
18
19 pub fn into_inner(self) -> DB {
21 self.0
22 }
23}
24
25impl<DB> core::fmt::Debug for StateProviderDatabase<DB> {
26 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
27 f.debug_struct("StateProviderDatabase").finish_non_exhaustive()
28 }
29}
30
31impl<DB> AsRef<DB> for StateProviderDatabase<DB> {
32 fn as_ref(&self) -> &DB {
33 self
34 }
35}
36
37impl<DB> Deref for StateProviderDatabase<DB> {
38 type Target = DB;
39
40 fn deref(&self) -> &Self::Target {
41 &self.0
42 }
43}
44
45impl<DB> DerefMut for StateProviderDatabase<DB> {
46 fn deref_mut(&mut self) -> &mut Self::Target {
47 &mut self.0
48 }
49}
50
51impl<DB: EvmStateProvider> Database for StateProviderDatabase<DB> {
52 type Error = ProviderError;
53
54 fn basic(&mut self, address: Address) -> Result<Option<AccountInfo>, Self::Error> {
59 self.basic_ref(address)
60 }
61
62 fn code_by_hash(&mut self, code_hash: B256) -> Result<Bytecode, Self::Error> {
66 self.code_by_hash_ref(code_hash)
67 }
68
69 fn storage(&mut self, address: Address, index: U256) -> Result<U256, Self::Error> {
73 self.storage_ref(address, index)
74 }
75
76 fn block_hash(&mut self, number: u64) -> Result<B256, Self::Error> {
81 self.block_hash_ref(number)
82 }
83}
84
85impl<DB: EvmStateProvider> DatabaseRef for StateProviderDatabase<DB> {
86 type Error = <Self as Database>::Error;
87
88 fn basic_ref(&self, address: Address) -> Result<Option<AccountInfo>, Self::Error> {
93 Ok(self.basic_account(&address)?.map(Into::into))
94 }
95
96 fn code_by_hash_ref(&self, code_hash: B256) -> Result<Bytecode, Self::Error> {
100 Ok(self.bytecode_by_hash(&code_hash)?.unwrap_or_default().0)
101 }
102
103 fn storage_ref(&self, address: Address, index: U256) -> Result<U256, Self::Error> {
107 Ok(self.0.storage(address, B256::new(index.to_be_bytes()))?.unwrap_or_default())
108 }
109
110 fn block_hash_ref(&self, number: u64) -> Result<B256, Self::Error> {
114 Ok(self.0.block_hash(number)?.unwrap_or_default())
116 }
117}
118
119#[derive(Clone)]
129pub struct DatabaseStateProvider<DB>(pub DB);
130
131impl<DB> DatabaseStateProvider<DB> {
132 pub const fn new(db: DB) -> Self {
134 Self(db)
135 }
136
137 pub fn into_inner(self) -> DB {
139 self.0
140 }
141
142 pub const fn inner(&self) -> &DB {
144 &self.0
145 }
146}
147
148impl<DB> core::fmt::Debug for DatabaseStateProvider<DB> {
149 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
150 f.debug_struct("DatabaseStateProvider").finish_non_exhaustive()
151 }
152}
153
154impl<DB> AccountReader for DatabaseStateProvider<DB>
155where
156 DB: DatabaseRef<Error = ProviderError>,
157{
158 fn basic_account(&self, address: &Address) -> ProviderResult<Option<Account>> {
159 Ok(self.0.basic_ref(*address)?.map(Into::into))
160 }
161}
162
163impl<DB> BytecodeReader for DatabaseStateProvider<DB>
164where
165 DB: DatabaseRef<Error = ProviderError>,
166{
167 fn bytecode_by_hash(
168 &self,
169 code_hash: &B256,
170 ) -> ProviderResult<Option<reth_primitives_traits::Bytecode>> {
171 Ok(Some(reth_primitives_traits::Bytecode(self.0.code_by_hash_ref(*code_hash)?)))
172 }
173}
174
175#[cfg(test)]
176mod tests {
177 use super::*;
178 use crate::cached::CachedReads;
179 use alloy_consensus::constants::KECCAK_EMPTY;
180 use alloy_primitives::Bytes;
181 use core::sync::atomic::{AtomicUsize, Ordering};
182 use std::sync::Arc;
183
184 #[derive(Clone)]
185 struct CountingDatabaseRef {
186 address: Address,
187 code_hash: B256,
188 account: Option<AccountInfo>,
189 bytecode: Bytecode,
190 account_reads: Arc<AtomicUsize>,
191 bytecode_reads: Arc<AtomicUsize>,
192 fail_account_reads: bool,
193 fail_bytecode_reads: bool,
194 }
195
196 impl CountingDatabaseRef {
197 fn new(address: Address, account: Option<AccountInfo>, bytecode: Bytecode) -> Self {
198 let code_hash = account.as_ref().map(|account| account.code_hash).unwrap_or_default();
199 Self {
200 address,
201 code_hash,
202 account,
203 bytecode,
204 account_reads: Arc::default(),
205 bytecode_reads: Arc::default(),
206 fail_account_reads: false,
207 fail_bytecode_reads: false,
208 }
209 }
210 }
211
212 impl DatabaseRef for CountingDatabaseRef {
213 type Error = ProviderError;
214
215 fn basic_ref(&self, address: Address) -> Result<Option<AccountInfo>, Self::Error> {
216 if self.fail_account_reads {
217 return Err(ProviderError::UnsupportedProvider)
218 }
219
220 self.account_reads.fetch_add(1, Ordering::Relaxed);
221 Ok((address == self.address).then(|| self.account.clone()).flatten())
222 }
223
224 fn code_by_hash_ref(&self, code_hash: B256) -> Result<Bytecode, Self::Error> {
225 if self.fail_bytecode_reads {
226 return Err(ProviderError::UnsupportedProvider)
227 }
228
229 self.bytecode_reads.fetch_add(1, Ordering::Relaxed);
230 Ok(if code_hash == self.code_hash {
231 self.bytecode.clone()
232 } else {
233 Bytecode::default()
234 })
235 }
236
237 fn storage_ref(&self, _address: Address, _index: U256) -> Result<U256, Self::Error> {
238 Ok(U256::ZERO)
239 }
240
241 fn block_hash_ref(&self, _number: u64) -> Result<B256, Self::Error> {
242 Ok(B256::ZERO)
243 }
244 }
245
246 #[test]
247 fn database_state_provider_maps_missing_account() {
248 let address = Address::repeat_byte(0x01);
249 let db = CountingDatabaseRef::new(address, None, Bytecode::default());
250 let provider = DatabaseStateProvider::new(db);
251
252 assert_eq!(provider.basic_account(&address).unwrap(), None);
253 }
254
255 #[test]
256 fn database_state_provider_maps_empty_code_hash() {
257 let address = Address::repeat_byte(0x01);
258 let account = AccountInfo {
259 nonce: 7,
260 balance: U256::from(42),
261 code_hash: KECCAK_EMPTY,
262 code: None,
263 ..Default::default()
264 };
265 let db = CountingDatabaseRef::new(address, Some(account), Bytecode::default());
266 let provider = DatabaseStateProvider::new(db);
267
268 assert_eq!(
269 provider.basic_account(&address).unwrap(),
270 Some(Account { nonce: 7, balance: U256::from(42), ..Default::default() })
271 );
272 }
273
274 #[test]
275 #[allow(clippy::needless_update)]
276 fn database_state_provider_maps_code_hash_and_bytecode() {
277 let address = Address::repeat_byte(0x01);
278 let code_hash = B256::repeat_byte(0x42);
279 let bytecode = Bytecode::new_raw(Bytes::from_static(&[0x60, 0x00]));
280 let account = AccountInfo::new(U256::from(42), 7, code_hash, bytecode.clone());
281 let db = CountingDatabaseRef::new(address, Some(account), bytecode.clone());
282 let provider = DatabaseStateProvider::new(db);
283
284 assert_eq!(
285 provider.basic_account(&address).unwrap(),
286 Some(Account {
287 nonce: 7,
288 balance: U256::from(42),
289 bytecode_hash: Some(code_hash),
290 ..Default::default()
291 })
292 );
293 assert_eq!(
294 provider.bytecode_by_hash(&code_hash).unwrap(),
295 Some(reth_primitives_traits::Bytecode(bytecode))
296 );
297 }
298
299 #[test]
300 fn database_state_provider_wraps_default_bytecode_for_unknown_hash() {
301 let address = Address::repeat_byte(0x01);
302 let unknown_hash = B256::repeat_byte(0x42);
303 let db = CountingDatabaseRef::new(address, None, Bytecode::default());
304 let provider = DatabaseStateProvider::new(db);
305
306 assert_eq!(
307 provider.bytecode_by_hash(&unknown_hash).unwrap(),
308 Some(reth_primitives_traits::Bytecode(Bytecode::default()))
309 );
310 }
311
312 #[test]
313 fn database_state_provider_propagates_database_errors() {
314 let address = Address::repeat_byte(0x01);
315 let code_hash = B256::repeat_byte(0x42);
316 let db = CountingDatabaseRef {
317 fail_account_reads: true,
318 fail_bytecode_reads: true,
319 ..CountingDatabaseRef::new(address, None, Bytecode::default())
320 };
321 let provider = DatabaseStateProvider::new(db);
322
323 assert!(matches!(
324 provider.basic_account(&address),
325 Err(ProviderError::UnsupportedProvider)
326 ));
327 assert!(matches!(
328 provider.bytecode_by_hash(&code_hash),
329 Err(ProviderError::UnsupportedProvider)
330 ));
331 }
332
333 #[test]
334 #[allow(clippy::needless_update)]
335 fn database_state_provider_uses_cached_reads() {
336 let address = Address::repeat_byte(0x01);
337 let code_hash = B256::repeat_byte(0x42);
338 let bytecode = Bytecode::new_raw(Bytes::from_static(&[0x60, 0x00]));
339 let account = AccountInfo::new(U256::from(42), 7, code_hash, bytecode.clone());
340 let db = CountingDatabaseRef::new(address, Some(account), bytecode.clone());
341 let account_reads = db.account_reads.clone();
342 let bytecode_reads = db.bytecode_reads.clone();
343 let mut cached_reads = CachedReads::default();
344 let provider = DatabaseStateProvider::new(cached_reads.as_db(db));
345
346 assert_eq!(
347 provider.basic_account(&address).unwrap(),
348 Some(Account {
349 nonce: 7,
350 balance: U256::from(42),
351 bytecode_hash: Some(code_hash),
352 ..Default::default()
353 })
354 );
355 assert_eq!(
356 provider.basic_account(&address).unwrap(),
357 Some(Account {
358 nonce: 7,
359 balance: U256::from(42),
360 bytecode_hash: Some(code_hash),
361 ..Default::default()
362 })
363 );
364 assert_eq!(account_reads.load(Ordering::Relaxed), 1);
365
366 assert_eq!(
367 provider.bytecode_by_hash(&code_hash).unwrap(),
368 Some(reth_primitives_traits::Bytecode(bytecode.clone()))
369 );
370 assert_eq!(
371 provider.bytecode_by_hash(&code_hash).unwrap(),
372 Some(reth_primitives_traits::Bytecode(bytecode))
373 );
374 assert_eq!(bytecode_reads.load(Ordering::Relaxed), 1);
375 }
376}