binius_hash_prover/sha256/
portable.rs1use std::array;
19
20#[cfg(all(
21 target_arch = "x86_64",
22 target_feature = "avx512f",
23 target_feature = "avx512bw"
24))]
25use super::avx512;
26#[cfg(all(target_arch = "aarch64", target_feature = "sha2"))]
27use super::neon;
28#[cfg(all(
29 target_arch = "x86_64",
30 target_feature = "sha",
31 target_feature = "sse2",
32 target_feature = "ssse3",
33 target_feature = "sse4.1"
34))]
35use super::sha_ni;
36use super::{BLOCK_LEN, DIGEST_LEN, IV, K, SINGLE_BLOCK_MAX_LEN};
37
38pub const LANES: usize =
51 if cfg!(all(target_arch = "x86_64", target_feature = "avx512f", target_feature = "avx512bw")) {
52 16
54 } else if cfg!(all(target_arch = "x86_64", target_feature = "sha")) {
55 8
56 } else if cfg!(all(target_arch = "aarch64", target_feature = "sha2")) {
57 4
58 } else {
59 8
61 };
62
63#[inline(always)]
76fn round<const N: usize>(state: [[u32; N]; 8], w: &[u32; N], k: u32) -> [[u32; N]; 8] {
77 let [a, b, c, d, e, f, g, h] = state;
78 let mut next_a = [0u32; N];
79 let mut next_e = [0u32; N];
80
81 for i in 0..N {
84 let sigma1 = e[i].rotate_right(6) ^ e[i].rotate_right(11) ^ e[i].rotate_right(25);
86 let ch = (e[i] & f[i]) ^ (!e[i] & g[i]);
88 let t1 = h[i]
89 .wrapping_add(sigma1)
90 .wrapping_add(ch)
91 .wrapping_add(k)
92 .wrapping_add(w[i]);
93 let sigma0 = a[i].rotate_right(2) ^ a[i].rotate_right(13) ^ a[i].rotate_right(22);
95 let maj = (a[i] & b[i]) ^ (a[i] & c[i]) ^ (b[i] & c[i]);
97
98 next_e[i] = d[i].wrapping_add(t1);
99 next_a[i] = t1.wrapping_add(sigma0.wrapping_add(maj));
100 }
101
102 [next_a, a, b, c, next_e, e, f, g]
107}
108
109#[inline(always)]
130fn extend_window<const N: usize>(w: &mut [[u32; N]; 16]) {
131 for j in 0..16 {
132 let m15 = w[(j + 1) & 15];
134 let m7 = w[(j + 9) & 15];
135 let m2 = w[(j + 14) & 15];
136 for i in 0..N {
137 let sigma0 = m15[i].rotate_right(7) ^ m15[i].rotate_right(18) ^ (m15[i] >> 3);
138 let sigma1 = m2[i].rotate_right(17) ^ m2[i].rotate_right(19) ^ (m2[i] >> 10);
139 w[j][i] = w[j][i]
140 .wrapping_add(sigma0)
141 .wrapping_add(m7[i])
142 .wrapping_add(sigma1);
143 }
144 }
145}
146
147#[inline(always)]
152fn load_block_words<const N: usize>(blocks: &[[u8; BLOCK_LEN]; N]) -> [[u32; N]; 16] {
153 #[cfg(all(
154 target_arch = "x86_64",
155 target_feature = "avx512f",
156 target_feature = "avx512bw"
157 ))]
158 if avx512::handles_lanes(N) {
159 return avx512::load_block_words(blocks);
160 }
161
162 load_block_words_portable(blocks)
163}
164
165#[inline(always)]
170fn load_block_words_portable<const N: usize>(blocks: &[[u8; BLOCK_LEN]; N]) -> [[u32; N]; 16] {
171 let mut w = [[0u32; N]; 16];
172 for lane in 0..N {
173 for (t, slot) in w.iter_mut().enumerate() {
174 let off = t * 4;
175 slot[lane] = u32::from_be_bytes([
177 blocks[lane][off],
178 blocks[lane][off + 1],
179 blocks[lane][off + 2],
180 blocks[lane][off + 3],
181 ]);
182 }
183 }
184 w
185}
186
187#[inline]
194pub fn compress256_multi<const N: usize>(
195 states: &mut [[u32; 8]; N],
196 blocks: &[[u8; BLOCK_LEN]; N],
197) {
198 #[cfg(all(
200 target_arch = "x86_64",
201 target_feature = "avx512f",
202 target_feature = "avx512bw"
203 ))]
204 if avx512::handles_lanes(N) {
205 unsafe { avx512::compress256_multi(states, blocks) };
207 return;
208 }
209
210 #[cfg(all(
211 target_arch = "x86_64",
212 target_feature = "sha",
213 target_feature = "sse2",
214 target_feature = "ssse3",
215 target_feature = "sse4.1"
216 ))]
217 if sha_ni::handles_lanes(N) {
218 unsafe { sha_ni::compress256_multi(states, blocks) };
220 return;
221 }
222
223 #[cfg(all(target_arch = "aarch64", target_feature = "sha2"))]
224 if neon::handles_lanes(N) {
225 unsafe { neon::compress256_multi(states, blocks) };
227 return;
228 }
229
230 compress256_multi_portable(states, blocks);
231}
232
233#[inline]
238pub fn compress256_multi_portable<const N: usize>(
239 states: &mut [[u32; 8]; N],
240 blocks: &[[u8; BLOCK_LEN]; N],
241) {
242 let mut w = load_block_words(blocks);
243
244 let mut state: [[u32; N]; 8] = array::from_fn(|t| array::from_fn(|lane| states[lane][t]));
246 let incoming = state;
247
248 for group in 0..4 {
251 if group > 0 {
252 extend_window(&mut w);
253 }
254 for t in 16 * group..16 * group + 16 {
255 state = round(state, &w[t & 15], K[t]);
256 }
257 }
258
259 for (lane, out) in states.iter_mut().enumerate() {
261 for (t, word) in out.iter_mut().enumerate() {
262 *word = state[t][lane].wrapping_add(incoming[t][lane]);
263 }
264 }
265}
266
267pub fn sha256_multi<const N: usize>(inputs: [&[u8]; N]) -> [[u8; DIGEST_LEN]; N] {
275 let len = inputs[0].len();
276 assert!(inputs.iter().all(|input| input.len() == len), "the inputs must have equal length");
277
278 let mut states = [IV; N];
279 let mut blocks = [[0u8; BLOCK_LEN]; N];
280
281 for b in 0..len / BLOCK_LEN {
283 for (block, input) in blocks.iter_mut().zip(inputs) {
284 block.copy_from_slice(&input[b * BLOCK_LEN..(b + 1) * BLOCK_LEN]);
285 }
286 compress256_multi(&mut states, &blocks);
287 }
288
289 let rem = len % BLOCK_LEN;
293 let n_tail = if rem <= SINGLE_BLOCK_MAX_LEN { 1 } else { 2 };
294 let mut tails = [[0u8; 2 * BLOCK_LEN]; N];
295 for (tail, input) in tails.iter_mut().zip(inputs) {
296 tail[..rem].copy_from_slice(&input[len - rem..]);
297 tail[rem] = 0x80;
298 tail[n_tail * BLOCK_LEN - 8..n_tail * BLOCK_LEN]
299 .copy_from_slice(&((len as u64) * 8).to_be_bytes());
300 }
301 for b in 0..n_tail {
302 for (block, tail) in blocks.iter_mut().zip(&tails) {
303 block.copy_from_slice(&tail[b * BLOCK_LEN..(b + 1) * BLOCK_LEN]);
304 }
305 compress256_multi(&mut states, &blocks);
306 }
307
308 states.map(|state| {
310 let mut digest = [0u8; DIGEST_LEN];
311 for (chunk, word) in digest.chunks_exact_mut(4).zip(state) {
312 chunk.copy_from_slice(&word.to_be_bytes());
313 }
314 digest
315 })
316}
317
318#[cfg(test)]
319mod tests {
320 use proptest::prelude::*;
321 use rand::{Rng, RngExt, SeedableRng, rngs::StdRng};
322 use sha2::{Digest, Sha256, block_api::compress256};
323
324 use super::*;
325
326 fn check_portable_against_sha2<const N: usize>(rng: &mut StdRng) {
329 let states: [[u32; 8]; N] = array::from_fn(|_| array::from_fn(|_| rng.random()));
332 let blocks: [[u8; BLOCK_LEN]; N] = array::from_fn(|_| array::from_fn(|_| rng.random()));
333
334 let mut got = states;
335 compress256_multi_portable(&mut got, &blocks);
336
337 for (lane, (got_lane, state)) in got.iter().zip(&states).enumerate() {
339 let mut want = *state;
340 compress256(&mut want, std::slice::from_ref(&blocks[lane]));
341 assert_eq!(*got_lane, want, "lane {lane} of {N} diverged from sha2");
342 }
343 }
344
345 #[test]
346 fn test_portable_matches_sha2() {
347 let mut rng = StdRng::seed_from_u64(0);
348 for _ in 0..16 {
351 check_portable_against_sha2::<1>(&mut rng);
352 check_portable_against_sha2::<4>(&mut rng);
353 check_portable_against_sha2::<8>(&mut rng);
354 check_portable_against_sha2::<16>(&mut rng);
355 }
356 }
357
358 #[test]
359 fn test_portable_matches_sha2_at_state_extremes() {
360 for state_word in [0u32, u32::MAX] {
362 for block_byte in [0u8, 0xff] {
363 let states = [[state_word; 8]; 4];
364 let blocks = [[block_byte; BLOCK_LEN]; 4];
365
366 let mut got = states;
367 compress256_multi_portable(&mut got, &blocks);
368
369 let mut want = states[0];
370 compress256(&mut want, &blocks[..1]);
371 for (lane, got_lane) in got.iter().enumerate() {
372 assert_eq!(*got_lane, want, "lane {lane}, state {state_word:#x}");
373 }
374 }
375 }
376 }
377
378 fn check_dispatch_against_portable<const N: usize>(rng: &mut StdRng) {
382 let states: [[u32; 8]; N] = array::from_fn(|_| array::from_fn(|_| rng.random()));
383 let blocks: [[u8; BLOCK_LEN]; N] = array::from_fn(|_| array::from_fn(|_| rng.random()));
384
385 let mut want = states;
386 compress256_multi_portable(&mut want, &blocks);
387
388 let mut got = states;
389 compress256_multi(&mut got, &blocks);
390
391 assert_eq!(got, want, "the dispatched kernel diverged from the lane loops at {N} lanes");
392 }
393
394 #[test]
395 fn test_dispatch_matches_portable() {
396 let mut rng = StdRng::seed_from_u64(1);
397 for _ in 0..16 {
399 check_dispatch_against_portable::<1>(&mut rng);
400 check_dispatch_against_portable::<2>(&mut rng);
401 check_dispatch_against_portable::<4>(&mut rng);
402 check_dispatch_against_portable::<8>(&mut rng);
403 check_dispatch_against_portable::<16>(&mut rng);
404 }
405 }
406
407 proptest! {
408 #[test]
409 fn dispatch_matches_portable_proptest(seed in any::<u64>()) {
410 let mut rng = StdRng::seed_from_u64(seed);
413 check_dispatch_against_portable::<LANES>(&mut rng);
414 check_dispatch_against_portable::<16>(&mut rng);
415 }
416 }
417
418 fn check_sha256_multi<const N: usize>(rng: &mut StdRng, len: usize) {
420 let messages: [Vec<u8>; N] = array::from_fn(|_| {
422 let mut m = vec![0u8; len];
423 rng.fill_bytes(&mut m);
424 m
425 });
426 let refs: [&[u8]; N] = array::from_fn(|i| messages[i].as_slice());
427
428 let got = sha256_multi(refs);
429 for (lane, (got_lane, message)) in got.iter().zip(&messages).enumerate() {
430 let want: [u8; DIGEST_LEN] = <Sha256 as Digest>::digest(message).into();
431 assert_eq!(*got_lane, want, "len {len}, lane {lane} of {N}");
432 }
433 }
434
435 #[test]
436 fn test_sha256_multi_matches_sha2_at_padding_boundaries() {
437 let mut rng = StdRng::seed_from_u64(2);
438 for len in [0, 1, 54, 55, 56, 63, 64, 65, 119, 120, 128, 256] {
446 check_sha256_multi::<1>(&mut rng, len);
447 check_sha256_multi::<4>(&mut rng, len);
448 check_sha256_multi::<LANES>(&mut rng, len);
449 }
450 }
451
452 proptest! {
453 #[test]
454 fn sha256_multi_matches_sha2_proptest(seed in any::<u64>(), len in 0..300usize) {
455 let mut rng = StdRng::seed_from_u64(seed);
456 check_sha256_multi::<LANES>(&mut rng, len);
457 }
458 }
459
460 #[test]
461 #[should_panic(expected = "the inputs must have equal length")]
462 fn test_sha256_multi_rejects_unequal_lengths() {
463 sha256_multi::<2>([&[0u8; 8], &[0u8; 9]]);
465 }
466
467 #[cfg(all(
469 target_arch = "x86_64",
470 target_feature = "avx512f",
471 target_feature = "avx512bw"
472 ))]
473 fn check_avx512_transpose<const N: usize>(rng: &mut StdRng) {
474 let blocks: [[u8; BLOCK_LEN]; N] = array::from_fn(|_| array::from_fn(|_| rng.random()));
476
477 assert_eq!(
479 avx512::load_block_words(&blocks),
480 load_block_words_portable(&blocks),
481 "the shuffle network diverged from the byte-wise loader at {N} lanes"
482 );
483 }
484
485 #[cfg(all(
486 target_arch = "x86_64",
487 target_feature = "avx512f",
488 target_feature = "avx512bw"
489 ))]
490 #[test]
491 fn test_avx512_transpose_places_every_word() {
492 let mut blocks = [[0u8; BLOCK_LEN]; 16];
504 for (lane, block) in blocks.iter_mut().enumerate() {
505 for (w, word) in block.chunks_exact_mut(4).enumerate() {
506 word.copy_from_slice(&((lane * 16 + w) as u32).to_be_bytes());
507 }
508 }
509
510 let m = avx512::load_block_words(&blocks);
511 for (w, row) in m.iter().enumerate() {
512 for (lane, got) in row.iter().enumerate() {
513 assert_eq!(*got, (lane * 16 + w) as u32, "row {w}, lane {lane}");
514 }
515 }
516
517 let mut rng = StdRng::seed_from_u64(3);
519 for _ in 0..64 {
520 check_avx512_transpose::<16>(&mut rng);
521 }
522 }
523
524 #[cfg(all(
526 target_arch = "x86_64",
527 target_feature = "avx512f",
528 target_feature = "avx512bw"
529 ))]
530 fn check_avx512_core(rng: &mut StdRng) {
531 let states: [[u32; 8]; 16] = array::from_fn(|_| array::from_fn(|_| rng.random()));
532 let blocks: [[u8; BLOCK_LEN]; 16] = array::from_fn(|_| array::from_fn(|_| rng.random()));
533
534 let mut want = states;
535 compress256_multi_portable(&mut want, &blocks);
536
537 let mut got = states;
538 unsafe { avx512::compress256_multi(&mut got, &blocks) };
540
541 assert_eq!(got, want, "the AVX-512 kernel diverged from the lane loops");
542 }
543
544 #[cfg(all(
545 target_arch = "x86_64",
546 target_feature = "avx512f",
547 target_feature = "avx512bw"
548 ))]
549 #[test]
550 fn test_avx512_core_matches_portable() {
551 let mut rng = StdRng::seed_from_u64(4);
554 for _ in 0..64 {
555 check_avx512_core(&mut rng);
556 }
557 }
558
559 #[cfg(all(
560 target_arch = "x86_64",
561 target_feature = "avx512f",
562 target_feature = "avx512bw"
563 ))]
564 proptest! {
565 #[test]
566 fn avx512_core_matches_portable_proptest(seed in any::<u64>()) {
567 let mut rng = StdRng::seed_from_u64(seed);
568 check_avx512_core(&mut rng);
569 check_avx512_transpose::<16>(&mut rng);
570 }
571 }
572
573 #[cfg(all(
575 target_arch = "x86_64",
576 target_feature = "sha",
577 target_feature = "sse2",
578 target_feature = "ssse3",
579 target_feature = "sse4.1"
580 ))]
581 fn check_sha_ni_core<const N: usize>(rng: &mut StdRng) {
582 let states: [[u32; 8]; N] = array::from_fn(|_| array::from_fn(|_| rng.random()));
583 let blocks: [[u8; BLOCK_LEN]; N] = array::from_fn(|_| array::from_fn(|_| rng.random()));
584
585 let mut want = states;
586 compress256_multi_portable(&mut want, &blocks);
587
588 let mut got = states;
589 unsafe { sha_ni::compress256_multi(&mut got, &blocks) };
591
592 assert_eq!(got, want, "the SHA-extension kernel diverged from the lane loops at {N} lanes");
593 }
594
595 #[cfg(all(
596 target_arch = "x86_64",
597 target_feature = "sha",
598 target_feature = "sse2",
599 target_feature = "ssse3",
600 target_feature = "sse4.1"
601 ))]
602 #[test]
603 fn test_sha_ni_core_matches_portable() {
604 let mut rng = StdRng::seed_from_u64(5);
605 for _ in 0..32 {
607 check_sha_ni_core::<1>(&mut rng);
608 check_sha_ni_core::<2>(&mut rng);
609 check_sha_ni_core::<4>(&mut rng);
610 check_sha_ni_core::<8>(&mut rng);
611 check_sha_ni_core::<16>(&mut rng);
612 }
613 }
614
615 #[cfg(all(
616 target_arch = "x86_64",
617 target_feature = "sha",
618 target_feature = "sse2",
619 target_feature = "ssse3",
620 target_feature = "sse4.1"
621 ))]
622 proptest! {
623 #[test]
624 fn sha_ni_core_matches_portable_proptest(seed in any::<u64>()) {
625 let mut rng = StdRng::seed_from_u64(seed);
626 check_sha_ni_core::<8>(&mut rng);
627 }
628 }
629
630 #[cfg(all(target_arch = "aarch64", target_feature = "sha2"))]
632 fn check_neon_core(rng: &mut StdRng) {
633 let states: [[u32; 8]; 4] = array::from_fn(|_| array::from_fn(|_| rng.random()));
634 let blocks: [[u8; BLOCK_LEN]; 4] = array::from_fn(|_| array::from_fn(|_| rng.random()));
635
636 let mut want = states;
637 compress256_multi_portable(&mut want, &blocks);
638
639 let mut got = states;
640 unsafe { neon::compress256_multi(&mut got, &blocks) };
642
643 assert_eq!(got, want, "the crypto-extension kernel diverged from the lane loops");
644 }
645
646 #[cfg(all(target_arch = "aarch64", target_feature = "sha2"))]
647 #[test]
648 fn test_neon_core_matches_portable() {
649 let mut rng = StdRng::seed_from_u64(6);
650 for _ in 0..64 {
651 check_neon_core(&mut rng);
652 }
653 }
654
655 #[cfg(all(target_arch = "aarch64", target_feature = "sha2"))]
656 proptest! {
657 #[test]
658 fn neon_core_matches_portable_proptest(seed in any::<u64>()) {
659 let mut rng = StdRng::seed_from_u64(seed);
660 check_neon_core(&mut rng);
661 }
662 }
663}