openvm_recursion_circuit/stacking/eq_base/
air.rs

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    // Proof index columns for continuations
39    pub proof_idx: F,
40    pub is_valid: F,
41    pub is_first: F,
42    pub is_last: F,
43
44    // Row index for the given proof, in {0, 1, ..., l_skip}
45    pub row_idx: F,
46
47    // Value of u^{2^row}, r^{2^row}, and (r * omega)^{2^row}
48    pub u_pow: [F; D_EF],
49    pub r_pow: [F; D_EF],
50    pub r_omega_pow: [F; D_EF],
51
52    // Running product of (u^{2^i} + r^{2^i}) from i to row
53    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    // Lookup multiplicity for eq_0(u, r) and k_rot_0(u, r)
59    pub mult: F,
60
61    // Value of u^{2^{l_skip + n}}, where we set n = -row_idx
62    pub u_pow_rev: [F; D_EF],
63
64    // Values of eq_n(u, r) and k_rot_n(u, r) for each negative n
65    pub eq_neg: [F; D_EF],
66    pub k_rot_neg: [F; D_EF],
67
68    // Value of in_n(u) for n in {-1, -2, ..., -l_skip}, i.e. the running product of
69    // each (u_pow_rev + 1)
70    pub in_prod: [F; D_EF],
71
72    // Lookup multiplicity for eq_n(u, r) and k_rot_n(u, r)
73    pub mult_neg: F,
74}
75
76#[derive(ColumnsAir)]
77#[columns_via(EqBaseCols<u8>)]
78pub struct EqBaseAir {
79    // External buses
80    pub constraint_randomness_bus: ConstraintSumcheckRandomnessBus,
81    pub whir_opening_point_bus: WhirOpeningPointBus,
82
83    // Internal buses
84    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    // Other fields
91    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        /*
151         * Constrain value of row_idx and send u^{2^row} to WhirOpeningPointBus when
152         * row_idx < l_skip.
153         */
154        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        /*
175         * Receive the values of u_0 and r_0 from the AIRs that sample them. Send u_0
176         * and r_0 to EqNegAir.
177         */
178        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        /*
212         * Constrain the running product of (u^{2^i} + r^{2^i}) from i to row, which is
213         * used to compute eq_0(u, r).
214         */
215        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        /*
240         * Constrain the running product that is used to compute eq_0(u, r * omega).
241         */
242        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        /*
272         * Constrain the running products that are used to compute eq_0(u, 1)
273         * and eq_0(r * omega, 1).
274         */
275        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        /*
305         * Compute eq_0(u, r) and eq_0(u, r * omega), which are sent to the lookup
306         * bus. Note that k_rot_0(u, r) = eq_0(u, r * omega).
307         */
308        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        /*
334         * Compute eq_0(u, 1) and eq_0(r * omega, 1) and send to SumcheckRoundsAir.
335         */
336        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        /*
351         * Compute eq_n(u, r), k_rot_n(u, r), and in_n(u), which are used to
352         * provide the eq and k_rot lookups for n < 0.
353         */
354        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}