1use std::{iter::once, marker::PhantomData};
2
3use ndarray::s;
4use openvm_circuit_primitives::{
5 bitwise_op_lookup::BitwiseOperationLookupBus, encoder::Encoder, utils::select, ColumnsAir,
6 SubAir,
7};
8use openvm_stark_backend::{
9 interaction::{BusIndex, InteractionBuilder, PermutationCheckBus},
10 p3_air::{AirBuilder, BaseAir},
11 p3_field::{Field, PrimeCharacteristicRing},
12 p3_matrix::Matrix,
13};
14
15use super::{
16 big_sig0_field, big_sig1_field, ch_field, compose, maj_field, small_sig0_field,
17 small_sig1_field,
18};
19use crate::{
20 constraint_word_addition, word_into_u16_limbs, Sha2BlockHasherSubairConfig, Sha2DigestColsRef,
21 Sha2RoundColsRef,
22};
23
24#[derive(Clone, Debug)]
26pub struct Sha2BlockHasherSubAir<C: Sha2BlockHasherSubairConfig> {
27 pub bitwise_lookup_bus: BitwiseOperationLookupBus,
28 pub row_idx_encoder: Encoder,
29 pub private_bus: PermutationCheckBus,
31 _phantom: PhantomData<C>,
32}
33
34impl<C: Sha2BlockHasherSubairConfig> ColumnsAir for Sha2BlockHasherSubAir<C> {}
37
38impl<C: Sha2BlockHasherSubairConfig> Sha2BlockHasherSubAir<C> {
39 pub fn new(bitwise_lookup_bus: BitwiseOperationLookupBus, private_bus_idx: BusIndex) -> Self {
40 Self {
41 bitwise_lookup_bus,
42 row_idx_encoder: Encoder::new(C::ROWS_PER_BLOCK + 1, 2, false), private_bus: PermutationCheckBus::new(private_bus_idx),
45 _phantom: PhantomData,
46 }
47 }
48}
49
50impl<F, C: Sha2BlockHasherSubairConfig> BaseAir<F> for Sha2BlockHasherSubAir<C> {
51 fn width(&self) -> usize {
52 C::SUBAIR_WIDTH
53 }
54}
55
56impl<AB: InteractionBuilder, C: Sha2BlockHasherSubairConfig> SubAir<AB>
57 for Sha2BlockHasherSubAir<C>
58{
59 type AirContext<'a>
61 = usize
62 where
63 Self: 'a,
64 AB: 'a,
65 <AB as AirBuilder>::Var: 'a,
66 <AB as AirBuilder>::Expr: 'a;
67
68 fn eval<'a>(&'a self, builder: &'a mut AB, start_col: Self::AirContext<'a>)
69 where
70 AB::Var: 'a,
71 AB::Expr: 'a,
72 {
73 self.eval_row(builder, start_col);
74 self.eval_transitions(builder, start_col);
75 }
76}
77
78impl<C: Sha2BlockHasherSubairConfig> Sha2BlockHasherSubAir<C> {
79 fn eval_row<AB: InteractionBuilder>(&self, builder: &mut AB, start_col: usize) {
82 let main = builder.main();
83 let local = main.row_slice(0).unwrap();
84
85 let local_cols: Sha2DigestColsRef<AB::Var> =
88 Sha2DigestColsRef::from::<C>(&local[start_col..start_col + C::SUBAIR_DIGEST_WIDTH]);
89 let flags = &local_cols.flags;
90 builder.assert_bool(*flags.is_round_row);
91 builder.assert_bool(*flags.is_first_4_rows);
92 builder.assert_bool(*flags.is_digest_row);
93 builder.assert_bool(*flags.is_round_row + *flags.is_digest_row);
94
95 self.row_idx_encoder
96 .eval(builder, local_cols.flags.row_idx.to_slice().unwrap());
97 builder.assert_one(self.row_idx_encoder.contains_flag_range::<AB>(
99 local_cols.flags.row_idx.to_slice().unwrap(),
100 0..=C::ROWS_PER_BLOCK,
101 ));
102 builder.assert_eq(
104 self.row_idx_encoder
105 .contains_flag_range::<AB>(local_cols.flags.row_idx.to_slice().unwrap(), 0..=3),
106 *flags.is_first_4_rows,
107 );
108 builder.assert_eq(
110 self.row_idx_encoder.contains_flag_range::<AB>(
111 local_cols.flags.row_idx.to_slice().unwrap(),
112 0..=C::ROUND_ROWS - 1,
113 ),
114 *flags.is_round_row,
115 );
116 builder.assert_eq(
118 self.row_idx_encoder.contains_flag::<AB>(
119 local_cols.flags.row_idx.to_slice().unwrap(),
120 &[C::ROUND_ROWS],
121 ),
122 *flags.is_digest_row,
123 );
124 builder.assert_eq(
126 self.row_idx_encoder.contains_flag::<AB>(
127 local_cols.flags.row_idx.to_slice().unwrap(),
128 &[C::ROWS_PER_BLOCK],
129 ),
130 flags.is_padding_row(),
131 );
132
133 for i in 0..C::ROUNDS_PER_ROW {
136 for j in 0..C::WORD_BITS {
137 builder.assert_bool(local_cols.hash.a[[i, j]]);
138 builder.assert_bool(local_cols.hash.e[[i, j]]);
139 }
140 }
141 }
142
143 fn eval_digest_row<AB: InteractionBuilder>(
145 &self,
146 builder: &mut AB,
147 local: Sha2RoundColsRef<AB::Var>,
148 next: Sha2DigestColsRef<AB::Var>,
149 ) {
150 for i in 0..C::HASH_WORDS {
154 let mut carry = AB::Expr::ZERO;
155 for j in 0..C::WORD_U16S {
156 let work_var_limb = if i < C::ROUNDS_PER_ROW {
157 compose::<AB::Expr>(
158 local
159 .work_vars
160 .a
161 .slice(s![C::ROUNDS_PER_ROW - 1 - i, j * 16..(j + 1) * 16])
162 .as_slice()
163 .unwrap(),
164 1,
165 )
166 } else {
167 compose::<AB::Expr>(
168 local
169 .work_vars
170 .e
171 .slice(s![C::ROUNDS_PER_ROW + 3 - i, j * 16..(j + 1) * 16])
172 .as_slice()
173 .unwrap(),
174 1,
175 )
176 };
177 let final_hash_limb = compose::<AB::Expr>(
178 next.final_hash
179 .slice(s![i, j * 2..(j + 1) * 2])
180 .as_slice()
181 .unwrap(),
182 8,
183 );
184
185 carry = AB::Expr::from(AB::F::from_u32(1 << 16).inverse())
186 * (next.prev_hash[[i, j]] + work_var_limb + carry - final_hash_limb);
187 builder
188 .when(*next.flags.is_digest_row)
189 .assert_bool(carry.clone());
190 }
191 for chunk in next.final_hash.row(i).as_slice().unwrap().chunks(2) {
194 self.bitwise_lookup_bus
195 .send_range(chunk[0], chunk[1])
196 .eval(builder, *next.flags.is_digest_row);
197 }
198 }
199 }
200
201 fn eval_transitions<AB: InteractionBuilder>(&self, builder: &mut AB, start_col: usize) {
202 let main = builder.main();
203 let local = main.row_slice(0).unwrap();
204 let next = main.row_slice(1).unwrap();
205
206 let local_cols: Sha2RoundColsRef<AB::Var> =
208 Sha2RoundColsRef::from::<C>(&local[start_col..start_col + C::SUBAIR_ROUND_WIDTH]);
209 let next_cols: Sha2RoundColsRef<AB::Var> =
210 Sha2RoundColsRef::from::<C>(&next[start_col..start_col + C::SUBAIR_ROUND_WIDTH]);
211
212 let local_is_padding_row = local_cols.flags.is_padding_row();
213 let next_is_padding_row = next_cols.flags.is_padding_row();
217
218 builder
220 .when(*local_cols.flags.is_round_row)
221 .assert_zero(next_is_padding_row.clone());
222 builder
224 .when_first_row()
225 .assert_one(*local_cols.flags.is_round_row);
226 builder
229 .when_last_row()
230 .assert_one(local_is_padding_row.clone());
231 builder
233 .when_transition()
234 .when(local_is_padding_row.clone())
235 .assert_one(next_is_padding_row.clone());
236 builder
238 .when(*local_cols.flags.is_digest_row)
239 .assert_zero(*next_cols.flags.is_digest_row);
240 let delta = *local_cols.flags.is_round_row * AB::Expr::ONE
248 + *local_cols.flags.is_digest_row
249 * *next_cols.flags.is_round_row
250 * AB::Expr::from_usize(C::ROUND_ROWS)
251 * AB::Expr::NEG_ONE
252 + *local_cols.flags.is_digest_row * next_is_padding_row.clone() * AB::Expr::ONE;
253
254 let local_row_idx = self.row_idx_encoder.flag_with_val::<AB>(
255 local_cols.flags.row_idx.to_slice().unwrap(),
256 &(0..=C::ROWS_PER_BLOCK).map(|i| (i, i)).collect::<Vec<_>>(),
257 );
258 let next_row_idx = self.row_idx_encoder.flag_with_val::<AB>(
259 next_cols.flags.row_idx.to_slice().unwrap(),
260 &(0..=C::ROWS_PER_BLOCK).map(|i| (i, i)).collect::<Vec<_>>(),
261 );
262
263 builder
264 .when_transition()
265 .assert_eq(local_row_idx.clone() + delta, next_row_idx.clone());
266 builder.when_first_row().assert_zero(local_row_idx);
267
268 builder
270 .when_first_row()
271 .assert_one(*local_cols.flags.global_block_idx);
272
273 builder.when(*local_cols.flags.is_round_row).assert_eq(
275 *local_cols.flags.global_block_idx,
276 *next_cols.flags.global_block_idx,
277 );
278 builder
280 .when_transition()
281 .when(*local_cols.flags.is_digest_row)
282 .assert_eq(
283 *local_cols.flags.global_block_idx + AB::Expr::ONE,
284 *next_cols.flags.global_block_idx,
285 );
286 builder
288 .when_transition()
289 .when(local_is_padding_row.clone())
290 .assert_eq(
291 *local_cols.flags.global_block_idx,
292 *next_cols.flags.global_block_idx,
293 );
294
295 for i in 0..C::ROUNDS_PER_ROW {
303 for j in 0..C::WORD_BITS {
304 builder.when(next_cols.flags.is_padding_row()).assert_eq(
305 local_cols.work_vars.a[[i, j]],
306 next_cols.work_vars.a[[i, j]],
307 );
308 builder.when(next_cols.flags.is_padding_row()).assert_eq(
309 local_cols.work_vars.e[[i, j]],
310 next_cols.work_vars.e[[i, j]],
311 );
312 }
313 }
314
315 self.eval_message_schedule(builder, local_cols.clone(), next_cols.clone());
316 self.eval_work_vars(builder, local_cols.clone(), next_cols);
317 let next: Sha2DigestColsRef<AB::Var> =
318 Sha2DigestColsRef::from::<C>(&next[start_col..start_col + C::SUBAIR_DIGEST_WIDTH]);
319 self.eval_digest_row(builder, local_cols, next);
320 let local_cols: Sha2DigestColsRef<AB::Var> =
321 Sha2DigestColsRef::from::<C>(&local[start_col..start_col + C::SUBAIR_DIGEST_WIDTH]);
322 self.eval_prev_hash(builder, local_cols, next_is_padding_row);
323 }
324
325 fn eval_prev_hash<AB: InteractionBuilder>(
328 &self,
329 builder: &mut AB,
330 local: Sha2DigestColsRef<AB::Var>,
331 is_last_block_of_trace: AB::Expr, ) {
334 let composed_hash = (0..C::HASH_WORDS)
336 .map(|i| {
337 let hash_bits = if i < C::ROUNDS_PER_ROW {
338 local
339 .hash
340 .a
341 .row(C::ROUNDS_PER_ROW - 1 - i)
342 .mapv(|x| x.into())
343 .to_vec()
344 } else {
345 local
346 .hash
347 .e
348 .row(C::ROUNDS_PER_ROW + 3 - i)
349 .mapv(|x| x.into())
350 .to_vec()
351 };
352 (0..C::WORD_U16S)
353 .map(|j| compose::<AB::Expr>(&hash_bits[j * 16..(j + 1) * 16], 1))
354 .collect::<Vec<_>>()
355 })
356 .collect::<Vec<_>>();
357 let next_global_block_idx = select(
359 is_last_block_of_trace,
360 AB::Expr::ONE,
361 *local.flags.global_block_idx + AB::Expr::ONE,
362 );
363 self.private_bus.send(
365 builder,
366 composed_hash
367 .into_iter()
368 .flatten()
369 .chain(once(next_global_block_idx)),
370 *local.flags.is_digest_row,
371 );
372
373 self.private_bus.receive(
374 builder,
375 local
376 .prev_hash
377 .flatten()
378 .mapv(|x| x.into())
379 .into_iter()
380 .chain(once((*local.flags.global_block_idx).into())),
381 *local.flags.is_digest_row,
382 );
383 }
384
385 fn eval_message_schedule<'a, AB: InteractionBuilder>(
390 &self,
391 builder: &mut AB,
392 local: Sha2RoundColsRef<'a, AB::Var>,
393 next: Sha2RoundColsRef<'a, AB::Var>,
394 ) {
395 let w = ndarray::concatenate(
397 ndarray::Axis(0),
398 &[local.message_schedule.w, next.message_schedule.w],
399 )
400 .unwrap();
401
402 for i in 0..C::ROUNDS_PER_ROW - 1 {
404 let w_3 = w.row(i + 1).mapv(|x| x.into()).to_vec();
407 let expected_w_3 = next.schedule_helper.w_3.row(i);
408 for j in 0..C::WORD_U16S {
409 let w_3_limb = compose::<AB::Expr>(&w_3[j * 16..(j + 1) * 16], 1);
410 builder
411 .when(*local.flags.is_round_row)
412 .assert_eq(w_3_limb, expected_w_3[j].into());
413 }
414 }
415
416 let is_row_intermed_12 = self.row_idx_encoder.contains_flag_range::<AB>(
421 next.flags.row_idx.to_slice().unwrap(),
422 3..=C::ROUND_ROWS - 2,
423 );
424 let is_row_intermed_8 = self.row_idx_encoder.contains_flag_range::<AB>(
427 next.flags.row_idx.to_slice().unwrap(),
428 2..=C::ROUND_ROWS - 3,
429 );
430 for i in 0..C::ROUNDS_PER_ROW {
431 let w_idx = w.row(i).mapv(|x| x.into()).to_vec();
433 let sig_w = small_sig0_field::<AB::Expr, C>(w.row(i + 1).as_slice().unwrap());
435 for j in 0..C::WORD_U16S {
436 let w_idx_limb = compose::<AB::Expr>(&w_idx[j * 16..(j + 1) * 16], 1);
437 let sig_w_limb = compose::<AB::Expr>(&sig_w[j * 16..(j + 1) * 16], 1);
438
439 builder.assert_eq(
444 next.schedule_helper.intermed_4[[i, j]],
445 w_idx_limb + sig_w_limb,
446 );
447
448 builder.when(is_row_intermed_8.clone()).assert_eq(
449 next.schedule_helper.intermed_8[[i, j]],
450 local.schedule_helper.intermed_4[[i, j]],
451 );
452
453 builder.when(is_row_intermed_12.clone()).assert_eq(
454 next.schedule_helper.intermed_12[[i, j]],
455 local.schedule_helper.intermed_8[[i, j]],
456 );
457 }
458 }
459
460 for i in 0..C::ROUNDS_PER_ROW {
462 let w_7 = if i < 3 {
465 local.schedule_helper.w_3.row(i).mapv(|x| x.into()).to_vec()
466 } else {
467 let w_3 = w.row(i - 3).mapv(|x| x.into()).to_vec();
468 (0..C::WORD_U16S)
469 .map(|j| compose::<AB::Expr>(&w_3[j * 16..(j + 1) * 16], 1))
470 .collect::<Vec<_>>()
471 };
472 let intermed_16 = local.schedule_helper.intermed_12.row(i).mapv(|x| x.into());
474
475 let carries = (0..C::WORD_U16S)
476 .map(|j| {
477 next.message_schedule.carry_or_buffer[[i, j * 2]]
478 + AB::Expr::TWO * next.message_schedule.carry_or_buffer[[i, j * 2 + 1]]
479 })
480 .collect::<Vec<_>>();
481
482 constraint_word_addition::<_, C>(
488 builder,
489 &[&small_sig1_field::<AB::Expr, C>(
490 w.row(i + 2).as_slice().unwrap(),
491 )],
492 &[&w_7, intermed_16.as_slice().unwrap()],
493 w.row(i + 4).as_slice().unwrap(),
494 &carries,
495 );
496
497 for j in 0..C::WORD_U16S {
498 let is_row_4_or_more = *next.flags.is_round_row - *next.flags.is_first_4_rows;
500 builder
501 .when(is_row_4_or_more.clone())
502 .assert_bool(next.message_schedule.carry_or_buffer[[i, j * 2]]);
503 builder
504 .when(is_row_4_or_more)
505 .assert_bool(next.message_schedule.carry_or_buffer[[i, j * 2 + 1]]);
506 }
507 for j in 0..C::WORD_BITS {
509 builder
510 .when(*next.flags.is_round_row)
511 .assert_bool(next.message_schedule.w[[i, j]]);
512 }
513 }
514 }
515
516 fn eval_work_vars<'a, AB: InteractionBuilder>(
519 &self,
520 builder: &mut AB,
521 local: Sha2RoundColsRef<'a, AB::Var>,
522 next: Sha2RoundColsRef<'a, AB::Var>,
523 ) {
524 let a =
525 ndarray::concatenate(ndarray::Axis(0), &[local.work_vars.a, next.work_vars.a]).unwrap();
526 let e =
527 ndarray::concatenate(ndarray::Axis(0), &[local.work_vars.e, next.work_vars.e]).unwrap();
528
529 for i in 0..C::ROUNDS_PER_ROW {
530 for j in 0..C::WORD_U16S {
531 self.bitwise_lookup_bus
535 .send_range(
536 local.work_vars.carry_a[[i, j]],
537 local.work_vars.carry_e[[i, j]],
538 )
539 .eval(builder, *local.flags.is_round_row);
540 }
541
542 let w_limbs = (0..C::WORD_U16S)
543 .map(|j| {
544 compose::<AB::Expr>(
545 next.message_schedule
546 .w
547 .slice(s![i, j * 16..(j + 1) * 16])
548 .as_slice()
549 .unwrap(),
550 1,
551 ) * *next.flags.is_round_row
552 })
553 .collect::<Vec<_>>();
554
555 let k_limbs = (0..C::WORD_U16S)
556 .map(|j| {
557 self.row_idx_encoder.flag_with_val::<AB>(
558 next.flags.row_idx.to_slice().unwrap(),
559 &(0..C::ROUND_ROWS)
560 .map(|rw_idx| {
561 (
562 rw_idx,
563 word_into_u16_limbs::<C>(
564 C::get_k()[rw_idx * C::ROUNDS_PER_ROW + i],
565 )[j] as usize,
566 )
567 })
568 .collect::<Vec<_>>(),
569 )
570 })
571 .collect::<Vec<_>>();
572
573 constraint_word_addition::<_, C>(
578 builder,
579 &[
580 e.row(i).mapv(|x| x.into()).as_slice().unwrap(), &big_sig1_field::<AB::Expr, C>(e.row(i + 3).as_slice().unwrap()), &ch_field::<AB::Expr>(
584 e.row(i + 3).as_slice().unwrap(),
585 e.row(i + 2).as_slice().unwrap(),
586 e.row(i + 1).as_slice().unwrap(),
587 ), &big_sig0_field::<AB::Expr, C>(a.row(i + 3).as_slice().unwrap()), &maj_field::<AB::Expr>(
590 a.row(i + 3).as_slice().unwrap(),
591 a.row(i + 2).as_slice().unwrap(),
592 a.row(i + 1).as_slice().unwrap(),
593 ), ],
595 &[&w_limbs, &k_limbs], a.row(i + 4).as_slice().unwrap(), next.work_vars.carry_a.row(i).as_slice().unwrap(), );
599
600 constraint_word_addition::<_, C>(
605 builder,
606 &[
607 a.row(i).mapv(|x| x.into()).as_slice().unwrap(), e.row(i).mapv(|x| x.into()).as_slice().unwrap(), &big_sig1_field::<AB::Expr, C>(e.row(i + 3).as_slice().unwrap()), &ch_field::<AB::Expr>(
612 e.row(i + 3).as_slice().unwrap(),
613 e.row(i + 2).as_slice().unwrap(),
614 e.row(i + 1).as_slice().unwrap(),
615 ), ],
617 &[&w_limbs, &k_limbs], e.row(i + 4).as_slice().unwrap(), next.work_vars.carry_e.row(i).as_slice().unwrap(), );
621 }
622 }
623}