Skip to main content

binius_hash_prover/sha256/
mod.rs

1// Copyright 2024-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Batched SHA-256 leaf hashing and inner-node compression.
5//!
6//! A round cannot start until the one before it lands, so a single chain stalls on itself.
7//! Both paths here hash many independent messages per call, which fills those stalls.
8//!
9//! The portable submodule holds the multi-lane kernel every batched path runs on.
10//! The architecture-specific kernels it dispatches to are private submodules:
11//!
12//! ```text
13//!     avx512   x86-64, sixteen lanes held transposed across 512-bit registers
14//!     sha_ni   x86-64, independent chains interleaved over the SHA extension
15//!     neon     aarch64, the same over the ARMv8 crypto extension
16//! ```
17//!
18//! Each is compiled in only when the target has the feature.
19//! Each is also pinned equal to the portable path by its tests.
20//!
21//! The spec constants live here, since every kernel reads them.
22//!
23//! Reference: FIPS 180-4, section 6.2.
24
25use 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/// Interleaved chains over the x86-64 SHA extension, at any lane count.
49#[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/// Transposed 16-lane kernel, the fastest measured path on x86-64.
59///
60/// It is also the only wide option on an AVX-512 machine with no SHA extension, such as
61/// Skylake-SP, Cascade Lake, or Cooper Lake.
62#[cfg(all(
63	target_arch = "x86_64",
64	target_feature = "avx512f",
65	target_feature = "avx512bw"
66))]
67mod avx512;
68
69/// Interleaved chains over the ARMv8 crypto extension, at four lanes.
70#[cfg(all(target_arch = "aarch64", target_feature = "sha2"))]
71mod neon;
72
73/// Bytes in one SHA-256 message block.
74const BLOCK_LEN: usize = 64;
75
76/// Bytes in one SHA-256 digest.
77const DIGEST_LEN: usize = 32;
78
79/// SHA-256 initial hash values, the starting state of an unkeyed hash.
80///
81/// The fractional parts of the square roots of the first eight primes, per FIPS 180-4 section
82/// 5.3.3.
83const IV: [u32; 8] = [
84	0x6a09e667, 0xbb67ae85, 0x3c6ef372, 0xa54ff53a, 0x510e527f, 0x9b05688c, 0x1f83d9ab, 0x5be0cd19,
85];
86
87/// The 64 SHA-256 round constants.
88///
89/// The fractional parts of the cube roots of the first 64 primes, per FIPS 180-4 section 4.2.2.
90const 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
101/// The largest message, in bytes, that still fits with its padding in a single 64-byte block.
102///
103/// The padding is one `0x80` terminator and the eight-byte big-endian bit length:
104///
105/// ```text
106///     64 - 1 - 8 = 55
107/// ```
108const SINGLE_BLOCK_MAX_LEN: usize = BLOCK_LEN - 1 - 8;
109
110/// Minimum batches per task, so one task still covers the work floor the cost model sets.
111///
112/// The floor counts compressions, but a batched loop hands out whole batches.
113/// Applying it to batches directly would inflate the task by the batch width.
114/// A layer smaller than that inflated floor would then run as a single task.
115#[inline]
116fn min_batches_per_task(batch: usize) -> usize {
117	min_len_for_work(WorkPerItem::HashCompression).div_ceil(batch)
118}
119
120/// Serializes eight state words into a standard SHA-256 digest.
121///
122/// FIPS 180-4 section 6.2.2 emits the state most significant byte first.
123#[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
132/// Batched SHA-256, on top of the sequential pair the verifier side defines.
133impl ParallelHashSuite for Sha256HashSuite {
134	type ParLeafHash = ParallelSha256Digest;
135	type ParCompression = ParallelSha256Compression;
136}
137
138/// Folds a batch of node pairs with a single block compression.
139///
140/// A pair is `left || right`, two 32-byte digests, so the message is exactly one block.
141///
142/// The output length sets how many pairs are folded and must not exceed the batch width.
143/// A partial batch leaves the unused high lanes zero, and nothing reads their output.
144#[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	// Pack each pair into one 64-byte message block: bytes 0..32 left child, 32..64 right child.
151	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	// Every lane starts from the same domain-separated state.
158	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/// Parallel SHA-256 two-to-one compression for the inner nodes of a Merkle tree.
167///
168/// Batches of independent compressions run through the multi-lane kernel.
169/// Every output byte equals compressing that node on its own with the scalar function.
170#[derive(Debug, Clone, Default)]
171pub struct ParallelSha256Compression {
172	/// The scalar two-to-one compression whose output the batched path reproduces exactly.
173	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		// One batch is `LANES` parent nodes, fed by `2 * LANES` child digests.
191		// A trailing batch shorter than `LANES` runs the kernel on its valid lanes only.
192		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
202/// Hashes leaves that fit, together with their padding, in a single 64-byte block.
203///
204/// Every leaf is the same length, so the padding suffix is shared and built once.
205/// Each leaf then overwrites only the message prefix.
206///
207/// A task reuses its own block buffers across batches, keeping the serialize pass off the
208/// allocator.
209///
210/// # Panics
211///
212/// Panics if a leaf does not serialize to its expected length.
213fn 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	// The `0x80` terminator right after the message, zeros, then the 64-bit big-endian bit length.
223	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			// Overwrite each lane's message prefix.
233			// The padding suffix stays untouched.
234			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			// A trailing batch leaves the high lanes holding an earlier leaf, whose digest is
244			// computed and then dropped, since `out_batch` is shorter than the batch.
245			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
253/// Hashes leaves too long for a single block, a batch at a time.
254///
255/// Every leaf is the same length, which is what lets one batch share a block count.
256///
257/// # Panics
258///
259/// Panics if a leaf does not serialize to its expected length.
260fn 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					// Reuse the capacity this task's earlier batches already grew.
279					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				// All lanes must share a length, so pad an unfilled high lane with zeros.
292				// Its digest is computed and then dropped.
293				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/// Batches fixed-length SHA-256 leaves through the multi-lane kernel.
306///
307/// With a fixed leaf length, the leaf size picks the route:
308///
309/// ```text
310///     up to 55 bytes : one batched block compression, padding folded into the block
311///     longer         : a batched multi-block hash, one shared block count
312/// ```
313///
314/// Without a fixed length a batch cannot share a block count, so that path falls back to
315/// hashing one leaf at a time.
316#[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		// Every leaf serializes to the same fixed byte length.
341		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		// Invariant: batching must not change how much work one task takes on.
365		//
366		// The floor is stated per compression, so widening the batch has to divide it.
367		// Applying it to batches unconverted is what silently serializes a whole layer.
368		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			// A task covers at least the floor, and overshoots by less than one batch.
374			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		// Invariant: the batched path equals per-node scalar compression byte for byte.
387		//
388		// Node counts crossing every regime of the batching, at LANES of 4, 8, or 16:
389		//
390		//     1, 2, 3     -> tail only (the top Merkle layers)
391		//     4, 8, 16    -> exactly one full batch at some tuned width
392		//     5, 9, 17    -> one full batch plus a tail
393		//     64, 1000    -> many batches, the second not a multiple of any width
394		for n_nodes in [1usize, 2, 3, 4, 5, 8, 9, 16, 17, 64, 1000] {
395			// Two random child digests per output node.
396			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			// Compress with the batched path.
405			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			// Compress every node one at a time through the scalar function as the reference.
411			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				// SAFETY: the compression calls above initialize every output slot.
419				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	/// Hashes a run of fixed-length leaves and pins every one to the reference digest.
427	fn check_const_len_leaves(rng: &mut StdRng, n_items: usize, n_leaves: usize) {
428		// `u128` serializes to 16 little-endian bytes, so the leaf is `16 * n_items` bytes.
429		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			// SAFETY: the digest call above initializes every output slot.
448			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		// Leaf lengths straddle the single-block boundary of 55 bytes:
462		//
463		//     1, 2, 3 items -> 16, 32, 48 bytes -> one block, padding folded in
464		//     4, 8 items    -> 64, 128 bytes    -> the batched multi-block route
465		//
466		// Leaf counts straddle every tuned batch width, and 1 leaves a batch mostly unfilled,
467		// which is the case that has to pad the idle lanes rather than read stale bytes.
468		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		// Without a fixed length a batch cannot share a block count.
479		// So this routes to the path that hashes one leaf at a time.
480		// Pin that path too, since callers reach it.
481		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			// SAFETY: the digest call above initializes every output slot.
498			assert_eq!(unsafe { slot.assume_init() }, <Sha256 as Digest>::digest(&bytes));
499		}
500	}
501}