Skip to main content

binius_hash_prover/blake3/
mod.rs

1// Copyright 2026 The Binius Developers
2
3//! Blake3 hash and compression functions for use in Merkle tree constructions.
4//!
5//! [`portable`] holds the multi-lane kernel both parallel paths run on.
6//! The arch-specific kernels it dispatches to are private submodules here:
7//! - `neon` — hand-written message transpose and block compression for ARM64.
8//! - `avx512` — hand-written message transpose for x86-64.
9//!
10//! Each is compiled in only when the target has the feature, and each is pinned
11//! equal to the portable path in [`portable`]'s tests.
12//!
13//! The Blake3 spec constants live here too, since every kernel reads them.
14
15use binius_hash::Blake3HashSuite;
16pub use portable::{PortableBlake3ParallelCompression, PortableBlake3ParallelDigest};
17
18use crate::suite::ParallelHashSuite;
19
20pub mod portable;
21
22/// Hand-written vector kernels for the message load and the block compression.
23///
24/// They stand in for the byte-wise loader and the lane loops in [`portable`].
25#[cfg(all(target_arch = "aarch64", target_feature = "neon"))]
26mod neon;
27
28/// Hand-written vector transpose for the message load, used in place of the byte-wise loader.
29#[cfg(all(target_arch = "x86_64", target_feature = "avx512f"))]
30mod avx512;
31
32/// Blake3 domain-separation flag marking the first block of a chunk.
33const CHUNK_START: u32 = 1 << 0;
34
35/// Blake3 domain-separation flag marking the last block of a chunk.
36const CHUNK_END: u32 = 1 << 1;
37
38/// Blake3 domain-separation flag marking the last block of the whole tree.
39const ROOT: u32 = 1 << 3;
40
41/// Blake3 initial chaining value: the eight IV words, identical to the SHA-256 IV.
42const IV: [u32; 8] = [
43	0x6a09e667, 0xbb67ae85, 0x3c6ef372, 0xa54ff53a, 0x510e527f, 0x9b05688c, 0x1f83d9ab, 0x5be0cd19,
44];
45
46/// Blake3 message permutation applied between rounds.
47///
48/// The single fixed schedule from section 2.2 of the Blake3 spec, Table 2.
49const MSG_PERMUTATION: [usize; 16] = [2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8];
50
51/// The 7-round count of the Blake3 keyed permutation.
52const N_ROUNDS: usize = 7;
53
54/// Message word that each of a round's 16 slots reads, one row per round.
55///
56/// Blake3 advances the message between rounds by one fixed permutation.
57/// Applying that permutation `r` times gives the words round `r` reads:
58///
59/// ```text
60///     row 0:   0  1  2  3  4  5  6  7  8  9 10 11 12 13 14 15    <- the message in order
61///     row 1:   2  6  3 10  7  0  4 13  1 11 12  5  9 14 15  8    <- permuted once
62///     row 2:   3  4 10 12 13  2  7 14  6  5  9  0 11 15  8  1    <- permuted twice
63///     ...
64/// ```
65///
66/// # Why compose it here
67///
68/// Permuting the words a round reads is the same as permuting the words themselves.
69/// Settling that at compile time leaves the 16 message words loaded once and never moved.
70const MSG_SCHEDULE: [[usize; 16]; N_ROUNDS] = {
71	let mut schedule = [[0usize; 16]; N_ROUNDS];
72
73	// Row 0: the first round consumes the message in its natural order.
74	let mut w = 0;
75	while w < 16 {
76		schedule[0][w] = w;
77		w += 1;
78	}
79
80	// Row r: reading the row above through the permutation applies it one more time.
81	let mut r = 1;
82	while r < N_ROUNDS {
83		let mut w = 0;
84		while w < 16 {
85			schedule[r][w] = schedule[r - 1][MSG_PERMUTATION[w]];
86			w += 1;
87		}
88		r += 1;
89	}
90
91	schedule
92};
93
94/// Batched Blake3, on top of the sequential pair the verifier side defines.
95///
96/// The batch width is 16 lanes for both paths:
97/// - the throughput sweet spot on NEON in the portable-kernel benchmark.
98/// - the width the AVX2 and AVX-512 vectorizers fill.
99/// - 4 and 8 lanes both measure slower.
100impl ParallelHashSuite for Blake3HashSuite {
101	type ParLeafHash = PortableBlake3ParallelDigest<16>;
102	type ParCompression = PortableBlake3ParallelCompression<16>;
103}
104
105#[cfg(test)]
106mod tests {
107	use std::{iter::repeat_with, mem::MaybeUninit};
108
109	use binius_hash::{Blake3Compression, CompressionFunction};
110	use binius_utils::rayon::iter::{IntoParallelRefIterator, ParallelIterator};
111	use digest::Output;
112	use rand::{RngExt, SeedableRng, rngs::StdRng};
113
114	use super::*;
115	use crate::{ParallelDigest, parallel_digest::ParallelDigestAdapter};
116
117	/// Checks that the compression function matches `blake3::hash` of the concatenated inputs.
118	#[test]
119	fn test_blake3_compression_matches_reference() {
120		let mut rng = StdRng::seed_from_u64(0);
121		let left: [u8; 32] = rng.random();
122		let right: [u8; 32] = rng.random();
123
124		let compressed = Blake3Compression.compress([left.into(), right.into()]);
125
126		let mut concatenated = [0u8; 64];
127		concatenated[..32].copy_from_slice(&left);
128		concatenated[32..].copy_from_slice(&right);
129		let expected = blake3::hash(&concatenated);
130
131		assert_eq!(compressed.as_slice(), expected.as_bytes());
132	}
133
134	/// Checks that the parallel leaf digest matches `blake3::hash` over the serialized leaf bytes.
135	#[test]
136	fn test_parallel_blake3_matches_serial() {
137		let mut rng = StdRng::seed_from_u64(0);
138		let n_leaves = 50;
139		// `u128` serializes to 16 little-endian bytes.
140		let leaves: Vec<Vec<u128>> = (0..n_leaves)
141			.map(|_| (0..4).map(|_| rng.random()).collect())
142			.collect();
143
144		let digest = <ParallelDigestAdapter<blake3::Hasher> as ParallelDigest>::new();
145		let mut results = repeat_with(MaybeUninit::<Output<blake3::Hasher>>::uninit)
146			.take(n_leaves)
147			.collect::<Vec<_>>();
148		digest.digest(leaves.par_iter().map(|leaf| leaf.iter().copied()), &mut results);
149
150		for (result, leaf) in results.into_iter().zip(&leaves) {
151			let mut bytes = Vec::new();
152			for &item in leaf {
153				bytes.extend_from_slice(&item.to_le_bytes());
154			}
155			let expected = blake3::hash(&bytes);
156			assert_eq!(unsafe { result.assume_init() }.as_slice(), expected.as_bytes());
157		}
158	}
159
160	#[test]
161	fn test_portable_leaf_hash_matches_scalar_reference() {
162		// The suite's parallel leaf path is the portable vectorized kernel.
163		//
164		// Pin it equal to the scalar adapter and to `blake3::hash` across the routing boundary.
165		let mut rng = StdRng::seed_from_u64(1);
166		let n_leaves = 50;
167
168		// Leaves are `u8` items (BYTE_SIZE = 1), so `leaf_len` bytes == `leaf_len` items.
169		let mut check = |leaf_len: usize| {
170			let leaves: Vec<Vec<u8>> = (0..n_leaves)
171				.map(|_| (0..leaf_len).map(|_| rng.random()).collect())
172				.collect();
173
174			// The scalar adapter that walks the Blake3 tree — the reference path.
175			let mut scalar = repeat_with(MaybeUninit::<Output<blake3::Hasher>>::uninit)
176				.take(n_leaves)
177				.collect::<Vec<_>>();
178			ParallelDigestAdapter::<blake3::Hasher>::default().digest_with_const_len(
179				leaf_len,
180				leaves.par_iter().map(|leaf| leaf.iter().copied()),
181				&mut scalar,
182			);
183
184			// The suite's parallel leaf path — the portable batch kernel.
185			let mut portable = repeat_with(MaybeUninit::<Output<blake3::Hasher>>::uninit)
186				.take(n_leaves)
187				.collect::<Vec<_>>();
188			<Blake3HashSuite as ParallelHashSuite>::ParLeafHash::default().digest_with_const_len(
189				leaf_len,
190				leaves.par_iter().map(|leaf| leaf.iter().copied()),
191				&mut portable,
192			);
193
194			// Invariant: both reproduce `blake3::hash`, so their leaf digests match.
195			for ((s, p), leaf) in scalar.into_iter().zip(portable).zip(&leaves) {
196				let expected = blake3::hash(leaf);
197				let (s, p) = unsafe { (s.assume_init(), p.assume_init()) };
198				assert_eq!(s.as_slice(), expected.as_bytes(), "scalar, leaf_len {leaf_len}");
199				assert_eq!(p.as_slice(), expected.as_bytes(), "portable, leaf_len {leaf_len}");
200			}
201		};
202
203		// Straddle the 1024-byte routing boundary:
204		// - 0, 1, 63, 100, 1000, 1024 : within one chunk -> portable batch route.
205		// - 1025, 4096                : multi-chunk       -> scalar adapter fallback.
206		for leaf_len in [0, 1, 63, 100, 1000, 1024, 1025, 4096] {
207			check(leaf_len);
208		}
209	}
210}