Skip to main content

reth_revm/
database.rs

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/// A [Database] and [`DatabaseRef`] implementation that uses [`EvmStateProvider`] as the underlying
9/// data source.
10#[derive(Clone)]
11pub struct StateProviderDatabase<DB>(pub DB);
12
13impl<DB> StateProviderDatabase<DB> {
14    /// Creates a database backed by an [`EvmStateProvider`].
15    pub const fn new(db: DB) -> Self {
16        Self(db)
17    }
18
19    /// Consumes the database and returns its inner provider.
20    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    /// Retrieves basic account information for a given address.
55    ///
56    /// Returns `Ok` with `Some(AccountInfo)` if the account exists,
57    /// `None` if it doesn't, or an error if encountered.
58    fn basic(&mut self, address: Address) -> Result<Option<AccountInfo>, Self::Error> {
59        self.basic_ref(address)
60    }
61
62    /// Retrieves the bytecode associated with a given code hash.
63    ///
64    /// Returns `Ok` with the bytecode if found, or the default bytecode otherwise.
65    fn code_by_hash(&mut self, code_hash: B256) -> Result<Bytecode, Self::Error> {
66        self.code_by_hash_ref(code_hash)
67    }
68
69    /// Retrieves the storage value at a specific index for a given address.
70    ///
71    /// Returns `Ok` with the storage value, or the default value if not found.
72    fn storage(&mut self, address: Address, index: U256) -> Result<U256, Self::Error> {
73        self.storage_ref(address, index)
74    }
75
76    /// Retrieves the block hash for a given block number.
77    ///
78    /// Returns `Ok` with the block hash if found, or the default hash otherwise.
79    /// Note: It safely casts the `number` to `u64`.
80    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    /// Retrieves basic account information for a given address.
89    ///
90    /// Returns `Ok` with `Some(AccountInfo)` if the account exists,
91    /// `None` if it doesn't, or an error if encountered.
92    fn basic_ref(&self, address: Address) -> Result<Option<AccountInfo>, Self::Error> {
93        Ok(self.basic_account(&address)?.map(Into::into))
94    }
95
96    /// Retrieves the bytecode associated with a given code hash.
97    ///
98    /// Returns `Ok` with the bytecode if found, or the default bytecode otherwise.
99    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    /// Retrieves the storage value at a specific index for a given address.
104    ///
105    /// Returns `Ok` with the storage value, or the default value if not found.
106    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    /// Retrieves the block hash for a given block number.
111    ///
112    /// Returns `Ok` with the block hash if found, or the default hash otherwise.
113    fn block_hash_ref(&self, number: u64) -> Result<B256, Self::Error> {
114        // Get the block hash or default hash with an attempt to convert U256 block number to u64
115        Ok(self.0.block_hash(number)?.unwrap_or_default())
116    }
117}
118
119/// A [`DatabaseRef`] backed account info reader.
120///
121/// This adapts account and bytecode reads from revm's database interface back into Reth's storage
122/// reader traits. It is intentionally not a full [`reth_storage_api::StateProvider`] because
123/// [`DatabaseRef`] does not expose roots or proofs.
124///
125/// Note: [`DatabaseRef::code_by_hash_ref`] returns [`Bytecode`] directly, so this adapter cannot
126/// distinguish missing bytecode from the database's default bytecode and wraps whatever the
127/// database returns in `Some`.
128#[derive(Clone)]
129pub struct DatabaseStateProvider<DB>(pub DB);
130
131impl<DB> DatabaseStateProvider<DB> {
132    /// Create a new database-backed state reader.
133    pub const fn new(db: DB) -> Self {
134        Self(db)
135    }
136
137    /// Consume self and return the inner database.
138    pub fn into_inner(self) -> DB {
139        self.0
140    }
141
142    /// Returns the inner database.
143    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}