1use std::{fmt, mem::MaybeUninit, slice};
5
6use binius_compute::{Allocator, GlobalAllocator, VecLike};
7use binius_field::Field;
8use binius_hash::CompressionFunction;
9use binius_utils::{
10 checked_arithmetics::checked_log_2,
11 rayon::{self, prelude::*, slice::ParallelSlice},
12};
13use digest::{
14 OutputSizeUser,
15 array::{Array, ArraySize},
16};
17
18use crate::{
19 parallel_compression::ParallelPseudoCompression, parallel_digest::ParallelDigest,
20 suite::ParallelHashSuite,
21};
22
23#[derive(Debug, thiserror::Error)]
28pub enum Error {
29 #[error("{n_values} values do not split into leaves of {batch_size}")]
31 IncorrectBatchSize {
32 n_values: usize,
34 batch_size: usize,
36 },
37 #[error("the leaf count {n_leaves} is not a power of two")]
39 PowerOfTwoLengthRequired {
40 n_leaves: usize,
42 },
43 #[error("layer depth {layer_depth} exceeds the tree depth {log_len}")]
45 IncorrectLayerDepth {
46 layer_depth: usize,
48 log_len: usize,
50 },
51 #[error("leaf index {index} is outside the 2^{log_len} leaves")]
53 IndexOutOfRange {
54 index: usize,
56 log_len: usize,
58 },
59}
60
61pub struct BinaryMerkleTree<D: Send, A: Allocator = GlobalAllocator> {
73 pub log_len: usize,
75 pub inner_nodes: A::Vec<D>,
77}
78
79impl<D: fmt::Debug + Send, A: Allocator> fmt::Debug for BinaryMerkleTree<D, A> {
82 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
83 f.debug_struct("BinaryMerkleTree")
84 .field("log_len", &self.log_len)
85 .field("inner_nodes", &&*self.inner_nodes)
86 .finish()
87 }
88}
89
90impl<N: ArraySize, A: Allocator> BinaryMerkleTree<Array<u8, N>, A> {
91 pub fn new<F, H>(elements: &[F], batch_size: usize, alloc: &A) -> Result<Self, Error>
104 where
105 F: Field,
106 H: ParallelHashSuite<LeafHash: OutputSizeUser<OutputSize = N>>,
107 {
108 if batch_size == 0 || !elements.len().is_multiple_of(batch_size) {
111 return Err(Error::IncorrectBatchSize {
112 n_values: elements.len(),
113 batch_size,
114 });
115 }
116
117 let len = elements.len() / batch_size;
119 if !len.is_power_of_two() {
120 return Err(Error::PowerOfTwoLengthRequired { n_leaves: len });
121 }
122
123 Ok(Self::from_leaves::<F, H, _>(
125 elements
126 .par_chunks(batch_size)
127 .map(|chunk| chunk.iter().copied()),
128 batch_size,
129 alloc,
130 ))
131 }
132
133 pub fn from_leaves<F, H, ParIter>(leaves: ParIter, n_items_per_input: usize, alloc: &A) -> Self
157 where
158 F: Field,
159 H: ParallelHashSuite<LeafHash: OutputSizeUser<OutputSize = N>>,
160 ParIter: IndexedParallelIterator<Item: IntoIterator<Item = F, IntoIter: Send>>,
161 {
162 let log_len = checked_log_2(leaves.len());
164
165 let total_length = (1 << (log_len + 1)) - 1;
168 let mut inner_nodes = alloc.alloc::<Array<u8, N>>(total_length);
169
170 {
172 let _span = tracing::debug_span!("hash_leaves").entered();
173
174 H::ParLeafHash::default().digest_with_const_len(
177 n_items_per_input,
178 leaves,
179 &mut inner_nodes.spare_capacity_mut()[..(1 << log_len)],
180 );
181 }
182
183 let (leaf_layer, remaining) = inner_nodes.spare_capacity_mut().split_at_mut(1 << log_len);
184
185 let leaves: &[Array<u8, N>] = unsafe {
186 leaf_layer.assume_init_mut()
188 };
189
190 if log_len > 0 {
194 let parallel_compression = H::ParCompression::default();
195 let ctx = BuildContext {
196 leaves,
197 log_len,
198 nodes: SendPtr(remaining.as_mut_ptr()),
199 compression: ¶llel_compression,
200 };
201 ctx.build_node(0, log_len);
202 }
203
204 unsafe {
205 inner_nodes.set_len(total_length);
209 }
210 Self {
211 log_len,
212 inner_nodes,
213 }
214 }
215}
216
217const SUBTREE_LOG_DEPTH: usize = 6;
224
225#[derive(Clone, Copy)]
229struct SendPtr<T>(*mut T);
230
231unsafe impl<T: Send> Send for SendPtr<T> {}
232unsafe impl<T: Sync> Sync for SendPtr<T> {}
233
234struct BuildContext<'a, D, C> {
236 leaves: &'a [D],
238 log_len: usize,
240 nodes: SendPtr<MaybeUninit<D>>,
242 compression: &'a C,
244}
245
246const fn node_offset(log_len: usize, leaves_start: usize, log_width: usize) -> usize {
251 let layer_depth = log_len - log_width;
253 let layer_start = (1 << log_len) - (1 << (layer_depth + 1));
254 layer_start + (leaves_start >> log_width)
255}
256
257impl<D: Clone, C: ParallelPseudoCompression<D, 2>> BuildContext<'_, D, C> {
258 fn fold_small_subtree(&self, leaves_start: usize, log_width: usize) -> D {
263 let mut cur = &self.leaves[leaves_start..leaves_start + (1 << log_width)];
264 for depth in 1..=log_width {
265 let width = 1 << (log_width - depth);
266 let offset = node_offset(self.log_len, leaves_start, depth);
267 let out = unsafe {
268 slice::from_raw_parts_mut(self.nodes.0.add(offset), width)
271 };
272 self.compression.parallel_compress(cur, out);
273 cur = unsafe {
274 out.assume_init_mut()
276 };
277 }
278 cur[0].clone()
279 }
280}
281
282impl<D: Clone + Send + Sync, C: ParallelPseudoCompression<D, 2> + Sync> BuildContext<'_, D, C> {
283 fn build_node(&self, leaves_start: usize, log_width: usize) -> D {
288 if log_width <= SUBTREE_LOG_DEPTH {
289 return self.fold_small_subtree(leaves_start, log_width);
290 }
291
292 let half = log_width - 1;
293 let mid = leaves_start + (1 << half);
294 let (left, right) =
295 rayon::join(|| self.build_node(leaves_start, half), || self.build_node(mid, half));
296
297 let node = self.compression.compression().compress([left, right]);
298
299 let offset = node_offset(self.log_len, leaves_start, log_width);
300 unsafe {
301 (*self.nodes.0.add(offset)).write(node.clone());
304 }
305 node
306 }
307}
308
309impl<D: Clone + Send, A: Allocator> BinaryMerkleTree<D, A> {
310 pub fn root(&self) -> D {
312 self.inner_nodes
313 .last()
314 .expect("MerkleTree inner nodes can't be empty")
315 .clone()
316 }
317
318 pub fn layer(&self, layer_depth: usize) -> Result<&[D], Error> {
324 if layer_depth > self.log_len {
325 return Err(Error::IncorrectLayerDepth {
326 layer_depth,
327 log_len: self.log_len,
328 });
329 }
330 let range_start = self.inner_nodes.len() + 1 - (1 << (layer_depth + 1));
331
332 Ok(&self.inner_nodes[range_start..range_start + (1 << layer_depth)])
333 }
334
335 pub fn branch(&self, index: usize, layer_depth: usize) -> Result<Vec<D>, Error> {
342 if index >= 1 << self.log_len {
343 return Err(Error::IndexOutOfRange {
344 index,
345 log_len: self.log_len,
346 });
347 }
348 if layer_depth > self.log_len {
349 return Err(Error::IncorrectLayerDepth {
350 layer_depth,
351 log_len: self.log_len,
352 });
353 }
354
355 let branch = (0..self.log_len - layer_depth)
356 .map(|j| {
357 let node_index = (((1 << j) - 1) << (self.log_len + 1 - j)) | (index >> j) ^ 1;
358 self.inner_nodes[node_index].clone()
359 })
360 .collect();
361
362 Ok(branch)
363 }
364}
365
366#[cfg(test)]
367mod tests {
368 use binius_field::Ghash128b as B128;
369 use binius_hash::{Sha256Compression, Sha256HashSuite};
370 use digest::Output;
371
372 use super::*;
373
374 fn commit(
376 n_values: usize,
377 batch_size: usize,
378 ) -> Result<BinaryMerkleTree<Output<sha2::Sha256>>, Error> {
379 let elements = (0..n_values)
380 .map(|i| B128::new(i as u128))
381 .collect::<Vec<_>>();
382 BinaryMerkleTree::new::<B128, Sha256HashSuite>(&elements, batch_size, &GlobalAllocator)
383 }
384
385 #[test]
386 fn test_new_rejects_a_ragged_final_leaf() {
387 let err = commit(7, 2).unwrap_err();
389
390 let Error::IncorrectBatchSize {
391 n_values,
392 batch_size,
393 } = err
394 else {
395 panic!("expected IncorrectBatchSize, got {err}");
396 };
397 assert_eq!(n_values, 7);
398 assert_eq!(batch_size, 2);
399 }
400
401 #[test]
402 fn test_new_rejects_an_empty_leaf() {
403 let err = commit(0, 0).unwrap_err();
405
406 let Error::IncorrectBatchSize {
407 n_values,
408 batch_size,
409 } = err
410 else {
411 panic!("expected IncorrectBatchSize, got {err}");
412 };
413 assert_eq!(n_values, 0);
414 assert_eq!(batch_size, 0);
415 }
416
417 #[test]
418 fn test_new_rejects_a_leaf_count_off_the_power_of_two_grid() {
419 let err = commit(6, 2).unwrap_err();
421
422 let Error::PowerOfTwoLengthRequired { n_leaves } = err else {
423 panic!("expected PowerOfTwoLengthRequired, got {err}");
424 };
425 assert_eq!(n_leaves, 3);
426 }
427
428 #[test]
429 fn test_layer_spans_the_root_down_to_the_leaves() {
430 let tree = commit(8, 2).unwrap();
432 assert_eq!(tree.log_len, 2);
433
434 assert_eq!(tree.layer(0).unwrap(), &[tree.root()]);
436 assert_eq!(tree.layer(1).unwrap().len(), 2);
437 assert_eq!(tree.layer(2).unwrap().len(), 4);
438
439 let err = tree.layer(3).unwrap_err();
441
442 let Error::IncorrectLayerDepth {
443 layer_depth,
444 log_len,
445 } = err
446 else {
447 panic!("expected IncorrectLayerDepth, got {err}");
448 };
449 assert_eq!(layer_depth, 3);
450 assert_eq!(log_len, 2);
451 }
452
453 #[test]
454 fn test_branch_opens_every_leaf_up_to_the_named_layer() {
455 let tree = commit(8, 2).unwrap();
456
457 for index in 0..4 {
459 assert_eq!(tree.branch(index, 0).unwrap().len(), 2);
460 assert_eq!(tree.branch(index, 1).unwrap().len(), 1);
461 assert_eq!(tree.branch(index, 2).unwrap().len(), 0);
462 }
463
464 assert_eq!(tree.branch(0, 0).unwrap()[1], tree.branch(1, 0).unwrap()[1]);
466 }
467
468 #[test]
469 fn test_branch_rejects_a_leaf_past_the_last_one() {
470 let tree = commit(8, 2).unwrap();
472 let err = tree.branch(4, 0).unwrap_err();
473
474 let Error::IndexOutOfRange { index, log_len } = err else {
475 panic!("expected IndexOutOfRange, got {err}");
476 };
477 assert_eq!(index, 4);
478 assert_eq!(log_len, 2);
479 }
480
481 #[test]
482 fn test_branch_blames_the_depth_when_the_leaf_is_valid() {
483 let tree = commit(8, 2).unwrap();
485 let err = tree.branch(0, 3).unwrap_err();
486
487 let Error::IncorrectLayerDepth {
488 layer_depth,
489 log_len,
490 } = err
491 else {
492 panic!("expected IncorrectLayerDepth, got {err}");
493 };
494 assert_eq!(layer_depth, 3);
495 assert_eq!(log_len, 2);
496 }
497
498 #[test]
499 fn test_new_matches_a_naive_layer_by_layer_fold_across_the_recursion_cutoff() {
500 for log_len in [0, 1, SUBTREE_LOG_DEPTH, 2 * SUBTREE_LOG_DEPTH] {
503 let tree = commit(1 << log_len, 1).unwrap();
504
505 let compression = Sha256Compression::default();
507 let mut layer = tree.layer(log_len).unwrap().to_vec();
508 let mut expected_layers = vec![layer.clone()];
509 while layer.len() > 1 {
510 layer = layer
511 .chunks_exact(2)
512 .map(|pair| compression.compress([pair[0], pair[1]]))
513 .collect();
514 expected_layers.push(layer.clone());
515 }
516
517 for depth in 0..=log_len {
518 assert_eq!(
519 tree.layer(depth).unwrap(),
520 expected_layers[log_len - depth].as_slice(),
521 "log_len {log_len}, depth {depth}"
522 );
523 }
524 }
525 }
526}