arbos/merkle_accumulator/
mod.rs

1use alloy_primitives::{B256, U256, keccak256};
2use arb_storage::{Storage, StorageBackedUint64, StorageBackend, SystemStateBackend};
3use revm::Database;
4
5mod error;
6pub use error::MerkleAccumulatorError;
7
8/// Event emitted when a Merkle tree node is updated during append.
9#[derive(Debug, Clone)]
10pub struct MerkleTreeNodeEvent {
11    pub level: u64,
12    pub num_leaves: u64,
13    pub hash: B256,
14}
15
16/// Storage-backed Merkle accumulator.
17pub struct MerkleAccumulator<'a, D> {
18    backing_storage: Storage<'a, D>,
19    size: StorageBackedUint64,
20}
21
22pub fn initialize_merkle_accumulator<D: Database>(_sto: &Storage<'_, D>) {
23    // no-op
24}
25
26pub fn open_merkle_accumulator<D>(sto: Storage<'_, D>) -> MerkleAccumulator<'_, D> {
27    let size = StorageBackedUint64::new(sto.base_key(), 0);
28    MerkleAccumulator {
29        backing_storage: sto,
30        size,
31    }
32}
33
34/// Returns the number of partial tree hashes needed for a given size.
35/// This is the bit-length of `size` (i.e. floor(log2(size)) + 1).
36pub fn calc_num_partials(size: u64) -> u64 {
37    if size == 0 {
38        return 0;
39    }
40    64 - size.leading_zeros() as u64
41}
42
43impl<D> MerkleAccumulator<'_, D> {
44    fn partial_slot(&self, level: u64) -> U256 {
45        self.backing_storage.new_slot(2 + level)
46    }
47
48    fn read_partial<B: SystemStateBackend>(
49        &self,
50        backend: &mut B,
51        level: u64,
52    ) -> Result<B256, MerkleAccumulatorError> {
53        let slot = self.partial_slot(level);
54        let value = backend
55            .sload_system(self.backing_storage.account(), slot)
56            .map_err(Into::into)?;
57        Ok(B256::from(value.to_be_bytes::<32>()))
58    }
59
60    fn write_partial<B: StorageBackend>(
61        &self,
62        backend: &mut B,
63        level: u64,
64        val: B256,
65    ) -> Result<(), MerkleAccumulatorError> {
66        let slot = self.partial_slot(level);
67        backend
68            .sstore(
69                self.backing_storage.account(),
70                slot,
71                U256::from_be_bytes(val.0),
72            )
73            .map_err(Into::into)?;
74        Ok(())
75    }
76
77    pub fn append<B: StorageBackend>(
78        &self,
79        backend: &mut B,
80        item_hash: B256,
81    ) -> Result<Vec<MerkleTreeNodeEvent>, MerkleAccumulatorError> {
82        let current_size = self.size.get(backend)?;
83        let new_size = current_size + 1;
84        self.size.set(backend, new_size)?;
85
86        let mut events = Vec::new();
87        let mut level = 0u64;
88        let mut so_far = keccak256(item_hash.as_slice());
89
90        loop {
91            if level == calc_num_partials(current_size) {
92                self.write_partial(backend, level, so_far)?;
93                return Ok(events);
94            }
95
96            let this_level = self.read_partial(backend, level)?;
97            if this_level == B256::ZERO {
98                self.write_partial(backend, level, so_far)?;
99                return Ok(events);
100            }
101
102            let mut combined = Vec::with_capacity(64);
103            combined.extend_from_slice(this_level.as_slice());
104            combined.extend_from_slice(so_far.as_slice());
105            so_far = keccak256(&combined);
106
107            self.write_partial(backend, level, B256::ZERO)?;
108
109            level += 1;
110            events.push(MerkleTreeNodeEvent {
111                level,
112                num_leaves: new_size - 1,
113                hash: so_far,
114            });
115        }
116    }
117
118    pub fn size<B: SystemStateBackend>(
119        &self,
120        backend: &mut B,
121    ) -> Result<u64, MerkleAccumulatorError> {
122        Ok(self.size.get(backend)?)
123    }
124
125    /// Read a single partial hash without re-reading the accumulator size.
126    pub fn partial_at<B: SystemStateBackend>(
127        &self,
128        backend: &mut B,
129        level: u64,
130    ) -> Result<B256, MerkleAccumulatorError> {
131        self.read_partial(backend, level)
132    }
133
134    pub fn root<B: SystemStateBackend>(
135        &self,
136        backend: &mut B,
137    ) -> Result<B256, MerkleAccumulatorError> {
138        let size = self.size.get(backend)?;
139        if size == 0 {
140            return Ok(B256::ZERO);
141        }
142
143        let mut hash_so_far: Option<B256> = None;
144        let mut capacity_in_hash = 0u64;
145        let mut capacity = 1u64;
146
147        for level in 0..calc_num_partials(size) {
148            let partial = self.read_partial(backend, level)?;
149            if partial != B256::ZERO {
150                if let Some(ref mut current) = hash_so_far {
151                    while capacity_in_hash < capacity {
152                        let mut combined = Vec::with_capacity(64);
153                        combined.extend_from_slice(current.as_slice());
154                        combined.extend_from_slice(&[0u8; 32]);
155                        *current = keccak256(&combined);
156                        capacity_in_hash *= 2;
157                    }
158
159                    let mut combined = Vec::with_capacity(64);
160                    combined.extend_from_slice(partial.as_slice());
161                    combined.extend_from_slice(current.as_slice());
162                    *current = keccak256(&combined);
163                    capacity_in_hash = 2 * capacity;
164                } else {
165                    hash_so_far = Some(partial);
166                    capacity_in_hash = capacity;
167                }
168            }
169            capacity *= 2;
170        }
171
172        Ok(hash_so_far.unwrap_or(B256::ZERO))
173    }
174
175    pub fn get_partials<B: SystemStateBackend>(
176        &self,
177        backend: &mut B,
178    ) -> Result<Vec<B256>, MerkleAccumulatorError> {
179        let size = self.size.get(backend)?;
180        let num = calc_num_partials(size);
181        let mut partials = Vec::with_capacity(num as usize);
182        for i in 0..num {
183            partials.push(self.read_partial(backend, i)?);
184        }
185        Ok(partials)
186    }
187
188    pub fn state_for_export<B: SystemStateBackend>(
189        &self,
190        backend: &mut B,
191    ) -> Result<(u64, B256, Vec<B256>), MerkleAccumulatorError> {
192        let root = self.root(backend)?;
193        let size = self.size.get(backend)?;
194        let partials = self.get_partials(backend)?;
195        Ok((size, root, partials))
196    }
197}
198
199/// In-memory (non-persistent) Merkle accumulator for export/import and testing.
200pub struct InMemoryMerkleAccumulator {
201    size: u64,
202    partials: Vec<B256>,
203}
204
205impl InMemoryMerkleAccumulator {
206    pub fn new() -> Self {
207        Self {
208            size: 0,
209            partials: Vec::new(),
210        }
211    }
212
213    pub fn from_partials(partials: Vec<B256>) -> Self {
214        let mut size = 0u64;
215        let mut level_size = 1u64;
216        for p in &partials {
217            if *p != B256::ZERO {
218                size += level_size;
219            }
220            level_size *= 2;
221        }
222        Self { size, partials }
223    }
224
225    pub fn size(&self) -> u64 {
226        self.size
227    }
228
229    fn get_partial(&self, level: u64) -> B256 {
230        self.partials
231            .get(level as usize)
232            .copied()
233            .unwrap_or(B256::ZERO)
234    }
235
236    fn set_partial(&mut self, level: u64, val: B256) {
237        let idx = level as usize;
238        if idx >= self.partials.len() {
239            self.partials.resize(idx + 1, B256::ZERO);
240        }
241        self.partials[idx] = val;
242    }
243
244    pub fn append(&mut self, item_hash: B256) -> Vec<MerkleTreeNodeEvent> {
245        let current_size = self.size;
246        self.size += 1;
247        let new_size = self.size;
248
249        let mut events = Vec::new();
250        let mut level = 0u64;
251        let mut so_far = keccak256(item_hash.as_slice());
252
253        loop {
254            if level == calc_num_partials(current_size) {
255                self.set_partial(level, so_far);
256                return events;
257            }
258
259            let this_level = self.get_partial(level);
260            if this_level == B256::ZERO {
261                self.set_partial(level, so_far);
262                return events;
263            }
264
265            let mut combined = Vec::with_capacity(64);
266            combined.extend_from_slice(this_level.as_slice());
267            combined.extend_from_slice(so_far.as_slice());
268            so_far = keccak256(&combined);
269
270            self.set_partial(level, B256::ZERO);
271
272            level += 1;
273            events.push(MerkleTreeNodeEvent {
274                level,
275                num_leaves: new_size - 1,
276                hash: so_far,
277            });
278        }
279    }
280
281    pub fn root(&self) -> B256 {
282        if self.size == 0 {
283            return B256::ZERO;
284        }
285
286        let mut hash_so_far: Option<B256> = None;
287        let mut capacity_in_hash = 0u64;
288        let mut capacity = 1u64;
289
290        for level in 0..calc_num_partials(self.size) {
291            let partial = self.get_partial(level);
292            if partial != B256::ZERO {
293                if let Some(ref mut current) = hash_so_far {
294                    while capacity_in_hash < capacity {
295                        let mut combined = Vec::with_capacity(64);
296                        combined.extend_from_slice(current.as_slice());
297                        combined.extend_from_slice(&[0u8; 32]);
298                        *current = keccak256(&combined);
299                        capacity_in_hash *= 2;
300                    }
301
302                    let mut combined = Vec::with_capacity(64);
303                    combined.extend_from_slice(partial.as_slice());
304                    combined.extend_from_slice(current.as_slice());
305                    *current = keccak256(&combined);
306                    capacity_in_hash = 2 * capacity;
307                } else {
308                    hash_so_far = Some(partial);
309                    capacity_in_hash = capacity;
310                }
311            }
312            capacity *= 2;
313        }
314
315        hash_so_far.unwrap_or(B256::ZERO)
316    }
317
318    pub fn get_partials(&self) -> Vec<B256> {
319        let num = calc_num_partials(self.size);
320        let mut partials = Vec::with_capacity(num as usize);
321        for i in 0..num {
322            partials.push(self.get_partial(i));
323        }
324        partials
325    }
326}
327
328impl Default for InMemoryMerkleAccumulator {
329    fn default() -> Self {
330        Self::new()
331    }
332}