1use alloc::{string::String, vec::Vec};
2use core::{fmt, slice};
3
4use super::{InnerNodeInfo, MerkleError, MerklePath, NodeIndex, Poseidon2, Word};
5use crate::utils::{assume_init_vec, uninit_vector, word_to_hex};
6
7#[derive(Debug, Clone, PartialEq, Eq)]
12#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
13pub struct MerkleTree {
14 nodes: Vec<Word>,
15}
16
17impl MerkleTree {
18 pub fn new<T>(leaves: T) -> Result<Self, MerkleError>
25 where
26 T: AsRef<[Word]>,
27 {
28 let leaves = leaves.as_ref();
29 let n = leaves.len();
30 if n <= 1 {
31 return Err(MerkleError::DepthTooSmall(n as u8));
32 } else if !n.is_power_of_two() {
33 return Err(MerkleError::NumLeavesNotPowerOfTwo(n));
34 }
35
36 let mut nodes = uninit_vector::<Word>(2 * n);
41 nodes[0].write(Word::default());
42
43 nodes[n..].iter_mut().zip(leaves).for_each(|(node, leaf)| {
45 node.write(*leaf);
46 });
47
48 for i in (1..n).rev() {
50 let left = unsafe { nodes[2 * i].assume_init_read() };
53 let right = unsafe { nodes[2 * i + 1].assume_init_read() };
54 nodes[i].write(Poseidon2::merge(&[left, right]));
55 }
56
57 let nodes = unsafe { assume_init_vec(nodes) };
59
60 Ok(Self { nodes })
61 }
62
63 pub fn root(&self) -> Word {
68 self.nodes[1]
69 }
70
71 pub fn depth(&self) -> u8 {
75 (self.nodes.len() / 2).ilog2() as u8
76 }
77
78 pub fn get_node(&self, index: NodeIndex) -> Result<Word, MerkleError> {
85 if index.is_root() {
86 return Err(MerkleError::DepthTooSmall(index.depth()));
87 } else if index.depth() > self.depth() {
88 return Err(MerkleError::DepthTooBig(index.depth() as u64));
89 }
90
91 let pos = index.to_scalar_index()? as usize;
92 Ok(self.nodes[pos])
93 }
94
95 pub fn get_path(&self, index: NodeIndex) -> Result<MerklePath, MerkleError> {
103 if index.is_root() {
104 return Err(MerkleError::DepthTooSmall(index.depth()));
105 } else if index.depth() > self.depth() {
106 return Err(MerkleError::DepthTooBig(index.depth() as u64));
107 }
108
109 Ok(MerklePath::from(Vec::from_iter(
110 index.proof_indices().map(|index| self.get_node(index).unwrap()),
111 )))
112 }
113
114 pub fn leaves(&self) -> impl Iterator<Item = (u64, &Word)> {
119 let leaves_start = self.nodes.len() / 2;
120 self.nodes.iter().skip(leaves_start).enumerate().map(|(i, v)| (i as u64, v))
121 }
122
123 pub fn inner_nodes(&self) -> InnerNodeIterator<'_> {
127 InnerNodeIterator {
128 nodes: &self.nodes,
129 index: 1, }
131 }
132
133 pub fn update_leaf<'a>(&'a mut self, index_value: u64, value: Word) -> Result<(), MerkleError> {
141 let mut index = NodeIndex::new(self.depth(), index_value)?;
142
143 debug_assert_eq!(self.nodes.len() & 1, 0);
152 let n = self.nodes.len() / 2;
153
154 let ptr = self.nodes.as_ptr() as *const [Word; 2];
160 let pairs: &'a [[Word; 2]] = unsafe { slice::from_raw_parts(ptr, n) };
161
162 let pos = index.to_scalar_index()? as usize;
164 self.nodes[pos] = value;
165
166 for _ in 0..index.depth() {
168 index.move_up();
169 let pos = index.to_scalar_index()? as usize;
170 let value = Poseidon2::merge(&pairs[pos]);
171 self.nodes[pos] = value;
172 }
173
174 Ok(())
175 }
176}
177
178impl TryFrom<&[Word]> for MerkleTree {
182 type Error = MerkleError;
183
184 fn try_from(value: &[Word]) -> Result<Self, Self::Error> {
185 MerkleTree::new(value)
186 }
187}
188
189pub struct InnerNodeIterator<'a> {
196 nodes: &'a Vec<Word>,
197 index: usize,
198}
199
200impl Iterator for InnerNodeIterator<'_> {
201 type Item = InnerNodeInfo;
202
203 fn next(&mut self) -> Option<Self::Item> {
204 if self.index < self.nodes.len() / 2 {
205 let value = self.index;
206 let left = self.index * 2;
207 let right = left + 1;
208
209 self.index += 1;
210
211 Some(InnerNodeInfo {
212 value: self.nodes[value],
213 left: self.nodes[left],
214 right: self.nodes[right],
215 })
216 } else {
217 None
218 }
219 }
220}
221
222pub fn tree_to_text(tree: &MerkleTree) -> Result<String, fmt::Error> {
227 let indent = " ";
228 let mut s = String::new();
229 s.push_str(&word_to_hex(&tree.root())?);
230 s.push('\n');
231 for d in 1..=tree.depth() {
232 let entries = 2u64.pow(d.into());
233 for i in 0..entries {
234 let index = NodeIndex::new(d, i).expect("The index must always be valid");
235 let node = tree.get_node(index).expect("The node must always be found");
236
237 for _ in 0..d {
238 s.push_str(indent);
239 }
240 s.push_str(&word_to_hex(&node)?);
241 s.push('\n');
242 }
243 }
244
245 Ok(s)
246}
247
248pub fn path_to_text(path: &MerklePath) -> Result<String, fmt::Error> {
250 let mut s = String::new();
251 s.push('[');
252
253 for el in path.iter() {
254 s.push_str(&word_to_hex(el)?);
255 s.push_str(", ");
256 }
257
258 if !path.is_empty() {
260 s.pop();
261 s.pop();
262 }
263 s.push(']');
264
265 Ok(s)
266}
267
268#[cfg(test)]
272mod tests {
273 use core::mem::size_of;
274
275 use proptest::prelude::*;
276
277 use super::*;
278 use crate::{
279 Felt,
280 merkle::{int_to_leaf, int_to_node},
281 };
282
283 const LEAVES4: [Word; Word::NUM_ELEMENTS] =
284 [int_to_node(1), int_to_node(2), int_to_node(3), int_to_node(4)];
285
286 const LEAVES8: [Word; 8] = [
287 int_to_node(1),
288 int_to_node(2),
289 int_to_node(3),
290 int_to_node(4),
291 int_to_node(5),
292 int_to_node(6),
293 int_to_node(7),
294 int_to_node(8),
295 ];
296
297 #[test]
298 fn build_merkle_tree() {
299 let tree = MerkleTree::new(LEAVES4).unwrap();
300 assert_eq!(8, tree.nodes.len());
301
302 for (a, b) in tree.nodes.iter().skip(4).zip(LEAVES4.iter()) {
304 assert_eq!(a, b);
305 }
306
307 let (root, node2, node3) = compute_internal_nodes();
308
309 assert_eq!(root, tree.nodes[1]);
310 assert_eq!(node2, tree.nodes[2]);
311 assert_eq!(node3, tree.nodes[3]);
312
313 assert_eq!(root, tree.root());
314 }
315
316 #[test]
317 fn get_leaf() {
318 let tree = MerkleTree::new(LEAVES4).unwrap();
319
320 assert_eq!(LEAVES4[0], tree.get_node(NodeIndex::make(2, 0)).unwrap());
322 assert_eq!(LEAVES4[1], tree.get_node(NodeIndex::make(2, 1)).unwrap());
323 assert_eq!(LEAVES4[2], tree.get_node(NodeIndex::make(2, 2)).unwrap());
324 assert_eq!(LEAVES4[3], tree.get_node(NodeIndex::make(2, 3)).unwrap());
325
326 let (_, node2, node3) = compute_internal_nodes();
328
329 assert_eq!(node2, tree.get_node(NodeIndex::make(1, 0)).unwrap());
330 assert_eq!(node3, tree.get_node(NodeIndex::make(1, 1)).unwrap());
331 }
332
333 #[test]
334 fn get_path() {
335 let tree = MerkleTree::new(LEAVES4).unwrap();
336
337 let (_, node2, node3) = compute_internal_nodes();
338
339 assert_eq!(vec![LEAVES4[1], node3], *tree.get_path(NodeIndex::make(2, 0)).unwrap());
341 assert_eq!(vec![LEAVES4[0], node3], *tree.get_path(NodeIndex::make(2, 1)).unwrap());
342 assert_eq!(vec![LEAVES4[3], node2], *tree.get_path(NodeIndex::make(2, 2)).unwrap());
343 assert_eq!(vec![LEAVES4[2], node2], *tree.get_path(NodeIndex::make(2, 3)).unwrap());
344
345 assert_eq!(vec![node3], *tree.get_path(NodeIndex::make(1, 0)).unwrap());
347 assert_eq!(vec![node2], *tree.get_path(NodeIndex::make(1, 1)).unwrap());
348 }
349
350 #[test]
351 fn update_leaf() {
352 let mut tree = MerkleTree::new(LEAVES8).unwrap();
353
354 let value = 3;
356 let new_node = int_to_leaf(9);
357 let mut expected_leaves = LEAVES8.to_vec();
358 expected_leaves[value as usize] = new_node;
359 let expected_tree = MerkleTree::new(expected_leaves.clone()).unwrap();
360
361 tree.update_leaf(value, new_node).unwrap();
362 assert_eq!(expected_tree.nodes, tree.nodes);
363
364 let value = 6;
366 let new_node = int_to_leaf(10);
367 expected_leaves[value as usize] = new_node;
368 let expected_tree = MerkleTree::new(expected_leaves.clone()).unwrap();
369
370 tree.update_leaf(value, new_node).unwrap();
371 assert_eq!(expected_tree.nodes, tree.nodes);
372 }
373
374 #[test]
375 fn nodes() -> Result<(), MerkleError> {
376 let tree = MerkleTree::new(LEAVES4).unwrap();
377 let root = tree.root();
378 let l1n0 = tree.get_node(NodeIndex::make(1, 0))?;
379 let l1n1 = tree.get_node(NodeIndex::make(1, 1))?;
380 let l2n0 = tree.get_node(NodeIndex::make(2, 0))?;
381 let l2n1 = tree.get_node(NodeIndex::make(2, 1))?;
382 let l2n2 = tree.get_node(NodeIndex::make(2, 2))?;
383 let l2n3 = tree.get_node(NodeIndex::make(2, 3))?;
384
385 let nodes: Vec<InnerNodeInfo> = tree.inner_nodes().collect();
386 let expected = vec![
387 InnerNodeInfo { value: root, left: l1n0, right: l1n1 },
388 InnerNodeInfo { value: l1n0, left: l2n0, right: l2n1 },
389 InnerNodeInfo { value: l1n1, left: l2n2, right: l2n3 },
390 ];
391 assert_eq!(nodes, expected);
392
393 Ok(())
394 }
395
396 proptest! {
397 #[test]
398 fn arbitrary_word_can_be_represented_as_digest(
399 a in prop::num::u64::ANY,
400 b in prop::num::u64::ANY,
401 c in prop::num::u64::ANY,
402 d in prop::num::u64::ANY,
403 ) {
404 let word = [Felt::new_unchecked(a), Felt::new_unchecked(b), Felt::new_unchecked(c), Felt::new_unchecked(d)];
410 let digest = Word::from(word);
411
412 let word_ptr = word.as_ptr() as *const u8;
414 let digest_ptr = digest.as_ptr() as *const u8;
415 assert_ne!(word_ptr, digest_ptr);
416
417 let word_bytes = unsafe { slice::from_raw_parts(word_ptr, size_of::<Word>()) };
419 let digest_bytes = unsafe { slice::from_raw_parts(digest_ptr, size_of::<Word>()) };
420 assert_eq!(word_bytes, digest_bytes);
421 }
422 }
423
424 fn compute_internal_nodes() -> (Word, Word, Word) {
428 let node2 = Poseidon2::hash_elements(&[*LEAVES4[0], *LEAVES4[1]].concat());
429 let node3 = Poseidon2::hash_elements(&[*LEAVES4[2], *LEAVES4[3]].concat());
430 let root = Poseidon2::merge(&[node2, node3]);
431
432 (root, node2, node3)
433 }
434}