1use binius_core::word::Word;
4use binius_frontend::{CircuitBuilder, Wire};
5
6use super::{
7 SHA3_256_DIGEST_WORDS, SHA3_256_RATE_BYTES, SHA3_384_DIGEST_WORDS, SHA3_384_RATE_BYTES,
8 SHA3_512_DIGEST_WORDS, SHA3_512_RATE_BYTES, SHA3_DELIMITER_BYTE,
9};
10use crate::{
11 fixed_byte_vec::ByteVec,
12 keccak::{N_WORDS_PER_STATE, permutation::keccak_f1600},
13 multiplexer::{multi_wire_multiplex, single_wire_multiplex},
14};
15
16fn sha3_varlen(
31 builder: &CircuitBuilder,
32 message: &ByteVec,
33 rate_bytes: usize,
34 digest_words: usize,
35) -> Vec<Wire> {
36 let len_bytes = message.len_bytes;
37 let data = &message.data;
38
39 let n_words_per_block = rate_bytes / 8;
40 let max_len_bytes = data.len() << 3;
41 let n_blocks = (max_len_bytes + 1).div_ceil(rate_bytes);
43 let n_words = n_blocks * n_words_per_block;
44
45 let too_long = builder.icmp_ugt(len_bytes, builder.add_constant_64(max_len_bytes as u64));
47 builder.assert_false("len_check", too_long);
48
49 let zero = builder.add_constant(Word::ZERO);
50 let msb_one = builder.add_constant(Word::MSB_ONE);
51
52 let w_bd = builder.shr(len_bytes, 3);
56 let len_mod_8 = builder.band(len_bytes, builder.add_constant_64(7));
57
58 let mut end_block_index = zero;
62 for block_no in 0..n_blocks {
63 let block_start = builder.add_constant_64((block_no * rate_bytes) as u64);
64 let block_end = builder.add_constant_64(((block_no + 1) * rate_bytes) as u64);
65 let gte_start = builder.icmp_ule(block_start, len_bytes);
66 let lt_end = builder.icmp_ult(len_bytes, block_end);
67 let is_final_block = builder.band(gte_start, lt_end);
68 end_block_index = builder.select(
69 is_final_block,
70 builder.add_constant_64(block_no as u64),
71 end_block_index,
72 );
73 }
74
75 let boundary_message_word = single_wire_multiplex(builder, data, w_bd);
82 let candidates: Vec<Wire> = (0..8)
83 .map(|i| {
84 let mask = builder.add_constant_64(0x00FFFFFFFFFFFFFF >> ((7 - i) << 3));
85 let delimiter = builder.add_constant_64(SHA3_DELIMITER_BYTE << (i << 3));
86 let message_low = builder.band(boundary_message_word, mask);
87 builder.bxor(message_low, delimiter)
88 })
89 .collect();
90 let boundary_word = single_wire_multiplex(builder, &candidates, len_mod_8);
92
93 let padded_message: Vec<Wire> = (0..n_words)
96 .map(|word_index| {
97 let block_index = word_index / n_words_per_block;
98 let column_index = word_index % n_words_per_block;
99 let word_idx_wire = builder.add_constant_64(word_index as u64);
100
101 let is_message_word = builder.icmp_ult(word_idx_wire, w_bd);
102 let is_boundary_word = builder.icmp_eq(word_idx_wire, w_bd);
103 let is_end_block =
104 builder.icmp_eq(builder.add_constant_64(block_index as u64), end_block_index);
105
106 let msg_word = if word_index < data.len() {
111 data[word_index]
112 } else {
113 zero
114 };
115
116 let delimiter = if column_index == n_words_per_block - 1 {
121 builder.select(is_end_block, msb_one, zero)
122 } else {
123 zero
124 };
125 let boundary_val = if column_index == n_words_per_block - 1 {
126 builder.bxor(boundary_word, delimiter)
127 } else {
128 boundary_word
129 };
130
131 let boundary_or_padding = builder.select(is_boundary_word, boundary_val, delimiter);
132 builder.select(is_message_word, msg_word, boundary_or_padding)
133 })
134 .collect();
135
136 let mut states: Vec<[Wire; N_WORDS_PER_STATE]> = Vec::with_capacity(n_blocks + 1);
143 states.push([zero; N_WORDS_PER_STATE]);
144 for block_no in 0..n_blocks {
145 let mut state = states[block_no];
146 for i in 0..n_words_per_block {
147 state[i] = builder.bxor(state[i], padded_message[block_no * n_words_per_block + i]);
148 }
149 keccak_f1600(builder, &mut state);
150 states.push(state);
151 }
152
153 let inputs: Vec<&[Wire]> = states[1..].iter().map(|s| &s[..]).collect();
155 let digest_vec = multi_wire_multiplex(builder, &inputs, end_block_index);
156 digest_vec[..digest_words].to_vec()
157}
158
159pub fn sha3_256_varlen(
168 builder: &CircuitBuilder,
169 message: &ByteVec,
170) -> [Wire; SHA3_256_DIGEST_WORDS] {
171 sha3_varlen(builder, message, SHA3_256_RATE_BYTES, SHA3_256_DIGEST_WORDS)
172 .try_into()
173 .unwrap()
174}
175
176pub fn sha3_384_varlen(
185 builder: &CircuitBuilder,
186 message: &ByteVec,
187) -> [Wire; SHA3_384_DIGEST_WORDS] {
188 sha3_varlen(builder, message, SHA3_384_RATE_BYTES, SHA3_384_DIGEST_WORDS)
189 .try_into()
190 .unwrap()
191}
192
193pub fn sha3_512_varlen(
202 builder: &CircuitBuilder,
203 message: &ByteVec,
204) -> [Wire; SHA3_512_DIGEST_WORDS] {
205 sha3_varlen(builder, message, SHA3_512_RATE_BYTES, SHA3_512_DIGEST_WORDS)
206 .try_into()
207 .unwrap()
208}
209
210#[cfg(test)]
211mod tests {
212 use binius_core::Word;
213 use binius_frontend::{CircuitBuilder, Wire};
214 use rand::prelude::*;
215 use rstest::rstest;
216 use sha3::Digest;
217
218 use super::*;
219
220 fn test_varlen<const DIGEST_WORDS: usize>(
223 message: &[u8],
224 max_message_len_bytes: usize,
225 hash_fn: impl FnOnce(&CircuitBuilder, &ByteVec) -> [Wire; DIGEST_WORDS],
226 reference: impl FnOnce(&[u8]) -> Vec<u8>,
227 ) {
228 assert!(message.len() <= max_message_len_bytes);
229
230 let expected_digest = reference(message);
231 assert_eq!(expected_digest.len(), DIGEST_WORDS * 8);
232
233 let b = CircuitBuilder::new();
234 let max_len_words = max_message_len_bytes.div_ceil(8);
235 let input = ByteVec::new_inout(&b, max_len_words);
236 let expected_digest_wires: [Wire; DIGEST_WORDS] = std::array::from_fn(|_| b.add_witness());
237
238 let computed_digest = hash_fn(&b, &input);
240 for i in 0..DIGEST_WORDS {
241 b.assert_eq(format!("digest[{i}]"), computed_digest[i], expected_digest_wires[i]);
242 }
243
244 let circuit = b.build();
245 let cs = circuit.constraint_system();
246 let mut witness = circuit.new_witness_filler();
247
248 input.populate_data(&mut witness, message);
250 input.populate_len_bytes(&mut witness, message.len());
251 for (i, bytes) in expected_digest.chunks(8).enumerate() {
253 witness[expected_digest_wires[i]] = Word(u64::from_le_bytes(bytes.try_into().unwrap()));
254 }
255
256 circuit
258 .populate_wire_witness(&mut witness)
259 .expect("Circuit should accept valid witness");
260 cs.verify(&witness.into_value_vec())
263 .expect("All constraints should be satisfied");
264 }
265
266 fn random_message(message_len_bytes: usize, max_message_len_bytes: usize) -> Vec<u8> {
269 let seed = ((message_len_bytes as u64) << 32) | (max_message_len_bytes as u64);
270 let mut rng = StdRng::seed_from_u64(seed);
271 let mut message = vec![0u8; message_len_bytes];
272 rng.fill_bytes(&mut message);
273 message
274 }
275
276 #[rstest]
277 #[case(0, 100)] #[case(1, 100)] #[case(1, 144)] #[case(135, 136)] #[case(136, 136)] #[case(137, 272)] #[case(271, 272)] #[case(272, 272)] fn test_sha3_256_varlen(
286 #[case] message_len_bytes: usize,
287 #[case] max_message_len_bytes: usize,
288 ) {
289 let message = random_message(message_len_bytes, max_message_len_bytes);
290 test_varlen(&message, max_message_len_bytes, sha3_256_varlen, |m| {
291 sha3::Sha3_256::digest(m).to_vec()
292 });
293 }
294
295 #[rstest]
296 #[case(0, 100)] #[case(1, 100)] #[case(103, 104)] #[case(104, 104)] #[case(105, 208)] fn test_sha3_384_varlen(
302 #[case] message_len_bytes: usize,
303 #[case] max_message_len_bytes: usize,
304 ) {
305 let message = random_message(message_len_bytes, max_message_len_bytes);
306 test_varlen(&message, max_message_len_bytes, sha3_384_varlen, |m| {
307 sha3::Sha3_384::digest(m).to_vec()
308 });
309 }
310
311 #[rstest]
312 #[case(0, 100)] #[case(1, 100)] #[case(71, 72)] #[case(72, 72)] #[case(73, 144)] fn test_sha3_512_varlen(
318 #[case] message_len_bytes: usize,
319 #[case] max_message_len_bytes: usize,
320 ) {
321 let message = random_message(message_len_bytes, max_message_len_bytes);
322 test_varlen(&message, max_message_len_bytes, sha3_512_varlen, |m| {
323 sha3::Sha3_512::digest(m).to_vec()
324 });
325 }
326}