1use crate::{
7 updates::{TrieUpdates, TrieUpdatesSorted},
8 HashedPostState, HashedPostStateSorted,
9};
10use alloc::sync::Arc;
11use core::fmt;
12use reth_primitives_traits::sync::OnceLock;
13
14#[derive(Clone, Debug, Default, PartialEq, Eq)]
19#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
20pub struct SortedTrieData {
21 pub hashed_state: Arc<HashedPostStateSorted>,
23 pub trie_updates: Arc<TrieUpdatesSorted>,
25}
26
27impl SortedTrieData {
28 pub const fn new(
30 hashed_state: Arc<HashedPostStateSorted>,
31 trie_updates: Arc<TrieUpdatesSorted>,
32 ) -> Self {
33 Self { hashed_state, trie_updates }
34 }
35}
36
37#[derive(Clone, Debug, Default, PartialEq, Eq)]
39pub struct ComputedTrieData {
40 pub sorted: SortedTrieData,
42}
43
44impl ComputedTrieData {
45 pub const fn new(
47 hashed_state: Arc<HashedPostStateSorted>,
48 trie_updates: Arc<TrieUpdatesSorted>,
49 ) -> Self {
50 Self { sorted: SortedTrieData::new(hashed_state, trie_updates) }
51 }
52}
53
54pub struct LazyTrieData {
67 data: Arc<OnceLock<ComputedTrieData>>,
69 mode: LazyTrieDataMode,
73}
74
75impl Clone for LazyTrieData {
76 fn clone(&self) -> Self {
77 Self { data: Arc::clone(&self.data), mode: self.mode.clone() }
78 }
79}
80
81impl fmt::Debug for LazyTrieData {
82 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
83 f.debug_struct("LazyTrieData")
84 .field("data", &if self.data.get().is_some() { "initialized" } else { "pending" })
85 .finish()
86 }
87}
88
89impl PartialEq for LazyTrieData {
90 fn eq(&self, other: &Self) -> bool {
91 self.get() == other.get()
92 }
93}
94
95impl Eq for LazyTrieData {}
96
97impl LazyTrieData {
98 pub fn ready(sorted: ComputedTrieData) -> Self {
100 Self { data: Arc::new(OnceLock::from(sorted)), mode: LazyTrieDataMode::Ready }
101 }
102
103 pub fn deferred(compute: impl Fn() -> ComputedTrieData + Send + Sync + 'static) -> Self {
108 Self {
109 data: Arc::new(OnceLock::new()),
110 mode: LazyTrieDataMode::Deferred(Arc::new(compute)),
111 }
112 }
113
114 #[cfg(feature = "std")]
116 pub fn pending(
117 hashed_state: Arc<HashedPostState>,
118 trie_updates: Arc<TrieUpdates>,
119 ) -> (Self, LazyTrieDataProducer) {
120 let value = Arc::new(OnceLock::new());
121 (
122 Self { data: Arc::clone(&value), mode: LazyTrieDataMode::Pending },
123 LazyTrieDataProducer { value, inputs: PendingInputs { hashed_state, trie_updates } },
124 )
125 }
126
127 pub fn get(&self) -> &ComputedTrieData {
133 match &self.mode {
134 LazyTrieDataMode::Ready => self.data.get().expect("LazyTrieData must be initialized"),
135 LazyTrieDataMode::Deferred(compute) => self.data.get_or_init(|| compute.as_ref()()),
136 #[cfg(feature = "std")]
137 LazyTrieDataMode::Pending => self.data.wait(),
138 }
139 }
140
141 pub fn hashed_state(&self) -> Arc<HashedPostStateSorted> {
145 Arc::clone(&self.get().sorted.hashed_state)
146 }
147
148 pub fn trie_updates(&self) -> Arc<TrieUpdatesSorted> {
152 Arc::clone(&self.get().sorted.trie_updates)
153 }
154
155 pub fn sorted_trie_data(&self) -> SortedTrieData {
159 self.get().sorted.clone()
160 }
161}
162
163#[cfg(feature = "serde")]
164impl serde::Serialize for LazyTrieData {
165 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
166 where
167 S: serde::Serializer,
168 {
169 self.get().sorted.serialize(serializer)
170 }
171}
172
173#[cfg(feature = "serde")]
174impl<'de> serde::Deserialize<'de> for LazyTrieData {
175 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
176 where
177 D: serde::Deserializer<'de>,
178 {
179 let data = SortedTrieData::deserialize(deserializer)?;
180 Ok(Self::ready(ComputedTrieData::new(data.hashed_state, data.trie_updates)))
181 }
182}
183
184#[derive(Clone)]
185enum LazyTrieDataMode {
186 Ready,
187 Deferred(Arc<dyn Fn() -> ComputedTrieData + Send + Sync>),
188 #[cfg(feature = "std")]
189 Pending,
190}
191
192impl fmt::Debug for LazyTrieDataMode {
193 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
194 match self {
195 Self::Ready => write!(f, "Ready"),
196 Self::Deferred(_) => write!(f, "Deferred(..)"),
197 #[cfg(feature = "std")]
198 Self::Pending => write!(f, "Pending"),
199 }
200 }
201}
202
203#[must_use = "LazyTrieDataProducer must be consumed with compute_and_publish to wake trie data waiters"]
205pub struct LazyTrieDataProducer {
206 value: Arc<OnceLock<ComputedTrieData>>,
208 inputs: PendingInputs,
210}
211
212impl fmt::Debug for LazyTrieDataProducer {
213 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
214 f.debug_struct("LazyTrieDataProducer").field("inputs", &self.inputs).finish_non_exhaustive()
215 }
216}
217
218impl LazyTrieDataProducer {
219 pub fn compute_and_publish(self) -> ComputedTrieData {
221 let Self { value, inputs } = self;
222 let computed = Self::sort(inputs.hashed_state, inputs.trie_updates);
223 let _ = value.set(computed.clone());
224 computed
225 }
226
227 pub fn sort(
229 hashed_state: Arc<HashedPostState>,
230 trie_updates: Arc<TrieUpdates>,
231 ) -> ComputedTrieData {
232 #[cfg(feature = "rayon")]
233 let (sorted_hashed_state, sorted_trie_updates) = rayon::join(
234 || match Arc::try_unwrap(hashed_state) {
235 Ok(state) => state.into_sorted(),
236 Err(arc) => arc.clone_into_sorted(),
237 },
238 || match Arc::try_unwrap(trie_updates) {
239 Ok(updates) => updates.into_sorted(),
240 Err(arc) => arc.clone_into_sorted(),
241 },
242 );
243
244 #[cfg(not(feature = "rayon"))]
245 let (sorted_hashed_state, sorted_trie_updates) = (
246 match Arc::try_unwrap(hashed_state) {
247 Ok(state) => state.into_sorted(),
248 Err(arc) => arc.clone_into_sorted(),
249 },
250 match Arc::try_unwrap(trie_updates) {
251 Ok(updates) => updates.into_sorted(),
252 Err(arc) => arc.clone_into_sorted(),
253 },
254 );
255
256 ComputedTrieData::new(Arc::new(sorted_hashed_state), Arc::new(sorted_trie_updates))
257 }
258}
259
260#[derive(Clone, Debug)]
262struct PendingInputs {
263 hashed_state: Arc<HashedPostState>,
265 trie_updates: Arc<TrieUpdates>,
267}
268
269#[cfg(test)]
270mod tests {
271 use crate::HashedStorage;
272
273 use super::*;
274 use alloy_primitives::{map::B256Map, B256, U256};
275 use reth_primitives_traits::Account;
276 use std::{
277 thread,
278 time::{Duration, Instant},
279 };
280
281 fn empty_pending() -> (LazyTrieData, LazyTrieDataProducer) {
282 LazyTrieData::pending(
283 Arc::new(HashedPostState::default()),
284 Arc::new(TrieUpdates::default()),
285 )
286 }
287
288 #[test]
289 fn test_lazy_ready_is_initialized() {
290 let lazy = LazyTrieData::ready(ComputedTrieData::default());
291 let _ = lazy.hashed_state();
292 let _ = lazy.trie_updates();
293 }
294
295 #[test]
296 fn test_lazy_clone_shares_state() {
297 let lazy1 = LazyTrieData::ready(ComputedTrieData::default());
298 let lazy2 = lazy1.clone();
299
300 assert!(Arc::ptr_eq(&lazy1.hashed_state(), &lazy2.hashed_state()));
302 assert!(Arc::ptr_eq(&lazy1.trie_updates(), &lazy2.trie_updates()));
303 }
304
305 #[test]
306 fn test_lazy_deferred() {
307 let lazy = LazyTrieData::deferred(ComputedTrieData::default);
308 assert!(lazy.hashed_state().is_empty());
309 assert!(lazy.trie_updates().is_empty());
310 }
311
312 #[test]
313 fn ready_returns_immediately() {
314 let bundle = ComputedTrieData::default();
315 let deferred = LazyTrieData::ready(bundle.clone());
316
317 let result = deferred.get();
318
319 assert_eq!(result.sorted.hashed_state.total_len(), bundle.sorted.hashed_state.total_len());
320 assert_eq!(result.sorted.trie_updates.total_len(), bundle.sorted.trie_updates.total_len());
321 }
322
323 #[test]
324 fn pending_waits_for_task_and_caches_result() {
325 let (deferred, task) = empty_pending();
326
327 let published = task.compute_and_publish();
328 let first = deferred.get();
329 let second = deferred.get();
330
331 assert!(Arc::ptr_eq(&published.sorted.hashed_state, &first.sorted.hashed_state));
332 assert!(Arc::ptr_eq(&published.sorted.trie_updates, &first.sorted.trie_updates));
333 assert!(Arc::ptr_eq(&first.sorted.hashed_state, &second.sorted.hashed_state));
334 assert!(Arc::ptr_eq(&first.sorted.trie_updates, &second.sorted.trie_updates));
335 }
336
337 #[test]
338 fn pending_wait_blocks_until_task_publishes() {
339 let (deferred, task) = empty_pending();
340
341 let handle = thread::spawn(move || deferred.get().clone());
342 thread::sleep(Duration::from_millis(20));
343 assert!(!handle.is_finished());
344
345 let published = task.compute_and_publish();
346 let result = handle.join().unwrap();
347
348 assert!(Arc::ptr_eq(&published.sorted.hashed_state, &result.sorted.hashed_state));
349 assert!(Arc::ptr_eq(&published.sorted.trie_updates, &result.sorted.trie_updates));
350 }
351
352 #[test]
353 fn concurrent_waits_share_published_result() {
354 let (deferred, task) = empty_pending();
355 let deferred2 = deferred.clone();
356
357 let handle = thread::spawn(move || deferred2.get().clone());
358 let published = task.compute_and_publish();
359 let result1 = deferred.get().clone();
360 let result2 = handle.join().unwrap();
361
362 assert!(Arc::ptr_eq(&published.sorted.hashed_state, &result1.sorted.hashed_state));
363 assert!(Arc::ptr_eq(&published.sorted.trie_updates, &result1.sorted.trie_updates));
364 assert!(Arc::ptr_eq(&result1.sorted.hashed_state, &result2.sorted.hashed_state));
365 assert!(Arc::ptr_eq(&result1.sorted.trie_updates, &result2.sorted.trie_updates));
366 }
367
368 #[test]
369 fn sorts_non_empty_inputs() {
370 let hashed_address = B256::with_last_byte(1);
371 let hashed_slot = B256::with_last_byte(2);
372 let hashed_state = HashedPostState::default()
373 .with_accounts([(hashed_address, Some(Account::default()))])
374 .with_storages([(
375 hashed_address,
376 HashedStorage::from_iter(false, [(hashed_slot, U256::from(1))]),
377 )]);
378
379 let (deferred, task) =
380 LazyTrieData::pending(Arc::new(hashed_state), Arc::new(TrieUpdates::default()));
381 let _ = task.compute_and_publish();
382 let result = deferred.get().clone();
383
384 assert_eq!(result.sorted.hashed_state.total_len(), 2);
385 assert_eq!(result.sorted.trie_updates.total_len(), 0);
386 }
387
388 #[test]
389 fn wait_does_not_block_after_first_compute() {
390 let mut accounts = B256Map::default();
391 for i in 0..100 {
392 accounts.insert(B256::with_last_byte(i), Some(Account::default()));
393 }
394 let (deferred, task) = LazyTrieData::pending(
395 Arc::new(HashedPostState { accounts, storages: Default::default() }),
396 Arc::new(TrieUpdates::default()),
397 );
398
399 let _ = task.compute_and_publish();
400 let _ = deferred.get().clone();
401 let start = Instant::now();
402 let _ = deferred.get().clone();
403
404 assert!(start.elapsed() < Duration::from_millis(10));
405 }
406}