Skip to main content

binius_hash_prover/
binary_merkle_tree.rs

1// Copyright 2024-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use 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/// Reason a Merkle tree operation rejects its arguments.
24///
25/// Every variant carries the numbers that failed the check alongside the bound they broke.
26/// A log line therefore says which call was malformed without the reader rerunning it.
27#[derive(Debug, thiserror::Error)]
28pub enum Error {
29	/// A committed batch of values does not split into whole leaves.
30	#[error("{n_values} values do not split into leaves of {batch_size}")]
31	IncorrectBatchSize {
32		/// The number of values handed over for commitment.
33		n_values: usize,
34		/// The number of values every leaf must hold.
35		batch_size: usize,
36	},
37	/// The leaves a commitment splits into do not span a binary tree.
38	#[error("the leaf count {n_leaves} is not a power of two")]
39	PowerOfTwoLengthRequired {
40		/// The number of leaves the values split into.
41		n_leaves: usize,
42	},
43	/// A layer is asked for below the leaves of the tree.
44	#[error("layer depth {layer_depth} exceeds the tree depth {log_len}")]
45	IncorrectLayerDepth {
46		/// The depth asked for, counted from the root.
47		layer_depth: usize,
48		/// The depth of the tree, which is its deepest valid layer.
49		log_len: usize,
50	},
51	/// A branch is asked for at a leaf the tree does not have.
52	#[error("leaf index {index} is outside the 2^{log_len} leaves")]
53	IndexOutOfRange {
54		/// The leaf position asked for.
55		index: usize,
56		/// Base-2 logarithm of the number of leaves the tree has.
57		log_len: usize,
58	},
59}
60
61/// A binary Merkle tree that commits batches of vectors.
62///
63/// # Overview
64///
65/// The entries sharing an index across a batch are hashed together into one leaf digest.
66/// A binary tree is then folded over those leaf digests.
67///
68/// All committed vectors must have the same length, and that length must be a power of two.
69///
70/// The nodes are drawn from an [`Allocator`], so a prover can back the whole tree with pooled
71/// memory instead of the global heap. The default is [`GlobalAllocator`], which is a plain [`Vec`].
72pub struct BinaryMerkleTree<D: Send, A: Allocator = GlobalAllocator> {
73	/// Base-2 logarithm of the number of leaves.
74	pub log_len: usize,
75	/// The inner nodes, arranged as a flattened array of layers with the root at the end.
76	pub inner_nodes: A::Vec<D>,
77}
78
79/// Written through the node slice rather than the buffer, so no allocator has to be [`Debug`]
80/// itself for a tree over it to be.
81impl<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	/// Commits a slice of values, cutting it into consecutive leaves.
92	///
93	/// # Arguments
94	///
95	/// * `elements` - the values to commit, in leaf order.
96	/// * `batch_size` - how many consecutive values are hashed together into one leaf.
97	/// * `alloc` - the allocator the tree's nodes are drawn from.
98	///
99	/// # Errors
100	///
101	/// * The value count is not a multiple of the batch size.
102	/// * The resulting leaf count is not a power of two.
103	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		// Every leaf holds the same number of values, so the split has to come out even.
109		// An empty leaf divides no value count, so it is rejected here rather than dividing by it.
110		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		// A binary tree only spans a power-of-two number of leaves.
118		let len = elements.len() / batch_size;
119		if !len.is_power_of_two() {
120			return Err(Error::PowerOfTwoLengthRequired { n_leaves: len });
121		}
122
123		// Hand the leaves over one contiguous chunk at a time.
124		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	/// Commits leaves drawn from a parallel iterator, one iterator item per leaf.
134	///
135	/// # Overview
136	///
137	/// The tree is laid out as one flat buffer of layers, widest first:
138	///
139	/// ```text
140	///     [ leaf digests | layer 1 | ... | root ]
141	///        2^log_len      2^(log_len-1)    1
142	/// ```
143	///
144	/// Each layer is written into the buffer's spare capacity.
145	/// It is then read back as the input to the layer above it.
146	///
147	/// # Arguments
148	///
149	/// * `leaves` - one iterator per leaf, each yielding that leaf's values.
150	/// * `n_items_per_input` - how many values every leaf iterator yields.
151	/// * `alloc` - the allocator the tree's nodes are drawn from.
152	///
153	/// # Panics
154	///
155	/// Panics unless the number of leaves is a power of two.
156	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		// Panics unless the leaf count is a power of two, which the binary layout requires.
163		let log_len = checked_log_2(leaves.len());
164
165		// A binary tree over 2^log_len leaves has 2^(log_len+1) - 1 nodes in total.
166		// The whole tree is one allocation, which the layer loop then fills front to back.
167		let total_length = (1 << (log_len + 1)) - 1;
168		let mut inner_nodes = alloc.alloc::<Array<u8, N>>(total_length);
169
170		// Fill the widest layer first, straight into uninitialized capacity.
171		{
172			let _span = tracing::debug_span!("hash_leaves").entered();
173
174			// Every leaf holds exactly the same number of values, so its byte length is a constant.
175			// Handing that length to the hasher lets it specialize for short leaves.
176			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			// SAFETY: the leaf digests were just written over the whole widest layer
187			leaf_layer.assume_init_mut()
188		};
189
190		// One synchronization point: the leaves are already written, so build every level
191		// above them as a single recursive tree of rayon tasks.
192		// A node is written the instant both its children are ready, with no other barrier.
193		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: &parallel_compression,
200			};
201			ctx.build_node(0, log_len);
202		}
203
204		unsafe {
205			// SAFETY: inner_nodes should be entirely initialized by now
206			// Note that we don't incrementally update inner_nodes.len() since
207			// that doesn't play well with using split_at_mut on spare capacity.
208			inner_nodes.set_len(total_length);
209		}
210		Self {
211			log_len,
212			inner_nodes,
213		}
214	}
215}
216
217/// Leaf-span depth below which a subtree folds without spawning further rayon tasks.
218///
219/// Below this depth, a subtree's whole span already fills one SIMD-batched compression
220/// call per level, so splitting it further would only add task overhead.
221///
222/// 6 was picked from per-layer costs measured on a 2^17-leaf SHA-256 tree.
223const SUBTREE_LOG_DEPTH: usize = 6;
224
225/// A raw pointer that can cross into another rayon task.
226///
227/// Every write through it targets a disjoint element, checked by hand at each write site.
228#[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
234/// State shared by every recursive call building the tree above the leaves.
235struct BuildContext<'a, D, C> {
236	/// The leaf digests, already written.
237	leaves: &'a [D],
238	/// Base-2 logarithm of the leaf count.
239	log_len: usize,
240	/// The first node above the leaf layer, addressed by [`node_offset`].
241	nodes: SendPtr<MaybeUninit<D>>,
242	/// Two-to-one compression, used both one pair at a time and in SIMD-batched groups.
243	compression: &'a C,
244}
245
246/// Flat-buffer offset of the node spanning `1 << log_width` leaves starting at `leaves_start`.
247///
248/// The offset is relative to the first node above the leaf layer, matching `nodes` in
249/// [`BuildContext`].
250const fn node_offset(log_len: usize, leaves_start: usize, log_width: usize) -> usize {
251	// Depth counted from the root, matching `BinaryMerkleTree::layer`.
252	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	/// Folds `1 << log_width` leaves down to their single root without spawning rayon tasks.
259	///
260	/// Every level still runs through the hash suite's SIMD-batched compression, since a small
261	/// subtree's whole level is one batch, not one task.
262	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				// SAFETY: this subtree owns a leaf range disjoint from every other subtree's, so
269				// its node offset at every depth is disjoint from every other subtree's.
270				slice::from_raw_parts_mut(self.nodes.0.add(offset), width)
271			};
272			self.compression.parallel_compress(cur, out);
273			cur = unsafe {
274				// SAFETY: the compression above wrote every slot of this layer
275				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	/// Builds every node from `1 << log_width` leaves starting at `leaves_start` up to their root.
284	///
285	/// Levels past [`SUBTREE_LOG_DEPTH`] recurse as two independent halves, run in parallel by
286	/// rayon, with no synchronization beyond a node waiting on its own two children.
287	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			// SAFETY: this subtree owns a leaf range disjoint from every other subtree's, so its
302			// node offset is disjoint from every other subtree's.
303			(*self.nodes.0.add(offset)).write(node.clone());
304		}
305		node
306	}
307}
308
309impl<D: Clone + Send, A: Allocator> BinaryMerkleTree<D, A> {
310	/// Clones the root digest, which sits last in the flattened layers.
311	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	/// Borrows one whole layer of the tree, counting depth from the root.
319	///
320	/// # Errors
321	///
322	/// * The depth lies below the leaves.
323	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	/// Collects the sibling digests opening a leaf up to the layer at `layer_depth`.
336	///
337	/// # Errors
338	///
339	/// * The leaf index lies outside the tree.
340	/// * The depth lies below the leaves.
341	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	/// Commits `n_values` distinct field elements in leaves of `batch_size` values each.
375	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		// 7 values leave a partial leaf behind once 3 whole leaves of 2 are cut.
388		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		// A leaf holding no value divides no value count, so the check runs before the division.
404		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		// 6 values in leaves of 2 split evenly, but 3 leaves do not span a binary tree.
420		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		// 8 values in leaves of 2 give 4 leaves, so the tree is 2 layers deep.
431		let tree = commit(8, 2).unwrap();
432		assert_eq!(tree.log_len, 2);
433
434		// Layer d holds 2^d digests, from the single root down to the 4 leaves.
435		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		// One layer past the leaves is the first depth the tree does not hold.
440		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		// Opening a leaf to the root names one sibling per layer crossed.
458		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		// Sibling leaves cross the same inner node, so their top sibling agrees.
465		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		// 4 leaves are indexed 0..=3, so 4 is the first index outside the tree.
471		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		// The index is in range, so the depth is the argument the diagnostic must name.
484		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		// Sizes chosen to land on both sides of the small-subtree cutoff:
501		// a single leaf, two leaves, exactly at the cutoff, and twice past it.
502		for log_len in [0, 1, SUBTREE_LOG_DEPTH, 2 * SUBTREE_LOG_DEPTH] {
503			let tree = commit(1 << log_len, 1).unwrap();
504
505			// Reference: fold the same leaves one pair at a time, one full layer at a time.
506			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}