1use binius_core::word::Word;
8use binius_frontend::{CircuitBuilder, Wire, WitnessFiller};
9
10use super::constants::{BLOCK_BYTES, IV, ROUNDS, SIGMA};
11
12pub struct Blake2bCircuit {
15 pub length: usize,
17
18 pub message: Vec<Wire>,
20
21 pub digest: [Wire; 8],
23}
24
25impl Blake2bCircuit {
26 pub fn new(builder: &CircuitBuilder) -> Self {
28 Self::new_with_length(builder, BLOCK_BYTES) }
30
31 pub fn new_with_length(builder: &CircuitBuilder, max_msg_len_bytes: usize) -> Self {
33 Self::new_with_params(builder, max_msg_len_bytes, 64)
34 }
35
36 pub fn new_with_params(
38 builder: &CircuitBuilder,
39 max_msg_len_bytes: usize,
40 outlen: usize,
41 ) -> Self {
42 assert!(outlen > 0 && outlen <= 64, "Output length must be 1-64 bytes");
43 let num_message_words = max_msg_len_bytes.div_ceil(8).max(1);
48 let message: Vec<Wire> = (0..num_message_words)
49 .map(|_| builder.add_witness())
50 .collect();
51
52 let digest = std::array::from_fn(|_| builder.add_witness());
54
55 Self::build_circuit(builder, max_msg_len_bytes, &message, digest, outlen);
57
58 Self {
59 length: max_msg_len_bytes,
60 message,
61 digest,
62 }
63 }
64
65 pub fn populate_message(&self, w: &mut WitnessFiller<'_>, message: &[u8]) {
67 assert!(message.len() <= self.length, "Message exceeds circuit capacity");
68
69 for (i, chunk) in message.chunks(8).enumerate() {
71 let mut word_value = 0u64;
72 for (j, &byte) in chunk.iter().enumerate() {
73 word_value |= (byte as u64) << (j * 8);
74 }
75 w[self.message[i]] = Word(word_value);
76 }
77
78 for i in message.len().div_ceil(8)..self.message.len() {
80 w[self.message[i]] = Word(0);
81 }
82 }
83
84 pub fn populate_digest(&self, w: &mut WitnessFiller<'_>, digest: &[u8; 64]) {
86 for i in 0..8 {
88 let mut word_value = 0u64;
89 for j in 0..8 {
90 word_value |= (digest[i * 8 + j] as u64) << (j * 8);
91 }
92 w[self.digest[i]] = Word(word_value);
93 }
94 }
95
96 fn build_circuit(
106 builder: &CircuitBuilder,
107 length: usize,
108 message: &[Wire],
109 expected_digest: [Wire; 8],
110 outlen: usize,
111 ) {
112 let num_blocks = if length == 0 {
114 1
115 } else {
116 length.div_ceil(BLOCK_BYTES)
117 };
118 let zero = builder.add_constant(Word::ZERO);
119
120 let param_block = 0x01010000 | (outlen as u64);
123
124 let init_state = [
125 builder.add_constant(Word(IV[0] ^ param_block)),
126 builder.add_constant(Word(IV[1])),
127 builder.add_constant(Word(IV[2])),
128 builder.add_constant(Word(IV[3])),
129 builder.add_constant(Word(IV[4])),
130 builder.add_constant(Word(IV[5])),
131 builder.add_constant(Word(IV[6])),
132 builder.add_constant(Word(IV[7])),
133 ];
134
135 let mut h = init_state;
136 let mut final_digest = [zero; 8];
137
138 for block_idx in 0..num_blocks {
140 let mut m = [zero; 16];
142
143 for word_idx in 0..16 {
145 let byte_start = block_idx * BLOCK_BYTES + word_idx * 8;
146
147 if byte_start < length {
148 let msg_word_idx = byte_start / 8;
150
151 if msg_word_idx < message.len() {
152 let msg_word = message[msg_word_idx];
153
154 if byte_start + 8 > length {
156 let valid_bytes = length - byte_start;
158 let mask = builder.add_constant(Word((1u64 << (valid_bytes * 8)) - 1));
159 m[word_idx] = builder.band(msg_word, mask);
160 } else {
161 m[word_idx] = msg_word;
162 }
163 }
164 }
165 }
167
168 let is_final_block = block_idx == num_blocks - 1;
170
171 let t_low = if is_final_block {
173 builder.add_constant(Word(length as u64))
174 } else {
175 builder.add_constant(Word(((block_idx + 1) * BLOCK_BYTES) as u64))
176 };
177 let t_high = zero; let last_flag = if is_final_block {
181 builder.add_constant(Word(0xFFFFFFFFFFFFFFFF))
182 } else {
183 zero
184 };
185
186 h = compress(builder, &h, &m, t_low, t_high, last_flag);
188
189 if is_final_block {
191 final_digest.copy_from_slice(&h);
192 }
193 }
194
195 for i in 0..8 {
197 builder.assert_eq(format!("digest[{}]", i), final_digest[i], expected_digest[i]);
198 }
199 }
200}
201
202fn compress(
204 builder: &CircuitBuilder,
205 h: &[Wire; 8],
206 m: &[Wire; 16],
207 t_low: Wire,
208 t_high: Wire,
209 last_block_flag: Wire,
210) -> [Wire; 8] {
211 let mut v = [builder.add_constant(Word::ZERO); 16];
213
214 v[0..8].copy_from_slice(h);
216
217 for i in 0..8 {
219 v[i + 8] = builder.add_constant(Word(IV[i]));
220 }
221
222 v[12] = builder.bxor(v[12], t_low);
224 v[13] = builder.bxor(v[13], t_high);
225
226 v[14] = builder.bxor(v[14], last_block_flag);
228
229 for round in 0..ROUNDS {
231 g_mixing(builder, &mut v, 0, 4, 8, 12, m[SIGMA[round][0]], m[SIGMA[round][1]]);
233 g_mixing(builder, &mut v, 1, 5, 9, 13, m[SIGMA[round][2]], m[SIGMA[round][3]]);
234 g_mixing(builder, &mut v, 2, 6, 10, 14, m[SIGMA[round][4]], m[SIGMA[round][5]]);
235 g_mixing(builder, &mut v, 3, 7, 11, 15, m[SIGMA[round][6]], m[SIGMA[round][7]]);
236
237 g_mixing(builder, &mut v, 0, 5, 10, 15, m[SIGMA[round][8]], m[SIGMA[round][9]]);
239 g_mixing(builder, &mut v, 1, 6, 11, 12, m[SIGMA[round][10]], m[SIGMA[round][11]]);
240 g_mixing(builder, &mut v, 2, 7, 8, 13, m[SIGMA[round][12]], m[SIGMA[round][13]]);
241 g_mixing(builder, &mut v, 3, 4, 9, 14, m[SIGMA[round][14]], m[SIGMA[round][15]]);
242 }
243
244 let mut h_new = [builder.add_constant(Word::ZERO); 8];
246 for i in 0..8 {
247 h_new[i] = builder.bxor_multi(&[h[i], v[i], v[i + 8]]);
248 }
249
250 h_new
251}
252
253#[allow(clippy::too_many_arguments)]
269pub fn g_mixing(
270 builder: &CircuitBuilder,
271 v: &mut [Wire; 16],
272 a: usize,
273 b: usize,
274 c: usize,
275 d: usize,
276 x: Wire,
277 y: Wire,
278) {
279 let (temp1, _) = builder.iadd(v[a], v[b]);
281 let (v_a_new1, _) = builder.iadd(temp1, x);
282 v[a] = v_a_new1;
283
284 let xor1 = builder.bxor(v[d], v[a]);
286 v[d] = builder.rotr(xor1, 32);
287
288 let (v_c_new1, _) = builder.iadd(v[c], v[d]);
290 v[c] = v_c_new1;
291
292 let xor2 = builder.bxor(v[b], v[c]);
294 v[b] = builder.rotr(xor2, 24);
295
296 let (temp2, _) = builder.iadd(v[a], v[b]);
298 let (v_a_new2, _) = builder.iadd(temp2, y);
299 v[a] = v_a_new2;
300
301 let xor3 = builder.bxor(v[d], v[a]);
303 v[d] = builder.rotr(xor3, 16);
304
305 let (v_c_new2, _) = builder.iadd(v[c], v[d]);
307 v[c] = v_c_new2;
308
309 let xor4 = builder.bxor(v[b], v[c]);
311 v[b] = builder.rotr(xor4, 63);
312}
313
314#[cfg(test)]
315mod tests {
316 use binius_core::word::Word;
317 use binius_frontend::CircuitBuilder;
318
319 use crate::blake2b::{circuit::g_mixing, reference};
320
321 #[test]
323 fn test_g_mixing_function() {
324 let builder = CircuitBuilder::new();
325
326 let mut v = core::array::from_fn(|_| builder.add_inout());
327 let x = builder.add_inout();
328 let y = builder.add_inout();
329
330 let expected: [_; 16] = core::array::from_fn(|_| builder.add_inout());
332
333 let v_initial = v;
335
336 g_mixing(&builder, &mut v, 0, 4, 8, 12, x, y);
338
339 for i in [0, 4, 8, 12] {
342 builder.assert_eq(format!("v[{}]", i), v[i], expected[i]);
343 }
344
345 let circuit = builder.build();
346
347 let mut w = circuit.new_witness_filler();
349
350 let initial_v = [
352 0x0000000000000001u64, 0x0000000000000002u64, 0x0000000000000003u64, 0x0000000000000004u64, 0x0000000000000005u64, 0x0000000000000006u64, 0x0000000000000007u64, 0x0000000000000008u64, 0x0000000000000009u64, 0x000000000000000Au64, 0x000000000000000Bu64, 0x000000000000000Cu64, 0x000000000000000Du64, 0x000000000000000Eu64, 0x000000000000000Fu64, 0x0000000000000010u64, ];
369
370 let x_val = 0x123456789ABCDEFu64;
371 let y_val = 0xFEDCBA9876543210u64;
372
373 for i in 0..16 {
374 w[v_initial[i]] = Word(initial_v[i]);
375 }
376 w[x] = Word(x_val);
377 w[y] = Word(y_val);
378
379 let mut expected_v = initial_v;
381 reference::g(&mut expected_v, 0, 4, 8, 12, x_val, y_val);
382
383 for i in [0, 4, 8, 12] {
384 w[expected[i]] = Word(expected_v[i]);
385 }
386
387 circuit.populate_wire_witness(&mut w).unwrap();
389
390 let cs = circuit.constraint_system();
391 cs.verify(&w.into_value_vec()).unwrap();
392 }
393}