binius_hash_prover/blake3/
mod.rs1use binius_hash::Blake3HashSuite;
16pub use portable::{PortableBlake3ParallelCompression, PortableBlake3ParallelDigest};
17
18use crate::suite::ParallelHashSuite;
19
20pub mod portable;
21
22#[cfg(all(target_arch = "aarch64", target_feature = "neon"))]
26mod neon;
27
28#[cfg(all(target_arch = "x86_64", target_feature = "avx512f"))]
30mod avx512;
31
32const CHUNK_START: u32 = 1 << 0;
34
35const CHUNK_END: u32 = 1 << 1;
37
38const ROOT: u32 = 1 << 3;
40
41const IV: [u32; 8] = [
43 0x6a09e667, 0xbb67ae85, 0x3c6ef372, 0xa54ff53a, 0x510e527f, 0x9b05688c, 0x1f83d9ab, 0x5be0cd19,
44];
45
46const MSG_PERMUTATION: [usize; 16] = [2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8];
50
51const N_ROUNDS: usize = 7;
53
54const MSG_SCHEDULE: [[usize; 16]; N_ROUNDS] = {
71 let mut schedule = [[0usize; 16]; N_ROUNDS];
72
73 let mut w = 0;
75 while w < 16 {
76 schedule[0][w] = w;
77 w += 1;
78 }
79
80 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
94impl 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 #[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 #[test]
136 fn test_parallel_blake3_matches_serial() {
137 let mut rng = StdRng::seed_from_u64(0);
138 let n_leaves = 50;
139 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 let mut rng = StdRng::seed_from_u64(1);
166 let n_leaves = 50;
167
168 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 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 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 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 for leaf_len in [0, 1, 63, 100, 1000, 1024, 1025, 4096] {
207 check(leaf_len);
208 }
209 }
210}