1use std::mem::MaybeUninit;
26
27use binius_hash::{Sha256Compression, Sha256HashSuite};
28use binius_utils::{
29 FixedSizeSerializeBytes, SerializeBytes,
30 rayon::{
31 iter::{IndexedParallelIterator, ParallelIterator},
32 slice::{ParallelSlice, ParallelSliceMut},
33 task_size::{WorkPerItem, min_len_for_work},
34 },
35};
36use bytemuck::must_cast;
37use portable::LANES;
38use sha2::{Sha256, digest::Output};
39
40use crate::{
41 parallel_compression::ParallelPseudoCompression,
42 parallel_digest::{ParallelDigest, ParallelDigestAdapter},
43 suite::ParallelHashSuite,
44};
45
46pub mod portable;
47
48#[cfg(all(
50 target_arch = "x86_64",
51 target_feature = "sha",
52 target_feature = "sse2",
53 target_feature = "ssse3",
54 target_feature = "sse4.1"
55))]
56mod sha_ni;
57
58#[cfg(all(
63 target_arch = "x86_64",
64 target_feature = "avx512f",
65 target_feature = "avx512bw"
66))]
67mod avx512;
68
69#[cfg(all(target_arch = "aarch64", target_feature = "sha2"))]
71mod neon;
72
73const BLOCK_LEN: usize = 64;
75
76const DIGEST_LEN: usize = 32;
78
79const IV: [u32; 8] = [
84 0x6a09e667, 0xbb67ae85, 0x3c6ef372, 0xa54ff53a, 0x510e527f, 0x9b05688c, 0x1f83d9ab, 0x5be0cd19,
85];
86
87const K: [u32; 64] = [
91 0x428a2f98, 0x71374491, 0xb5c0fbcf, 0xe9b5dba5, 0x3956c25b, 0x59f111f1, 0x923f82a4, 0xab1c5ed5,
92 0xd807aa98, 0x12835b01, 0x243185be, 0x550c7dc3, 0x72be5d74, 0x80deb1fe, 0x9bdc06a7, 0xc19bf174,
93 0xe49b69c1, 0xefbe4786, 0x0fc19dc6, 0x240ca1cc, 0x2de92c6f, 0x4a7484aa, 0x5cb0a9dc, 0x76f988da,
94 0x983e5152, 0xa831c66d, 0xb00327c8, 0xbf597fc7, 0xc6e00bf3, 0xd5a79147, 0x06ca6351, 0x14292967,
95 0x27b70a85, 0x2e1b2138, 0x4d2c6dfc, 0x53380d13, 0x650a7354, 0x766a0abb, 0x81c2c92e, 0x92722c85,
96 0xa2bfe8a1, 0xa81a664b, 0xc24b8b70, 0xc76c51a3, 0xd192e819, 0xd6990624, 0xf40e3585, 0x106aa070,
97 0x19a4c116, 0x1e376c08, 0x2748774c, 0x34b0bcb5, 0x391c0cb3, 0x4ed8aa4a, 0x5b9cca4f, 0x682e6ff3,
98 0x748f82ee, 0x78a5636f, 0x84c87814, 0x8cc70208, 0x90befffa, 0xa4506ceb, 0xbef9a3f7, 0xc67178f2,
99];
100
101const SINGLE_BLOCK_MAX_LEN: usize = BLOCK_LEN - 1 - 8;
109
110#[inline]
116fn min_batches_per_task(batch: usize) -> usize {
117 min_len_for_work(WorkPerItem::HashCompression).div_ceil(batch)
118}
119
120#[inline]
124fn be_digest(state: &[u32; 8]) -> Output<Sha256> {
125 let mut digest = Output::<Sha256>::default();
126 for (chunk, word) in digest.chunks_exact_mut(4).zip(state) {
127 chunk.copy_from_slice(&word.to_be_bytes());
128 }
129 digest
130}
131
132impl ParallelHashSuite for Sha256HashSuite {
134 type ParLeafHash = ParallelSha256Digest;
135 type ParCompression = ParallelSha256Compression;
136}
137
138#[inline]
145fn compress_node_pairs<const N: usize>(
146 initial_state: &[u32; 8],
147 inputs: &[Output<Sha256>],
148 out: &mut [MaybeUninit<Output<Sha256>>],
149) {
150 let mut blocks = [[0u8; BLOCK_LEN]; N];
152 for (block, pair) in blocks.iter_mut().zip(inputs.chunks_exact(2)) {
153 block[..DIGEST_LEN].copy_from_slice(&pair[0]);
154 block[DIGEST_LEN..].copy_from_slice(&pair[1]);
155 }
156
157 let mut states = [*initial_state; N];
159 portable::compress256_multi(&mut states, &blocks);
160
161 for (slot, state) in out.iter_mut().zip(states) {
162 slot.write(must_cast::<[u32; 8], [u8; DIGEST_LEN]>(state).into());
163 }
164}
165
166#[derive(Debug, Clone, Default)]
171pub struct ParallelSha256Compression {
172 compression: Sha256Compression,
174}
175
176impl ParallelPseudoCompression<Output<Sha256>, 2> for ParallelSha256Compression {
177 type Compression = Sha256Compression;
178
179 fn compression(&self) -> &Self::Compression {
180 &self.compression
181 }
182
183 fn parallel_compress(
184 &self,
185 inputs: &[Output<Sha256>],
186 out: &mut [MaybeUninit<Output<Sha256>>],
187 ) {
188 assert_eq!(inputs.len(), 2 * out.len(), "Input length must be N * output length");
189
190 inputs
193 .par_chunks(2 * LANES)
194 .zip(out.par_chunks_mut(LANES))
195 .with_min_len(min_batches_per_task(LANES))
196 .for_each(|(pairs, out_batch)| {
197 compress_node_pairs::<LANES>(self.compression.initial_state(), pairs, out_batch);
198 });
199 }
200}
201
202fn digest_single_block_leaves<const N: usize, I>(
214 leaf_len: usize,
215 source: impl IndexedParallelIterator<Item = I>,
216 out: &mut [MaybeUninit<Output<Sha256>>],
217) where
218 I: IntoIterator<Item: SerializeBytes>,
219{
220 debug_assert!(leaf_len <= SINGLE_BLOCK_MAX_LEN, "pre-condition: the leaf fits in one block");
221
222 let mut template = [0u8; BLOCK_LEN];
224 template[leaf_len] = 0x80;
225 template[BLOCK_LEN - 8..].copy_from_slice(&((leaf_len as u64) * 8).to_be_bytes());
226
227 source
228 .chunks(N)
229 .zip(out.par_chunks_mut(N))
230 .with_min_len(min_batches_per_task(N))
231 .for_each_with([template; N], |blocks, (leaves, out_batch)| {
232 for (block, items) in blocks.iter_mut().zip(leaves) {
235 let mut cursor = &mut block[..leaf_len];
236 for item in items {
237 item.serialize(&mut cursor)
238 .expect("pre-condition: items must serialize without error");
239 }
240 debug_assert!(cursor.is_empty(), "pre-condition: each leaf serializes to leaf_len");
241 }
242
243 let mut states = [IV; N];
246 portable::compress256_multi(&mut states, blocks);
247 for (slot, state) in out_batch.iter_mut().zip(states) {
248 slot.write(be_digest(&state));
249 }
250 });
251}
252
253fn digest_multi_block_leaves<const N: usize, I>(
261 leaf_len: usize,
262 source: impl IndexedParallelIterator<Item = I>,
263 out: &mut [MaybeUninit<Output<Sha256>>],
264) where
265 I: IntoIterator<Item: SerializeBytes>,
266{
267 use bytes::BytesMut;
268
269 source
270 .chunks(N)
271 .zip(out.par_chunks_mut(N))
272 .with_min_len(min_batches_per_task(N))
273 .for_each_with(
274 std::array::from_fn::<_, N, _>(|_| BytesMut::new()),
275 |bufs, (leaves, out_batch)| {
276 let n_leaves = leaves.len();
277 for (buf, items) in bufs.iter_mut().zip(leaves) {
278 buf.clear();
280 for item in items {
281 item.serialize(&mut *buf)
282 .expect("pre-condition: items must serialize without error");
283 }
284 assert_eq!(
285 buf.len(),
286 leaf_len,
287 "pre-condition: each leaf serializes to leaf_len"
288 );
289 }
290
291 for buf in &mut bufs[n_leaves..] {
294 buf.resize(leaf_len, 0);
295 }
296
297 let inputs: [&[u8]; N] = std::array::from_fn(|i| bufs[i].as_ref());
298 for (slot, digest) in out_batch.iter_mut().zip(portable::sha256_multi(inputs)) {
299 slot.write(digest.into());
300 }
301 },
302 );
303}
304
305#[derive(Debug, Clone, Default)]
317pub struct ParallelSha256Digest;
318
319impl ParallelDigest for ParallelSha256Digest {
320 type Digest = Sha256;
321
322 fn new() -> Self {
323 Self
324 }
325
326 fn digest<I: IntoIterator<Item: SerializeBytes>>(
327 &self,
328 source: impl IndexedParallelIterator<Item = I>,
329 out: &mut [MaybeUninit<Output<Sha256>>],
330 ) {
331 ParallelDigestAdapter::<Sha256>::new().digest(source, out);
332 }
333
334 fn digest_with_const_len<I: IntoIterator<Item: FixedSizeSerializeBytes>>(
335 &self,
336 n_items_per_input: usize,
337 source: impl IndexedParallelIterator<Item = I>,
338 out: &mut [MaybeUninit<Output<Sha256>>],
339 ) {
340 let leaf_len = n_items_per_input * <I::Item as FixedSizeSerializeBytes>::BYTE_SIZE;
342
343 if leaf_len <= SINGLE_BLOCK_MAX_LEN {
344 digest_single_block_leaves::<LANES, I>(leaf_len, source, out);
345 } else {
346 digest_multi_block_leaves::<LANES, I>(leaf_len, source, out);
347 }
348 }
349}
350
351#[cfg(test)]
352mod tests {
353 use std::iter::repeat_with;
354
355 use binius_utils::rayon::iter::{IntoParallelRefIterator, ParallelIterator};
356 use digest::Digest;
357 use rand::{Rng, RngExt, SeedableRng, rngs::StdRng};
358
359 use super::*;
360 use crate::parallel_compression::ParallelCompressionAdaptor;
361
362 #[test]
363 fn test_batch_floor_counts_compressions_not_batches() {
364 let per_compression = min_len_for_work(WorkPerItem::HashCompression);
369 assert_eq!(min_batches_per_task(1), per_compression);
370
371 for batch in [2usize, 4, 8, 16] {
372 let batches = min_batches_per_task(batch);
373 assert!(batches * batch >= per_compression, "batch {batch} undershoots the floor");
375 assert!(
376 (batches - 1) * batch < per_compression,
377 "batch {batch} overshoots by a whole batch"
378 );
379 }
380 }
381
382 #[test]
383 fn test_parallel_compression_matches_adaptor() {
384 let mut rng = StdRng::seed_from_u64(0);
385
386 for n_nodes in [1usize, 2, 3, 4, 5, 8, 9, 16, 17, 64, 1000] {
395 let inputs: Vec<Output<Sha256>> = repeat_with(|| {
397 let mut digest = Output::<Sha256>::default();
398 rng.fill_bytes(&mut digest);
399 digest
400 })
401 .take(2 * n_nodes)
402 .collect();
403
404 let mut got = repeat_with(MaybeUninit::<Output<Sha256>>::uninit)
406 .take(n_nodes)
407 .collect::<Vec<_>>();
408 ParallelSha256Compression::default().parallel_compress(&inputs, &mut got);
409
410 let mut want = repeat_with(MaybeUninit::<Output<Sha256>>::uninit)
412 .take(n_nodes)
413 .collect::<Vec<_>>();
414 ParallelCompressionAdaptor::new(Sha256Compression::default())
415 .parallel_compress(&inputs, &mut want);
416
417 for (i, (got_i, want_i)) in got.iter().zip(&want).enumerate() {
418 let (got_i, want_i) =
420 unsafe { (got_i.assume_init_ref(), want_i.assume_init_ref()) };
421 assert_eq!(got_i, want_i, "mismatch at node {i} of {n_nodes}");
422 }
423 }
424 }
425
426 fn check_const_len_leaves(rng: &mut StdRng, n_items: usize, n_leaves: usize) {
428 let leaves: Vec<Vec<u128>> = (0..n_leaves)
430 .map(|_| (0..n_items).map(|_| rng.random()).collect())
431 .collect();
432
433 let mut got = repeat_with(MaybeUninit::<Output<Sha256>>::uninit)
434 .take(n_leaves)
435 .collect::<Vec<_>>();
436 ParallelSha256Digest::new().digest_with_const_len(
437 n_items,
438 leaves.par_iter().map(|leaf| leaf.iter().copied()),
439 &mut got,
440 );
441
442 for (i, (slot, leaf)) in got.into_iter().zip(&leaves).enumerate() {
443 let mut bytes = Vec::new();
444 for &item in leaf {
445 bytes.extend_from_slice(&item.to_le_bytes());
446 }
447 let got_i = unsafe { slot.assume_init() };
449 assert_eq!(
450 got_i,
451 <Sha256 as Digest>::digest(&bytes),
452 "leaf {i} of {n_leaves}, {n_items} items"
453 );
454 }
455 }
456
457 #[test]
458 fn test_const_len_leaves_match_serial() {
459 let mut rng = StdRng::seed_from_u64(0);
460
461 for n_items in [1, 2, 3, 4, 8] {
469 for n_leaves in [1, 4, 7, 8, 9, 16, 17, 48, 50] {
470 check_const_len_leaves(&mut rng, n_items, n_leaves);
471 }
472 }
473 }
474
475 #[test]
476 fn test_variable_len_leaves_match_serial() {
477 let mut rng = StdRng::seed_from_u64(1);
478 let n_leaves = 50;
482 let leaves: Vec<Vec<u128>> = (0..n_leaves)
483 .map(|_| (0..4).map(|_| rng.random()).collect())
484 .collect();
485
486 let mut got = repeat_with(MaybeUninit::<Output<Sha256>>::uninit)
487 .take(n_leaves)
488 .collect::<Vec<_>>();
489 ParallelSha256Digest::new()
490 .digest(leaves.par_iter().map(|leaf| leaf.iter().copied()), &mut got);
491
492 for (slot, leaf) in got.into_iter().zip(&leaves) {
493 let mut bytes = Vec::new();
494 for &item in leaf {
495 bytes.extend_from_slice(&item.to_le_bytes());
496 }
497 assert_eq!(unsafe { slot.assume_init() }, <Sha256 as Digest>::digest(&bytes));
499 }
500 }
501}