use alloy::primitives::map::{AddressHashMap, B256HashMap};
use alloy::primitives::Address;
use parking_lot::RwLock;
use revm::bytecode::Bytecode;
use revm::database_interface::DatabaseRef;
use revm::primitives::{StorageKey, StorageValue, B256};
use revm::state::AccountInfo;
use std::sync::Arc;
pub const WARM_CODE_CACHE_TTL_BLOCKS: u64 = 10_000;
pub struct WarmCodeCacheInner {
accounts: AddressHashMap<(u64, Option<AccountInfo>)>,
bytecode: B256HashMap<(u64, Bytecode)>,
ttl_blocks: u64,
}
impl WarmCodeCacheInner {
#[must_use]
fn new(ttl_blocks: u64) -> Self {
Self {
accounts: AddressHashMap::default(),
bytecode: B256HashMap::default(),
ttl_blocks,
}
}
#[must_use]
pub fn shared_default() -> Arc<RwLock<Self>> {
Arc::new(RwLock::new(Self::new(WARM_CODE_CACHE_TTL_BLOCKS)))
}
#[must_use]
pub fn shared_with_ttl(ttl_blocks: u64) -> Arc<RwLock<Self>> {
Arc::new(RwLock::new(Self::new(ttl_blocks)))
}
fn is_fresh(&self, loaded_block: u64, block: u64) -> bool {
block.saturating_sub(loaded_block) <= self.ttl_blocks
}
}
pub struct WarmCodeCache<Db>
where
Db: DatabaseRef,
{
cache: Arc<RwLock<WarmCodeCacheInner>>,
block: u64,
db: Db,
}
impl<Db> WarmCodeCache<Db>
where
Db: DatabaseRef,
{
#[must_use]
pub fn with_owner(cache: Arc<RwLock<WarmCodeCacheInner>>, block: u64, db: Db) -> Self {
Self { cache, block, db }
}
}
impl<Db> DatabaseRef for WarmCodeCache<Db>
where
Db: DatabaseRef,
{
type Error = Db::Error;
fn basic_ref(&self, address: Address) -> Result<Option<AccountInfo>, Self::Error> {
{
let guard = self.cache.read();
if let Some((loaded_block, info)) = guard.accounts.get(&address) {
if guard.is_fresh(*loaded_block, self.block) {
return Ok(info.clone());
}
}
}
let info = self.db.basic_ref(address)?;
let mut guard = self.cache.write();
guard.accounts.insert(address, (self.block, info.clone()));
Ok(info)
}
fn code_by_hash_ref(&self, code_hash: B256) -> Result<Bytecode, Self::Error> {
{
let guard = self.cache.read();
if let Some((loaded_block, code)) = guard.bytecode.get(&code_hash) {
if guard.is_fresh(*loaded_block, self.block) {
return Ok(code.clone());
}
}
}
let code = self.db.code_by_hash_ref(code_hash)?;
let mut guard = self.cache.write();
guard.bytecode.insert(code_hash, (self.block, code.clone()));
Ok(code)
}
fn storage_ref(
&self,
address: Address,
index: StorageKey,
) -> Result<StorageValue, Self::Error> {
self.db.storage_ref(address, index)
}
fn block_hash_ref(&self, number: u64) -> Result<B256, Self::Error> {
self.db.block_hash_ref(number)
}
}
#[expect(clippy::expect_used)]
#[cfg(test)]
mod tests {
use super::*;
use alloy::primitives::{address, U256};
use revm::database_interface::EmptyDB;
use std::cell::Cell;
#[expect(clippy::struct_field_names)]
struct CountingMockDb {
basic_calls: Cell<u64>,
code_by_hash_calls: Cell<u64>,
storage_calls: Cell<u64>,
block_hash_calls: Cell<u64>,
}
impl CountingMockDb {
const EMPTY_CODE_HASH: B256 = B256::ZERO;
fn new() -> Self {
Self {
basic_calls: Cell::new(0),
code_by_hash_calls: Cell::new(0),
storage_calls: Cell::new(0),
block_hash_calls: Cell::new(0),
}
}
}
impl DatabaseRef for CountingMockDb {
type Error = core::convert::Infallible;
fn basic_ref(&self, _address: Address) -> Result<Option<AccountInfo>, Self::Error> {
self.basic_calls.set(self.basic_calls.get() + 1);
Ok(Some(AccountInfo::new(
U256::from(1),
0,
Self::EMPTY_CODE_HASH,
revm::bytecode::Bytecode::new_legacy(alloy::primitives::Bytes::from_static(&[
0x00,
])),
)))
}
fn code_by_hash_ref(&self, _code_hash: B256) -> Result<Bytecode, Self::Error> {
self.code_by_hash_calls
.set(self.code_by_hash_calls.get() + 1);
Ok(revm::bytecode::Bytecode::new_legacy(
alloy::primitives::Bytes::new(),
))
}
fn storage_ref(
&self,
_address: Address,
_index: StorageKey,
) -> Result<StorageValue, Self::Error> {
self.storage_calls.set(self.storage_calls.get() + 1);
Ok(StorageValue::ZERO)
}
fn block_hash_ref(&self, _number: u64) -> Result<B256, Self::Error> {
self.block_hash_calls.set(self.block_hash_calls.get() + 1);
Ok(B256::ZERO)
}
}
const ADDR: Address = address!("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa");
fn view(
cache: &Arc<RwLock<WarmCodeCacheInner>>,
block: u64,
db: CountingMockDb,
) -> WarmCodeCache<CountingMockDb> {
WarmCodeCache::with_owner(Arc::clone(cache), block, db)
}
#[test]
fn basic_ref_second_call_same_block_hits_cache() {
let inner = WarmCodeCacheInner::shared_with_ttl(10);
let v1 = view(&inner, 100, CountingMockDb::new());
let _ = v1.basic_ref(ADDR);
let v2 = view(&inner, 100, CountingMockDb::new());
let _ = v2.basic_ref(ADDR);
let guard = inner.read();
assert_eq!(guard.accounts.len(), 1, "exactly one cached account");
let (loaded_block, info) = guard.accounts.get(&ADDR).expect("cached");
assert_eq!(*loaded_block, 100);
assert!(info.is_some(), "cached value is Some");
}
#[test]
fn basic_ref_first_call_hits_the_db() {
let inner = WarmCodeCacheInner::shared_with_ttl(10);
let db = CountingMockDb::new();
let v = view(&inner, 100, db);
let _ = v.basic_ref(ADDR);
assert!(inner.read().accounts.contains_key(&ADDR));
}
#[test]
fn ttl_expiry_refetches_on_stale_read() {
let inner = WarmCodeCacheInner::shared_with_ttl(5);
let v = view(&inner, 10, CountingMockDb::new());
let info10 = v.basic_ref(ADDR).unwrap();
let loaded_block_10 = inner.read().accounts.get(&ADDR).expect("cached").0;
assert_eq!(loaded_block_10, 10);
let v15 = view(&inner, 15, CountingMockDb::new());
let info15 = v15.basic_ref(ADDR).unwrap();
assert_eq!(info15, info10, "fresh read returns the cached value");
let loaded_block_15 = inner.read().accounts.get(&ADDR).expect("cached").0;
assert_eq!(
loaded_block_15, 10,
"fresh read does NOT rewrite loaded_block"
);
let v16 = view(&inner, 16, CountingMockDb::new());
let _ = v16.basic_ref(ADDR).unwrap();
let loaded_block_16 = inner.read().accounts.get(&ADDR).expect("cached").0;
assert_eq!(
loaded_block_16, 16,
"stale read re-fetches + rewrites loaded_block"
);
}
#[test]
fn storage_ref_is_never_cached() {
let inner = WarmCodeCacheInner::shared_with_ttl(100);
let db = CountingMockDb::new();
let v = view(&inner, 100, db);
let slot = StorageKey::ZERO;
let _ = v.storage_ref(ADDR, slot);
let _ = v.storage_ref(ADDR, slot);
assert_eq!(
v.db.storage_calls.get(),
2,
"both storage reads hit the inner db (never cached)"
);
assert!(
inner.read().accounts.is_empty(),
"storage reads never populate the accounts map"
);
}
#[test]
fn block_hash_ref_is_never_cached() {
let inner = WarmCodeCacheInner::shared_with_ttl(100);
let db = CountingMockDb::new();
let v = view(&inner, 100, db);
let _ = v.block_hash_ref(50);
let _ = v.block_hash_ref(50);
assert_eq!(
v.db.block_hash_calls.get(),
2,
"both block_hash reads hit the inner db (never cached)"
);
}
#[test]
fn code_by_hash_ref_caches_and_expires() {
let inner = WarmCodeCacheInner::shared_with_ttl(5);
let hash = B256::repeat_byte(0xab);
let v = view(&inner, 20, CountingMockDb::new());
let code20 = v.code_by_hash_ref(hash).unwrap();
assert_eq!(v.db.code_by_hash_calls.get(), 1, "first call hits the db");
let loaded_block_20 = inner.read().bytecode.get(&hash).expect("cached").0;
assert_eq!(loaded_block_20, 20);
let v25 = view(&inner, 25, CountingMockDb::new());
let code25 = v25.code_by_hash_ref(hash).unwrap();
assert_eq!(
code25.hash_slow(),
code20.hash_slow(),
"fresh returns cached"
);
let loaded_block_25 = inner.read().bytecode.get(&hash).expect("cached").0;
assert_eq!(loaded_block_25, 20, "fresh read does NOT rewrite");
let v26 = view(&inner, 26, CountingMockDb::new());
let _ = v26.code_by_hash_ref(hash).unwrap();
let loaded_block_26 = inner.read().bytecode.get(&hash).expect("cached").0;
assert_eq!(loaded_block_26, 26, "stale re-fetches + rewrites");
}
#[test]
fn basic_ref_none_is_cached() {
struct NoneDb {
calls: Cell<u64>,
}
impl DatabaseRef for NoneDb {
type Error = core::convert::Infallible;
fn basic_ref(&self, _address: Address) -> Result<Option<AccountInfo>, Self::Error> {
self.calls.set(self.calls.get() + 1);
Ok(None)
}
fn code_by_hash_ref(&self, _code_hash: B256) -> Result<Bytecode, Self::Error> {
Ok(Bytecode::new_legacy(alloy::primitives::Bytes::new()))
}
fn storage_ref(
&self,
_address: Address,
_index: StorageKey,
) -> Result<StorageValue, Self::Error> {
Ok(StorageValue::ZERO)
}
fn block_hash_ref(&self, _number: u64) -> Result<B256, Self::Error> {
Ok(B256::ZERO)
}
}
let inner = WarmCodeCacheInner::shared_with_ttl(100);
let db = NoneDb {
calls: Cell::new(0),
};
let v = WarmCodeCache::with_owner(Arc::clone(&inner), 100, db);
let info1 = v.basic_ref(ADDR).unwrap();
assert!(info1.is_none(), "first call returns None");
assert_eq!(v.db.calls.get(), 1, "first call hit the db");
let info2 = v.basic_ref(ADDR).unwrap();
assert!(info2.is_none(), "second call returns cached None");
assert_eq!(
v.db.calls.get(),
1,
"second call hit the cache (None is cached, not re-queried)"
);
}
#[test]
fn warm_code_cache_compiles_over_empty_db() {
let inner = WarmCodeCacheInner::shared_default();
let v = WarmCodeCache::with_owner(Arc::clone(&inner), 1, EmptyDB::default());
let info = v.basic_ref(ADDR).unwrap();
assert!(info.is_none(), "EmptyDB returns no account");
assert_eq!(
inner.read().accounts.len(),
1,
"Ok(None) IS cached (the existence-negative decision)"
);
}
#[test]
fn warm_code_cache_module_has_no_pyo3_dependency() {
}
}