binius_circuits/sha3/
fixed_length.rs1use 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::keccak::{N_WORDS_PER_STATE, permutation::keccak_f1600};
11
12fn sha3_fixed(
32 builder: &CircuitBuilder,
33 message: &[Wire],
34 len_bytes: usize,
35 rate_bytes: usize,
36 digest_words: usize,
37) -> Vec<Wire> {
38 assert_eq!(
40 message.len(),
41 len_bytes.div_ceil(8),
42 "message.len() ({}) must equal len_bytes.div_ceil(8) ({})",
43 message.len(),
44 len_bytes.div_ceil(8)
45 );
46
47 let n_words_per_block = rate_bytes / 8;
48 let n_blocks = (len_bytes + 1).div_ceil(rate_bytes);
50 let n_padded_words = n_blocks * n_words_per_block;
51
52 let mut padded_message = Vec::with_capacity(n_padded_words);
53
54 if len_bytes.is_multiple_of(8) {
60 padded_message.extend_from_slice(message);
62 padded_message.push(builder.add_constant(Word(SHA3_DELIMITER_BYTE)));
63 } else {
64 padded_message.extend_from_slice(&message[..message.len() - 1]);
66
67 let last_idx = message.len() - 1;
68 let byte_in_word = len_bytes % 8;
69
70 let mask = (1u64 << (byte_in_word * 8)) - 1;
72 let masked_word = builder.band(message[last_idx], builder.add_constant(Word(mask)));
73
74 let padding_bit = SHA3_DELIMITER_BYTE << (byte_in_word * 8);
76 let boundary_word = builder.bxor(masked_word, builder.add_constant(Word(padding_bit)));
77 padded_message.push(boundary_word);
78 }
79
80 let zero = builder.add_constant(Word::ZERO);
82 padded_message.resize(n_padded_words, zero);
83
84 let last_byte_mask = 0x80u64 << 56;
88 let last_idx = n_padded_words - 1;
89 padded_message[last_idx] =
90 builder.bxor(padded_message[last_idx], builder.add_constant(Word(last_byte_mask)));
91
92 let mut state = [zero; N_WORDS_PER_STATE];
96 for block in padded_message.chunks(n_words_per_block) {
97 for (i, &word) in block.iter().enumerate() {
98 state[i] = builder.bxor(state[i], word);
99 }
100 keccak_f1600(builder, &mut state);
101 }
102
103 state[..digest_words].to_vec()
107}
108
109pub fn sha3_256(
123 builder: &CircuitBuilder,
124 message: &[Wire],
125 len_bytes: usize,
126) -> [Wire; SHA3_256_DIGEST_WORDS] {
127 sha3_fixed(builder, message, len_bytes, SHA3_256_RATE_BYTES, SHA3_256_DIGEST_WORDS)
129 .try_into()
130 .unwrap()
131}
132
133pub fn sha3_384(
147 builder: &CircuitBuilder,
148 message: &[Wire],
149 len_bytes: usize,
150) -> [Wire; SHA3_384_DIGEST_WORDS] {
151 sha3_fixed(builder, message, len_bytes, SHA3_384_RATE_BYTES, SHA3_384_DIGEST_WORDS)
153 .try_into()
154 .unwrap()
155}
156
157pub fn sha3_512(
171 builder: &CircuitBuilder,
172 message: &[Wire],
173 len_bytes: usize,
174) -> [Wire; SHA3_512_DIGEST_WORDS] {
175 sha3_fixed(builder, message, len_bytes, SHA3_512_RATE_BYTES, SHA3_512_DIGEST_WORDS)
177 .try_into()
178 .unwrap()
179}
180
181#[cfg(test)]
182mod tests {
183 use binius_frontend::CircuitBuilder;
184 use rand::prelude::*;
185 use rstest::rstest;
186 use sha3::Digest;
187
188 use super::*;
189
190 fn test_fixed<const DIGEST_WORDS: usize>(
193 message_len_bytes: usize,
194 hash_fn: impl FnOnce(&CircuitBuilder, &[Wire], usize) -> [Wire; DIGEST_WORDS],
195 reference: impl FnOnce(&[u8]) -> Vec<u8>,
196 ) {
197 let seed = message_len_bytes as u64;
199 let mut rng = StdRng::seed_from_u64(seed);
200 let mut message = vec![0u8; message_len_bytes];
201 rng.fill_bytes(&mut message);
202
203 let expected_digest = reference(&message);
204 assert_eq!(expected_digest.len(), DIGEST_WORDS * 8);
205
206 let builder = CircuitBuilder::new();
207 let n_words = message_len_bytes.div_ceil(8);
208 let message_wires: Vec<_> = (0..n_words).map(|_| builder.add_witness()).collect();
209 let expected_digest_wires: [Wire; DIGEST_WORDS] =
210 std::array::from_fn(|_| builder.add_witness());
211
212 let computed_digest = hash_fn(&builder, &message_wires, message_len_bytes);
214 for i in 0..DIGEST_WORDS {
215 builder.assert_eq(format!("digest[{i}]"), computed_digest[i], expected_digest_wires[i]);
216 }
217
218 let circuit = builder.build();
219 let cs = circuit.constraint_system();
220 let mut witness = circuit.new_witness_filler();
221
222 for (i, chunk) in message.chunks(8).enumerate() {
224 let mut word_bytes = [0u8; 8];
225 word_bytes[..chunk.len()].copy_from_slice(chunk);
226 witness[message_wires[i]] = Word(u64::from_le_bytes(word_bytes));
227 }
228 for (i, chunk) in expected_digest.chunks(8).enumerate() {
230 witness[expected_digest_wires[i]] = Word(u64::from_le_bytes(chunk.try_into().unwrap()));
231 }
232
233 circuit.populate_wire_witness(&mut witness).unwrap();
235 cs.verify(&witness.into_value_vec())
238 .expect("Circuit constraints should be satisfied");
239 }
240
241 #[rstest]
242 #[case(0)] #[case(1)] #[case(8)] #[case(135)] #[case(136)] #[case(137)] #[case(272)] #[case(500)] fn test_sha3_256(#[case] message_len_bytes: usize) {
251 test_fixed(message_len_bytes, sha3_256, |m| sha3::Sha3_256::digest(m).to_vec());
252 }
253
254 #[rstest]
255 #[case(0)] #[case(1)] #[case(103)] #[case(104)] #[case(105)] #[case(500)] fn test_sha3_384(#[case] message_len_bytes: usize) {
262 test_fixed(message_len_bytes, sha3_384, |m| sha3::Sha3_384::digest(m).to_vec());
263 }
264
265 #[rstest]
266 #[case(0)] #[case(1)] #[case(71)] #[case(72)] #[case(73)] #[case(500)] fn test_sha3_512(#[case] message_len_bytes: usize) {
273 test_fixed(message_len_bytes, sha3_512, |m| sha3::Sha3_512::digest(m).to_vec());
274 }
275
276 #[test]
277 #[should_panic(expected = "message.len() (1) must equal len_bytes.div_ceil(8) (2)")]
278 fn test_sha3_256_wrong_wire_count() {
279 let builder = CircuitBuilder::new();
281 let message_wires = vec![builder.add_witness()];
282 sha3_256(&builder, &message_wires, 10);
283 }
284}