1use crate::{
2 hashed_cursor::HashedCursorFactory, prefix_set::TriePrefixSetsMut, proof::Proof, proof_v2,
3 trie_cursor::TrieCursorFactory, TRIE_ACCOUNT_RLP_MAX_SIZE,
4};
5use alloy_primitives::{
6 keccak256,
7 map::{B256Map, HashMap},
8 Bytes, B256,
9};
10use alloy_rlp::{Encodable, EMPTY_STRING_CODE};
11use alloy_trie::{nodes::BranchNodeRef, EMPTY_ROOT_HASH};
12use reth_execution_errors::{SparseStateTrieErrorKind, StateProofError, TrieWitnessError};
13use reth_trie_common::{
14 DecodedMultiProofV2, ExecutionWitnessMode, HashedPostState, MultiProofTargetsV2, ProofV2Target,
15 TrieNodeV2,
16};
17use reth_trie_sparse::{LeafUpdate, SparseStateTrie, SparseTrie as _, TrieNodeEpoch};
18
19#[derive(Debug)]
21pub struct TrieWitness<T, H> {
22 trie_cursor_factory: T,
24 hashed_cursor_factory: H,
26 prefix_sets: TriePrefixSetsMut,
28 always_include_root_node: bool,
33 mode: ExecutionWitnessMode,
35 witness: B256Map<Bytes>,
37}
38
39impl<T, H> TrieWitness<T, H> {
40 pub fn new(trie_cursor_factory: T, hashed_cursor_factory: H) -> Self {
42 Self {
43 trie_cursor_factory,
44 hashed_cursor_factory,
45 prefix_sets: TriePrefixSetsMut::default(),
46 always_include_root_node: false,
47 mode: ExecutionWitnessMode::Legacy,
48 witness: HashMap::default(),
49 }
50 }
51
52 pub fn with_trie_cursor_factory<TF>(self, trie_cursor_factory: TF) -> TrieWitness<TF, H> {
54 TrieWitness {
55 trie_cursor_factory,
56 hashed_cursor_factory: self.hashed_cursor_factory,
57 prefix_sets: self.prefix_sets,
58 always_include_root_node: self.always_include_root_node,
59 mode: self.mode,
60 witness: self.witness,
61 }
62 }
63
64 pub fn with_hashed_cursor_factory<HF>(self, hashed_cursor_factory: HF) -> TrieWitness<T, HF> {
66 TrieWitness {
67 trie_cursor_factory: self.trie_cursor_factory,
68 hashed_cursor_factory,
69 prefix_sets: self.prefix_sets,
70 always_include_root_node: self.always_include_root_node,
71 mode: self.mode,
72 witness: self.witness,
73 }
74 }
75
76 pub fn with_prefix_sets_mut(mut self, prefix_sets: TriePrefixSetsMut) -> Self {
78 self.prefix_sets = prefix_sets;
79 self
80 }
81
82 pub const fn always_include_root_node(mut self) -> Self {
86 self.always_include_root_node = true;
87 self
88 }
89
90 pub const fn with_execution_witness_mode(mut self, mode: ExecutionWitnessMode) -> Self {
92 self.mode = mode;
93 self
94 }
95}
96
97impl<T, H> TrieWitness<T, H>
98where
99 T: TrieCursorFactory + Clone,
100 H: HashedCursorFactory + Clone,
101{
102 #[allow(clippy::clone_on_copy)]
109 pub fn compute(mut self, state: HashedPostState) -> Result<B256Map<Bytes>, TrieWitnessError> {
110 let is_state_empty = state.is_empty();
111 if is_state_empty && !self.always_include_root_node {
112 return Ok(Default::default())
113 }
114
115 let proof_targets = if is_state_empty {
116 MultiProofTargetsV2 {
117 account_targets: vec![ProofV2Target::new(B256::ZERO)],
118 ..Default::default()
119 }
120 } else {
121 Self::get_proof_targets(&state)
122 };
123 let multiproof =
124 Proof::new(self.trie_cursor_factory.clone(), self.hashed_cursor_factory.clone())
125 .with_prefix_sets_mut(self.prefix_sets.clone())
126 .multiproof_v2(proof_targets)?;
127
128 if is_state_empty {
131 let (root_hash, root_node) = if let Some(root_node) =
132 multiproof.account_proofs.into_iter().find(|n| n.path.is_empty())
133 {
134 let bytes = Bytes::from(alloy_rlp::encode(&root_node.node));
135 (keccak256(&bytes), bytes)
136 } else {
137 (EMPTY_ROOT_HASH, Bytes::from([EMPTY_STRING_CODE]))
138 };
139 return Ok(B256Map::from_iter([(root_hash, root_node)]))
140 }
141
142 self.record_multiproof_nodes(&multiproof);
144
145 let mut sparse_trie = SparseStateTrie::new();
146 sparse_trie.reveal_decoded_multiproof_v2(multiproof)?;
147
148 let mut storage_removals: B256Map<B256Map<LeafUpdate>> = B256Map::default();
155 let mut storage_upserts: B256Map<B256Map<LeafUpdate>> = B256Map::default();
156 for (hashed_address, storage) in &state.storages {
157 for (&hashed_slot, value) in &storage.storage {
158 if value.is_zero() {
159 storage_removals
160 .entry(*hashed_address)
161 .or_default()
162 .insert(hashed_slot, LeafUpdate::Changed(vec![]));
163 } else {
164 storage_upserts.entry(*hashed_address).or_default().insert(
165 hashed_slot,
166 LeafUpdate::Changed(alloy_rlp::encode_fixed_size(value).to_vec()),
167 );
168 }
169 }
170 }
171
172 let storage_update_sets = if self.mode.is_canonical() {
173 [&mut storage_upserts, &mut storage_removals]
174 } else {
175 [&mut storage_removals, &mut storage_upserts]
176 };
177
178 for storage_updates in storage_update_sets {
180 loop {
181 let mut targets = MultiProofTargetsV2::default();
182
183 for (&hashed_address, slot_updates) in storage_updates.iter_mut() {
184 if slot_updates.is_empty() {
185 continue;
186 }
187 let storage_trie = sparse_trie
188 .storage_trie_mut(&hashed_address)
189 .expect("storage trie was revealed from multiproof");
190 storage_trie
191 .update_leaves(slot_updates, |key, parent| {
192 targets
193 .storage_targets
194 .entry(hashed_address)
195 .or_default()
196 .push(ProofV2Target::new(key).with_parent(parent));
197 })
198 .map_err(|err| {
199 SparseStateTrieErrorKind::SparseStorageTrie(
200 hashed_address,
201 err.into_kind(),
202 )
203 })?;
204 }
205
206 if targets.is_empty() {
207 break;
208 }
209
210 let multiproof = Proof::new(
211 self.trie_cursor_factory.clone(),
212 self.hashed_cursor_factory.clone(),
213 )
214 .with_prefix_sets_mut(self.prefix_sets.clone())
215 .multiproof_v2(targets)?;
216 self.record_multiproof_nodes(&multiproof);
217 sparse_trie.reveal_decoded_multiproof_v2(multiproof)?;
218 }
219 }
220
221 let mut account_removals: B256Map<LeafUpdate> = B256Map::default();
226 let mut account_upserts: B256Map<LeafUpdate> = B256Map::default();
227 for &hashed_address in state.accounts.keys().chain(state.storages.keys()) {
228 if account_removals.contains_key(&hashed_address) ||
229 account_upserts.contains_key(&hashed_address)
230 {
231 continue;
232 }
233
234 let account = state
235 .accounts
236 .get(&hashed_address)
237 .ok_or(TrieWitnessError::MissingAccount(hashed_address))?
238 .clone()
239 .unwrap_or_default();
240
241 let storage_root =
242 if let Some(storage_trie) = sparse_trie.storage_trie_mut(&hashed_address) {
243 storage_trie.root(TrieNodeEpoch::UNMODIFIED)
244 } else {
245 let record_root_node = !self.mode.is_canonical() ||
246 state
247 .storages
248 .get(&hashed_address)
249 .is_some_and(|storage| !storage.storage.is_empty());
250 self.account_storage_root(hashed_address, record_root_node)?
251 };
252
253 if account.is_empty() && storage_root == EMPTY_ROOT_HASH {
254 account_removals.insert(hashed_address, LeafUpdate::Changed(vec![]));
255 } else {
256 let mut rlp = Vec::with_capacity(TRIE_ACCOUNT_RLP_MAX_SIZE);
257 account.into_trie_account(storage_root).encode(&mut rlp);
258 account_upserts.insert(hashed_address, LeafUpdate::Changed(rlp));
259 }
260 }
261
262 let account_update_sets = if self.mode.is_canonical() {
263 [&mut account_upserts, &mut account_removals]
264 } else {
265 [&mut account_removals, &mut account_upserts]
266 };
267
268 for account_updates in account_update_sets {
270 loop {
271 let mut targets = MultiProofTargetsV2::default();
272
273 sparse_trie
274 .trie_mut()
275 .update_leaves(account_updates, |key, parent| {
276 targets.account_targets.push(ProofV2Target::new(key).with_parent(parent));
277 })
278 .map_err(SparseStateTrieErrorKind::from)?;
279
280 if targets.is_empty() {
281 break;
282 }
283
284 let multiproof = Proof::new(
285 self.trie_cursor_factory.clone(),
286 self.hashed_cursor_factory.clone(),
287 )
288 .with_prefix_sets_mut(self.prefix_sets.clone())
289 .multiproof_v2(targets)?;
290 self.record_multiproof_nodes(&multiproof);
291 sparse_trie.reveal_decoded_multiproof_v2(multiproof)?;
292 }
293 }
294
295 if self.mode.is_canonical() {
296 self.witness.retain(|_, value| value.as_ref() != [EMPTY_STRING_CODE]);
299 }
300
301 Ok(self.witness)
302 }
303
304 fn record_multiproof_nodes(&mut self, multiproof: &DecodedMultiProofV2) {
306 let mut encoded = Vec::new();
307 for proof_node in &multiproof.account_proofs {
308 self.record_witness_node(&proof_node.node, &mut encoded);
309 }
310 for proof_nodes in multiproof.storage_proofs.values() {
311 for proof_node in proof_nodes {
312 self.record_witness_node(&proof_node.node, &mut encoded);
313 }
314 }
315 }
316
317 fn record_witness_node(&mut self, node: &TrieNodeV2, encoded: &mut Vec<u8>) {
319 encoded.clear();
320 node.encode(encoded);
321 let hash = keccak256(encoded.as_slice());
322 self.witness.entry(hash).or_insert_with(|| Bytes::copy_from_slice(encoded));
323
324 if let TrieNodeV2::Branch(branch) = node &&
325 !branch.key.is_empty()
326 {
327 encoded.clear();
328 BranchNodeRef::new(&branch.stack, branch.state_mask).encode(encoded);
329 let hash = keccak256(encoded.as_slice());
330 self.witness.entry(hash).or_insert_with(|| Bytes::copy_from_slice(encoded));
331 }
332 }
333
334 fn account_storage_root(
337 &mut self,
338 hashed_address: B256,
339 record_root_node: bool,
340 ) -> Result<B256, TrieWitnessError> {
341 let storage_trie_cursor = self
342 .trie_cursor_factory
343 .storage_trie_cursor(hashed_address)
344 .map_err(StateProofError::from)?;
345 let hashed_storage_cursor = self
346 .hashed_cursor_factory
347 .hashed_storage_cursor(hashed_address)
348 .map_err(StateProofError::from)?;
349 let mut calculator = proof_v2::StorageProofCalculator::new_storage(
350 storage_trie_cursor,
351 hashed_storage_cursor,
352 );
353 if let Some(prefix_set) = self.prefix_sets.storage_prefix_sets.get(&hashed_address) {
354 calculator = calculator.with_prefix_set(prefix_set.clone().freeze());
355 }
356 let root_node = calculator.storage_root_node(hashed_address)?;
357 let root_hash = calculator
358 .compute_root_hash(core::slice::from_ref(&root_node))?
359 .unwrap_or(EMPTY_ROOT_HASH);
360 drop(calculator);
361 if record_root_node {
362 let mut encoded = Vec::new();
363 self.record_witness_node(&root_node.node, &mut encoded);
364 }
365 Ok(root_hash)
366 }
367
368 fn get_proof_targets(state: &HashedPostState) -> MultiProofTargetsV2 {
371 let mut targets = MultiProofTargetsV2::default();
372 for &hashed_address in state.accounts.keys() {
373 targets.account_targets.push(ProofV2Target::new(hashed_address));
374 }
375 for (&hashed_address, storage) in &state.storages {
376 if !state.accounts.contains_key(&hashed_address) {
377 targets.account_targets.push(ProofV2Target::new(hashed_address));
378 }
379 if storage.storage.is_empty() {
382 continue;
383 }
384 let storage_keys = storage.storage.keys().map(|k| ProofV2Target::new(*k)).collect();
385 targets.storage_targets.insert(hashed_address, storage_keys);
386 }
387 targets
388 }
389}