1use alloy_primitives::{
4 B256, Bytes, U256, keccak256,
5 map::{AddressMap, B256Map, HashSet, U256Map},
6};
7use alloy_rlp::Encodable;
8use alloy_trie::{
9 EMPTY_ROOT_HASH, HashBuilder, Nibbles, TrieMask,
10 nodes::{BranchNodeRef, ExtensionNodeRef, LeafNodeRef, RlpNode},
11};
12use revm::{
13 database::{AccountState, DbAccount},
14 state::{Account, AccountInfo},
15};
16use std::{array, mem};
17
18#[derive(Debug, Default)]
25pub struct StateRootCache {
26 trie: Option<IncrementalStateTrie>,
27 dirty: AddressMap<DirtyAccount>,
28 rlp_buf: Vec<u8>,
30}
31
32#[derive(Debug, Default)]
33struct DirtyAccount {
34 storage: HashSet<U256>,
35 reset_storage: bool,
36}
37
38impl StateRootCache {
39 pub fn record_changes(&mut self, changes: &AddressMap<Account>) {
41 for (address, account) in changes {
42 if !account.is_touched() {
43 continue;
44 }
45
46 let dirty = self.dirty.entry(*address).or_default();
47 dirty.reset_storage |= account.is_created() || account.is_selfdestructed();
48 dirty.storage.extend(account.changed_storage_slots().map(|(slot, _)| *slot));
49 }
50 }
51
52 pub fn record_account(&mut self, address: alloy_primitives::Address) {
54 self.dirty.entry(address).or_default();
55 }
56
57 pub fn record_storage(&mut self, address: alloy_primitives::Address, slot: U256) {
59 self.dirty.entry(address).or_default().storage.insert(slot);
60 }
61
62 pub fn record_overlay(
68 &mut self,
69 accounts: &AddressMap<DbAccount>,
70 previous: &AddressMap<DbAccount>,
71 ) -> bool {
72 let incremental = previous.iter().all(|(address, previous)| {
73 accounts.get(address).is_some_and(|account| {
74 matches!(
75 account.account_state,
76 AccountState::StorageCleared | AccountState::NotExisting
77 ) || (!matches!(
78 previous.account_state,
79 AccountState::StorageCleared | AccountState::NotExisting
80 ) && previous.storage.keys().all(|slot| account.storage.contains_key(slot)))
81 })
82 });
83 if !incremental {
84 self.invalidate();
85 return false;
86 }
87
88 for (address, account) in accounts {
89 let dirty = self.dirty.entry(*address).or_default();
90 if account.account_state.is_storage_cleared() {
91 dirty.reset_storage = true;
92 } else {
93 dirty.storage.extend(account.storage.keys().copied());
94 }
95 }
96 true
97 }
98
99 pub fn invalidate(&mut self) {
101 self.trie = None;
102 self.dirty.clear();
103 }
104
105 pub fn root(&mut self, accounts: &AddressMap<DbAccount>) -> B256 {
109 let Self { trie, dirty, rlp_buf } = self;
110 if trie.is_none() {
111 *trie = Some(IncrementalStateTrie::from_accounts(accounts, rlp_buf));
112 dirty.clear();
113 return trie.as_mut().unwrap().root(rlp_buf);
114 }
115
116 let trie = trie.as_mut().unwrap();
117 for (address, dirty) in mem::take(dirty) {
118 let hashed_address = keccak256(address);
119 let Some(account) = accounts
120 .get(&address)
121 .filter(|account| account.account_state != AccountState::NotExisting)
122 else {
123 trie.accounts.remove(hashed_address);
124 trie.storage.remove(&hashed_address);
125 continue;
126 };
127
128 let storage_trie = if dirty.reset_storage {
129 trie.storage
130 .entry(hashed_address)
131 .insert_entry(IncrementalTrie::from_storage(&account.storage))
132 .into_mut()
133 } else {
134 let storage_trie = trie.storage.entry(hashed_address).or_default();
135 for slot in dirty.storage {
136 let key = keccak256(slot.to_be_bytes::<32>());
137 if let Some(value) = account.storage.get(&slot).filter(|value| !value.is_zero())
138 {
139 storage_trie.insert(key, alloy_rlp::encode(value));
140 } else {
141 storage_trie.remove(key);
142 }
143 }
144 storage_trie
145 };
146 let storage_root = storage_trie.root_with_buf(rlp_buf);
147 trie.accounts.insert(
148 hashed_address,
149 trie_account_rlp_with_storage_root(&account.info, storage_root),
150 );
151 }
152
153 trie.root(rlp_buf)
154 }
155}
156
157#[derive(Debug, Default)]
158struct IncrementalStateTrie {
159 accounts: IncrementalTrie,
160 storage: B256Map<IncrementalTrie>,
161}
162
163impl IncrementalStateTrie {
164 fn from_accounts(accounts: &AddressMap<DbAccount>, rlp_buf: &mut Vec<u8>) -> Self {
165 let mut trie = Self::default();
166 for (address, account) in accounts {
167 if account.account_state == AccountState::NotExisting {
168 continue;
169 }
170 let hashed_address = keccak256(address);
171 let mut storage_trie = IncrementalTrie::from_storage(&account.storage);
172 let storage_root = storage_trie.root_with_buf(rlp_buf);
173 trie.accounts.insert(
174 hashed_address,
175 trie_account_rlp_with_storage_root(&account.info, storage_root),
176 );
177 trie.storage.insert(hashed_address, storage_trie);
178 }
179 trie
180 }
181
182 fn root(&mut self, rlp_buf: &mut Vec<u8>) -> B256 {
183 self.accounts.root_with_buf(rlp_buf)
184 }
185}
186
187#[derive(Debug, Default)]
189struct IncrementalTrie {
190 root: TrieNode,
191}
192
193impl IncrementalTrie {
194 fn from_storage(storage: &U256Map<U256>) -> Self {
195 let mut trie = Self::default();
196 for (slot, value) in storage.iter().filter(|(_, value)| !value.is_zero()) {
197 trie.insert(keccak256(slot.to_be_bytes::<32>()), alloy_rlp::encode(value));
198 }
199 trie
200 }
201
202 fn insert(&mut self, key: B256, value: Vec<u8>) {
203 self.root.insert(Nibbles::unpack(key), value);
204 }
205
206 fn remove(&mut self, key: B256) {
207 self.root.remove(Nibbles::unpack(key));
208 }
209
210 #[cfg(test)]
211 fn root(&mut self) -> B256 {
212 self.root_with_buf(&mut Vec::new())
213 }
214
215 fn root_with_buf(&mut self, rlp_buf: &mut Vec<u8>) -> B256 {
216 let Some(root) = self.root.rlp(rlp_buf) else { return EMPTY_ROOT_HASH };
217 root.as_hash().unwrap_or_else(|| keccak256(root.as_ref()))
218 }
219}
220
221#[derive(Debug, Default)]
222struct TrieNode {
223 kind: TrieNodeKind,
224 rlp: Option<RlpNode>,
225}
226
227#[derive(Debug, Default)]
228enum TrieNodeKind {
229 #[default]
230 Empty,
231 Leaf {
232 path: Nibbles,
233 value: Vec<u8>,
234 },
235 Extension {
236 path: Nibbles,
237 child: Box<TrieNode>,
238 },
239 Branch {
240 children: [Option<Box<TrieNode>>; 16],
241 },
242}
243
244impl TrieNode {
245 const fn leaf(path: Nibbles, value: Vec<u8>) -> Self {
246 Self { kind: TrieNodeKind::Leaf { path, value }, rlp: None }
247 }
248
249 fn extension(path: Nibbles, child: Self) -> Self {
250 debug_assert!(!path.is_empty());
251 Self { kind: TrieNodeKind::Extension { path, child: Box::new(child) }, rlp: None }
252 }
253
254 const fn branch(children: [Option<Box<Self>>; 16]) -> Self {
255 Self { kind: TrieNodeKind::Branch { children }, rlp: None }
256 }
257
258 fn empty_children() -> [Option<Box<Self>>; 16] {
259 Default::default()
260 }
261
262 fn insert(&mut self, key: Nibbles, value: Vec<u8>) {
263 let kind = mem::take(&mut self.kind);
264 self.rlp = None;
265 *self = match kind {
266 TrieNodeKind::Empty => Self::leaf(key, value),
267 TrieNodeKind::Leaf { path, value: old_value } => {
268 let common = path.common_prefix_length(&key);
269 if common == path.len() {
270 debug_assert_eq!(common, key.len());
271 Self::leaf(path, value)
272 } else {
273 let mut children = Self::empty_children();
274 let old_index = path.get(common).unwrap() as usize;
275 let new_index = key.get(common).unwrap() as usize;
276 children[old_index] =
277 Some(Box::new(Self::leaf(path.slice(common + 1..), old_value)));
278 children[new_index] =
279 Some(Box::new(Self::leaf(key.slice(common + 1..), value)));
280 let branch = Self::branch(children);
281 if common == 0 { branch } else { Self::extension(path.slice(..common), branch) }
282 }
283 }
284 TrieNodeKind::Extension { path, mut child } => {
285 let common = path.common_prefix_length(&key);
286 if common == path.len() {
287 child.insert(key.slice(common..), value);
288 Self::extension(path, *child)
289 } else {
290 let mut children = Self::empty_children();
291 let old_index = path.get(common).unwrap() as usize;
292 let old_path = path.slice(common + 1..);
293 let old_child = if old_path.is_empty() {
294 *child
295 } else {
296 Self::extension(old_path, *child)
297 };
298 children[old_index] = Some(Box::new(old_child));
299
300 let new_index = key.get(common).unwrap() as usize;
301 children[new_index] =
302 Some(Box::new(Self::leaf(key.slice(common + 1..), value)));
303 let branch = Self::branch(children);
304 if common == 0 { branch } else { Self::extension(path.slice(..common), branch) }
305 }
306 }
307 TrieNodeKind::Branch { mut children } => {
308 let index = key.first().expect("trie keys have equal lengths") as usize;
309 children[index]
310 .get_or_insert_with(|| Box::new(Self::default()))
311 .insert(key.slice(1..), value);
312 Self::branch(children)
313 }
314 };
315 }
316
317 fn remove(&mut self, key: Nibbles) {
318 let kind = mem::take(&mut self.kind);
319 self.rlp = None;
320 *self = match kind {
321 TrieNodeKind::Empty => Self::default(),
322 TrieNodeKind::Leaf { path, value } => {
323 if path == key {
324 Self::default()
325 } else {
326 Self::leaf(path, value)
327 }
328 }
329 TrieNodeKind::Extension { path, mut child } => {
330 if key.starts_with(&path) {
331 child.remove(key.slice(path.len()..));
332 Self::normalize_extension(path, *child)
333 } else {
334 Self::extension(path, *child)
335 }
336 }
337 TrieNodeKind::Branch { mut children } => {
338 let index = key.first().expect("trie keys have equal lengths") as usize;
339 if let Some(child) = &mut children[index] {
340 child.remove(key.slice(1..));
341 if matches!(child.kind, TrieNodeKind::Empty) {
342 children[index] = None;
343 }
344 }
345 Self::normalize_branch(children)
346 }
347 };
348 }
349
350 fn normalize_extension(path: Nibbles, child: Self) -> Self {
351 match child.kind {
352 TrieNodeKind::Empty => Self::default(),
353 TrieNodeKind::Leaf { path: child_path, value } => {
354 Self::leaf(path.join(&child_path), value)
355 }
356 TrieNodeKind::Extension { path: child_path, child } => {
357 Self::extension(path.join(&child_path), *child)
358 }
359 TrieNodeKind::Branch { children } => Self::extension(path, Self::branch(children)),
360 }
361 }
362
363 fn normalize_branch(mut children: [Option<Box<Self>>; 16]) -> Self {
364 let mut indexes =
365 children.iter().enumerate().filter_map(|(index, child)| child.as_ref().map(|_| index));
366 let Some(index) = indexes.next() else { return Self::default() };
367 if indexes.next().is_some() {
368 return Self::branch(children);
369 }
370
371 let child = *children[index].take().unwrap();
372 let prefix = Nibbles::from_nibbles([index as u8]);
373 match child.kind {
374 TrieNodeKind::Empty => unreachable!("empty branch children are removed"),
375 TrieNodeKind::Leaf { path, value } => Self::leaf(prefix.join(&path), value),
376 TrieNodeKind::Extension { path, child } => Self::extension(prefix.join(&path), *child),
377 TrieNodeKind::Branch { children } => Self::extension(prefix, Self::branch(children)),
378 }
379 }
380
381 fn rlp(&mut self, out: &mut Vec<u8>) -> Option<RlpNode> {
382 if let Some(rlp) = &self.rlp {
383 return Some(rlp.clone());
384 }
385
386 let rlp = self.encode(out)?;
387 self.rlp = Some(rlp.clone());
388 Some(rlp)
389 }
390
391 fn encode(&mut self, out: &mut Vec<u8>) -> Option<RlpNode> {
393 let rlp = match &mut self.kind {
394 TrieNodeKind::Empty => return None,
395 TrieNodeKind::Leaf { path, value } => {
396 out.clear();
397 LeafNodeRef::new(path, value).rlp(out)
398 }
399 TrieNodeKind::Extension { path, child } => {
400 let child = child.rlp(out).expect("extension nodes have a child");
401 out.clear();
402 ExtensionNodeRef::new(path, child.as_ref()).rlp(out)
403 }
404 TrieNodeKind::Branch { children } => {
405 let mut stack: [RlpNode; 16] = array::from_fn(|_| RlpNode::default());
406 let mut stack_len = 0;
407 let mut state_mask = TrieMask::default();
408 for (index, child) in children.iter_mut().enumerate() {
409 if let Some(child) = child {
410 stack[stack_len] = child.rlp(out).expect("branch children are not empty");
411 stack_len += 1;
412 state_mask.set_bit(index as u8);
413 }
414 }
415 out.clear();
416 BranchNodeRef::new(&stack[..stack_len], state_mask).rlp(out)
417 }
418 };
419 Some(rlp)
420 }
421
422 fn collect_nodes(&mut self, rlp_buf: &mut Vec<u8>, nodes: &mut Vec<Bytes>) {
424 match &mut self.kind {
425 TrieNodeKind::Empty => return,
426 TrieNodeKind::Leaf { .. } => {}
427 TrieNodeKind::Extension { child, .. } => child.collect_nodes(rlp_buf, nodes),
428 TrieNodeKind::Branch { children } => {
429 for child in children.iter_mut().flatten() {
430 child.collect_nodes(rlp_buf, nodes);
431 }
432 }
433 }
434 self.encode(rlp_buf);
435 nodes.push(Bytes::copy_from_slice(rlp_buf));
436 }
437}
438
439pub fn state_trie_witness(accounts: &AddressMap<DbAccount>) -> (B256, Vec<Bytes>) {
445 let mut rlp_buf = Vec::new();
446 let mut account_trie = IncrementalTrie::default();
447 let mut nodes = Vec::new();
448 for (address, account) in accounts {
449 if account.account_state == AccountState::NotExisting {
450 continue;
451 }
452
453 let mut storage_trie = IncrementalTrie::from_storage(&account.storage);
455 let storage_root = storage_trie.root_with_buf(&mut rlp_buf);
456 storage_trie.root.collect_nodes(&mut rlp_buf, &mut nodes);
457 account_trie.insert(
458 keccak256(address),
459 trie_account_rlp_with_storage_root(&account.info, storage_root),
460 );
461 }
462 account_trie.root.collect_nodes(&mut rlp_buf, &mut nodes);
463 let root = account_trie.root_with_buf(&mut rlp_buf);
464 nodes.sort_unstable();
465 nodes.dedup();
466 (root, nodes)
467}
468
469pub fn build_root(values: impl IntoIterator<Item = (Nibbles, Vec<u8>)>) -> B256 {
470 let mut builder = HashBuilder::default();
471 for (key, value) in values {
472 builder.add_leaf(key, value.as_ref());
473 }
474 builder.root()
475}
476
477pub fn state_root(accounts: &AddressMap<DbAccount>) -> B256 {
479 build_root(trie_accounts(accounts))
480}
481
482pub fn storage_root(storage: &U256Map<U256>) -> B256 {
484 build_root(trie_storage(storage))
485}
486
487pub fn trie_storage(storage: &U256Map<U256>) -> Vec<(Nibbles, Vec<u8>)> {
489 let mut storage = storage
490 .iter()
491 .filter(|(_, value)| !value.is_zero())
492 .map(|(key, value)| {
493 let data = alloy_rlp::encode(value);
494 (Nibbles::unpack(keccak256(key.to_be_bytes::<32>())), data)
495 })
496 .collect::<Vec<_>>();
497 storage.sort_by_key(|(key, _)| *key);
498
499 storage
500}
501
502pub fn trie_accounts(accounts: &AddressMap<DbAccount>) -> Vec<(Nibbles, Vec<u8>)> {
504 let mut accounts: Vec<(Nibbles, Vec<u8>)> = accounts
505 .iter()
506 .filter(|(_, account)| account.account_state != AccountState::NotExisting)
507 .map(|(address, account)| {
508 let data = trie_account_rlp(&account.info, &account.storage);
509 (Nibbles::unpack(keccak256(*address)), data)
510 })
511 .collect();
512 accounts.sort_by_key(|(key, _)| *key);
513
514 accounts
515}
516
517pub fn trie_account_rlp(info: &AccountInfo, storage: &U256Map<U256>) -> Vec<u8> {
519 trie_account_rlp_with_storage_root(info, storage_root(storage))
520}
521
522fn trie_account_rlp_with_storage_root(info: &AccountInfo, storage_root: B256) -> Vec<u8> {
524 let mut out: Vec<u8> = Vec::new();
525 let list: [&dyn Encodable; 4] = [&info.nonce, &info.balance, &storage_root, &info.code_hash()];
526
527 alloy_rlp::encode_list::<_, dyn Encodable>(&list, &mut out);
528
529 out
530}
531
532#[cfg(test)]
533mod tests {
534 use super::*;
535
536 #[test]
537 fn canonical_roots_omit_zero_storage_and_non_existing_accounts() {
538 let mut storage = U256Map::default();
539 storage.insert(U256::ONE, U256::ZERO);
540 assert_eq!(storage_root(&storage), EMPTY_ROOT_HASH);
541
542 let mut accounts = AddressMap::default();
543 accounts
544 .insert(alloy_primitives::Address::with_last_byte(1), DbAccount::new_not_existing());
545 assert_eq!(state_root(&accounts), EMPTY_ROOT_HASH);
546 assert_eq!(StateRootCache::default().root(&accounts), EMPTY_ROOT_HASH);
547 }
548
549 #[test]
550 fn incremental_roots_preserve_uncached_state_and_replace_cleared_storage() {
551 let address = alloy_primitives::Address::with_last_byte(1);
552 let untouched = alloy_primitives::Address::with_last_byte(2);
553 let [one, two, three, four] = [1u64, 2, 3, 4].map(U256::from);
554 let original = DbAccount {
555 info: AccountInfo { balance: U256::from(100), ..Default::default() },
556 account_state: AccountState::None,
557 storage: [(one, U256::from(11)), (two, U256::from(22))].into_iter().collect(),
558 };
559 let mut full = AddressMap::from_iter([(address, original.clone()), (untouched, original)]);
560 let mut cache = StateRootCache::default();
561 assert_eq!(cache.root(&full), state_root(&full));
562 let mut previous = AddressMap::default();
563 let mut check = |overlay: &AddressMap<DbAccount>, full: &AddressMap<DbAccount>| {
564 assert!(cache.record_overlay(overlay, &previous));
565 assert_eq!(cache.root(overlay), state_root(full));
566 previous.clone_from(overlay);
567 };
568
569 let mut overlay = AddressMap::from_iter([(
571 address,
572 DbAccount {
573 info: full[&address].info.clone(),
574 account_state: AccountState::Touched,
575 storage: [(one, U256::from(33))].into_iter().collect(),
576 },
577 )]);
578 full.get_mut(&address).unwrap().storage.insert(one, U256::from(33));
579 check(&overlay, &full);
580
581 overlay.get_mut(&address).unwrap().storage.extend([(two, U256::ZERO), (three, four)]);
583 full.get_mut(&address).unwrap().storage.remove(&two);
584 full.get_mut(&address).unwrap().storage.insert(three, four);
585 check(&overlay, &full);
586
587 overlay.get_mut(&address).unwrap().info.balance = U256::from(101);
589 full.get_mut(&address).unwrap().info.balance = U256::from(101);
590 check(&overlay, &full);
591
592 overlay.insert(address, DbAccount::new_not_existing());
594 full.remove(&address);
595 check(&overlay, &full);
596
597 for slot in [three, four] {
600 let account = DbAccount {
601 info: AccountInfo { balance: U256::from(102), ..Default::default() },
602 account_state: AccountState::StorageCleared,
603 storage: [(slot, U256::from(55))].into_iter().collect(),
604 };
605 overlay.insert(address, account.clone());
606 full.insert(address, account);
607 check(&overlay, &full);
608 }
609 }
610
611 #[test]
612 fn overlay_roots_rebuild_when_base_storage_is_exposed_again() {
613 let address = alloy_primitives::Address::with_last_byte(1);
614 let [one, two] = [1u64, 2].map(U256::from);
615 let base = DbAccount {
616 info: AccountInfo { balance: U256::from(100), ..Default::default() },
617 storage: [(one, U256::from(11)), (two, U256::from(22))].into_iter().collect(),
618 ..Default::default()
619 };
620 for account_state in
621 [AccountState::Touched, AccountState::StorageCleared, AccountState::NotExisting]
622 {
623 let storage = if account_state == AccountState::NotExisting {
624 U256Map::default()
625 } else {
626 [(one, U256::from(33))].into_iter().collect()
627 };
628 let previous = AddressMap::from_iter([(
629 address,
630 DbAccount { account_state: account_state.clone(), storage, ..base.clone() },
631 )]);
632 let storage = if account_state.is_storage_cleared() {
635 previous[&address].storage.clone()
636 } else {
637 U256Map::default()
638 };
639 let mut merged = base.clone();
640 merged.storage.extend(storage.clone());
641 let overlay = AddressMap::from_iter([(
642 address,
643 DbAccount { account_state: AccountState::Touched, storage, ..base.clone() },
644 )]);
645 let mut cache = StateRootCache::default();
646 cache.root(&previous);
647 assert!(!cache.record_overlay(&overlay, &previous));
648 let full = AddressMap::from_iter([(address, merged)]);
649 assert_eq!(cache.root(&full), state_root(&full));
650 assert!(!cache.record_overlay(&AddressMap::default(), &overlay));
651 }
652 }
653
654 fn rebuilt_root(values: &B256Map<Vec<u8>>) -> B256 {
655 let mut leaves = values
656 .iter()
657 .map(|(key, value)| (Nibbles::unpack(*key), value.clone()))
658 .collect::<Vec<_>>();
659 leaves.sort_by_key(|(key, _)| *key);
660 build_root(leaves)
661 }
662
663 #[test]
664 fn incremental_trie_matches_full_rebuild() {
665 let mut trie = IncrementalTrie::default();
666 let mut values = B256Map::default();
667 assert_eq!(trie.root(), EMPTY_ROOT_HASH);
668
669 for index in 0..128 {
670 let key = keccak256(U256::from(index).to_be_bytes::<32>());
671 let value = alloy_rlp::encode(U256::from(index + 1));
672 trie.insert(key, value.clone());
673 values.insert(key, value);
674 assert_eq!(trie.root(), rebuilt_root(&values));
675 }
676
677 for index in (0..128).step_by(3) {
678 let key = keccak256(U256::from(index).to_be_bytes::<32>());
679 let value = alloy_rlp::encode(U256::from(index + 1_000));
680 trie.insert(key, value.clone());
681 values.insert(key, value);
682 assert_eq!(trie.root(), rebuilt_root(&values));
683 }
684
685 for index in (0..128).rev() {
686 let key = keccak256(U256::from(index).to_be_bytes::<32>());
687 trie.remove(key);
688 values.remove(&key);
689 assert_eq!(trie.root(), rebuilt_root(&values));
690 }
691 }
692}