1use binius_core::word::Word;
21use binius_frontend::{CircuitBuilder, Wire};
22
23use crate::{
24 fixed_byte_vec::ByteVec,
25 util::{clear_high_bits, zeroed_u32_words},
26};
27
28pub mod compress;
29
30pub use compress::{
31 Blake3Compress2x, blake3_compress, blake3_compress_2x, blake3_compress_2x_seq, ref_compress,
32};
33
34pub const IV: [u32; 8] = [
36 0x6a09e667, 0xbb67ae85, 0x3c6ef372, 0xa54ff53a, 0x510e527f, 0x9b05688c, 0x1f83d9ab, 0x5be0cd19,
37];
38
39pub const MSG_SCHEDULE: [[usize; 16]; 7] = [
45 [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15],
46 [2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8],
47 [3, 4, 10, 12, 13, 2, 7, 14, 6, 5, 9, 0, 11, 15, 8, 1],
48 [10, 7, 12, 9, 14, 3, 13, 15, 4, 0, 11, 2, 5, 8, 1, 6],
49 [12, 13, 9, 11, 15, 10, 14, 8, 7, 2, 5, 3, 0, 1, 6, 4],
50 [9, 14, 11, 5, 8, 12, 15, 1, 13, 3, 0, 10, 2, 6, 4, 7],
51 [11, 15, 5, 0, 1, 9, 8, 6, 14, 10, 2, 12, 3, 4, 7, 13],
52];
53
54pub const CHUNK_START: u32 = 1 << 0;
56pub const CHUNK_END: u32 = 1 << 1;
57pub const PARENT: u32 = 1 << 2;
58pub const ROOT: u32 = 1 << 3;
59pub const KEYED_HASH: u32 = 1 << 4;
60pub const DERIVE_KEY_CONTEXT: u32 = 1 << 5;
61pub const DERIVE_KEY_MATERIAL: u32 = 1 << 6;
62
63pub const BLOCK_BYTES: usize = 64;
65
66pub const CHUNK_BYTES: usize = 1024;
68
69pub const KEY_BYTES: usize = 32;
71
72fn init_cv(builder: &CircuitBuilder, key: Option<[Wire; 8]>) -> [Wire; 8] {
75 key.unwrap_or_else(|| std::array::from_fn(|i| builder.add_constant(Word(IV[i] as u64))))
76}
77
78const fn key_flag(key: Option<[Wire; 8]>) -> u32 {
84 if key.is_some() { KEYED_HASH } else { 0 }
85}
86
87fn pack_lanes(builder: &CircuitBuilder, lo: Wire, hi: Wire) -> Wire {
91 builder.bxor(lo, builder.shl(hi, 32))
92}
93
94fn dup32(builder: &CircuitBuilder, value: u32) -> Wire {
96 let value = value as u64;
97 builder.add_constant(Word(value | (value << 32)))
98}
99
100pub fn blake3_chunk(
121 builder: &CircuitBuilder,
122 key: Option<[Wire; 8]>,
123 blocks: &[[Wire; 16]],
124 block_lens: &[Wire],
125 counter: u64,
126 last_flags_extra: u32,
127) -> [Wire; 8] {
128 let n_blocks = blocks.len();
129 assert!((1..=16).contains(&n_blocks), "blake3_chunk: n_blocks ({n_blocks}) must be in 1..=16",);
130 assert_eq!(
131 block_lens.len(),
132 n_blocks,
133 "blake3_chunk: block_lens.len() ({}) must equal blocks.len() ({n_blocks})",
134 block_lens.len(),
135 );
136
137 let counter = builder.add_constant_64(counter);
138
139 let flags: Vec<Wire> = (0..n_blocks)
140 .map(|j| {
141 let start = if j == 0 { CHUNK_START } else { 0 };
142 let end = if j + 1 == n_blocks {
143 CHUNK_END | last_flags_extra
144 } else {
145 0
146 };
147 builder.add_constant(Word((start | end | key_flag(key)) as u64))
148 })
149 .collect();
150
151 let mut cv = init_cv(builder, key);
152
153 let n_pairs = n_blocks / 2;
164 for pair in 0..n_pairs {
165 let (lo, hi) = (2 * pair, 2 * pair + 1);
166 cv = blake3_compress_2x_seq(
168 &builder.subcircuit(format!("blake3_chunk_compress[{pair}]")),
169 cv,
170 [blocks[lo], blocks[hi]],
171 counter,
172 [block_lens[lo], block_lens[hi]],
173 [flags[lo], flags[hi]],
174 );
175 }
176
177 if n_blocks % 2 == 1 {
188 let last = n_blocks - 1;
189 cv = blake3_compress(
190 &builder.subcircuit("blake3_chunk_compress[last]"),
191 cv,
192 blocks[last],
193 counter,
194 block_lens[last],
195 flags[last],
196 );
197 }
198
199 std::array::from_fn(|i| clear_high_bits(builder, cv[i], 32))
205}
206
207fn blake3_parent(
216 builder: &CircuitBuilder,
217 key: Option<[Wire; 8]>,
218 left: [Wire; 8],
219 right: [Wire; 8],
220 is_root: bool,
221) -> [Wire; 8] {
222 let cv = init_cv(builder, key);
223 let block: [Wire; 16] = std::array::from_fn(|i| if i < 8 { left[i] } else { right[i - 8] });
224 let counter = builder.add_constant(Word::ZERO);
225 let block_len = builder.add_constant(Word(BLOCK_BYTES as u64));
226 let root_flag = if is_root { ROOT } else { 0 };
227 let flags = builder.add_constant(Word((PARENT | root_flag | key_flag(key)) as u64));
228 let out = blake3_compress(builder, cv, block, counter, block_len, flags);
229 std::array::from_fn(|i| clear_high_bits(builder, out[i], 32))
230}
231
232fn blake3_parent_pair(
240 builder: &CircuitBuilder,
241 key: Option<[Wire; 8]>,
242 a: ([Wire; 8], [Wire; 8]),
243 b: ([Wire; 8], [Wire; 8]),
244) -> ([Wire; 8], [Wire; 8]) {
245 let cv: [Wire; 8] = key.map_or_else(
248 || std::array::from_fn(|i| dup32(builder, IV[i])),
249 |key| std::array::from_fn(|i| pack_lanes(builder, key[i], key[i])),
250 );
251 let block: [Wire; 16] = std::array::from_fn(|i| {
252 if i < 8 {
253 pack_lanes(builder, a.0[i], b.0[i])
254 } else {
255 pack_lanes(builder, a.1[i - 8], b.1[i - 8])
256 }
257 });
258 let zero = builder.add_constant(Word::ZERO);
259 let block_len = dup32(builder, BLOCK_BYTES as u32);
260 let flags = dup32(builder, PARENT | key_flag(key));
261 let out = blake3_compress_2x(builder, cv, block, zero, zero, block_len, flags);
262 let cv_a: [Wire; 8] = std::array::from_fn(|i| clear_high_bits(builder, out[i], 32));
263 let cv_b: [Wire; 8] = std::array::from_fn(|i| builder.shr(out[i], 32));
264 (cv_a, cv_b)
265}
266
267fn blake3_tree_root(
277 builder: &CircuitBuilder,
278 key: Option<[Wire; 8]>,
279 chunk_cvs: Vec<[Wire; 8]>,
280) -> [Wire; 8] {
281 assert!(chunk_cvs.len() >= 2, "blake3_tree_root: needs at least two chunks");
282
283 let mut level = chunk_cvs;
284 let mut depth = 0;
285 loop {
286 if level.len() == 2 {
288 return blake3_parent(
289 &builder.subcircuit("blake3_tree_root"),
290 key,
291 level[0],
292 level[1],
293 true,
294 );
295 }
296
297 let sub = builder.subcircuit(format!("blake3_tree_level[{depth}]"));
298 let n = level.len();
299 let n_pairs = n / 2;
300 let mut next: Vec<[Wire; 8]> = Vec::with_capacity(n.div_ceil(2));
301
302 let mut p = 0;
304 while p + 1 < n_pairs {
305 let (cv_a, cv_b) = blake3_parent_pair(
306 &sub,
307 key,
308 (level[2 * p], level[2 * p + 1]),
309 (level[2 * p + 2], level[2 * p + 3]),
310 );
311 next.push(cv_a);
312 next.push(cv_b);
313 p += 2;
314 }
315 if p < n_pairs {
317 next.push(blake3_parent(&sub, key, level[2 * p], level[2 * p + 1], false));
318 }
319 if n % 2 == 1 {
321 next.push(level[n - 1]);
322 }
323
324 level = next;
325 depth += 1;
326 }
327}
328
329pub fn blake3_fixed(builder: &CircuitBuilder, message: &[Wire], len_bytes: usize) -> [Wire; 8] {
352 blake3_hash_fixed(builder, None, message, len_bytes)
353}
354
355pub fn blake3_keyed_fixed(
379 builder: &CircuitBuilder,
380 message: &[Wire],
381 len_bytes: usize,
382 key: &ByteVec,
383) -> [Wire; 8] {
384 blake3_hash_fixed(builder, Some(key_words(builder, key)), message, len_bytes)
385}
386
387fn key_words(builder: &CircuitBuilder, key: &ByteVec) -> [Wire; 8] {
399 assert_eq!(
400 key.len_range.start(),
401 key.len_range.end(),
402 "BLAKE3: the key length must be fixed at circuit construction time, but len_range is {:?}",
403 key.len_range,
404 );
405 let key_len = *key.len_range.start();
406 assert!(key_len <= KEY_BYTES, "BLAKE3: key length ({key_len}) exceeds {KEY_BYTES}");
407
408 let words = zeroed_u32_words(builder, &key.data, key_len, 8);
409 std::array::from_fn(|i| words[i])
410}
411
412fn padded_message_words(
423 builder: &CircuitBuilder,
424 message: &[Wire],
425 len_bytes: usize,
426 mask_high_halves: bool,
427) -> Vec<Wire> {
428 assert_eq!(
429 message.len(),
430 len_bytes.div_ceil(4),
431 "blake3: message.len() ({}) must equal len_bytes.div_ceil(4) ({})",
432 message.len(),
433 len_bytes.div_ceil(4),
434 );
435
436 let n_padded_words = len_bytes.div_ceil(BLOCK_BYTES).max(1) * 16;
437 let n_whole_words = len_bytes / 4;
438 let boundary_bytes = len_bytes % 4;
439
440 let mut padded: Vec<Wire> = Vec::with_capacity(n_padded_words);
441 padded.extend(message[..n_whole_words].iter().map(|&w| {
442 if mask_high_halves {
443 clear_high_bits(builder, w, 32)
444 } else {
445 w
446 }
447 }));
448 if boundary_bytes > 0 {
449 let mask_value = (1u64 << (boundary_bytes * 8)) - 1;
452 let mask = builder.add_constant(Word(mask_value));
453 padded.push(builder.band(message[n_whole_words], mask));
454 }
455 padded.resize(n_padded_words, builder.add_constant(Word::ZERO));
457 padded
458}
459
460fn block_len_bytes(len_bytes: usize, j: usize) -> usize {
462 (len_bytes - j * BLOCK_BYTES).min(BLOCK_BYTES)
463}
464
465fn blake3_hash_fixed(
469 builder: &CircuitBuilder,
470 key: Option<[Wire; 8]>,
471 message: &[Wire],
472 len_bytes: usize,
473) -> [Wire; 8] {
474 let n_blocks = len_bytes.div_ceil(BLOCK_BYTES).max(1);
475 let padded = padded_message_words(builder, message, len_bytes, false);
477
478 let block = |j: usize| -> [Wire; 16] { std::array::from_fn(|i| padded[j * 16 + i]) };
479 let block_len =
480 |j: usize| -> Wire { builder.add_constant(Word(block_len_bytes(len_bytes, j) as u64)) };
481
482 let n_chunks = len_bytes.div_ceil(CHUNK_BYTES).max(1);
484 let blocks_per_chunk = CHUNK_BYTES / BLOCK_BYTES;
485 let chunk_cvs: Vec<[Wire; 8]> = (0..n_chunks)
486 .map(|c| {
487 let block_start = c * blocks_per_chunk;
488 let block_end = ((c + 1) * blocks_per_chunk).min(n_blocks);
489 let blocks: Vec<[Wire; 16]> = (block_start..block_end).map(block).collect();
490 let block_lens: Vec<Wire> = (block_start..block_end).map(block_len).collect();
491 let last_flags_extra = if n_chunks == 1 { ROOT } else { 0 };
494 blake3_chunk(
495 &builder.subcircuit(format!("blake3_chunk[{c}]")),
496 key,
497 &blocks,
498 &block_lens,
499 c as u64,
500 last_flags_extra,
501 )
502 })
503 .collect();
504
505 if n_chunks == 1 {
507 chunk_cvs[0]
508 } else {
509 blake3_tree_root(builder, key, chunk_cvs)
510 }
511}
512
513pub fn blake3_keyed_fixed_2x(
551 builder: &CircuitBuilder,
552 messages: [&[Wire]; 2],
553 len_bytes: usize,
554 keys: [&ByteVec; 2],
555) -> [[Wire; 8]; 2] {
556 let lanes = keys.map(|key| key_words(builder, key));
558 let key_2x: [Wire; 8] = std::array::from_fn(|i| pack_lanes(builder, lanes[0][i], lanes[1][i]));
559
560 let padded = messages.map(|message| padded_message_words(builder, message, len_bytes, true));
562 let n_blocks = len_bytes.div_ceil(BLOCK_BYTES).max(1);
563 let n_content_words = len_bytes.div_ceil(4);
564 let zero = builder.add_constant(Word::ZERO);
565 let block = |j: usize| -> [Wire; 16] {
566 std::array::from_fn(|i| {
567 let k = j * 16 + i;
568 if k >= n_content_words {
570 zero
571 } else {
572 pack_lanes(builder, padded[0][k], padded[1][k])
573 }
574 })
575 };
576 let block_len = |j: usize| -> Wire { dup32(builder, block_len_bytes(len_bytes, j) as u32) };
578
579 let n_chunks = len_bytes.div_ceil(CHUNK_BYTES).max(1);
580 let blocks_per_chunk = CHUNK_BYTES / BLOCK_BYTES;
581 let chunk_cvs: Vec<[Wire; 8]> = (0..n_chunks)
582 .map(|c| {
583 let block_start = c * blocks_per_chunk;
584 let block_end = ((c + 1) * blocks_per_chunk).min(n_blocks);
585 let blocks: Vec<[Wire; 16]> = (block_start..block_end).map(block).collect();
586 let block_lens: Vec<Wire> = (block_start..block_end).map(block_len).collect();
587 let last_flags_extra = if n_chunks == 1 { ROOT } else { 0 };
590 blake3_chunk_2x(
591 &builder.subcircuit(format!("blake3_chunk_2x[{c}]")),
592 key_2x,
593 &blocks,
594 &block_lens,
595 c as u64,
596 last_flags_extra,
597 )
598 })
599 .collect();
600
601 let root = if n_chunks == 1 {
602 chunk_cvs[0]
603 } else {
604 blake3_tree_root_2x(builder, key_2x, chunk_cvs)
605 };
606
607 [
609 std::array::from_fn(|i| clear_high_bits(builder, root[i], 32)),
610 std::array::from_fn(|i| builder.shr(root[i], 32)),
611 ]
612}
613
614fn blake3_chunk_2x(
634 builder: &CircuitBuilder,
635 key_2x: [Wire; 8],
636 blocks: &[[Wire; 16]],
637 block_lens: &[Wire],
638 counter: u64,
639 last_flags_extra: u32,
640) -> [Wire; 8] {
641 let n_blocks = blocks.len();
642 assert!(
643 (1..=16).contains(&n_blocks),
644 "blake3_chunk_2x: n_blocks ({n_blocks}) must be in 1..=16",
645 );
646 assert_eq!(
647 block_lens.len(),
648 n_blocks,
649 "blake3_chunk_2x: block_lens.len() ({}) must equal blocks.len() ({n_blocks})",
650 block_lens.len(),
651 );
652
653 let counter_lo = dup32(builder, counter as u32);
654 let counter_hi = dup32(builder, (counter >> 32) as u32);
655
656 let mut cv = key_2x;
657 for (j, block) in blocks.iter().enumerate() {
658 let start = if j == 0 { CHUNK_START } else { 0 };
659 let end = if j + 1 == n_blocks {
660 CHUNK_END | last_flags_extra
661 } else {
662 0
663 };
664 let flags = dup32(builder, start | end | KEYED_HASH);
665 cv = blake3_compress_2x(
666 &builder.subcircuit(format!("blake3_chunk_2x_compress[{j}]")),
667 cv,
668 *block,
669 counter_lo,
670 counter_hi,
671 block_lens[j],
672 flags,
673 );
674 }
675
676 cv
678}
679
680fn blake3_parent_2x(
684 builder: &CircuitBuilder,
685 key_2x: [Wire; 8],
686 left: [Wire; 8],
687 right: [Wire; 8],
688 is_root: bool,
689) -> [Wire; 8] {
690 let block: [Wire; 16] = std::array::from_fn(|i| if i < 8 { left[i] } else { right[i - 8] });
691 let zero = builder.add_constant(Word::ZERO);
692 let block_len = dup32(builder, BLOCK_BYTES as u32);
693 let root_flag = if is_root { ROOT } else { 0 };
694 let flags = dup32(builder, PARENT | root_flag | KEYED_HASH);
695 blake3_compress_2x(builder, key_2x, block, zero, zero, block_len, flags)
696}
697
698fn blake3_tree_root_2x(
707 builder: &CircuitBuilder,
708 key_2x: [Wire; 8],
709 chunk_cvs: Vec<[Wire; 8]>,
710) -> [Wire; 8] {
711 assert!(chunk_cvs.len() >= 2, "blake3_tree_root_2x: needs at least two chunks");
712
713 let mut level = chunk_cvs;
714 let mut depth = 0;
715 loop {
716 if level.len() == 2 {
718 return blake3_parent_2x(
719 &builder.subcircuit("blake3_tree_root_2x"),
720 key_2x,
721 level[0],
722 level[1],
723 true,
724 );
725 }
726
727 let sub = builder.subcircuit(format!("blake3_tree_level_2x[{depth}]"));
728 let n = level.len();
729 let mut next: Vec<[Wire; 8]> = Vec::with_capacity(n.div_ceil(2));
730 for p in 0..n / 2 {
731 next.push(blake3_parent_2x(&sub, key_2x, level[2 * p], level[2 * p + 1], false));
732 }
733 if n % 2 == 1 {
735 next.push(level[n - 1]);
736 }
737
738 level = next;
739 depth += 1;
740 }
741}
742
743#[cfg(test)]
744mod tests {
745 use binius_frontend::CircuitStat;
746 use hex_literal::hex;
747 use proptest::prelude::*;
748
749 use super::*;
750
751 fn bytes_to_le_words(bytes: &[u8]) -> Vec<u64> {
754 let n_words = bytes.len().div_ceil(4);
755 (0..n_words)
756 .map(|i| {
757 let mut buf = [0u8; 4];
758 let start = i * 4;
759 let end = (start + 4).min(bytes.len());
760 buf[..end - start].copy_from_slice(&bytes[start..end]);
761 u32::from_le_bytes(buf) as u64
762 })
763 .collect()
764 }
765
766 fn check_digest(input: &[u8], expected: [u8; 32]) {
773 let builder = CircuitBuilder::new();
774 let message: Vec<Wire> = (0..input.len().div_ceil(4))
776 .map(|_| builder.add_witness())
777 .collect();
778 let digest = blake3_fixed(&builder, &message, input.len());
779 let digest_out: [Wire; 8] = std::array::from_fn(|_| builder.add_inout());
781 for i in 0..8 {
782 builder.assert_eq("digest_match", digest[i], digest_out[i]);
783 }
784
785 let circuit = builder.build();
786 let mut w = circuit.new_witness_filler();
787 for (wire, word) in message.iter().zip(bytes_to_le_words(input)) {
788 w[*wire] = Word(word);
789 }
790 for i in 0..8 {
792 let bytes: [u8; 4] = expected[i * 4..i * 4 + 4].try_into().unwrap();
793 w[digest_out[i]] = Word(u32::from_le_bytes(bytes) as u64);
794 }
795 circuit
796 .populate_wire_witness(&mut w)
797 .unwrap_or_else(|e| panic!("digest disagreed with the specification vector: {e:?}"));
798 }
799
800 #[test]
801 fn draft_b1_digest_matches_spec() {
802 check_digest(
805 b"IETF",
806 hex!("83a2de1ee6f4e6ab686889248f4ec0cf4cc5709446a682ffd1cbb4d6165181e2"),
807 );
808 }
809
810 #[test]
811 fn draft_b2_digest_matches_spec() {
812 let mut input = vec![0xaau8; CHUNK_BYTES];
823 input.extend_from_slice(&[0xbbu8; CHUNK_BYTES]);
824 check_digest(
825 &input,
826 hex!("e79d2838915accd3b21bb0ba76b5edf8dc08d3d78d0db65b713f0f37ec58c346"),
827 );
828 }
829
830 proptest! {
831 #![proptest_config(ProptestConfig::with_cases(12))]
835
836 #[test]
837 fn fixed_matches_blake3_crate(input in prop::collection::vec(any::<u8>(), 0..=600)) {
838 check(&input);
840 }
841 }
842
843 fn check(input: &[u8]) {
845 let builder = CircuitBuilder::new();
846 let message: Vec<Wire> = (0..input.len().div_ceil(4))
847 .map(|_| builder.add_witness())
848 .collect();
849 let digest = blake3_fixed(&builder, &message, input.len());
850 let digest_out: [Wire; 8] = std::array::from_fn(|_| builder.add_inout());
851 for i in 0..8 {
852 builder.assert_eq("digest_match", digest[i], digest_out[i]);
853 }
854
855 let circuit = builder.build();
856 let mut w = circuit.new_witness_filler();
857 let words = bytes_to_le_words(input);
858 for (wire, word) in message.iter().zip(words.iter()) {
859 w[*wire] = Word(*word);
860 }
861
862 let expected = blake3::hash(input);
863 let expected_words: [u32; 8] = std::array::from_fn(|i| {
864 u32::from_le_bytes(expected.as_bytes()[i * 4..i * 4 + 4].try_into().unwrap())
865 });
866 for i in 0..8 {
867 w[digest_out[i]] = Word(expected_words[i] as u64);
868 }
869 circuit
870 .populate_wire_witness(&mut w)
871 .unwrap_or_else(|e| panic!("blake3_fixed failed for len_bytes={}: {e:?}", input.len()));
872 }
873
874 #[test]
875 fn empty() {
876 check(b"");
877 }
878
879 #[test]
880 fn one_byte() {
881 check(&[0x5a]);
882 }
883
884 #[test]
885 fn abc() {
886 check(b"abc");
887 }
888
889 #[test]
890 fn block_boundaries() {
891 for &len in &[
894 1usize, 63, 64, 65, 127, 128, 129, 192, 256, 257, 320, 448, 511, 512, 1023, 1024,
895 ] {
896 let input: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
897 check(&input);
898 }
899 }
900
901 fn check_keyed(input: &[u8], key: &[u8], garbage_padding: bool) {
907 let builder = CircuitBuilder::new();
908 let message: Vec<Wire> = (0..input.len().div_ceil(4))
909 .map(|_| builder.add_witness())
910 .collect();
911 let key_data: Vec<Wire> = (0..key.len().div_ceil(8))
912 .map(|_| builder.add_witness())
913 .collect();
914 let key_vec = ByteVec::new_const_len(&builder, key_data, key.len());
915 let digest = blake3_keyed_fixed(&builder, &message, input.len(), &key_vec);
916 let digest_out: [Wire; 8] = std::array::from_fn(|_| builder.add_inout());
917 for i in 0..8 {
918 builder.assert_eq("digest_match", digest[i], digest_out[i]);
919 }
920
921 let circuit = builder.build();
922 let mut w = circuit.new_witness_filler();
923 for (wire, word) in message.iter().zip(bytes_to_le_words(input)) {
924 w[*wire] = Word(word);
925 }
926 let mut key_bytes = key.to_vec();
928 key_bytes.resize(key.len().next_multiple_of(8), if garbage_padding { 0xff } else { 0 });
929 for (i, chunk) in key_bytes.chunks(8).enumerate() {
930 w[key_vec.data[i]] = Word(u64::from_le_bytes(chunk.try_into().unwrap()));
931 }
932
933 let mut padded_key = [0u8; KEY_BYTES];
934 padded_key[..key.len()].copy_from_slice(key);
935 let expected = blake3::keyed_hash(&padded_key, input);
936 for i in 0..8 {
937 let bytes: [u8; 4] = expected.as_bytes()[i * 4..i * 4 + 4].try_into().unwrap();
938 w[digest_out[i]] = Word(u32::from_le_bytes(bytes) as u64);
939 }
940 circuit.populate_wire_witness(&mut w).unwrap_or_else(|e| {
941 panic!(
942 "blake3_keyed_fixed failed for len_bytes={}, key_len={}: {e:?}",
943 input.len(),
944 key.len()
945 )
946 });
947 }
948
949 #[test]
950 fn keyed_full_length_key() {
951 let key: Vec<u8> = (0..KEY_BYTES).map(|i| (i * 7 + 3) as u8).collect();
954 for &len in &[0usize, 1, 64, 65, 192, 1024, 1025, 3072] {
955 let input: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
956 check_keyed(&input, &key, false);
957 }
958 }
959
960 #[test]
961 fn keyed_short_keys() {
962 for key_len in [0usize, 1, 4, 5, 8, 9, 16, 23, 31] {
967 let key: Vec<u8> = (0..key_len).map(|i| (i * 11 + 5) as u8).collect();
968 check_keyed(b"abc", &key, false);
969 }
970 }
971
972 #[test]
973 fn keyed_ignores_bytes_past_the_key_length() {
974 for key_len in [0usize, 1, 5, 9, 23, 31] {
977 let key: Vec<u8> = (0..key_len).map(|i| (i * 11 + 5) as u8).collect();
978 check_keyed(b"abc", &key, true);
979 }
980 }
981
982 proptest! {
983 #![proptest_config(ProptestConfig::with_cases(12))]
985
986 #[test]
987 fn keyed_matches_blake3_crate(
988 input in prop::collection::vec(any::<u8>(), 0..=300),
989 key in prop::collection::vec(any::<u8>(), 0..=KEY_BYTES),
990 ) {
991 check_keyed(&input, &key, true);
992 }
993 }
994
995 #[test]
996 #[should_panic(expected = "key length (33) exceeds 32")]
997 fn keyed_rejects_an_oversized_key() {
998 let builder = CircuitBuilder::new();
999 let data: Vec<Wire> = (0..5).map(|_| builder.add_witness()).collect();
1000 let key = ByteVec::new_const_len(&builder, data, KEY_BYTES + 1);
1001 blake3_keyed_fixed(&builder, &[], 0, &key);
1002 }
1003
1004 #[test]
1005 #[should_panic(expected = "key length must be fixed at circuit construction time")]
1006 fn keyed_rejects_a_runtime_length_key() {
1007 let builder = CircuitBuilder::new();
1008 let key = ByteVec::new_witness(&builder, 4);
1011 blake3_keyed_fixed(&builder, &[], 0, &key);
1012 }
1013
1014 fn check_keyed_2x(inputs: [&[u8]; 2], keys: [&[u8]; 2], garbage: bool) {
1020 let len_bytes = inputs[0].len();
1021 assert_eq!(len_bytes, inputs[1].len(), "the two messages must have equal length");
1022
1023 let builder = CircuitBuilder::new();
1024 let messages: [Vec<Wire>; 2] = std::array::from_fn(|_| {
1025 (0..len_bytes.div_ceil(4))
1026 .map(|_| builder.add_witness())
1027 .collect()
1028 });
1029 let key_vecs: [ByteVec; 2] = std::array::from_fn(|l| {
1030 let data = (0..keys[l].len().div_ceil(8))
1031 .map(|_| builder.add_witness())
1032 .collect();
1033 ByteVec::new_const_len(&builder, data, keys[l].len())
1034 });
1035 let digests = blake3_keyed_fixed_2x(
1036 &builder,
1037 [&messages[0], &messages[1]],
1038 len_bytes,
1039 [&key_vecs[0], &key_vecs[1]],
1040 );
1041 let digests_out: [[Wire; 8]; 2] =
1043 std::array::from_fn(|_| std::array::from_fn(|_| builder.add_inout()));
1044 for l in 0..2 {
1045 for i in 0..8 {
1046 builder.assert_eq("digest_match_2x", digests[l][i], digests_out[l][i]);
1047 }
1048 }
1049
1050 let circuit = builder.build();
1051 let mut w = circuit.new_witness_filler();
1052 let dirty_high = if garbage { 0xFFFF_FFFF_0000_0000 } else { 0 };
1054 for l in 0..2 {
1055 for (wire, word) in messages[l].iter().zip(bytes_to_le_words(inputs[l])) {
1056 w[*wire] = Word(word | dirty_high);
1057 }
1058 let mut key_bytes = keys[l].to_vec();
1060 key_bytes.resize(keys[l].len().next_multiple_of(8), if garbage { 0xff } else { 0 });
1061 for (i, chunk) in key_bytes.chunks(8).enumerate() {
1062 w[key_vecs[l].data[i]] = Word(u64::from_le_bytes(chunk.try_into().unwrap()));
1063 }
1064
1065 let mut padded_key = [0u8; KEY_BYTES];
1066 padded_key[..keys[l].len()].copy_from_slice(keys[l]);
1067 let expected = blake3::keyed_hash(&padded_key, inputs[l]);
1068 for i in 0..8 {
1069 let bytes: [u8; 4] = expected.as_bytes()[i * 4..i * 4 + 4].try_into().unwrap();
1070 w[digests_out[l][i]] = Word(u32::from_le_bytes(bytes) as u64);
1071 }
1072 }
1073 circuit.populate_wire_witness(&mut w).unwrap_or_else(|e| {
1074 panic!(
1075 "blake3_keyed_fixed_2x failed for len_bytes={len_bytes}, key_lens=({}, {}): {e:?}",
1076 keys[0].len(),
1077 keys[1].len()
1078 )
1079 });
1080 }
1081
1082 fn distinct_pair(len: usize) -> ([Vec<u8>; 2], [Vec<u8>; 2]) {
1084 let messages =
1085 std::array::from_fn(|l| (0..len).map(|i| (i * 37 + 1 + l * 91) as u8).collect());
1086 let keys =
1087 std::array::from_fn(|l| (0..KEY_BYTES).map(|i| (i * 7 + 3 + l * 53) as u8).collect());
1088 (messages, keys)
1089 }
1090
1091 #[test]
1092 fn keyed_2x_block_boundaries() {
1093 for &len in &[
1096 0usize, 1, 3, 63, 64, 65, 127, 128, 129, 192, 320, 511, 512, 1023, 1024,
1097 ] {
1098 let (messages, keys) = distinct_pair(len);
1099 check_keyed_2x([&messages[0], &messages[1]], [&keys[0], &keys[1]], false);
1100 }
1101 }
1102
1103 #[test]
1104 fn keyed_2x_multi_chunk() {
1105 for &len in &[1025usize, 2048, 2049, 3072, 5121, 7168, 8192, 9217] {
1109 let (messages, keys) = distinct_pair(len);
1110 check_keyed_2x([&messages[0], &messages[1]], [&keys[0], &keys[1]], false);
1111 }
1112 }
1113
1114 #[test]
1115 fn keyed_2x_lanes_take_different_key_lengths() {
1116 for (len0, len1) in [(0usize, KEY_BYTES), (1, 31), (5, 9), (23, 4), (16, 16)] {
1119 let key0: Vec<u8> = (0..len0).map(|i| (i * 11 + 5) as u8).collect();
1120 let key1: Vec<u8> = (0..len1).map(|i| (i * 13 + 2) as u8).collect();
1121 check_keyed_2x([b"abc", b"xyz"], [&key0, &key1], false);
1122 }
1123 }
1124
1125 #[test]
1126 fn keyed_2x_lanes_are_independent() {
1127 for &len in &[1usize, 64, 65, 1025] {
1133 let (messages, keys) = distinct_pair(len);
1134 check_keyed_2x([&messages[0], &messages[1]], [&keys[0], &keys[1]], true);
1135 }
1136 }
1137
1138 #[test]
1139 fn keyed_2x_lanes_hashing_the_same_input_agree() {
1140 let key: Vec<u8> = (0..KEY_BYTES).map(|i| (i * 7 + 3) as u8).collect();
1143 let message: Vec<u8> = (0..200).map(|i| (i * 37 + 1) as u8).collect();
1144 check_keyed_2x([&message, &message], [&key, &key], true);
1145 }
1146
1147 proptest! {
1148 #![proptest_config(ProptestConfig::with_cases(12))]
1150
1151 #[test]
1152 fn keyed_2x_matches_blake3_crate(
1153 len in 0usize..=300,
1154 bytes0 in prop::collection::vec(any::<u8>(), 300),
1155 bytes1 in prop::collection::vec(any::<u8>(), 300),
1156 key0 in prop::collection::vec(any::<u8>(), 0..=KEY_BYTES),
1157 key1 in prop::collection::vec(any::<u8>(), 0..=KEY_BYTES),
1158 ) {
1159 check_keyed_2x([&bytes0[..len], &bytes1[..len]], [&key0, &key1], true);
1161 }
1162 }
1163
1164 #[test]
1165 #[should_panic(expected = "key length (33) exceeds 32")]
1166 fn keyed_2x_rejects_an_oversized_key() {
1167 let builder = CircuitBuilder::new();
1168 let good = ByteVec::new_const_len(&builder, vec![], 0);
1169 let data: Vec<Wire> = (0..5).map(|_| builder.add_witness()).collect();
1170 let oversized = ByteVec::new_const_len(&builder, data, KEY_BYTES + 1);
1171 blake3_keyed_fixed_2x(&builder, [&[], &[]], 0, [&good, &oversized]);
1172 }
1173
1174 #[test]
1175 #[should_panic(expected = "key length must be fixed at circuit construction time")]
1176 fn keyed_2x_rejects_a_runtime_length_key() {
1177 let builder = CircuitBuilder::new();
1178 let good = ByteVec::new_const_len(&builder, vec![], 0);
1179 let runtime = ByteVec::new_witness(&builder, 4);
1182 blake3_keyed_fixed_2x(&builder, [&[], &[]], 0, [&good, &runtime]);
1183 }
1184
1185 fn and_counts(len_bytes: usize) -> (usize, usize) {
1189 let key = |builder: &CircuitBuilder| {
1190 let data: Vec<Wire> = (0..KEY_BYTES / 8).map(|_| builder.add_witness()).collect();
1191 ByteVec::new_const_len(builder, data, KEY_BYTES)
1192 };
1193 let message = |builder: &CircuitBuilder| -> Vec<Wire> {
1194 (0..len_bytes.div_ceil(4))
1195 .map(|_| builder.add_witness())
1196 .collect()
1197 };
1198 let expose = |builder: &CircuitBuilder, digest: [Wire; 8]| {
1200 for word in digest {
1201 let out = builder.add_inout();
1202 builder.assert_eq("digest", word, out);
1203 }
1204 };
1205
1206 let single = {
1207 let builder = CircuitBuilder::new();
1208 for _ in 0..2 {
1209 let digest =
1210 blake3_keyed_fixed(&builder, &message(&builder), len_bytes, &key(&builder));
1211 expose(&builder, digest);
1212 }
1213 CircuitStat::collect(&builder.build()).n_and_constraints
1214 };
1215 let paired = {
1216 let builder = CircuitBuilder::new();
1217 let (m0, m1) = (message(&builder), message(&builder));
1218 let (k0, k1) = (key(&builder), key(&builder));
1219 for digest in blake3_keyed_fixed_2x(&builder, [&m0, &m1], len_bytes, [&k0, &k1]) {
1220 expose(&builder, digest);
1221 }
1222 CircuitStat::collect(&builder.build()).n_and_constraints
1223 };
1224 (single, paired)
1225 }
1226
1227 #[test]
1228 fn keyed_2x_never_costs_more_than_two_single_hashes() {
1229 for &len in &[0usize, 1, 64, 65, 128, 192, 1024, 2048] {
1236 let (single, paired) = and_counts(len);
1237 assert!(
1238 paired <= single,
1239 "blake3_keyed_fixed_2x costs {paired} AND constraints at len={len}, \
1240 against {single} for two single hashes"
1241 );
1242 }
1243
1244 let (single, paired) = and_counts(BLOCK_BYTES);
1249 assert!(
1250 paired < single,
1251 "a one-block pair costs {paired} AND constraints, no less than the {single} two \
1252 single hashes cost"
1253 );
1254 }
1255
1256 #[test]
1257 #[should_panic(expected = "message.len() (2) must equal len_bytes.div_ceil(4) (3)")]
1258 fn keyed_2x_rejects_a_message_of_the_wrong_length() {
1259 let builder = CircuitBuilder::new();
1261 let key = ByteVec::new_const_len(&builder, vec![], 0);
1262 let long: Vec<Wire> = (0..3).map(|_| builder.add_witness()).collect();
1263 let short: Vec<Wire> = (0..2).map(|_| builder.add_witness()).collect();
1264 blake3_keyed_fixed_2x(&builder, [&long, &short], 9, [&key, &key]);
1265 }
1266
1267 #[test]
1268 fn multi_chunk() {
1269 for &len in &[
1273 1025usize, 2048, 2049, 3072, 4096, 5121, 7168, 8192, 9217, 10240, ] {
1284 let input: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
1285 check(&input);
1286 }
1287 }
1288
1289 fn check_with_compress_chip(input: &[u8]) {
1295 let builder = CircuitBuilder::new();
1296 builder.register_chip(Blake3Compress2x, &[]);
1297
1298 let message: Vec<Wire> = (0..input.len().div_ceil(4))
1299 .map(|_| builder.add_witness())
1300 .collect();
1301 let digest = blake3_fixed(&builder, &message, input.len());
1302 let digest_out: [Wire; 8] = std::array::from_fn(|_| builder.add_inout());
1303 for i in 0..8 {
1304 builder.assert_eq("digest_match", digest[i], digest_out[i]);
1305 }
1306
1307 let circuit = builder.build_m4();
1308 circuit.validate().unwrap();
1309 let cs = circuit.to_constraint_system();
1310 cs.validate().unwrap();
1311
1312 let expected = blake3::hash(input);
1313 let expected_words: [u32; 8] = std::array::from_fn(|i| {
1314 u32::from_le_bytes(expected.as_bytes()[i * 4..i * 4 + 4].try_into().unwrap())
1315 });
1316
1317 let witness = circuit
1318 .generate_witness(|w| {
1319 for (wire, word) in message.iter().zip(bytes_to_le_words(input)) {
1320 w[*wire] = Word(word);
1321 }
1322 for i in 0..8 {
1323 w[digest_out[i]] = Word(expected_words[i] as u64);
1324 }
1325 })
1326 .unwrap_or_else(|e| panic!("blake3_fixed failed for len_bytes={}: {e:?}", input.len()));
1327
1328 witness.verify(&cs).unwrap();
1329 }
1330
1331 #[test]
1339 fn a_registered_chip_serves_every_compression() {
1340 for &len in &[128usize, 320, 1025, 5121] {
1341 let input: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
1342 check_with_compress_chip(&input);
1343 }
1344 }
1345
1346 #[test]
1349 fn a_chip_no_compression_reaches_leaves_an_uncalled_chip() {
1350 let builder = CircuitBuilder::new();
1351 builder.register_chip(Blake3Compress2x, &[]);
1352 blake3_fixed(&builder, &[builder.add_witness()], 4);
1353
1354 let error = builder.build_m4().validate().unwrap_err();
1355 assert!(matches!(error, binius_frontend::CircuitM4Error::NeverCalled { .. }), "{error:?}");
1356 }
1357}