1pub mod fixed_length;
4pub mod permutation;
5
6use binius_core::word::Word;
7use binius_frontend::{CircuitBuilder, Wire};
8pub use permutation::{KeccakF1600, keccak_f1600, ref_keccak_f1600};
9
10use crate::{
11 fixed_byte_vec::ByteVec,
12 multiplexer::{multi_wire_multiplex, single_wire_multiplex},
13};
14
15pub const N_WORDS_PER_DIGEST: usize = 4;
16pub const N_WORDS_PER_STATE: usize = 25;
17pub const RATE_BYTES: usize = 136;
18pub const N_WORDS_PER_BLOCK: usize = RATE_BYTES / 8;
19
20pub fn keccak256_varlen(builder: &CircuitBuilder, message: &ByteVec) -> [Wire; N_WORDS_PER_DIGEST] {
52 let len_bytes = message.len_bytes;
53 let data = &message.data;
54
55 let max_len_bytes = data.len() << 3;
56 let n_blocks = (max_len_bytes + 1).div_ceil(RATE_BYTES);
58 let n_words = n_blocks * N_WORDS_PER_BLOCK;
59
60 let too_long = builder.icmp_ugt(len_bytes, builder.add_constant_64(max_len_bytes as u64));
62 builder.assert_false("len_check", too_long);
63
64 let zero = builder.add_constant(Word::ZERO);
65 let msb_one = builder.add_constant(Word::MSB_ONE);
66
67 let w_bd = builder.shr(len_bytes, 3);
70 let len_mod_8 = builder.band(len_bytes, builder.add_constant_64(7));
71
72 let mut end_block_index = zero;
76 for block_no in 0..n_blocks {
77 let block_start = builder.add_constant_64((block_no * RATE_BYTES) as u64);
78 let block_end = builder.add_constant_64(((block_no + 1) * RATE_BYTES) as u64);
79 let gte_start = builder.icmp_ule(block_start, len_bytes);
80 let lt_end = builder.icmp_ult(len_bytes, block_end);
81 let is_final_block = builder.band(gte_start, lt_end);
82 end_block_index = builder.select(
83 is_final_block,
84 builder.add_constant_64(block_no as u64),
85 end_block_index,
86 );
87 }
88
89 let boundary_message_word = single_wire_multiplex(builder, data, w_bd);
93 let candidates: Vec<Wire> = (0..8)
94 .map(|i| {
95 let mask = builder.add_constant_64(0x00FFFFFFFFFFFFFF >> ((7 - i) << 3));
96 let delimiter = builder.add_constant_64(1u64 << (i << 3));
97 let message_low = builder.band(boundary_message_word, mask);
98 builder.bxor(message_low, delimiter)
99 })
100 .collect();
101 let boundary_word = single_wire_multiplex(builder, &candidates, len_mod_8);
102
103 let padded_message: Vec<Wire> = (0..n_words)
105 .map(|word_index| {
106 let block_index = word_index / N_WORDS_PER_BLOCK;
107 let column_index = word_index % N_WORDS_PER_BLOCK;
108 let word_idx_wire = builder.add_constant_64(word_index as u64);
109
110 let is_message_word = builder.icmp_ult(word_idx_wire, w_bd);
111 let is_boundary_word = builder.icmp_eq(word_idx_wire, w_bd);
112 let is_end_block =
113 builder.icmp_eq(builder.add_constant_64(block_index as u64), end_block_index);
114
115 let msg_word = if word_index < data.len() {
119 data[word_index]
120 } else {
121 zero
122 };
123
124 let delimiter = if column_index == N_WORDS_PER_BLOCK - 1 {
128 builder.select(is_end_block, msb_one, zero)
129 } else {
130 zero
131 };
132 let boundary_val = if column_index == N_WORDS_PER_BLOCK - 1 {
133 builder.bxor(boundary_word, delimiter)
134 } else {
135 boundary_word
136 };
137
138 let boundary_or_padding = builder.select(is_boundary_word, boundary_val, delimiter);
139 builder.select(is_message_word, msg_word, boundary_or_padding)
140 })
141 .collect();
142
143 let mut states: Vec<[Wire; N_WORDS_PER_STATE]> = Vec::with_capacity(n_blocks + 1);
145 states.push([zero; N_WORDS_PER_STATE]);
146 for block_no in 0..n_blocks {
147 let mut state = states[block_no];
148 for i in 0..N_WORDS_PER_BLOCK {
149 state[i] = builder.bxor(state[i], padded_message[block_no * N_WORDS_PER_BLOCK + i]);
150 }
151 keccak_f1600(builder, &mut state);
152 states.push(state);
153 }
154
155 let inputs: Vec<&[Wire]> = states[1..].iter().map(|s| &s[..]).collect();
157 let digest_vec = multi_wire_multiplex(builder, &inputs, end_block_index);
158 digest_vec[..N_WORDS_PER_DIGEST].try_into().unwrap()
159}
160
161#[cfg(test)]
162mod tests {
163 use binius_core::Word;
164 use binius_frontend::{CircuitBuilder, Wire, WitnessFiller};
165 use rand::prelude::*;
166 use rstest::rstest;
167 use sha3::Digest;
168
169 use super::*;
170 use crate::fixed_byte_vec::ByteVec;
171
172 fn test_keccak_varlen_with_input(message: &[u8], max_message_len_bytes: usize) {
176 assert!(
177 message.len() <= max_message_len_bytes,
178 "Message length {} exceeds max capacity {} bytes",
179 message.len(),
180 max_message_len_bytes
181 );
182
183 let mut hasher = sha3::Keccak256::new();
185 hasher.update(message);
186 let expected_digest: [u8; 32] = hasher.finalize().into();
187
188 let b = CircuitBuilder::new();
189 let max_len_words = max_message_len_bytes.div_ceil(8);
190 let input = ByteVec::new_inout(&b, max_len_words);
191 let expected_digest_wires: [Wire; N_WORDS_PER_DIGEST] =
192 std::array::from_fn(|_| b.add_witness());
193
194 let computed_digest = keccak256_varlen(&b, &input);
195 for i in 0..N_WORDS_PER_DIGEST {
196 b.assert_eq(format!("digest[{i}]"), computed_digest[i], expected_digest_wires[i]);
197 }
198
199 let circuit = b.build();
200 let cs = circuit.constraint_system();
201 let mut witness = circuit.new_witness_filler();
202
203 input.populate_data(&mut witness, message);
204 input.populate_len_bytes(&mut witness, message.len());
205 for (i, bytes) in expected_digest.chunks(8).enumerate() {
206 witness[expected_digest_wires[i]] = Word(u64::from_le_bytes(bytes.try_into().unwrap()));
207 }
208
209 circuit
210 .populate_wire_witness(&mut witness)
211 .expect("Circuit should accept valid witness");
212 cs.verify(&witness.into_value_vec())
213 .expect("All constraints should be satisfied");
214 }
215
216 #[rstest]
217 #[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_keccak_varlen(#[case] message_len_bytes: usize, #[case] max_message_len_bytes: usize) {
226 let seed = ((message_len_bytes as u64) << 32) | (max_message_len_bytes as u64);
228 let mut rng = StdRng::seed_from_u64(seed);
229 let mut message = vec![0u8; message_len_bytes];
230 rng.fill_bytes(&mut message);
231
232 test_keccak_varlen_with_input(&message, max_message_len_bytes);
233 }
234
235 fn check_chip_serves(
243 b: CircuitBuilder,
244 computed_digest: [Wire; N_WORDS_PER_DIGEST],
245 expected: [u8; 32],
246 n_permutations: usize,
247 fill: impl FnOnce(&mut WitnessFiller<'_>),
248 ) {
249 let digest_out: [Wire; N_WORDS_PER_DIGEST] = std::array::from_fn(|_| b.add_inout());
250 for i in 0..N_WORDS_PER_DIGEST {
251 b.assert_eq(format!("digest[{i}]"), computed_digest[i], digest_out[i]);
252 }
253
254 let circuit = b.build_m4();
255 circuit.validate().unwrap();
256 assert_eq!(circuit.chips.len(), 1, "the permutation is the system's only chip");
257 assert_eq!(
258 circuit.chips[0].1, n_permutations,
259 "every permutation the sponge runs has to reach the chip"
260 );
261
262 let cs = circuit.to_constraint_system();
263 cs.validate().unwrap();
264
265 let witness = circuit
266 .generate_witness(|w| {
267 fill(w);
268 for (i, bytes) in expected.chunks(8).enumerate() {
269 w[digest_out[i]] = Word(u64::from_le_bytes(bytes.try_into().unwrap()));
270 }
271 })
272 .unwrap();
273
274 witness.verify(&cs).unwrap();
275 }
276
277 #[test]
284 fn a_registered_chip_serves_every_fixed_length_permutation() {
285 for &len in &[0usize, 1, 135, 136, 137, 271, 272, 500] {
286 let message: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
287
288 let b = CircuitBuilder::new();
289 b.register_chip(KeccakF1600, &[]);
290 let message_wires: Vec<Wire> = (0..len.div_ceil(8)).map(|_| b.add_witness()).collect();
291 let digest = fixed_length::keccak256(&b, &message_wires, len);
292
293 check_chip_serves(
294 b,
295 digest,
296 sha3::Keccak256::digest(&message).into(),
297 (len + 1).div_ceil(RATE_BYTES),
298 |w| {
299 for (wire, chunk) in std::iter::zip(&message_wires, message.chunks(8)) {
300 let mut bytes = [0u8; 8];
301 bytes[..chunk.len()].copy_from_slice(chunk);
302 w[*wire] = Word(u64::from_le_bytes(bytes));
303 }
304 },
305 );
306 }
307 }
308
309 #[test]
314 fn a_registered_chip_serves_every_variable_length_permutation() {
315 for &(len, max_len) in &[
316 (0usize, 100usize),
317 (1, 144),
318 (135, 136),
319 (137, 272),
320 (272, 272),
321 ] {
322 let message: Vec<u8> = (0..len).map(|i| (i * 37 + 1) as u8).collect();
323
324 let b = CircuitBuilder::new();
325 b.register_chip(KeccakF1600, &[]);
326 let max_len_words = max_len.div_ceil(8);
327 let input = ByteVec::new_witness(&b, max_len_words);
328 let digest = keccak256_varlen(&b, &input);
329
330 check_chip_serves(
331 b,
332 digest,
333 sha3::Keccak256::digest(&message).into(),
334 ((max_len_words << 3) + 1).div_ceil(RATE_BYTES),
335 |w| {
336 input.populate_data(w, &message);
337 input.populate_len_bytes(w, message.len());
338 },
339 );
340 }
341 }
342}