1use 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
17pub trait MultiDigest<const N: usize>: Clone {
25 type Digest: Digest;
27
28 fn new() -> Self;
30
31 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 fn update(&mut self, data: [&[u8]; N]);
41
42 #[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 fn finalize_into(self, out: &mut [MaybeUninit<Output<Self::Digest>>; N]);
52
53 fn finalize_into_reset(&mut self, out: &mut [MaybeUninit<Output<Self::Digest>>; N]);
55
56 fn reset(&mut self);
58
59 fn digest(data: [&[u8]; N], out: &mut [MaybeUninit<Output<Self::Digest>>; N]);
65}
66
67pub trait ParallelDigest: Send {
68 type Digest: Digest;
70
71 fn new() -> Self;
73
74 fn digest<I: IntoIterator<Item: SerializeBytes>>(
85 &self,
86 source: impl IndexedParallelIterator<Item = I>,
87 out: &mut [MaybeUninit<Output<Self::Digest>>],
88 );
89
90 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#[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
166pub 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}