1use std::{array, mem::MaybeUninit};
25
26use binius_utils::{
27 FixedSizeSerializeBytes, SerializeBytes,
28 rayon::{
29 iter::{IndexedParallelIterator, ParallelIterator},
30 slice::{ParallelSlice, ParallelSliceMut},
31 },
32};
33use blake3::{BLOCK_LEN, CHUNK_LEN, OUT_LEN};
34use digest::Output;
35
36use super::{
37 blake3::Blake3Compression,
38 parallel_compression::ParallelPseudoCompression,
39 parallel_digest::{
40 MultiDigest, ParallelDigest, ParallelDigestAdapter, ParallelMultidigestImpl,
41 },
42};
43
44const CHUNK_START: u32 = 1 << 0;
46
47const CHUNK_END: u32 = 1 << 1;
49
50const ROOT: u32 = 1 << 3;
52
53const IV: [u32; 8] = [
55 0x6a09e667, 0xbb67ae85, 0x3c6ef372, 0xa54ff53a, 0x510e527f, 0x9b05688c, 0x1f83d9ab, 0x5be0cd19,
56];
57
58const MSG_PERMUTATION: [usize; 16] = [2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8];
62
63const N_ROUNDS: usize = 7;
65
66#[inline(always)]
71fn quarter_round<const N: usize>(
72 v: &mut [[u32; N]; 16],
73 a: usize,
74 b: usize,
75 c: usize,
76 d: usize,
77 mx: &[u32; N],
78 my: &[u32; N],
79) {
80 for i in 0..N {
82 v[a][i] = v[a][i].wrapping_add(v[b][i]).wrapping_add(mx[i]);
83 v[d][i] = (v[d][i] ^ v[a][i]).rotate_right(16);
84 v[c][i] = v[c][i].wrapping_add(v[d][i]);
85 v[b][i] = (v[b][i] ^ v[c][i]).rotate_right(12);
86 v[a][i] = v[a][i].wrapping_add(v[b][i]).wrapping_add(my[i]);
87 v[d][i] = (v[d][i] ^ v[a][i]).rotate_right(8);
88 v[c][i] = v[c][i].wrapping_add(v[d][i]);
89 v[b][i] = (v[b][i] ^ v[c][i]).rotate_right(7);
90 }
91}
92
93#[inline(always)]
97fn round<const N: usize>(v: &mut [[u32; N]; 16], m: &[[u32; N]; 16]) {
98 quarter_round(v, 0, 4, 8, 12, &m[0], &m[1]);
100 quarter_round(v, 1, 5, 9, 13, &m[2], &m[3]);
101 quarter_round(v, 2, 6, 10, 14, &m[4], &m[5]);
102 quarter_round(v, 3, 7, 11, 15, &m[6], &m[7]);
103 quarter_round(v, 0, 5, 10, 15, &m[8], &m[9]);
105 quarter_round(v, 1, 6, 11, 12, &m[10], &m[11]);
106 quarter_round(v, 2, 7, 8, 13, &m[12], &m[13]);
107 quarter_round(v, 3, 4, 9, 14, &m[14], &m[15]);
108}
109
110#[inline(always)]
112fn permute<const N: usize>(m: &mut [[u32; N]; 16]) {
113 let permuted: [[u32; N]; 16] = array::from_fn(|i| m[MSG_PERMUTATION[i]]);
115 *m = permuted;
116}
117
118#[inline(always)]
120fn load_block_words<const N: usize>(block: &[[u8; BLOCK_LEN]; N]) -> [[u32; N]; 16] {
121 let mut m = [[0u32; N]; 16];
122 for lane in 0..N {
123 for (w, slot) in m.iter_mut().enumerate() {
124 let off = w * 4;
125 slot[lane] = u32::from_le_bytes([
126 block[lane][off],
127 block[lane][off + 1],
128 block[lane][off + 2],
129 block[lane][off + 3],
130 ]);
131 }
132 }
133 m
134}
135
136#[inline(always)]
141fn compress_block<const N: usize>(
142 cv: &mut [[u32; N]; 8],
143 block: &[[u32; N]; 16],
144 counter: u64,
145 block_len: u32,
146 flags: u32,
147) {
148 let counter_lo = counter as u32;
150 let counter_hi = (counter >> 32) as u32;
151
152 let mut v: [[u32; N]; 16] = [
154 cv[0],
155 cv[1],
156 cv[2],
157 cv[3],
158 cv[4],
159 cv[5],
160 cv[6],
161 cv[7],
162 [IV[0]; N],
163 [IV[1]; N],
164 [IV[2]; N],
165 [IV[3]; N],
166 [counter_lo; N],
167 [counter_hi; N],
168 [block_len; N],
169 [flags; N],
170 ];
171
172 let mut m = *block;
174 for r in 0..N_ROUNDS {
175 round(&mut v, &m);
176 if r < N_ROUNDS - 1 {
177 permute(&mut m);
178 }
179 }
180
181 for i in 0..8 {
183 for lane in 0..N {
184 cv[i][lane] = v[i][lane] ^ v[i + 8][lane];
185 }
186 }
187}
188
189#[inline(always)]
191fn broadcast_iv<const N: usize>() -> [[u32; N]; 8] {
192 array::from_fn(|w| [IV[w]; N])
193}
194
195#[inline(always)]
197fn serialize_cv_lane<const N: usize>(cv: &[[u32; N]; 8], lane: usize) -> [u8; OUT_LEN] {
198 let mut digest = [0u8; OUT_LEN];
199 for (w, chunk) in digest.chunks_exact_mut(4).enumerate() {
200 chunk.copy_from_slice(&cv[w][lane].to_le_bytes());
201 }
202 digest
203}
204
205#[derive(Clone)]
213pub struct PortableBlake3MultiDigest<const N: usize> {
214 cv: [[u32; N]; 8],
216 block: [[u8; BLOCK_LEN]; N],
218 block_len: usize,
220 blocks_compressed: usize,
222}
223
224impl<const N: usize> Default for PortableBlake3MultiDigest<N> {
225 fn default() -> Self {
226 Self {
228 cv: broadcast_iv(),
229 block: [[0u8; BLOCK_LEN]; N],
230 block_len: 0,
231 blocks_compressed: 0,
232 }
233 }
234}
235
236impl<const N: usize> PortableBlake3MultiDigest<N> {
237 fn compress_full_block(&mut self) {
239 let flags = if self.blocks_compressed == 0 {
241 CHUNK_START
242 } else {
243 0
244 };
245 let m = load_block_words(&self.block);
246 compress_block(&mut self.cv, &m, 0, BLOCK_LEN as u32, flags);
247 self.blocks_compressed += 1;
248 self.block_len = 0;
249 }
250
251 fn write_root(&self, out: &mut [MaybeUninit<Output<blake3::Hasher>>; N]) {
255 let mut cv = self.cv;
256 let mut block = self.block;
257 for lane in 0..N {
259 block[lane][self.block_len..].fill(0);
260 }
261 let start = if self.blocks_compressed == 0 {
263 CHUNK_START
264 } else {
265 0
266 };
267 let m = load_block_words(&block);
268 compress_block(&mut cv, &m, 0, self.block_len as u32, start | CHUNK_END | ROOT);
269
270 for lane in 0..N {
272 out[lane].write(serialize_cv_lane(&cv, lane).into());
273 }
274 }
275}
276
277impl<const N: usize> MultiDigest<N> for PortableBlake3MultiDigest<N> {
278 type Digest = blake3::Hasher;
279
280 fn new() -> Self {
281 Self::default()
282 }
283
284 fn update(&mut self, data: [&[u8]; N]) {
285 let mut consumed = [0usize; N];
287 loop {
288 let remaining = (0..N)
290 .map(|i| data[i].len() - consumed[i])
291 .max()
292 .unwrap_or(0);
293 if remaining == 0 {
294 break;
295 }
296 if self.block_len == BLOCK_LEN {
298 self.compress_full_block();
299 }
300 let take = (BLOCK_LEN - self.block_len).min(remaining);
302 for lane in 0..N {
303 let avail = data[lane].len() - consumed[lane];
304 let n = take.min(avail);
305 self.block[lane][self.block_len..self.block_len + n]
306 .copy_from_slice(&data[lane][consumed[lane]..consumed[lane] + n]);
307 consumed[lane] += n;
308 }
309 self.block_len += take;
310 }
311 }
312
313 fn finalize_into(self, out: &mut [MaybeUninit<Output<Self::Digest>>; N]) {
314 self.write_root(out);
315 }
316
317 fn finalize_into_reset(&mut self, out: &mut [MaybeUninit<Output<Self::Digest>>; N]) {
318 self.write_root(out);
319 self.reset();
320 }
321
322 fn reset(&mut self) {
323 self.cv = broadcast_iv();
326 self.block_len = 0;
327 self.blocks_compressed = 0;
328 }
329
330 fn digest(data: [&[u8]; N], out: &mut [MaybeUninit<Output<Self::Digest>>; N]) {
331 let mut hasher = Self::new();
332 hasher.update(data);
333 hasher.finalize_into(out);
334 }
335}
336
337#[derive(Debug, Clone, Default)]
344pub struct PortableBlake3ParallelDigest<const LANES: usize>;
345
346impl<const LANES: usize> ParallelDigest for PortableBlake3ParallelDigest<LANES> {
347 type Digest = blake3::Hasher;
348
349 fn new() -> Self {
350 Self
351 }
352
353 fn digest<I: IntoIterator<Item: SerializeBytes>>(
354 &self,
355 source: impl IndexedParallelIterator<Item = I>,
356 out: &mut [MaybeUninit<Output<Self::Digest>>],
357 ) {
358 ParallelDigestAdapter::<blake3::Hasher>::new().digest(source, out);
361 }
362
363 fn digest_with_const_len<I: IntoIterator<Item: FixedSizeSerializeBytes>>(
364 &self,
365 n_items_per_input: usize,
366 source: impl IndexedParallelIterator<Item = I>,
367 out: &mut [MaybeUninit<Output<Self::Digest>>],
368 ) {
369 let leaf_len = n_items_per_input * I::Item::BYTE_SIZE;
371
372 if leaf_len <= CHUNK_LEN {
373 ParallelMultidigestImpl::<PortableBlake3MultiDigest<LANES>, LANES>::new()
375 .digest(source, out);
376 } else {
377 ParallelDigestAdapter::<blake3::Hasher>::new().digest(source, out);
379 }
380 }
381}
382
383#[inline]
397fn compress_node_pairs<const N: usize>(
398 inputs: &[Output<blake3::Hasher>],
399 out: &mut [MaybeUninit<Output<blake3::Hasher>>],
400) {
401 let mut blocks = [[0u8; BLOCK_LEN]; N];
403 for (lane, block) in blocks.iter_mut().enumerate().take(out.len()) {
404 block[..OUT_LEN].copy_from_slice(inputs[2 * lane].as_slice());
405 block[OUT_LEN..].copy_from_slice(inputs[2 * lane + 1].as_slice());
406 }
407
408 let m = load_block_words(&blocks);
410 let mut cv = broadcast_iv::<N>();
411 compress_block(&mut cv, &m, 0, BLOCK_LEN as u32, CHUNK_START | CHUNK_END | ROOT);
412
413 for (lane, slot) in out.iter_mut().enumerate() {
414 slot.write(serialize_cv_lane(&cv, lane).into());
415 }
416}
417
418#[derive(Debug, Clone, Default)]
427pub struct PortableBlake3ParallelCompression<const LANES: usize> {
428 compression: Blake3Compression,
430}
431
432impl<const LANES: usize> ParallelPseudoCompression<Output<blake3::Hasher>, 2>
433 for PortableBlake3ParallelCompression<LANES>
434{
435 type Compression = Blake3Compression;
436
437 fn compression(&self) -> &Self::Compression {
438 &self.compression
439 }
440
441 fn parallel_compress(
442 &self,
443 inputs: &[Output<blake3::Hasher>],
444 out: &mut [MaybeUninit<Output<blake3::Hasher>>],
445 ) {
446 assert_eq!(inputs.len(), 2 * out.len(), "Input length must be 2 * output length");
447
448 inputs
451 .par_chunks(2 * LANES)
452 .zip(out.par_chunks_mut(LANES))
453 .for_each(|(in_batch, out_batch)| compress_node_pairs::<LANES>(in_batch, out_batch));
454 }
455}
456
457#[cfg(test)]
458mod tests {
459 use std::iter::repeat_with;
460
461 use binius_utils::rayon::iter::{IntoParallelRefIterator, ParallelIterator};
462 use proptest::prelude::*;
463 use rand::{Rng, SeedableRng, rngs::StdRng};
464
465 use super::{super::compress::CompressionFunction, *};
466
467 fn check_parallel_compression<const N: usize>(pairs: &[[[u8; OUT_LEN]; 2]]) {
470 let inputs: Vec<Output<blake3::Hasher>> = pairs
472 .iter()
473 .flat_map(|[l, r]| [(*l).into(), (*r).into()])
474 .collect();
475 let mut out = repeat_with(MaybeUninit::<Output<blake3::Hasher>>::uninit)
476 .take(pairs.len())
477 .collect::<Vec<_>>();
478
479 PortableBlake3ParallelCompression::<N>::default().parallel_compress(&inputs, &mut out);
480
481 for (slot, [l, r]) in out.into_iter().zip(pairs) {
483 let expected = Blake3Compression.compress([(*l).into(), (*r).into()]);
484 assert_eq!(unsafe { slot.assume_init() }.as_slice(), expected.as_slice());
485 }
486 }
487
488 #[test]
489 fn test_parallel_compression_boundaries() {
490 let zero = [0u8; OUT_LEN];
492 let ones = [0xffu8; OUT_LEN];
493 check_parallel_compression::<16>(&[[zero, zero], [ones, ones], [zero, ones], [ones, zero]]);
494
495 let mut rng = StdRng::seed_from_u64(7);
497 for count in [0usize, 1, 15, 16, 17, 33] {
498 let pairs: Vec<[[u8; OUT_LEN]; 2]> = (0..count)
499 .map(|_| {
500 let mut pair = [[0u8; OUT_LEN]; 2];
501 rng.fill_bytes(&mut pair[0]);
502 rng.fill_bytes(&mut pair[1]);
503 pair
504 })
505 .collect();
506 check_parallel_compression::<16>(&pairs);
507 }
508 }
509
510 proptest! {
511 #[test]
512 fn parallel_compression_matches_scalar(
513 pairs in prop::collection::vec(
514 (prop::array::uniform32(any::<u8>()), prop::array::uniform32(any::<u8>())),
515 0..40usize,
516 ),
517 ) {
518 let pairs: Vec<[[u8; OUT_LEN]; 2]> = pairs.into_iter().map(|(l, r)| [l, r]).collect();
520 check_parallel_compression::<4>(&pairs);
521 check_parallel_compression::<8>(&pairs);
522 check_parallel_compression::<16>(&pairs);
523 }
524 }
525
526 fn check_portable_batch<const N: usize>(rng: &mut StdRng, len: usize) {
529 let messages: [Vec<u8>; N] = array::from_fn(|_| {
531 let mut m = vec![0u8; len];
532 rng.fill_bytes(&mut m);
533 m
534 });
535 let refs: [&[u8]; N] = array::from_fn(|i| messages[i].as_slice());
536 let mut out = array::from_fn::<_, N, _>(|_| MaybeUninit::uninit());
537 PortableBlake3MultiDigest::<N>::digest(refs, &mut out);
538
539 for (o, message) in out.iter().zip(messages.iter()) {
541 let got = unsafe { o.assume_init_ref() };
542 assert_eq!(got.as_slice(), blake3::hash(message).as_bytes(), "len = {len}, N = {N}");
543 }
544 }
545
546 #[test]
547 fn test_portable_lengths_match_reference() {
548 let mut rng = StdRng::seed_from_u64(0);
549
550 for len in [0, 1, 31, 63, 64, 65, 100, 127, 128, 1000, 1024] {
558 check_portable_batch::<4>(&mut rng, len);
559 check_portable_batch::<8>(&mut rng, len);
560 check_portable_batch::<16>(&mut rng, len);
561 }
562 }
563
564 #[test]
565 fn test_portable_chained_update() {
566 let mut rng = StdRng::seed_from_u64(2);
567 let messages: [Vec<u8>; 4] = array::from_fn(|_| {
569 let mut m = vec![0u8; 200];
570 rng.fill_bytes(&mut m);
571 m
572 });
573
574 let mut hasher = PortableBlake3MultiDigest::<4>::new();
577 hasher.update(array::from_fn(|i| &messages[i][..50]));
578 hasher.update(array::from_fn(|i| &messages[i][50..]));
579 let mut out = array::from_fn::<_, 4, _>(|_| MaybeUninit::uninit());
580 hasher.finalize_into(&mut out);
581
582 for (o, message) in out.iter().zip(messages.iter()) {
583 assert_eq!(unsafe { o.assume_init_ref() }.as_slice(), blake3::hash(message).as_bytes());
584 }
585 }
586
587 #[test]
588 fn test_portable_routing_matches_reference() {
589 let mut rng = StdRng::seed_from_u64(3);
590 let mut check = |leaf_len: usize| {
592 let leaves: Vec<Vec<u8>> = (0..50)
593 .map(|_| {
594 let mut m = vec![0u8; leaf_len];
595 rng.fill_bytes(&mut m);
596 m
597 })
598 .collect();
599 let digest = PortableBlake3ParallelDigest::<8>::new();
600 let mut results = repeat_with(MaybeUninit::<Output<blake3::Hasher>>::uninit)
601 .take(50)
602 .collect::<Vec<_>>();
603 digest.digest_with_const_len(
604 leaf_len,
605 leaves.par_iter().map(|leaf| leaf.iter().copied()),
606 &mut results,
607 );
608 for (result, leaf) in results.into_iter().zip(&leaves) {
609 let got = unsafe { result.assume_init() };
610 assert_eq!(got.as_slice(), blake3::hash(leaf).as_bytes(), "leaf_len {leaf_len}");
611 }
612 };
613
614 for leaf_len in [0, 1, 63, 64, 65, 100, 1000, 1024, 1025, 2048] {
620 check(leaf_len);
621 }
622 }
623}