1use std::borrow::Borrow;
2
3use openvm_circuit_primitives::{
4 utils::{and, assert_array_eq, not, or},
5 ColumnsAir, StructReflection, StructReflectionHelper, SubAir,
6};
7use openvm_recursion_circuit_derive::AlignedBorrow;
8use openvm_stark_backend::{
9 interaction::InteractionBuilder, BaseAirWithPublicValues, PartitionedBaseAir,
10};
11use openvm_stark_sdk::config::baby_bear_poseidon2::{D_EF, F};
12use p3_air::{Air, AirBuilder, BaseAir};
13use p3_field::{
14 extension::BinomiallyExtendable, Field, PrimeCharacteristicRing, PrimeField32, TwoAdicField,
15};
16use p3_matrix::Matrix;
17
18use crate::{
19 bus::{
20 ConstraintSumcheckRandomness, ConstraintSumcheckRandomnessBus, EqNegBaseRandBus,
21 EqNegBaseRandMessage, EqNegResultBus, EqNegResultMessage, WhirOpeningPointBus,
22 WhirOpeningPointMessage,
23 },
24 stacking::bus::{
25 EqBaseBus, EqBaseMessage, EqKernelLookupBus, EqKernelLookupMessage, EqRandValuesLookupBus,
26 EqRandValuesLookupMessage,
27 },
28 subairs::nested_for_loop::{NestedForLoopIoCols, NestedForLoopSubAir},
29 utils::{
30 assert_one_ext, ext_field_add, ext_field_add_scalar, ext_field_multiply,
31 ext_field_multiply_scalar, ext_field_subtract,
32 },
33};
34
35#[repr(C)]
36#[derive(AlignedBorrow, StructReflection)]
37pub struct EqBaseCols<F> {
38 pub proof_idx: F,
40 pub is_valid: F,
41 pub is_first: F,
42 pub is_last: F,
43
44 pub row_idx: F,
46
47 pub u_pow: [F; D_EF],
49 pub r_pow: [F; D_EF],
50 pub r_omega_pow: [F; D_EF],
51
52 pub prod_u_r: [F; D_EF],
54 pub prod_u_r_omega: [F; D_EF],
55 pub prod_u_1: [F; D_EF],
56 pub prod_r_omega_1: [F; D_EF],
57
58 pub mult: F,
60
61 pub u_pow_rev: [F; D_EF],
63
64 pub eq_neg: [F; D_EF],
66 pub k_rot_neg: [F; D_EF],
67
68 pub in_prod: [F; D_EF],
71
72 pub mult_neg: F,
74}
75
76#[derive(ColumnsAir)]
77#[columns_via(EqBaseCols<u8>)]
78pub struct EqBaseAir {
79 pub constraint_randomness_bus: ConstraintSumcheckRandomnessBus,
81 pub whir_opening_point_bus: WhirOpeningPointBus,
82
83 pub eq_base_bus: EqBaseBus,
85 pub eq_rand_values_bus: EqRandValuesLookupBus,
86 pub eq_kernel_lookup_bus: EqKernelLookupBus,
87 pub eq_neg_base_rand_bus: EqNegBaseRandBus,
88 pub eq_neg_result_bus: EqNegResultBus,
89
90 pub l_skip: usize,
92}
93
94impl BaseAirWithPublicValues<F> for EqBaseAir {}
95impl PartitionedBaseAir<F> for EqBaseAir {}
96
97impl<F> BaseAir<F> for EqBaseAir {
98 fn width(&self) -> usize {
99 EqBaseCols::<F>::width()
100 }
101}
102
103impl<AB: AirBuilder + InteractionBuilder> Air<AB> for EqBaseAir
104where
105 AB::F: PrimeField32 + TwoAdicField,
106 <AB::Expr as PrimeCharacteristicRing>::PrimeSubfield: BinomiallyExtendable<{ D_EF }>,
107{
108 fn eval(&self, builder: &mut AB) {
109 let main = builder.main();
110 let (local, next) = (
111 main.row_slice(0).expect("window should have two elements"),
112 main.row_slice(1).expect("window should have two elements"),
113 );
114
115 let local: &EqBaseCols<AB::Var> = (*local).borrow();
116 let next: &EqBaseCols<AB::Var> = (*next).borrow();
117
118 NestedForLoopSubAir::<1> {}.eval(
119 builder,
120 (
121 NestedForLoopIoCols {
122 is_enabled: local.is_valid,
123 counter: [local.proof_idx],
124 is_first: [local.is_first],
125 }
126 .map_into(),
127 NestedForLoopIoCols {
128 is_enabled: next.is_valid,
129 counter: [next.proof_idx],
130 is_first: [next.is_first],
131 }
132 .map_into(),
133 ),
134 );
135
136 builder.when(local.is_valid).assert_eq(
137 local.is_last,
138 NestedForLoopSubAir::<1>::local_is_last(local.is_valid, next.is_valid, next.is_first),
139 );
140
141 builder.assert_bool(local.is_last);
142 builder
143 .when(and(local.is_valid, local.is_last))
144 .assert_zero((local.proof_idx + AB::F::ONE - next.proof_idx) * next.proof_idx);
145 builder
146 .when(and(not(local.is_valid), local.is_last))
147 .assert_zero(next.proof_idx);
148 builder.assert_zero(local.is_first * local.is_last);
149
150 let is_valid_transition = and(local.is_valid, not(local.is_last));
155
156 builder.when(local.is_first).assert_zero(local.row_idx);
157 builder
158 .when(is_valid_transition.clone())
159 .assert_eq(local.row_idx + AB::F::ONE, next.row_idx);
160 builder
161 .when(and(local.is_valid, local.is_last))
162 .assert_eq(local.row_idx, AB::F::from_usize(self.l_skip));
163
164 self.whir_opening_point_bus.send(
165 builder,
166 local.proof_idx,
167 WhirOpeningPointMessage {
168 idx: local.row_idx,
169 value: local.u_pow,
170 },
171 is_valid_transition,
172 );
173
174 self.constraint_randomness_bus.receive(
179 builder,
180 local.proof_idx,
181 ConstraintSumcheckRandomness {
182 idx: AB::Expr::ZERO,
183 challenge: local.r_pow.map(Into::into),
184 },
185 local.is_first,
186 );
187
188 self.eq_rand_values_bus.lookup_key(
189 builder,
190 local.proof_idx,
191 EqRandValuesLookupMessage {
192 idx: AB::Expr::ZERO,
193 u: ext_field_add(
194 ext_field_multiply_scalar(local.u_pow, local.is_first),
195 ext_field_multiply_scalar(local.u_pow_rev, local.is_last),
196 ),
197 },
198 and(local.is_valid, local.is_first + local.is_last),
199 );
200
201 self.eq_neg_base_rand_bus.send(
202 builder,
203 local.proof_idx,
204 EqNegBaseRandMessage {
205 u: local.u_pow,
206 r: local.r_pow,
207 },
208 local.is_first,
209 );
210
211 assert_array_eq(
216 &mut builder.when(local.is_first),
217 local.prod_u_r,
218 ext_field_multiply(local.u_pow, ext_field_add(local.u_pow, local.r_pow)),
219 );
220
221 assert_array_eq(
222 &mut builder.when(not(local.is_last)),
223 ext_field_multiply(local.u_pow, local.u_pow),
224 next.u_pow,
225 );
226
227 assert_array_eq(
228 &mut builder.when(not(local.is_last)),
229 ext_field_multiply(local.r_pow, local.r_pow),
230 next.r_pow,
231 );
232
233 assert_array_eq(
234 &mut builder.when(not(local.is_last)),
235 ext_field_multiply(local.prod_u_r, ext_field_add(next.u_pow, next.r_pow)),
236 next.prod_u_r,
237 );
238
239 let omega = AB::F::two_adic_generator(self.l_skip);
243
244 assert_array_eq(
245 &mut builder.when(local.is_first),
246 local.prod_u_r_omega,
247 ext_field_multiply(local.u_pow, ext_field_add(local.u_pow, local.r_omega_pow)),
248 );
249
250 assert_array_eq(
251 &mut builder.when(local.is_first),
252 local.r_omega_pow,
253 ext_field_multiply_scalar(local.r_pow, omega),
254 );
255
256 assert_array_eq(
257 &mut builder.when(not(local.is_last)),
258 ext_field_multiply(local.r_omega_pow, local.r_omega_pow),
259 next.r_omega_pow,
260 );
261
262 assert_array_eq(
263 &mut builder.when(not(local.is_last)),
264 ext_field_multiply(
265 local.prod_u_r_omega,
266 ext_field_add(next.u_pow, next.r_omega_pow),
267 ),
268 next.prod_u_r_omega,
269 );
270
271 let ef_one = [AB::F::ONE, AB::F::ZERO, AB::F::ZERO, AB::F::ZERO];
276
277 assert_array_eq(
278 &mut builder.when(local.is_first),
279 local.prod_u_1,
280 ext_field_add(local.u_pow, ef_one),
281 );
282
283 assert_array_eq(
284 &mut builder.when(local.is_first),
285 local.prod_r_omega_1,
286 ext_field_add(local.r_omega_pow, ef_one),
287 );
288
289 assert_array_eq(
290 &mut builder.when(not(local.is_last)),
291 ext_field_multiply(local.prod_u_1, ext_field_add(next.u_pow, ef_one)),
292 next.prod_u_1,
293 );
294
295 assert_array_eq(
296 &mut builder.when(not(local.is_last)),
297 ext_field_multiply(
298 local.prod_r_omega_1,
299 ext_field_add(next.r_omega_pow, ef_one),
300 ),
301 next.prod_r_omega_1,
302 );
303
304 let omega_pow_inv = AB::F::from_usize(1 << self.l_skip).inverse();
309
310 let eq_u_r = ext_field_multiply_scalar(
311 ext_field_add::<AB::Expr>(ext_field_subtract(local.prod_u_r, next.u_pow), ef_one),
312 omega_pow_inv,
313 );
314
315 let eq_u_r_omega = ext_field_multiply_scalar(
316 ext_field_add::<AB::Expr>(ext_field_subtract(local.prod_u_r_omega, next.u_pow), ef_one),
317 omega_pow_inv,
318 );
319
320 builder.when(next.mult).assert_one(next.is_last);
321
322 self.eq_kernel_lookup_bus.add_key_with_lookups(
323 builder,
324 local.proof_idx,
325 EqKernelLookupMessage {
326 n: AB::Expr::ZERO,
327 eq_in: eq_u_r.clone(),
328 k_rot_in: eq_u_r_omega.clone(),
329 },
330 next.is_valid * next.mult,
331 );
332
333 let eq_u_1 = ext_field_multiply_scalar(local.prod_u_1, omega_pow_inv);
337 let eq_r_omega_1 = ext_field_multiply_scalar(local.prod_r_omega_1, omega_pow_inv);
338
339 self.eq_base_bus.send(
340 builder,
341 local.proof_idx,
342 EqBaseMessage {
343 eq_u_r,
344 eq_u_r_omega,
345 eq_u_r_prod: ext_field_multiply(eq_u_1, eq_r_omega_1),
346 },
347 and(next.is_last, next.is_valid),
348 );
349
350 self.eq_neg_result_bus.receive(
355 builder,
356 local.proof_idx,
357 EqNegResultMessage {
358 n: AB::Expr::ZERO - local.row_idx,
359 eq: local.eq_neg.map(Into::into),
360 k_rot: local.k_rot_neg.map(Into::into),
361 },
362 and(
363 local.is_valid,
364 not::<AB::Expr>(or(local.is_first, local.is_last)),
365 ),
366 );
367
368 assert_one_ext(
369 &mut builder.when(and(local.is_valid, local.is_last)),
370 local.eq_neg,
371 );
372
373 assert_one_ext(
374 &mut builder.when(and(local.is_valid, local.is_last)),
375 local.k_rot_neg,
376 );
377
378 assert_array_eq(
379 &mut builder.when(not(local.is_last)),
380 ext_field_multiply(next.u_pow_rev, next.u_pow_rev),
381 local.u_pow_rev,
382 );
383
384 assert_one_ext(&mut builder.when(local.is_first), local.in_prod);
385
386 assert_array_eq(
387 &mut builder.when(not(local.is_last)),
388 ext_field_multiply(
389 local.in_prod,
390 ext_field_add_scalar(next.u_pow_rev, AB::F::ONE),
391 ),
392 next.in_prod,
393 );
394
395 builder
396 .when(not(local.is_valid))
397 .assert_zero(local.mult_neg);
398 builder.when(local.is_first).assert_zero(local.mult_neg);
399
400 let in_n = ext_field_multiply_scalar::<AB::Expr>(local.in_prod, omega_pow_inv);
401 self.eq_kernel_lookup_bus.add_key_with_lookups(
402 builder,
403 local.proof_idx,
404 EqKernelLookupMessage {
405 n: AB::Expr::ZERO - local.row_idx,
406 eq_in: ext_field_multiply(in_n.clone(), local.eq_neg),
407 k_rot_in: ext_field_multiply(in_n, local.k_rot_neg),
408 },
409 local.mult_neg,
410 );
411 }
412}