1use core::{
3 error::Error,
4 fmt::Display,
5 ops::{Deref, DerefMut},
6};
7use primitives::{Address, StorageKey, StorageValue, B256};
8use state::{
9 bal::{alloy::AlloyBal, Bal, BalAccountLookup, BalError, BlockAccessIndex},
10 Account, AccountId, AccountInfo, Bytecode, EvmState,
11};
12use std::sync::Arc;
13
14use crate::{DBErrorMarker, Database, DatabaseCommit};
15
16#[derive(Clone, Default, Debug, PartialEq, Eq)]
18#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
19pub struct BalState {
20 pub bal: Option<Arc<Bal>>,
22 pub bal_builder: Option<Bal>,
25 pub bal_index: BlockAccessIndex,
28 #[cfg_attr(feature = "serde", serde(default))]
36 pub allow_db_fallback: bool,
37}
38
39impl BalState {
40 #[inline]
42 pub fn new() -> Self {
43 Self::default()
44 }
45
46 #[inline]
48 pub const fn reset_bal_index(&mut self) {
49 self.bal_index = BlockAccessIndex::PRE_EXECUTION;
50 }
51
52 #[inline]
54 pub const fn bump_bal_index(&mut self) {
55 self.bal_index.increment();
56 }
57
58 #[inline]
60 pub const fn bal_index(&self) -> BlockAccessIndex {
61 self.bal_index
62 }
63
64 #[inline]
66 pub fn bal(&self) -> Option<Arc<Bal>> {
67 self.bal.clone()
68 }
69
70 #[inline]
72 pub fn bal_builder(&self) -> Option<Bal> {
73 self.bal_builder.clone()
74 }
75
76 #[inline]
78 pub fn with_bal(mut self, bal: Arc<Bal>) -> Self {
79 self.bal = Some(bal);
80 self
81 }
82
83 #[inline]
85 pub fn with_bal_builder(mut self) -> Self {
86 self.bal_builder = Some(Bal::new());
87 self
88 }
89
90 #[inline]
94 pub const fn with_allow_db_fallback(mut self, allow: bool) -> Self {
95 self.allow_db_fallback = allow;
96 self
97 }
98
99 #[inline]
103 pub const fn set_allow_db_fallback(&mut self, allow: bool) {
104 self.allow_db_fallback = allow;
105 }
106
107 #[inline]
109 pub const fn take_built_bal(&mut self) -> Option<Bal> {
110 self.reset_bal_index();
111 self.bal_builder.take()
112 }
113
114 #[inline]
116 pub fn take_built_alloy_bal(&mut self) -> Option<AlloyBal> {
117 self.take_built_bal().map(|bal| bal.into_alloy_bal())
118 }
119
120 #[inline]
128 pub fn get_account_id(&self, address: &Address) -> Result<Option<AccountId>, BalError> {
129 let Some(bal) = self.bal.as_ref() else {
130 return Ok(None);
131 };
132 match bal.accounts.get_full(address) {
133 Some(i) => Ok(Some(AccountId::new(i.0).expect("too many bals"))),
134 None if self.allow_db_fallback => Ok(None),
135 None => Err(BalError::AccountNotFound { address: *address }),
136 }
137 }
138
139 #[inline]
145 pub fn basic(
146 &self,
147 address: Address,
148 basic: &mut Option<AccountInfo>,
149 ) -> Result<bool, BalError> {
150 let Some(account_id) = self.get_account_id(&address)? else {
151 return Ok(false);
152 };
153 self.basic_by_account_id(account_id, basic)
154 }
155
156 #[inline]
158 pub fn basic_by_account_id(
159 &self,
160 account_id: AccountId,
161 basic: &mut Option<AccountInfo>,
162 ) -> Result<bool, BalError> {
163 let Some(bal) = &self.bal else {
164 return Ok(false);
165 };
166 let is_none = basic.is_none();
167 let mut bal_basic = core::mem::take(basic).unwrap_or_default();
168 let changed = bal.populate_account_info(account_id, self.bal_index, &mut bal_basic)?;
169
170 if !changed && is_none {
172 return Ok(true);
173 }
174
175 *basic = Some(bal_basic);
176 Ok(true)
177 }
178
179 #[inline]
200 pub fn get_bal_account_info(&self, address: &Address) -> Result<BalAccountLookup, BalError> {
201 let Some(bal) = &self.bal else {
202 return Ok(BalAccountLookup::NotCovered);
203 };
204 let Some(bal_account) = bal.accounts.get(address) else {
205 if self.allow_db_fallback {
206 return Ok(BalAccountLookup::NotCovered);
207 }
208 return Err(BalError::AccountNotFound { address: *address });
209 };
210 Ok(bal_account.account_info.account_info_lookup(self.bal_index))
211 }
212
213 #[inline]
221 pub fn storage(
222 &self,
223 account: &Address,
224 storage_key: StorageKey,
225 ) -> Result<Option<StorageValue>, BalError> {
226 let Some(bal) = &self.bal else {
227 return Ok(None);
228 };
229
230 let Some(bal_account) = bal.accounts.get(account) else {
231 if self.allow_db_fallback {
232 return Ok(None);
233 }
234 return Err(BalError::AccountNotFound { address: *account });
235 };
236
237 match bal_account.storage.get_bal_writes(account, storage_key) {
238 Ok(writes) => Ok(writes.get(self.bal_index)),
239 Err(BalError::SlotNotFound { .. }) if self.allow_db_fallback => Ok(None),
240 Err(err) => Err(err),
241 }
242 }
243
244 #[inline]
252 pub fn storage_by_account_id(
253 &self,
254 account_id: AccountId,
255 storage_key: StorageKey,
256 ) -> Result<Option<StorageValue>, BalError> {
257 let Some(bal) = &self.bal else {
258 return Ok(None);
259 };
260
261 let Some((address, bal_account)) = bal.accounts.get_index(account_id.get()) else {
262 return Err(BalError::InvalidAccountId { account_id });
263 };
264
265 match bal_account.storage.get_bal_writes(address, storage_key) {
266 Ok(writes) => Ok(writes.get(self.bal_index)),
267 Err(BalError::SlotNotFound { .. }) if self.allow_db_fallback => Ok(None),
268 Err(err) => Err(err),
269 }
270 }
271
272 #[inline]
274 pub fn commit(&mut self, changes: &EvmState) {
275 if let Some(bal_builder) = &mut self.bal_builder {
276 for (address, account) in changes.iter() {
277 bal_builder.update_account(self.bal_index, *address, account);
278 }
279 }
280 }
281
282 #[inline]
284 pub fn commit_one(&mut self, address: Address, account: &Account) {
285 if let Some(bal_builder) = &mut self.bal_builder {
286 bal_builder.update_account(self.bal_index, address, account);
287 }
288 }
289}
290
291#[derive(Clone, Debug, PartialEq, Eq)]
293#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
294pub struct BalDatabase<DB> {
295 pub bal_state: BalState,
297 pub db: DB,
299}
300
301impl<DB> Deref for BalDatabase<DB> {
302 type Target = DB;
303
304 fn deref(&self) -> &Self::Target {
305 &self.db
306 }
307}
308
309impl<DB> DerefMut for BalDatabase<DB> {
310 fn deref_mut(&mut self) -> &mut Self::Target {
311 &mut self.db
312 }
313}
314
315impl<DB> BalDatabase<DB> {
316 #[inline]
318 pub fn new(db: DB) -> Self {
319 Self {
320 bal_state: BalState::default(),
321 db,
322 }
323 }
324
325 #[inline]
327 pub fn with_bal_option(self, bal: Option<Arc<Bal>>) -> Self {
328 Self {
329 bal_state: BalState {
330 bal,
331 ..self.bal_state
332 },
333 ..self
334 }
335 }
336
337 #[inline]
339 pub fn with_bal_builder(self) -> Self {
340 Self {
341 bal_state: self.bal_state.with_bal_builder(),
342 ..self
343 }
344 }
345
346 #[inline]
350 pub const fn with_allow_bal_db_fallback(mut self, allow: bool) -> Self {
351 self.bal_state.allow_db_fallback = allow;
352 self
353 }
354
355 #[inline]
357 pub const fn reset_bal_index(mut self) -> Self {
358 self.bal_state.reset_bal_index();
359 self
360 }
361
362 #[inline]
364 pub const fn bump_bal_index(&mut self) {
365 self.bal_state.bump_bal_index();
366 }
367}
368
369#[derive(Clone, Debug, PartialEq, Eq)]
371#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
372pub enum EvmDatabaseError<ERROR> {
373 Bal(BalError),
375 Database(ERROR),
377}
378
379impl<ERROR> From<BalError> for EvmDatabaseError<ERROR> {
380 fn from(error: BalError) -> Self {
381 Self::Bal(error)
382 }
383}
384
385impl<ERROR: core::error::Error + Send + Sync + 'static> DBErrorMarker for EvmDatabaseError<ERROR> {
386 fn is_fatal(&self) -> bool {
387 match self {
388 Self::Bal(_) => false,
389 Self::Database(_) => true,
390 }
391 }
392}
393
394impl<ERROR: Display> Display for EvmDatabaseError<ERROR> {
395 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
396 match self {
397 Self::Bal(error) => write!(f, "Bal error: {error}"),
398 Self::Database(error) => write!(f, "Database error: {error}"),
399 }
400 }
401}
402
403impl<ERROR: Error> Error for EvmDatabaseError<ERROR> {}
404
405impl<ERROR> EvmDatabaseError<ERROR> {
406 pub fn into_external_error(self) -> ERROR {
410 match self {
411 Self::Bal(_) => panic!("Expected database error, got BAL error"),
412 Self::Database(error) => error,
413 }
414 }
415}
416
417impl<DB: Database> Database for BalDatabase<DB> {
418 type Error = EvmDatabaseError<DB::Error>;
419
420 #[inline]
421 fn basic(&mut self, address: Address) -> Result<Option<AccountInfo>, Self::Error> {
422 let account_id = self.bal_state.get_account_id(&address)?;
423
424 let mut account = self.db.basic(address).map_err(EvmDatabaseError::Database)?;
425
426 if let Some(account_id) = account_id {
427 self.bal_state
428 .basic_by_account_id(account_id, &mut account)?;
429 }
430
431 Ok(account)
432 }
433
434 #[inline]
435 fn code_by_hash(&mut self, code_hash: B256) -> Result<Bytecode, Self::Error> {
436 self.db
437 .code_by_hash(code_hash)
438 .map_err(EvmDatabaseError::Database)
439 }
440
441 #[inline]
442 fn storage(&mut self, address: Address, key: StorageKey) -> Result<StorageValue, Self::Error> {
443 if let Some(storage) = self.bal_state.storage(&address, key)? {
444 return Ok(storage);
445 }
446
447 self.db
448 .storage(address, key)
449 .map_err(EvmDatabaseError::Database)
450 }
451
452 #[inline]
453 fn storage_by_account_id(
454 &mut self,
455 address: Address,
456 account_id: AccountId,
457 storage_key: StorageKey,
458 ) -> Result<StorageValue, Self::Error> {
459 if let Some(value) = self
460 .bal_state
461 .storage_by_account_id(account_id, storage_key)?
462 {
463 return Ok(value);
464 }
465
466 self.db
467 .storage(address, storage_key)
468 .map_err(EvmDatabaseError::Database)
469 }
470
471 fn block_hash(&mut self, number: u64) -> Result<B256, Self::Error> {
472 self.db
473 .block_hash(number)
474 .map_err(EvmDatabaseError::Database)
475 }
476}
477
478impl<DB: DatabaseCommit> DatabaseCommit for BalDatabase<DB> {
479 fn commit(&mut self, changes: EvmState) {
480 self.bal_state.commit(&changes);
481 self.db.commit(changes);
482 }
483
484 fn commit_iter(&mut self, changes: &mut dyn Iterator<Item = (Address, Account)>) {
485 let bal_state = &mut self.bal_state;
486 let mut changes = changes.map(|(address, account)| {
487 bal_state.commit_one(address, &account);
488 (address, account)
489 });
490 self.db.commit_iter(&mut changes);
491 }
492}
493
494#[cfg(test)]
495mod tests {
496 use super::*;
497 use primitives::U256;
498 use state::bal::{AccountBal, BalAccountInfo, BalWrites};
499
500 fn bal_with_account(address: Address, slot: StorageKey) -> Arc<Bal> {
501 let mut account = AccountBal::default();
502 account.storage.storage.insert(
503 slot,
504 BalWrites::new(vec![(BlockAccessIndex::new(1), StorageValue::from(42u64))]),
505 );
506 Arc::new(Bal::from_iter([(address, account)]))
507 }
508
509 #[test]
510 fn bal_misses_error_without_fallback() {
511 let address = Address::with_last_byte(1);
512 let missing = Address::with_last_byte(2);
513 let slot = U256::from(1);
514 let missing_slot = U256::from(2);
515 let bal_state = BalState::new().with_bal(bal_with_account(address, slot));
516
517 assert_eq!(
518 bal_state.get_account_id(&missing),
519 Err(BalError::AccountNotFound { address: missing })
520 );
521 assert_eq!(
522 bal_state.storage(&missing, slot),
523 Err(BalError::AccountNotFound { address: missing })
524 );
525 assert_eq!(
526 bal_state.storage(&address, missing_slot),
527 Err(BalError::SlotNotFound {
528 address,
529 slot: missing_slot
530 })
531 );
532 }
533
534 #[test]
535 fn bal_misses_fall_back_to_database_with_fallback() {
536 let address = Address::with_last_byte(1);
537 let missing = Address::with_last_byte(2);
538 let slot = U256::from(1);
539 let missing_slot = U256::from(2);
540 let mut bal_state = BalState::new()
541 .with_bal(bal_with_account(address, slot))
542 .with_allow_db_fallback(true);
543
544 assert_eq!(bal_state.get_account_id(&missing), Ok(None));
546 assert_eq!(bal_state.storage(&missing, slot), Ok(None));
547 assert_eq!(bal_state.storage(&address, missing_slot), Ok(None));
548
549 bal_state.bal_index = BlockAccessIndex::new(2);
551 assert!(bal_state.get_account_id(&address).unwrap().is_some());
552 assert_eq!(
553 bal_state.storage(&address, slot),
554 Ok(Some(StorageValue::from(42u64)))
555 );
556 }
557
558 #[test]
559 fn account_info_lookup_obeys_coverage_and_fallback() {
560 let address = Address::with_last_byte(1);
561 let mut bal_state = BalState::new();
562 assert_eq!(
563 bal_state.get_bal_account_info(&address),
564 Ok(BalAccountLookup::NotCovered)
565 );
566
567 bal_state.bal = Some(Arc::new(Bal::new()));
568 assert_eq!(
569 bal_state.get_bal_account_info(&address),
570 Err(BalError::AccountNotFound { address })
571 );
572 bal_state.set_allow_db_fallback(true);
573 assert_eq!(
574 bal_state.get_bal_account_info(&address),
575 Ok(BalAccountLookup::NotCovered)
576 );
577
578 bal_state.bal = Some(Arc::new(Bal::from_iter([(address, AccountBal::default())])));
580 bal_state.set_allow_db_fallback(false);
581 assert_eq!(
582 bal_state.get_bal_account_info(&address),
583 Ok(BalAccountLookup::Partial(BalAccountInfo::default()))
584 );
585 }
586}