1pub mod compress;
3
4use binius_core::word::Word;
5use binius_frontend::{CircuitBuilder, Wire};
6pub use compress::{Sha512Compress, State, compress, pack_message_block, ref_compress};
7
8use crate::{
9 bytes::swap_bytes,
10 fixed_byte_vec::ByteVec,
11 multiplexer::{multi_wire_multiplex, single_wire_multiplex},
12};
13
14pub fn sha512_fixed(builder: &CircuitBuilder, message: &[Wire], len_bytes: usize) -> [Wire; 8] {
49 assert_eq!(
51 message.len(),
52 len_bytes.div_ceil(8),
53 "message.len() ({}) must equal len_bytes.div_ceil(8) ({})",
54 message.len(),
55 len_bytes.div_ceil(8)
56 );
57
58 assert!(
60 (len_bytes as u64).checked_mul(8).is_some(),
61 "Message length in bits must fit in 64 bits"
62 );
63
64 let n_blocks = (len_bytes + 17).div_ceil(128);
69 let n_padded_words = n_blocks * 16; let mut padded_message = Vec::with_capacity(n_padded_words);
73 if len_bytes.is_multiple_of(8) {
74 padded_message.extend_from_slice(message);
76 padded_message.push(builder.add_constant(Word(0x8000000000000000)));
78 } else {
79 padded_message.extend_from_slice(&message[..message.len() - 1]);
81
82 let last_idx = message.len() - 1;
84 let boundary_byte_in_word = len_bytes % 8;
85
86 let shift_amount = (8 - boundary_byte_in_word) * 8;
89 let shifted_right = builder.shr(message[last_idx], shift_amount as u32);
90 let shifted_back = builder.shl(shifted_right, shift_amount as u32);
91
92 let delimiter_shift = (7 - boundary_byte_in_word) * 8;
94 let delimiter = builder.add_constant(Word(0x80u64 << delimiter_shift));
95 let boundary_word = builder.bxor(shifted_back, delimiter);
96 padded_message.push(boundary_word);
97 }
98
99 let zero = builder.add_constant(Word::ZERO);
101 padded_message.resize(n_padded_words - 2, zero);
102
103 padded_message.push(zero);
105
106 let bitlen = (len_bytes as u64) * 8;
107 padded_message.push(builder.add_constant(Word(bitlen))); let state_out = padded_message.chunks(16).enumerate().fold(
111 State::iv(builder),
112 |state, (block_idx, block)| {
113 let block_message: [Wire; 16] = block
114 .try_into()
115 .expect("padded_message.len() must be divisible by 16");
116 compress(
117 &builder.subcircuit(format!("sha512_fixed_compress[{}]", block_idx)),
118 state,
119 block_message,
120 )
121 },
122 );
123
124 state_out.0
126}
127
128pub fn sha512_varlen(builder: &CircuitBuilder, message: &ByteVec) -> [Wire; 8] {
154 let len_bytes = message.len_bytes;
160 assert!(
161 message.data.len() << Word::LOG_BITS <= u64::MAX as usize,
162 "length of message in bits must fit in 64-bit wire"
163 );
164
165 let max_len_bytes = message.data.len() << Word::LOG_BYTES;
166 let n_blocks = (message.data.len() + 3).div_ceil(16);
167 let n_words: usize = n_blocks << 4; let too_long = builder.icmp_ugt(len_bytes, builder.add_constant_64(max_len_bytes as u64));
170 builder.assert_false("len_check", too_long);
171
172 let message_be: Vec<Wire> = message
174 .data
175 .iter()
176 .map(|&word| swap_bytes(builder, word))
177 .collect();
178
179 let zero = builder.add_constant(Word::ZERO);
181 let w_bd = builder.shr(len_bytes, 3);
182 let len_mod_8 = builder.band(len_bytes, builder.add_constant_zx_8(7));
183 let bitlen = builder.shl(len_bytes, 3);
184
185 let (sum, _carry) = builder.iadd(len_bytes, builder.add_constant_64(16));
187 let end_block_index = builder.shr(sum, 7);
188
189 let boundary_message_word = single_wire_multiplex(builder, &message_be, w_bd);
197 let candidates: Vec<Wire> = (0..8)
198 .map(|i| {
199 let mask = builder.add_constant_64(0xFFFFFFFFFFFFFF00 << ((7 - i) << 3));
200 let padding_byte = builder.add_constant_64(0x8000000000000000 >> (i << 3));
201 let message_low = builder.band(boundary_message_word, mask);
202 builder.bxor(message_low, padding_byte)
203 })
204 .collect();
205 let boundary_word = single_wire_multiplex(builder, &candidates, len_mod_8);
206
207 let padded_message: Vec<Wire> = (0..n_words)
215 .map(|word_index| {
216 let block_index = word_index >> 4;
217 let column_index = word_index & 15;
218
219 let is_message_word =
220 builder.icmp_ult(builder.add_constant_64(word_index as u64), w_bd);
221 let is_boundary_word =
222 builder.icmp_eq(builder.add_constant_64(word_index as u64), w_bd);
223 let is_length_block =
224 builder.icmp_eq(builder.add_constant_64(block_index as u64), end_block_index);
225
226 let msg_word = if word_index < message_be.len() {
230 message_be[word_index]
231 } else {
232 zero
233 };
234
235 let past_word = if column_index == 15 {
239 builder.select(is_length_block, bitlen, zero)
240 } else {
241 zero
242 };
243
244 let boundary_or_past = builder.select(is_boundary_word, boundary_word, past_word);
245 builder.select(is_message_word, msg_word, boundary_or_past)
246 })
247 .collect();
248
249 let mut states = Vec::with_capacity(n_blocks + 1);
253 states.push(State::iv(builder));
254 for block_no in 0..n_blocks {
255 let m: [Wire; 16] = padded_message[block_no << 4..(block_no + 1) << 4]
256 .try_into()
257 .unwrap();
258 let state_out =
259 compress(&builder.subcircuit(format!("compress[{block_no}]")), states[block_no], m);
260 states.push(state_out);
261 }
262
263 let inputs: Vec<&[Wire]> = states[1..].iter().map(|s| &s.0[..]).collect();
267 let final_digest_vec = multi_wire_multiplex(builder, &inputs, end_block_index);
268 final_digest_vec.try_into().unwrap()
269}
270
271#[cfg(test)]
272mod tests {
273 use binius_core::Word;
274 use binius_frontend::{CircuitBuilder, Wire};
275 use hex_literal::hex;
276 use sha2::Digest;
277
278 use super::{Sha512Compress, sha512_fixed, sha512_varlen};
279 use crate::fixed_byte_vec::ByteVec;
280
281 fn test_sha512_fixed_with_input(message_bytes: &[u8], expected_digest: [u8; 64]) {
285 let builder = CircuitBuilder::new();
286
287 let n_words = message_bytes.len().div_ceil(8);
289 let message_wires: Vec<Wire> = (0..n_words).map(|_| builder.add_witness()).collect();
290
291 let expected_digest_wires: [Wire; 8] = std::array::from_fn(|_| builder.add_witness());
293
294 let computed_digest = sha512_fixed(&builder, &message_wires, message_bytes.len());
296
297 for i in 0..8 {
299 builder.assert_eq(format!("digest[{i}]"), computed_digest[i], expected_digest_wires[i]);
300 }
301
302 let circuit = builder.build();
303 let cs = circuit.constraint_system();
304 let mut w = circuit.new_witness_filler();
305
306 for (i, wire) in message_wires.iter().enumerate() {
308 let byte_start = i * 8;
309 let byte_end = ((i + 1) * 8).min(message_bytes.len());
310
311 let mut word = 0u64;
312 for j in byte_start..byte_end {
313 word |= (message_bytes[j] as u64) << (56 - (j - byte_start) * 8);
314 }
315 w[*wire] = Word(word);
316 }
317
318 for (i, bytes) in expected_digest.chunks(8).enumerate() {
320 let word = u64::from_be_bytes(bytes.try_into().unwrap());
321 w[expected_digest_wires[i]] = Word(word);
322 }
323
324 circuit.populate_wire_witness(&mut w).unwrap();
325 cs.verify(&w.into_value_vec()).unwrap();
326 }
327
328 #[test]
329 #[should_panic(expected = "message.len() (1) must equal len_bytes.div_ceil(8) (2)")]
330 fn test_sha512_fixed_with_insufficient_wires() {
331 let builder = CircuitBuilder::new();
332
333 let message_wires: Vec<Wire> = vec![builder.add_witness()];
335
336 sha512_fixed(&builder, &message_wires, 10);
338 }
339
340 #[test]
341 fn test_sha512_fixed_exact_wire_count() {
342 let builder = CircuitBuilder::new();
343
344 let empty: Vec<Wire> = vec![];
348 let _ = sha512_fixed(&builder, &empty, 0);
349
350 let one_wire: Vec<Wire> = vec![builder.add_witness()];
352 let _ = sha512_fixed(&builder, &one_wire, 8);
353
354 let two_wires: Vec<Wire> = vec![builder.add_witness(), builder.add_witness()];
356 let _ = sha512_fixed(&builder, &two_wires, 10);
357
358 let two_wires_full: Vec<Wire> = vec![builder.add_witness(), builder.add_witness()];
360 let _ = sha512_fixed(&builder, &two_wires_full, 16);
361
362 let three_wires: Vec<Wire> = vec![
364 builder.add_witness(),
365 builder.add_witness(),
366 builder.add_witness(),
367 ];
368 let _ = sha512_fixed(&builder, &three_wires, 17);
369 }
370
371 #[test]
372 fn test_sha512_fixed_various_sizes() {
373 use rand::prelude::*;
374
375 let sizes = vec![
377 0, 1, 7, 8, 9, 63, 64, 65, 111, 112, 127, 128, 129, 239, 240, 256, ];
394
395 let mut rng = StdRng::seed_from_u64(0);
396
397 for size in sizes {
398 let mut message = vec![0u8; size];
400 rng.fill(&mut message[..]);
401
402 let expected = sha2::Sha512::digest(&message);
404 let expected_bytes: [u8; 64] = expected.into();
405
406 test_sha512_fixed_with_input(&message, expected_bytes);
408 }
409 }
410
411 fn test_sha512_varlen_with_input(
417 message_bytes: &[u8],
418 expected_digest: [u8; 64],
419 max_len_bytes: usize,
420 ) {
421 assert!(message_bytes.len() <= max_len_bytes);
422
423 let builder = CircuitBuilder::new();
424 let max_len_words = max_len_bytes.div_ceil(8);
425 let input = ByteVec::new_inout(&builder, max_len_words);
426 let expected_digest_wires: [Wire; 8] = std::array::from_fn(|_| builder.add_witness());
427
428 let computed_digest = sha512_varlen(&builder, &input);
429 for i in 0..8 {
430 builder.assert_eq(format!("digest[{i}]"), computed_digest[i], expected_digest_wires[i]);
431 }
432
433 let circuit = builder.build();
434 let cs = circuit.constraint_system();
435 let mut w = circuit.new_witness_filler();
436
437 input.populate_data(&mut w, message_bytes);
438 input.populate_len_bytes(&mut w, message_bytes.len());
439
440 for (i, bytes) in expected_digest.chunks(8).enumerate() {
441 let word = u64::from_be_bytes(bytes.try_into().unwrap());
442 w[expected_digest_wires[i]] = Word(word);
443 }
444
445 circuit.populate_wire_witness(&mut w).unwrap();
446 cs.verify(&w.into_value_vec()).unwrap();
447 }
448
449 #[test]
450 fn test_sha512_varlen_empty() {
451 test_sha512_varlen_with_input(
452 b"",
453 hex!(
454 "cf83e1357eefb8bdf1542850d66d8007d620e4050b5715dc83f4a921d36ce9ce47d0d13c5d85f2b0ff8318d2877eec2f63b931bd47417a81a538327af927da3e"
455 ),
456 128,
457 );
458 }
459
460 #[test]
461 fn test_sha512_varlen_abc() {
462 test_sha512_varlen_with_input(
463 b"abc",
464 hex!(
465 "ddaf35a193617abacc417349ae20413112e6fa4e89a97ea20a9eeee64b55d39a2192992a274fc1a836ba3c23a3feebbd454d4423643ce80e2a9ac94fa54ca49f"
466 ),
467 128,
468 );
469 }
470
471 #[test]
472 fn test_sha512_varlen_various_sizes() {
473 use rand::prelude::*;
474
475 let sizes: Vec<usize> = vec![
477 0, 1, 7, 8, 9, 63, 64, 65, 111, 112, 127, 128, 129, 239, 240, 256,
478 ];
479 let max_len_bytes = 320;
481
482 let mut rng = StdRng::seed_from_u64(0);
483 for size in sizes {
484 let mut message = vec![0u8; size];
485 rng.fill(&mut message[..]);
486
487 let expected = sha2::Sha512::digest(&message);
488 let expected_bytes: [u8; 64] = expected.into();
489
490 test_sha512_varlen_with_input(&message, expected_bytes, max_len_bytes);
491 }
492 }
493
494 fn check_fixed_with_compress_chip(message: &[u8]) {
500 let builder = CircuitBuilder::new();
501 builder.register_chip(Sha512Compress, &[]);
502
503 let n_words = message.len().div_ceil(8);
504 let message_wires: Vec<Wire> = (0..n_words).map(|_| builder.add_witness()).collect();
505 let computed_digest = sha512_fixed(&builder, &message_wires, message.len());
506 let digest_out: [Wire; 8] = std::array::from_fn(|_| builder.add_inout());
507 for i in 0..8 {
508 builder.assert_eq(format!("digest[{i}]"), computed_digest[i], digest_out[i]);
509 }
510
511 let circuit = builder.build_m4();
512 circuit.validate().unwrap();
513 let cs = circuit.to_constraint_system();
514 cs.validate().unwrap();
515
516 let expected: [u8; 64] = sha2::Sha512::digest(message).into();
517
518 let witness = circuit
519 .generate_witness(|w| {
520 for (i, wire) in message_wires.iter().enumerate() {
521 let byte_start = i * 8;
522 let byte_end = ((i + 1) * 8).min(message.len());
523
524 let mut word = 0u64;
525 for j in byte_start..byte_end {
526 word |= (message[j] as u64) << (56 - (j - byte_start) * 8);
527 }
528 w[*wire] = Word(word);
529 }
530 for (i, bytes) in expected.chunks(8).enumerate() {
531 w[digest_out[i]] = Word(u64::from_be_bytes(bytes.try_into().unwrap()));
532 }
533 })
534 .unwrap_or_else(|e| {
535 panic!("sha512_fixed failed for len_bytes={}: {e:?}", message.len())
536 });
537
538 witness.verify(&cs).unwrap();
539 }
540
541 #[test]
547 fn a_registered_chip_serves_every_compression() {
548 for &len in &[10usize, 112, 300, 1024, 5000] {
549 let message: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
550 check_fixed_with_compress_chip(&message);
551 }
552 }
553}