1pub mod compress;
4
5use binius_core::word::Word;
6use binius_frontend::{CircuitBuilder, Wire};
7pub use compress::{
8 Sha256Compress2x, State, populate_message_block, ref_compress, sha256_compress,
9 sha256_compress_2x, sha256_compress_2x_seq,
10};
11
12use crate::{
13 bytes::swap_bytes_32,
14 fixed_byte_vec::ByteVec,
15 multiplexer::{multi_wire_multiplex, single_wire_multiplex},
16 util::clear_high_bits,
17};
18
19pub fn sha256_fixed(builder: &CircuitBuilder, message: &[Wire], len_bytes: usize) -> [Wire; 8] {
54 assert_eq!(
56 message.len(),
57 len_bytes.div_ceil(4),
58 "message.len() ({}) must equal len_bytes.div_ceil(4) ({})",
59 message.len(),
60 len_bytes.div_ceil(4)
61 );
62
63 assert!(
65 (len_bytes as u64)
66 .checked_mul(8)
67 .is_some_and(|bits| bits <= u32::MAX as u64),
68 "Message length in bits must fit in 32 bits"
69 );
70
71 let n_blocks = (len_bytes + 9).div_ceil(64);
76 let n_padded_words = n_blocks * 16; let mut padded_message = Vec::with_capacity(n_padded_words);
80
81 let n_message_words = len_bytes / 4;
83 let boundary_bytes = len_bytes % 4;
84
85 padded_message.extend_from_slice(&message[0..n_message_words]);
87
88 if boundary_bytes > 0 {
90 let last_word = message[n_message_words];
92
93 let shift_amount = (4 - boundary_bytes) * 8;
95 let mask = builder.add_constant(Word((0xFFFFFFFFu64 >> shift_amount) << shift_amount));
96 let masked = builder.band(last_word, mask);
97
98 let delimiter_shift = (3 - boundary_bytes) * 8;
100 let delimiter = builder.add_constant(Word(0x80u64 << delimiter_shift));
101 let boundary_word = builder.bxor(masked, delimiter);
102
103 padded_message.push(boundary_word);
104 } else {
105 padded_message.push(builder.add_constant(Word(0x80000000)));
107 }
108
109 let zero = builder.add_constant(Word::ZERO);
111 padded_message.resize(n_padded_words - 2, zero);
112
113 padded_message.push(zero); let bitlen = (len_bytes as u64) * 8;
116 padded_message.push(builder.add_constant(Word(bitlen)));
117
118 let blocks: Vec<[Wire; 16]> = padded_message
123 .chunks_exact(16)
124 .map(|block| block.try_into().unwrap())
125 .collect();
126 let n_blocks = blocks.len();
127
128 let mut state = State::iv(builder);
129 let mut block_idx = 0;
130 while block_idx + 1 < n_blocks {
138 state = sha256_compress_2x_seq(
140 &builder.subcircuit(format!("sha256_fixed_compress[{block_idx}..{}]", block_idx + 2)),
141 state,
142 [blocks[block_idx], blocks[block_idx + 1]],
143 );
144 block_idx += 2;
145 }
146 if block_idx < n_blocks {
147 let sub = builder.subcircuit(format!("sha256_fixed_compress[{block_idx}]"));
153 state = sha256_compress(&sub, state, blocks[block_idx]);
154 }
155
156 if n_blocks % 2 == 1 {
164 return state.0;
165 }
166 std::array::from_fn(|i| clear_high_bits(builder, state.0[i], 32))
167}
168
169pub fn sha256_varlen(builder: &CircuitBuilder, message: &ByteVec) -> [Wire; 4] {
203 let len_bytes = message.len_bytes;
209 assert!(
210 message.data.len() << Word::LOG_BITS <= u32::MAX as usize,
211 "length of message in bits must fit within 32 bits"
212 );
213
214 let max_len_bytes = message.data.len() << Word::LOG_BYTES;
215 let n_blocks = (message.data.len() + 2).div_ceil(8);
216 let n_words: usize = n_blocks << 4; let too_long = builder.icmp_ugt(len_bytes, builder.add_constant_64(max_len_bytes as u64));
219 builder.assert_false("len_check", too_long);
220
221 let mut message_be: Vec<Wire> = Vec::with_capacity(message.data.len() * 2);
226 for &word in &message.data {
227 let swapped = swap_bytes_32(builder, word);
228 message_be.push(clear_high_bits(builder, swapped, 32));
229 message_be.push(builder.shr(swapped, 32));
230 }
231
232 let zero = builder.add_constant(Word::ZERO);
234 let w_bd = builder.shr(len_bytes, 2);
235 let len_mod_4 = builder.band(len_bytes, builder.add_constant_zx_8(3));
236 let bitlen = builder.shl(len_bytes, 3);
237
238 let (sum, _carry) = builder.iadd(len_bytes, builder.add_constant_64(8));
240 let end_block_index = builder.shr(sum, 6);
241
242 let boundary_message_word = single_wire_multiplex(builder, &message_be, w_bd);
250 let candidates: Vec<Wire> = (0..4)
251 .map(|i| {
252 let mask = builder.add_constant_64((0xFFFFFFFFu64 << ((4 - i) << 3)) & 0xFFFFFFFF);
253 let padding_byte = builder.add_constant_64(0x80000000u64 >> (i << 3));
254 let message_low = builder.band(boundary_message_word, mask);
255 builder.bxor(message_low, padding_byte)
256 })
257 .collect();
258 let boundary_word = single_wire_multiplex(builder, &candidates, len_mod_4);
259
260 let padded_message: Vec<Wire> = (0..n_words)
268 .map(|word_index| {
269 let block_index = word_index >> 4;
270 let column_index = word_index & 15;
271
272 let is_message_word =
273 builder.icmp_ult(builder.add_constant_64(word_index as u64), w_bd);
274 let is_boundary_word =
275 builder.icmp_eq(builder.add_constant_64(word_index as u64), w_bd);
276 let is_length_block =
277 builder.icmp_eq(builder.add_constant_64(block_index as u64), end_block_index);
278
279 let msg_word = if word_index < message_be.len() {
283 message_be[word_index]
284 } else {
285 zero
286 };
287
288 let past_word = if column_index == 15 {
292 builder.select(is_length_block, bitlen, zero)
293 } else {
294 zero
295 };
296
297 let boundary_or_past = builder.select(is_boundary_word, boundary_word, past_word);
298 builder.select(is_message_word, msg_word, boundary_or_past)
299 })
300 .collect();
301
302 let mut states = Vec::with_capacity(n_blocks + 1);
311 states.push(State::iv(builder));
312 let mk_m = |block_no: usize| -> [Wire; 16] {
313 padded_message[block_no << 4..(block_no + 1) << 4]
314 .try_into()
315 .unwrap()
316 };
317 let mut block_no = 0;
318 while block_no + 1 < n_blocks {
319 let out = sha256_compress_2x_seq(
320 &builder.subcircuit(format!("compress[{block_no}..{}]", block_no + 2)),
321 states[block_no],
322 [mk_m(block_no), mk_m(block_no + 1)],
323 );
324 let state_first = State::new(std::array::from_fn(|i| builder.shr(out.0[i], 32)));
327 let state_second =
328 State::new(std::array::from_fn(|i| clear_high_bits(builder, out.0[i], 32)));
329 states.push(state_first);
330 states.push(state_second);
331 block_no += 2;
332 }
333 if block_no < n_blocks {
335 let state_out = sha256_compress(
336 &builder.subcircuit(format!("compress[{block_no}]")),
337 states[block_no],
338 mk_m(block_no),
339 );
340 states.push(state_out);
341 }
342
343 let block_digests: Vec<[Wire; 4]> = states[1..].iter().map(|s| s.pack_4x64b(builder)).collect();
349 let inputs: Vec<&[Wire]> = block_digests.iter().map(|d| &d[..]).collect();
350 let final_digest_vec = multi_wire_multiplex(builder, &inputs, end_block_index);
351 final_digest_vec.try_into().unwrap()
352}
353
354#[cfg(test)]
355mod tests {
356 use std::array;
357
358 use binius_core::Word;
359 use binius_frontend::{CircuitBuilder, CircuitStat, Wire};
360 use hex_literal::hex;
361 use sha2::Digest;
362
363 use super::*;
364
365 fn test_sha256_varlen_with_input(
371 message_bytes: &[u8],
372 expected_digest: [u8; 32],
373 max_len_bytes: usize,
374 ) {
375 assert!(message_bytes.len() <= max_len_bytes);
376
377 let builder = CircuitBuilder::new();
378 let max_len_words = max_len_bytes.div_ceil(8);
379 let input = ByteVec::new_inout(&builder, max_len_words);
380 let expected_digest_wires: [Wire; 4] = array::from_fn(|_| builder.add_witness());
381
382 let computed_digest = sha256_varlen(&builder, &input);
383 for i in 0..4 {
384 builder.assert_eq(format!("digest[{i}]"), computed_digest[i], expected_digest_wires[i]);
385 }
386
387 let circuit = builder.build();
388 let cs = circuit.constraint_system();
389 let mut w = circuit.new_witness_filler();
390
391 input.populate_data(&mut w, message_bytes);
392 input.populate_len_bytes(&mut w, message_bytes.len());
393
394 for (i, bytes) in expected_digest.chunks(8).enumerate() {
395 let word = u64::from_be_bytes(bytes.try_into().unwrap());
396 w[expected_digest_wires[i]] = Word(word);
397 }
398
399 circuit.populate_wire_witness(&mut w).unwrap();
400 cs.verify(&w.into_value_vec()).unwrap();
401 }
402
403 #[test]
404 fn test_sha256_varlen_empty() {
405 test_sha256_varlen_with_input(
406 b"",
407 hex!("e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"),
408 64,
409 );
410 }
411
412 #[test]
413 fn test_sha256_varlen_abc() {
414 test_sha256_varlen_with_input(
415 b"abc",
416 hex!("ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad"),
417 64,
418 );
419 }
420
421 #[test]
422 fn test_sha256_varlen_two_block_boundary() {
423 test_sha256_varlen_with_input(
425 &[b'a'; 56],
426 hex!("b35439a4ac6f0948b6d6f9e3c6af0f5f590ce20f1bde7090ef7970686ec6738a"),
427 128,
428 );
429 }
430
431 #[test]
432 fn test_sha256_varlen_various_sizes() {
433 use rand::prelude::*;
434
435 let sizes: Vec<usize> = vec![
438 0, 1, 3, 4, 5, 31, 32, 33, 55, 56, 63, 64, 65, 119, 120, 128, 256,
439 ];
440 let max_len_bytes = 320;
442
443 let mut rng = StdRng::seed_from_u64(0);
444 for size in sizes {
445 let mut message = vec![0u8; size];
446 rng.fill(&mut message[..]);
447
448 let expected = sha2::Sha256::digest(&message);
449 let expected_bytes: [u8; 32] = expected.into();
450
451 test_sha256_varlen_with_input(&message, expected_bytes, max_len_bytes);
452 }
453 }
454
455 #[test]
456 fn test_sha256_varlen_length_exceeds_max_rejection() {
457 let builder = CircuitBuilder::new();
460 let max_len_bytes = 64usize;
461 let max_len_words = max_len_bytes.div_ceil(8);
462 let input = ByteVec::new_inout(&builder, max_len_words);
463 let _ = sha256_varlen(&builder, &input);
464
465 let circuit = builder.build();
466 let mut w = circuit.new_witness_filler();
467 input.populate_data(&mut w, b"");
468 w[input.len_bytes] = Word(max_len_bytes as u64 + 1);
470 assert!(circuit.populate_wire_witness(&mut w).is_err());
471 }
472
473 fn test_sha256_fixed_with_input(message: &[u8], expected_bytes: [u8; 32]) {
475 let b = CircuitBuilder::new();
476
477 let n_words = message.len().div_ceil(4);
479 let mut message_wires = Vec::new();
480
481 for word_idx in 0..n_words {
482 let mut packed = 0u32;
483 for i in 0..4 {
484 let byte_idx = word_idx * 4 + i;
485 if byte_idx < message.len() {
486 packed |= (message[byte_idx] as u32) << (24 - i * 8);
487 }
488 }
489 message_wires.push(b.add_constant(Word(packed as u64)));
490 }
491
492 let expected_digest_wires = array::from_fn::<_, 8, _>(|_| b.add_inout());
494
495 let computed_digest = sha256_fixed(&b, &message_wires, message.len());
497
498 for i in 0..8 {
500 b.assert_eq(format!("digest[{}]", i), computed_digest[i], expected_digest_wires[i]);
501 }
502
503 let circuit = b.build();
504 let cs = circuit.constraint_system();
505 let mut w = circuit.new_witness_filler();
506
507 for i in 0..8 {
509 let mut word = 0u32;
510 for j in 0..4 {
511 word |= (expected_bytes[i * 4 + j] as u32) << (24 - j * 8);
512 }
513 w[expected_digest_wires[i]] = Word(word as u64);
514 }
515
516 circuit.populate_wire_witness(&mut w).unwrap();
517 cs.verify(&w.into_value_vec()).unwrap();
518 }
519
520 #[test]
521 #[should_panic(expected = "message.len() (1) must equal len_bytes.div_ceil(4) (2)")]
522 fn test_sha256_fixed_with_insufficient_wires() {
523 use super::sha256_fixed;
524 let builder = CircuitBuilder::new();
525
526 let message_wires: Vec<Wire> = vec![builder.add_witness()];
528
529 sha256_fixed(&builder, &message_wires, 5);
531 }
532
533 #[test]
534 fn test_sha256_fixed_various_sizes() {
535 use rand::prelude::*;
536
537 let sizes = vec![
539 0, 1, 3, 4, 5, 31, 32, 33, 55, 56, 63, 64, 65, 119, 120, 128, 256, ];
557
558 let mut rng = StdRng::seed_from_u64(0);
559
560 for size in sizes {
561 let mut message = vec![0u8; size];
563 rng.fill(&mut message[..]);
564
565 let expected = sha2::Sha256::digest(&message);
567 let expected_bytes: [u8; 32] = expected.into();
568
569 test_sha256_fixed_with_input(&message, expected_bytes);
571 }
572 }
573
574 #[test]
575 fn every_block_costs_the_same() {
576 const AND_PER_BLOCK: usize = 364;
578
579 for (len_bytes, n_blocks) in [(32, 1), (64, 2), (128, 3), (192, 4), (256, 5)] {
581 let b = CircuitBuilder::new();
582 let message: Vec<Wire> = (0..len_bytes / 4).map(|_| b.add_inout()).collect();
583
584 for wire in sha256_fixed(&b, &message, len_bytes) {
586 let public = b.add_inout();
587 b.assert_eq("digest", wire, public);
588 }
589
590 let stat = CircuitStat::collect(&b.build());
594 assert_eq!(stat.n_and_constraints, n_blocks * AND_PER_BLOCK, "{len_bytes} bytes");
595 }
596 }
597
598 fn check_fixed_with_compress_chip(message: &[u8]) {
604 let b = CircuitBuilder::new();
605 b.register_chip(Sha256Compress2x, &[]);
606
607 let n_words = message.len().div_ceil(4);
608 let message_wires: Vec<Wire> = (0..n_words).map(|_| b.add_witness()).collect();
609 let computed_digest = sha256_fixed(&b, &message_wires, message.len());
610 let digest_out: [Wire; 8] = array::from_fn(|_| b.add_inout());
611 for i in 0..8 {
612 b.assert_eq(format!("digest[{i}]"), computed_digest[i], digest_out[i]);
613 }
614
615 let circuit = b.build_m4();
616 circuit.validate().unwrap();
617 let cs = circuit.to_constraint_system();
618 cs.validate().unwrap();
619
620 let expected: [u8; 32] = sha2::Sha256::digest(message).into();
621
622 let witness = circuit
623 .generate_witness(|w| {
624 for (word_idx, wire) in message_wires.iter().enumerate() {
625 let mut packed = 0u32;
626 for i in 0..4 {
627 let byte_idx = word_idx * 4 + i;
628 if byte_idx < message.len() {
629 packed |= (message[byte_idx] as u32) << (24 - i * 8);
630 }
631 }
632 w[*wire] = Word(packed as u64);
633 }
634 for i in 0..8 {
635 let mut word = 0u32;
636 for j in 0..4 {
637 word |= (expected[i * 4 + j] as u32) << (24 - j * 8);
638 }
639 w[digest_out[i]] = Word(word as u64);
640 }
641 })
642 .unwrap_or_else(|e| {
643 panic!("sha256_fixed failed for len_bytes={}: {e:?}", message.len())
644 });
645
646 witness.verify(&cs).unwrap();
647 }
648
649 #[test]
654 fn a_registered_chip_serves_every_paired_compression() {
655 for &len in &[64usize, 128, 192, 300] {
656 let message: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
657 check_fixed_with_compress_chip(&message);
658 }
659 }
660
661 #[test]
664 fn a_chip_no_paired_compression_reaches_leaves_an_uncalled_chip() {
665 let b = CircuitBuilder::new();
666 b.register_chip(Sha256Compress2x, &[]);
667 sha256_fixed(&b, &[b.add_witness()], 4);
668
669 let error = b.build_m4().validate().unwrap_err();
670 assert!(matches!(error, binius_frontend::CircuitM4Error::NeverCalled { .. }), "{error:?}");
671 }
672}