openvm_recursion_circuit/stacking/univariate/
air.rs

1use std::borrow::Borrow;
2
3use openvm_circuit_primitives::{
4    utils::{and, assert_array_eq, not},
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::{extension::BinomiallyExtendable, PrimeCharacteristicRing, PrimeField32};
14use p3_matrix::Matrix;
15
16use crate::{
17    bus::{TranscriptBus, TranscriptBusMessage},
18    stacking::bus::{
19        EqKernelLookupBus, EqRandValuesLookupBus, EqRandValuesLookupMessage, StackingModuleTidxBus,
20        StackingModuleTidxMessage, SumcheckClaimsBus, SumcheckClaimsMessage,
21    },
22    subairs::nested_for_loop::{NestedForLoopIoCols, NestedForLoopSubAir},
23    utils::{assert_one_ext, ext_field_add, ext_field_multiply, ext_field_multiply_scalar},
24};
25
26#[repr(C)]
27#[derive(AlignedBorrow, StructReflection)]
28pub struct UnivariateRoundCols<F> {
29    // Proof index columns for continuations
30    pub proof_idx: F,
31    pub is_valid: F,
32    pub is_first: F,
33    pub is_last: F,
34
35    // Sampled transcript values
36    pub tidx: F,
37    pub u_0: [F; D_EF],
38    pub u_0_pow: [F; D_EF],
39
40    // Coefficients of univariate round (s_0) polynomial
41    pub coeff: [F; D_EF],
42
43    // Columns to compute s_0(z) sum over all z in D
44    pub coeff_idx: F,
45    pub coeff_is_d: F,
46    pub s_0_sum_over_d: [F; D_EF],
47
48    // Evaluation of s_0 polynomial at u_0
49    pub poly_rand_eval: [F; D_EF],
50}
51
52#[derive(ColumnsAir)]
53#[columns_via(UnivariateRoundCols<u8>)]
54pub struct UnivariateRoundAir {
55    // External buses
56    pub transcript_bus: TranscriptBus,
57
58    // Internal buses
59    pub stacking_tidx_bus: StackingModuleTidxBus,
60    pub sumcheck_claims_bus: SumcheckClaimsBus,
61    pub eq_rand_values_bus: EqRandValuesLookupBus,
62    pub eq_kernel_lookup_bus: EqKernelLookupBus,
63
64    // Other fields
65    pub l_skip: usize,
66}
67
68impl BaseAirWithPublicValues<F> for UnivariateRoundAir {}
69impl PartitionedBaseAir<F> for UnivariateRoundAir {}
70
71impl<F> BaseAir<F> for UnivariateRoundAir {
72    fn width(&self) -> usize {
73        UnivariateRoundCols::<F>::width()
74    }
75}
76
77impl<AB: AirBuilder + InteractionBuilder> Air<AB> for UnivariateRoundAir
78where
79    AB::F: PrimeField32,
80    <AB::Expr as PrimeCharacteristicRing>::PrimeSubfield: BinomiallyExtendable<{ D_EF }>,
81{
82    fn eval(&self, builder: &mut AB) {
83        let main = builder.main();
84        let (local, next) = (
85            main.row_slice(0).expect("window should have two elements"),
86            main.row_slice(1).expect("window should have two elements"),
87        );
88
89        let local: &UnivariateRoundCols<AB::Var> = (*local).borrow();
90        let next: &UnivariateRoundCols<AB::Var> = (*next).borrow();
91
92        NestedForLoopSubAir::<1> {}.eval(
93            builder,
94            (
95                NestedForLoopIoCols {
96                    is_enabled: local.is_valid,
97                    counter: [local.proof_idx],
98                    is_first: [local.is_first],
99                }
100                .map_into(),
101                NestedForLoopIoCols {
102                    is_enabled: next.is_valid,
103                    counter: [next.proof_idx],
104                    is_first: [next.is_first],
105                }
106                .map_into(),
107            ),
108        );
109
110        builder.when(local.is_valid).assert_eq(
111            local.is_last,
112            NestedForLoopSubAir::<1>::local_is_last(local.is_valid, next.is_valid, next.is_first),
113        );
114
115        builder.assert_bool(local.is_last);
116        builder
117            .when(and(local.is_valid, local.is_last))
118            .assert_zero((local.proof_idx + AB::F::ONE - next.proof_idx) * next.proof_idx);
119        builder
120            .when(and(not(local.is_valid), local.is_last))
121            .assert_zero(next.proof_idx);
122
123        /*
124         * Constrain that the sum of s_0(z) for z in D via interaction equals the RLC of column
125         * claims from OpeningClaimsAir. We use the properties of D to do this efficiently -
126         * since D is a multiplicative subgroup, it turns out this sum is |D| * (a_0 + a_{|D|}).
127         */
128        let d_card = 1usize << self.l_skip;
129
130        builder.when(local.is_first).assert_zero(local.coeff_idx);
131        builder
132            .when(and(local.is_last, local.is_valid))
133            .assert_eq(local.coeff_idx, AB::F::from_usize(2 * (d_card - 1)));
134        builder
135            .when(and(not(local.is_last), local.is_valid))
136            .assert_one(next.coeff_idx - local.coeff_idx);
137
138        builder.assert_bool(local.coeff_is_d);
139        builder.when(local.coeff_is_d).assert_one(local.is_valid);
140        builder
141            .when(local.coeff_is_d)
142            .assert_eq(local.coeff_idx, AB::F::from_usize(d_card));
143
144        assert_array_eq(
145            &mut builder.when(local.is_first),
146            ext_field_multiply_scalar(local.coeff, AB::F::from_usize(d_card)),
147            local.s_0_sum_over_d,
148        );
149
150        assert_array_eq(
151            &mut builder.when(next.coeff_is_d),
152            ext_field_add(
153                local.s_0_sum_over_d,
154                ext_field_multiply_scalar(next.coeff, AB::F::from_usize(d_card)),
155            ),
156            next.s_0_sum_over_d,
157        );
158
159        assert_array_eq(
160            &mut builder.when(and::<AB::Expr>(not(next.coeff_is_d), not(local.is_last))),
161            local.s_0_sum_over_d,
162            next.s_0_sum_over_d,
163        );
164
165        self.sumcheck_claims_bus.receive(
166            builder,
167            next.proof_idx,
168            SumcheckClaimsMessage {
169                module_idx: AB::Expr::ZERO,
170                value: next.s_0_sum_over_d.map(Into::into),
171            },
172            next.coeff_is_d,
173        );
174
175        /*
176         * Compute evaluation of polynomial s_0(u_0) and send it to SumcheckRoundsAir, where
177         * it'll be used to constrain the correctness of s_1(0).
178         */
179        assert_one_ext(&mut builder.when(local.is_first), local.u_0_pow);
180        assert_array_eq(&mut builder.when(not(local.is_last)), local.u_0, next.u_0);
181
182        assert_array_eq(
183            &mut builder.when(not(local.is_last)),
184            ext_field_multiply(local.u_0, local.u_0_pow),
185            next.u_0_pow,
186        );
187
188        assert_array_eq(
189            &mut builder.when(local.is_first),
190            ext_field_multiply(local.coeff, local.u_0_pow),
191            local.poly_rand_eval,
192        );
193
194        assert_array_eq(
195            &mut builder.when(not(local.is_last)),
196            ext_field_add(
197                local.poly_rand_eval,
198                ext_field_multiply(next.coeff, next.u_0_pow),
199            ),
200            next.poly_rand_eval,
201        );
202
203        self.sumcheck_claims_bus.send(
204            builder,
205            local.proof_idx,
206            SumcheckClaimsMessage {
207                module_idx: AB::Expr::ONE,
208                value: local.poly_rand_eval.map(Into::into),
209            },
210            and(local.is_last, local.is_valid),
211        );
212
213        /*
214         * Because we sample u_0 from the transcript here, we send u_0 to other AIRs that
215         * need to use it.
216         */
217        self.eq_rand_values_bus.add_key_with_lookups(
218            builder,
219            local.proof_idx,
220            EqRandValuesLookupMessage {
221                idx: AB::Expr::ZERO,
222                u: local.u_0.map(Into::into),
223            },
224            and(local.is_last, local.is_valid) * AB::F::TWO,
225        );
226
227        /*
228         * Constrain transcript operations and send the final tidx to SumcheckRoundsAir.
229         */
230        let mut when_same_proof = builder.when(and(local.is_valid, not(local.is_last)));
231        when_same_proof.assert_one(next.is_valid);
232        when_same_proof.assert_eq(local.tidx + AB::Expr::from_usize(D_EF), next.tidx);
233
234        self.stacking_tidx_bus.receive(
235            builder,
236            local.proof_idx,
237            StackingModuleTidxMessage {
238                module_idx: AB::Expr::ZERO,
239                tidx: local.tidx.into(),
240            },
241            local.is_first,
242        );
243
244        for i in 0..D_EF {
245            self.transcript_bus.receive(
246                builder,
247                local.proof_idx,
248                TranscriptBusMessage {
249                    tidx: AB::Expr::from_usize(i) + local.tidx,
250                    value: local.coeff[i].into(),
251                    is_sample: AB::Expr::ZERO,
252                },
253                local.is_valid,
254            );
255
256            self.transcript_bus.receive(
257                builder,
258                local.proof_idx,
259                TranscriptBusMessage {
260                    tidx: AB::Expr::from_usize(i + D_EF) + local.tidx,
261                    value: local.u_0[i].into(),
262                    is_sample: AB::Expr::ONE,
263                },
264                and(local.is_last, local.is_valid),
265            );
266        }
267
268        self.stacking_tidx_bus.send(
269            builder,
270            local.proof_idx,
271            StackingModuleTidxMessage {
272                module_idx: AB::Expr::ONE,
273                tidx: AB::Expr::from_usize(2 * D_EF) + local.tidx,
274            },
275            and(local.is_last, local.is_valid),
276        );
277    }
278}