Skip to main content

binius_hash_prover/
parallel_digest.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{array, marker::PhantomData, mem::MaybeUninit};
5
6use binius_hash::HashBuffer;
7use binius_utils::{
8	FixedSizeSerializeBytes, SerializeBytes,
9	rayon::{
10		iter::{IndexedParallelIterator, IntoParallelRefMutIterator, ParallelIterator},
11		slice::ParallelSliceMut,
12	},
13};
14use bytes::BytesMut;
15use digest::{Digest, FixedOutputReset, Output, block_api::BlockSizeUser};
16
17/// An object that efficiently computes `N` instances of a cryptographic hash function
18/// in parallel.
19///
20/// This trait is useful when there is a more efficient way of calculating multiple digests at once,
21/// e.g. using SIMD instructions. It is supposed that this trait is implemented directly for some
22/// digest and some fixed `N` and passed as an implementation of the `ParallelDigest` trait which
23/// hides the `N` value.
24pub trait MultiDigest<const N: usize>: Clone {
25	/// The corresponding non-parallelized hash function.
26	type Digest: Digest;
27
28	/// Create new hasher instance with empty state.
29	fn new() -> Self;
30
31	/// Create new hasher instance which has processed the provided data.
32	fn new_with_prefix(data: impl AsRef<[u8]>) -> Self {
33		let mut hasher = Self::new();
34		hasher.update([data.as_ref(); N]);
35		hasher
36	}
37
38	/// Process data, updating the internal state.
39	/// The number of rows in `data` must be equal to `parallel_instances()`.
40	fn update(&mut self, data: [&[u8]; N]);
41
42	/// Process input data in a chained manner.
43	#[must_use]
44	fn chain_update(self, data: [&[u8]; N]) -> Self {
45		let mut hasher = self;
46		hasher.update(data);
47		hasher
48	}
49
50	/// Write result into provided array and consume the hasher instance.
51	fn finalize_into(self, out: &mut [MaybeUninit<Output<Self::Digest>>; N]);
52
53	/// Write result into provided array and reset the hasher instance.
54	fn finalize_into_reset(&mut self, out: &mut [MaybeUninit<Output<Self::Digest>>; N]);
55
56	/// Reset hasher instance to its initial state.
57	fn reset(&mut self);
58
59	/// Compute hash of `data`.
60	/// All slices in the `data` must have the same length.
61	///
62	/// # Panics
63	/// Panics if data contains slices of different lengths.
64	fn digest(data: [&[u8]; N], out: &mut [MaybeUninit<Output<Self::Digest>>; N]);
65}
66
67pub trait ParallelDigest: Send {
68	/// The corresponding non-parallelized hash function.
69	type Digest: Digest;
70
71	/// Create new hasher instance with empty state.
72	fn new() -> Self;
73
74	/// Calculate the digest of multiple hashes by processing a parallel iterator of iterators.
75	///
76	/// The source parameter provides a parallel iterator where:
77	/// - Each element of the outer iterator maps to one leaf/digest in the output
78	/// - Each element contains an inner iterator of items that will be serialized and concatenated
79	///   to form that leaf's content
80	///
81	/// # Panics
82	/// All items must be able to serialize with SerializationMode::Native without error, or this
83	/// method will panic.
84	fn digest<I: IntoIterator<Item: SerializeBytes>>(
85		&self,
86		source: impl IndexedParallelIterator<Item = I>,
87		out: &mut [MaybeUninit<Output<Self::Digest>>],
88	);
89
90	/// Like [`digest`](Self::digest), but specialized for the case where every leaf is built from
91	/// exactly `n_items_per_input` items of a [`FixedSizeSerializeBytes`] type, so that each leaf
92	/// has the same, compile-time-derivable byte length.
93	///
94	/// This extra structure lets implementations skip per-leaf length bookkeeping (and, for short
95	/// leaves, the message padding) that [`digest`](Self::digest) must redo every time. The default
96	/// implementation simply forwards to [`digest`](Self::digest).
97	///
98	/// # Panics
99	/// Each iterator in `source` must yield exactly `n_items_per_input` items, and all items must
100	/// serialize without error, or this method may panic.
101	fn digest_with_const_len<I: IntoIterator<Item: FixedSizeSerializeBytes>>(
102		&self,
103		n_items_per_input: usize,
104		source: impl IndexedParallelIterator<Item = I>,
105		out: &mut [MaybeUninit<Output<Self::Digest>>],
106	) {
107		let _ = n_items_per_input;
108		self.digest(source, out);
109	}
110}
111
112/// A wrapper that implements the `ParallelDigest` trait for a `MultiDigest` implementation.
113#[derive(Clone)]
114pub struct ParallelMultidigestImpl<D: MultiDigest<N>, const N: usize>(D);
115
116impl<D: MultiDigest<N> + Default, const N: usize> Default for ParallelMultidigestImpl<D, N> {
117	fn default() -> Self {
118		Self(D::default())
119	}
120}
121
122impl<D: MultiDigest<N, Digest: Send> + Send + Sync, const N: usize> ParallelDigest
123	for ParallelMultidigestImpl<D, N>
124{
125	type Digest = D::Digest;
126
127	fn new() -> Self {
128		Self(D::new())
129	}
130
131	fn digest<I: IntoIterator<Item: SerializeBytes>>(
132		&self,
133		source: impl IndexedParallelIterator<Item = I>,
134		out: &mut [MaybeUninit<Output<Self::Digest>>],
135	) {
136		let buffers = array::from_fn::<_, N, _>(|_| BytesMut::new());
137		source.chunks(N).zip(out.par_chunks_mut(N)).for_each_with(
138			buffers,
139			|buffers, (data, out_chunk)| {
140				let mut hasher = self.0.clone();
141				for (mut buf, chunk) in buffers.iter_mut().zip(data) {
142					buf.clear();
143					for item in chunk {
144						item.serialize(&mut buf)
145							.expect("pre-condition: items must serialize without error");
146					}
147				}
148				let data = array::from_fn(|i| buffers[i].as_ref());
149				hasher.update(data);
150
151				if out_chunk.len() == N {
152					hasher
153						.finalize_into_reset(out_chunk.try_into().expect("chunk size is correct"));
154				} else {
155					let mut result = array::from_fn::<_, N, _>(|_| MaybeUninit::uninit());
156					hasher.finalize_into(&mut result);
157					for (out, res) in out_chunk.iter_mut().zip(result) {
158						out.write(unsafe { res.assume_init() });
159					}
160				}
161			},
162		);
163	}
164}
165
166/// Adapts a sequential [`Digest`] into a [`ParallelDigest`] that hashes one leaf per element of a
167/// parallel iterator.
168///
169/// Each Rayon work-item is seeded with a single hasher (via `for_each_with`) which is recycled in
170/// place with `finalize_reset` between leaves, rather than cloning a fresh hasher per leaf. This
171/// requires `D: FixedOutputReset`.
172pub struct ParallelDigestAdapter<D>(PhantomData<D>);
173
174impl<D> Default for ParallelDigestAdapter<D> {
175	fn default() -> Self {
176		Self(PhantomData)
177	}
178}
179
180impl<D> ParallelDigest for ParallelDigestAdapter<D>
181where
182	D: Digest + FixedOutputReset + BlockSizeUser + Send + Sync + Clone,
183{
184	type Digest = D;
185
186	fn new() -> Self {
187		Self(PhantomData)
188	}
189
190	fn digest<I: IntoIterator<Item: SerializeBytes>>(
191		&self,
192		source: impl IndexedParallelIterator<Item = I>,
193		out: &mut [MaybeUninit<Output<Self::Digest>>],
194	) {
195		source
196			.zip(out.par_iter_mut())
197			.for_each_with(D::new(), |hasher, (items, out)| {
198				{
199					let mut buffer = HashBuffer::new(hasher);
200					for item in items {
201						item.serialize(&mut buffer)
202							.expect("pre-condition: items must serialize without error");
203					}
204				}
205				out.write(hasher.finalize_reset());
206			});
207	}
208}
209
210#[cfg(test)]
211mod tests {
212	use std::iter::repeat_with;
213
214	use binius_utils::rayon::iter::IntoParallelRefIterator;
215	use digest::{
216		FixedOutput, HashMarker, OutputSizeUser, Reset, Update,
217		consts::{U1, U32},
218	};
219	use itertools::izip;
220	use rand::prelude::*;
221
222	use super::*;
223
224	#[derive(Clone, Default)]
225	struct MockDigest {
226		state: u8,
227	}
228
229	impl HashMarker for MockDigest {}
230
231	impl Update for MockDigest {
232		fn update(&mut self, data: &[u8]) {
233			for &byte in data {
234				self.state ^= byte;
235			}
236		}
237	}
238
239	impl Reset for MockDigest {
240		fn reset(&mut self) {
241			self.state = 0;
242		}
243	}
244
245	impl OutputSizeUser for MockDigest {
246		type OutputSize = U32;
247	}
248
249	impl BlockSizeUser for MockDigest {
250		type BlockSize = U1;
251	}
252
253	impl FixedOutput for MockDigest {
254		fn finalize_into(self, out: &mut Output<Self>) {
255			out[0] = self.state;
256			for byte in &mut out[1..] {
257				*byte = 0;
258			}
259		}
260	}
261
262	#[derive(Clone, Default)]
263	struct MockMultiDigest {
264		digests: [MockDigest; 4],
265	}
266
267	impl MultiDigest<4> for MockMultiDigest {
268		type Digest = MockDigest;
269
270		fn new() -> Self {
271			Self::default()
272		}
273
274		fn update(&mut self, data: [&[u8]; 4]) {
275			for (digest, &chunk) in self.digests.iter_mut().zip(data.iter()) {
276				digest::Digest::update(digest, chunk);
277			}
278		}
279
280		fn finalize_into(self, out: &mut [MaybeUninit<Output<Self::Digest>>; 4]) {
281			for (digest, out) in self.digests.into_iter().zip(out.iter_mut()) {
282				let mut output = digest::Output::<Self::Digest>::default();
283				digest::Digest::finalize_into(digest, &mut output);
284				*out = MaybeUninit::new(output);
285			}
286		}
287
288		fn finalize_into_reset(&mut self, out: &mut [MaybeUninit<Output<Self::Digest>>; 4]) {
289			for (digest, out) in self.digests.iter_mut().zip(out.iter_mut()) {
290				let mut digest_copy = MockDigest::default();
291				std::mem::swap(digest, &mut digest_copy);
292				*out = MaybeUninit::new(digest_copy.finalize());
293			}
294			self.reset();
295		}
296
297		fn reset(&mut self) {
298			for digest in &mut self.digests {
299				*digest = MockDigest::default();
300			}
301		}
302
303		fn digest(data: [&[u8]; 4], out: &mut [MaybeUninit<Output<Self::Digest>>; 4]) {
304			let mut hasher = Self::default();
305			hasher.update(data);
306			hasher.finalize_into(out);
307		}
308	}
309
310	fn generate_mock_data(n_hashes: usize, chunk_size: usize) -> Vec<Vec<u8>> {
311		let mut rng = StdRng::seed_from_u64(0);
312
313		(0..n_hashes)
314			.map(|_| {
315				let mut chunk = vec![0; chunk_size];
316				rng.fill_bytes(&mut chunk);
317				chunk
318			})
319			.collect()
320	}
321
322	fn check_parallel_digest_consistency<
323		D: ParallelDigest<Digest: BlockSizeUser + Send + Sync + Clone>,
324	>(
325		data: &[Vec<u8>],
326	) {
327		let parallel_digest = D::new();
328		let mut parallel_results = repeat_with(MaybeUninit::<Output<D::Digest>>::uninit)
329			.take(data.len())
330			.collect::<Vec<_>>();
331		parallel_digest.digest(data.par_iter(), &mut parallel_results);
332
333		let serial_results = data.iter().map(<D::Digest as Digest>::digest);
334
335		for (parallel, serial) in izip!(parallel_results, serial_results) {
336			assert_eq!(unsafe { parallel.assume_init() }, serial);
337		}
338	}
339
340	#[test]
341	fn test_empty_data() {
342		let data = generate_mock_data(0, 16);
343		check_parallel_digest_consistency::<ParallelMultidigestImpl<MockMultiDigest, 4>>(&data);
344	}
345
346	#[test]
347	fn test_non_empty_data() {
348		for n_hashes in [1, 2, 4, 8, 9] {
349			let data = generate_mock_data(n_hashes, 16);
350			check_parallel_digest_consistency::<ParallelMultidigestImpl<MockMultiDigest, 4>>(&data);
351		}
352	}
353
354	#[test]
355	fn test_adapter_matches_serial_sha256() {
356		use sha2::Sha256;
357
358		for n_hashes in [0, 1, 2, 4, 8, 9, 100] {
359			let data = generate_mock_data(n_hashes, 16);
360
361			let adapter = ParallelDigestAdapter::<Sha256>::new();
362			let mut results = repeat_with(MaybeUninit::<Output<Sha256>>::uninit)
363				.take(data.len())
364				.collect::<Vec<_>>();
365			adapter.digest(data.par_iter(), &mut results);
366
367			for (result, leaf) in results.into_iter().zip(&data) {
368				assert_eq!(unsafe { result.assume_init() }, <Sha256 as Digest>::digest(leaf));
369			}
370		}
371	}
372}