use crate::zkevm_opcode_defs::system_params::MAX_PUBDATA_COST_PER_QUERY;
use zk_evm_abstractions::aux::{PubdataCost, Timestamp};
use zk_evm_abstractions::vm::{Storage, StorageAccessRefund};
use zk_evm_abstractions::zkevm_opcode_defs::system_params::STORAGE_AUX_BYTE;
use super::ApplicationData;
use super::*;
#[derive(Debug, Clone)]
pub struct InMemoryStorage {
pub inner: [HashMap<Address, HashMap<U256, U256>>; NUM_SHARDS],
pub inner_transient: [HashMap<Address, HashMap<U256, U256>>; NUM_SHARDS],
pub cold_warm_markers: [HashMap<Address, HashSet<U256>>; NUM_SHARDS],
pub transient_cold_warm_markers: [HashMap<Address, HashSet<U256>>; NUM_SHARDS], pub frames_stack: Vec<ApplicationData<LogQuery>>,
}
impl InMemoryStorage {
pub fn new() -> Self {
Self {
inner: [(); NUM_SHARDS].map(|_| HashMap::default()),
inner_transient: [(); NUM_SHARDS].map(|_| HashMap::default()),
cold_warm_markers: [(); NUM_SHARDS].map(|_| HashMap::default()),
transient_cold_warm_markers: [(); NUM_SHARDS].map(|_| HashMap::default()),
frames_stack: vec![ApplicationData::empty()],
}
}
pub fn populate(&mut self, elements: Vec<(u8, Address, U256, U256)>) {
for (shard_id, address, key, value) in elements.into_iter() {
let shard_level_map = &mut self.inner[shard_id as usize];
let address_level_map = shard_level_map.entry(address).or_default();
address_level_map.insert(key, value);
}
}
pub fn flatten_and_net_history(
mut self,
) -> (Vec<LogQuery>, HashMap<(u8, Address, U256), Vec<LogQuery>>) {
assert_eq!(
self.frames_stack.len(),
1,
"there must exist an initial keeper frame"
);
let full_history = self.frames_stack.pop().unwrap();
let ApplicationData {
forward,
rollbacks: _,
} = full_history;
let history = forward.clone();
let mut tmp = HashMap::<(u8, Address, U256), Vec<LogQuery>>::with_capacity(forward.len());
for el in forward.into_iter() {
if el.aux_byte != STORAGE_AUX_BYTE {
continue;
}
let LogQuery {
timestamp,
shard_id,
address,
key,
rollback,
..
} = ⪙
let entry = tmp.entry((*shard_id, *address, *key)).or_insert(vec![]);
if let Some(last) = entry.last() {
if !rollback {
assert!(timestamp.0 > last.timestamp.0);
}
}
entry.push(el);
}
(history, tmp)
}
}
impl Storage for InMemoryStorage {
#[track_caller]
fn get_access_refund(
&mut self, _monotonic_cycle_counter: u32,
_partial_query: &LogQuery,
) -> StorageAccessRefund {
StorageAccessRefund::Cold
}
#[track_caller]
fn execute_partial_query(
&mut self,
_monotonic_cycle_counter: u32,
mut query: LogQuery,
) -> (LogQuery, PubdataCost) {
let aux_byte = query.aux_byte;
let shard_level_map = if aux_byte == STORAGE_AUX_BYTE {
&mut self.inner[query.shard_id as usize]
} else {
&mut self.inner_transient[query.shard_id as usize]
};
let shard_level_warm_map = if aux_byte == STORAGE_AUX_BYTE {
&mut self.cold_warm_markers[query.shard_id as usize]
} else {
&mut self.transient_cold_warm_markers[query.shard_id as usize]
};
let frame_data = self.frames_stack.last_mut().expect("frame must be started");
assert!(!query.rollback);
if query.rw_flag {
let address_level_map = shard_level_map.entry(query.address).or_default();
let current_value = address_level_map
.get(&query.key)
.copied()
.unwrap_or(U256::zero());
address_level_map.insert(query.key, query.written_value);
let address_level_warm_map = shard_level_warm_map.entry(query.address).or_default();
let warm = address_level_warm_map.contains(&query.key);
if !warm {
address_level_warm_map.insert(query.key);
}
query.read_value = current_value;
frame_data.forward.push(query);
query.rollback = true;
frame_data.rollbacks.push(query);
query.rollback = false;
let pubdata_cost = if aux_byte == STORAGE_AUX_BYTE {
PubdataCost(MAX_PUBDATA_COST_PER_QUERY)
} else {
PubdataCost(0i32)
};
(query, pubdata_cost)
} else {
let address_level_map = shard_level_map.entry(query.address).or_default();
let current_value = address_level_map
.get(&query.key)
.copied()
.unwrap_or(U256::zero());
let address_level_warm_map = shard_level_warm_map.entry(query.address).or_default();
let warm = address_level_warm_map.contains(&query.key);
if !warm {
address_level_warm_map.insert(query.key);
}
query.read_value = current_value;
frame_data.forward.push(query);
(query, PubdataCost(0i32))
}
}
#[track_caller]
fn start_frame(&mut self, _timestamp: Timestamp) {
let new = ApplicationData::empty();
self.frames_stack.push(new);
}
#[track_caller]
fn finish_frame(&mut self, _timestamp: Timestamp, panicked: bool) {
let current_frame = self
.frames_stack
.pop()
.expect("frame must be started before finishing");
let ApplicationData { forward, rollbacks } = current_frame;
let parent_data = self
.frames_stack
.last_mut()
.expect("parent_frame_must_exist");
if panicked {
for query in rollbacks.iter().rev() {
let LogQuery {
shard_id,
address,
key,
read_value,
written_value,
aux_byte,
..
} = *query;
let shard_level_map = if aux_byte == STORAGE_AUX_BYTE {
&mut self.inner[shard_id as usize]
} else {
&mut self.inner_transient[shard_id as usize]
};
let address_level_map = shard_level_map
.get_mut(&address)
.expect("must always exist on rollback");
let current_value_ref = address_level_map
.get_mut(&key)
.expect("must always exist on rollback");
assert_eq!(*current_value_ref, written_value); *current_value_ref = read_value; }
parent_data.forward.extend(forward);
parent_data.forward.extend(rollbacks.into_iter().rev());
} else {
parent_data.forward.extend(forward);
parent_data.rollbacks.extend(rollbacks);
}
}
#[track_caller]
fn start_new_tx(&mut self, _: Timestamp) {
for transient in self.inner_transient.iter_mut() {
transient.clear();
}
}
}