openvm_recursion_circuit/batch_constraint/expr_eval/interactions_folding/
air.rs

1use std::borrow::Borrow;
2
3use openvm_circuit_primitives::{
4    utils::{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;
12use p3_air::{Air, AirBuilder, BaseAir};
13use p3_field::{extension::BinomiallyExtendable, PrimeCharacteristicRing};
14use p3_matrix::Matrix;
15
16use crate::{
17    batch_constraint::bus::{
18        Eq3bBus, Eq3bMessage, ExpressionClaimBus, ExpressionClaimMessage, InteractionsFoldingBus,
19        InteractionsFoldingMessage,
20    },
21    bus::{
22        AirShapeBus, AirShapeBusMessage, AirShapeProperty, InteractionsFoldingInputBus,
23        InteractionsFoldingInputMessage, TranscriptBus,
24    },
25    subairs::nested_for_loop::{NestedForLoopIoCols, NestedForLoopSubAir},
26    utils::{assert_zeros, ext_field_add, ext_field_multiply},
27};
28
29#[derive(AlignedBorrow, Copy, Clone, StructReflection)]
30#[repr(C)]
31pub struct InteractionsFoldingCols<T> {
32    pub is_valid: T,
33    pub is_first: T,
34    pub proof_idx: T,
35
36    pub beta_tidx: T,
37
38    pub air_idx: T,
39    pub sort_idx: T,
40    pub interaction_idx: T,
41
42    pub has_interactions: T,
43
44    pub is_first_in_air: T,
45    /// It's true for the num row, which doesn't need to be beta folded.
46    pub is_first_in_message: T, // aka "is_mult"
47    // the second in message is the first denom, and it's cur_sum is the folded denom
48    pub is_second_in_message: T,
49    pub is_bus_index: T,
50
51    pub idx_in_message: T,
52    pub value: [T; D_EF],
53    /// Current sum for doing beta folding. This is the value for one interaction.
54    /// When local.is_first_in_message, next.cur_sum should be the folded denom.
55    /// (because local row is for the num row, which doesn't need to be beta folded)
56    /// It doesn't multiply with eq_3b yet.
57    pub cur_sum: [T; D_EF],
58    pub beta: [T; D_EF],
59    pub eq_3b: [T; D_EF],
60
61    /// The summed num and denom for all interactions.
62    /// It's summed over all the interactions in the AIR: cur_sum * eq_3b when is_first_in_message
63    pub final_acc_num: [T; D_EF],
64    pub final_acc_denom: [T; D_EF],
65}
66
67#[derive(ColumnsAir)]
68#[columns_via(InteractionsFoldingCols<u8>)]
69pub struct InteractionsFoldingAir {
70    pub interaction_bus: InteractionsFoldingBus,
71    pub interactions_folding_input_bus: InteractionsFoldingInputBus,
72    pub air_shape_bus: AirShapeBus,
73    pub transcript_bus: TranscriptBus,
74    pub expression_claim_bus: ExpressionClaimBus,
75    pub eq_3b_bus: Eq3bBus,
76}
77
78impl<F> BaseAirWithPublicValues<F> for InteractionsFoldingAir {}
79impl<F> PartitionedBaseAir<F> for InteractionsFoldingAir {}
80
81impl<F> BaseAir<F> for InteractionsFoldingAir {
82    fn width(&self) -> usize {
83        InteractionsFoldingCols::<F>::width()
84    }
85}
86
87impl<AB: AirBuilder + InteractionBuilder> Air<AB> for InteractionsFoldingAir
88where
89    <AB::Expr as PrimeCharacteristicRing>::PrimeSubfield: BinomiallyExtendable<{ D_EF }>,
90{
91    fn eval(&self, builder: &mut AB) {
92        let main = builder.main();
93        let (local, next) = (
94            main.row_slice(0).expect("window should have two elements"),
95            main.row_slice(1).expect("window should have two elements"),
96        );
97
98        let local: &InteractionsFoldingCols<AB::Var> = (*local).borrow();
99        let next: &InteractionsFoldingCols<AB::Var> = (*next).borrow();
100
101        type LoopSubAir = NestedForLoopSubAir<3>;
102        LoopSubAir {}.eval(
103            builder,
104            (
105                NestedForLoopIoCols {
106                    is_enabled: local.is_valid,
107                    counter: [local.proof_idx, local.sort_idx, local.interaction_idx],
108                    is_first: [
109                        local.is_first,
110                        local.is_first_in_air,
111                        local.is_first_in_message,
112                    ],
113                }
114                .map_into(),
115                NestedForLoopIoCols {
116                    is_enabled: next.is_valid,
117                    counter: [next.proof_idx, next.sort_idx, next.interaction_idx],
118                    is_first: [
119                        next.is_first,
120                        next.is_first_in_air,
121                        next.is_first_in_message,
122                    ],
123                }
124                .map_into(),
125            ),
126        );
127
128        builder.when_first_row().assert_zero(local.proof_idx);
129
130        builder.assert_bool(local.has_interactions);
131        builder.assert_bool(local.is_bus_index);
132        builder
133            .when(local.is_bus_index)
134            .assert_one(local.has_interactions);
135        builder
136            .when(local.has_interactions + local.is_bus_index)
137            .assert_one(local.is_valid);
138        let is_same_proof = next.is_valid - next.is_first;
139        let is_same_air = next.is_valid - next.is_first_in_air;
140        let is_same_message = next.is_valid - next.is_first_in_message;
141        let next_is_first_in_air_or_invalid =
142            next.is_first_in_air + (AB::Expr::ONE - next.is_valid);
143        let next_is_first_in_message_or_invalid =
144            next.is_first_in_message + (AB::Expr::ONE - next.is_valid);
145
146        // =========================== indices consistency ===============================
147        // When we are within one proof, sort_idx increases by 0/1
148        builder
149            .when(is_same_proof.clone())
150            .assert_bool(next.sort_idx - local.sort_idx);
151        // When we are within one AIR, interaction_idx increases by 0/1 as well
152        builder
153            .when(is_same_air.clone())
154            .assert_bool(next.interaction_idx - local.interaction_idx);
155        // First AIR within a proof is zero, and first interaction within an AIR is also zero
156        builder.when(local.is_first).assert_zero(local.sort_idx);
157        builder
158            .when(not::<AB::Expr>(is_same_air.clone()))
159            .assert_zero(next.interaction_idx);
160
161        builder.assert_bool(local.is_first_in_message + local.is_second_in_message);
162        builder
163            .when(local.is_first_in_message + local.is_second_in_message)
164            .assert_zero(local.idx_in_message);
165        builder
166            .when(is_same_message.clone())
167            .assert_one(next.idx_in_message - local.idx_in_message + local.is_first_in_message);
168
169        // // =========================== general consistency ================================
170        // air_idx is the same within an air
171        builder
172            .when(is_same_air.clone())
173            .assert_eq(local.air_idx, next.air_idx);
174        // The row describes an AIR without interactions iff it's first and last in the message,
175        // unless the row is invalid
176        builder.when(local.is_valid).assert_eq(
177            local.is_first_in_message * next_is_first_in_message_or_invalid.clone(),
178            not(local.has_interactions),
179        );
180        // If we have interactions, then the row is valid
181        builder
182            .when(local.has_interactions)
183            .assert_one(local.is_valid);
184        // If we don't have interactions and the row is valid, then it's first and last _within AIR_
185        builder
186            .when(not(local.has_interactions))
187            .when(local.is_valid)
188            .assert_one(local.is_first_in_air);
189        builder
190            .when(not(local.has_interactions))
191            .when(local.is_valid)
192            .assert_one(next_is_first_in_air_or_invalid.clone());
193        // If the row is valid, then this is the bus index iff the next one is first in message
194        // or invalid
195        builder.when(local.has_interactions).assert_eq(
196            local.is_bus_index,
197            next_is_first_in_message_or_invalid.clone(),
198        );
199        // An interaction has at least two fields (mult and bus index)
200        builder
201            .when(local.has_interactions)
202            .assert_bool(local.is_bus_index + local.is_first_in_message);
203        // final_acc_num only changes when it's first in message
204        assert_array_eq(
205            &mut builder
206                .when(not(local.is_first_in_message) * local.is_valid * is_same_air.clone()),
207            local.final_acc_num,
208            next.final_acc_num,
209        );
210        assert_array_eq(
211            &mut builder.when(local.is_first_in_message * local.has_interactions),
212            local.final_acc_num,
213            ext_field_add(
214                next.final_acc_num,
215                ext_field_multiply(local.cur_sum, local.eq_3b),
216            ),
217        );
218        assert_zeros(
219            &mut builder
220                .when(local.is_first_in_message * (local.is_valid - local.has_interactions)),
221            local.final_acc_num,
222        );
223        assert_zeros(
224            &mut builder
225                .when(local.is_first_in_message * (local.is_valid - local.has_interactions)),
226            local.final_acc_denom,
227        );
228        // final_acc_denom only changes when it's second in message
229        assert_array_eq(
230            &mut builder.when(
231                (not(local.is_second_in_message) + not(local.has_interactions))
232                    * local.is_valid
233                    * is_same_air.clone(),
234            ),
235            local.final_acc_denom,
236            next.final_acc_denom,
237        );
238        assert_array_eq(
239            &mut builder.when(local.is_second_in_message * local.is_valid),
240            local.final_acc_denom,
241            ext_field_add(
242                next.final_acc_denom,
243                ext_field_multiply(local.cur_sum, local.eq_3b),
244            ),
245        );
246        assert_array_eq(
247            &mut builder.when(is_same_message.clone()),
248            local.eq_3b,
249            next.eq_3b,
250        );
251        // the running sums are zero on the last row of the proof
252        assert_zeros(
253            &mut builder.when(LoopSubAir::local_is_last(
254                local.is_valid,
255                next.is_valid,
256                next.is_first_in_air,
257            )),
258            local.final_acc_num,
259        );
260        assert_zeros(
261            &mut builder.when(LoopSubAir::local_is_last(
262                local.is_valid,
263                next.is_valid,
264                next.is_first_in_air,
265            )),
266            local.final_acc_denom,
267        );
268        // Constraint is_second_in_message
269        builder.assert_bool(local.is_second_in_message);
270        builder
271            .when(local.is_first_in_message * local.has_interactions)
272            .assert_one(next.is_second_in_message);
273        builder
274            .when(next.is_second_in_message)
275            .assert_one(local.is_first_in_message);
276
277        // ======================== beta and cur sum consistency ============================
278        assert_array_eq(&mut builder.when(is_same_proof), local.beta, next.beta);
279        assert_array_eq(
280            &mut builder.when(is_same_message * not(local.is_first_in_message)),
281            local.cur_sum,
282            ext_field_add(
283                local.value,
284                ext_field_multiply::<AB::Expr>(local.beta, next.cur_sum),
285            ),
286        );
287        // numerator and the last element of the message are just the corresponding values
288        assert_array_eq(
289            &mut builder.when(next_is_first_in_message_or_invalid + local.is_first_in_message),
290            local.cur_sum,
291            local.value,
292        );
293
294        self.expression_claim_bus.send(
295            builder,
296            local.proof_idx,
297            ExpressionClaimMessage {
298                is_interaction: AB::Expr::ONE,
299                idx: local.sort_idx * AB::Expr::TWO,
300                // value: ext_field_multiply(local.cur_sum, local.eq_3b),
301                value: local.final_acc_num.map(Into::into),
302            },
303            local.is_first_in_air * local.is_valid,
304        );
305        self.expression_claim_bus.send(
306            builder,
307            local.proof_idx,
308            ExpressionClaimMessage {
309                is_interaction: AB::Expr::ONE,
310                idx: local.sort_idx * AB::Expr::TWO + AB::Expr::ONE,
311                // value: ext_field_multiply(next.cur_sum, next.eq_3b),
312                value: local.final_acc_denom.map(Into::into),
313            },
314            local.is_first_in_air * local.is_valid,
315        );
316        self.interaction_bus.receive(
317            builder,
318            local.proof_idx,
319            InteractionsFoldingMessage {
320                air_idx: local.air_idx.into(),
321                interaction_idx: local.interaction_idx.into(),
322                is_mult: AB::Expr::ZERO,
323                idx_in_message: local.idx_in_message.into(),
324                value: local.value.map(Into::into),
325            },
326            local.has_interactions
327                * (AB::Expr::ONE - local.is_first_in_message - local.is_bus_index),
328        );
329        self.interaction_bus.receive(
330            builder,
331            local.proof_idx,
332            InteractionsFoldingMessage {
333                air_idx: local.air_idx.into(),
334                interaction_idx: local.interaction_idx.into(),
335                is_mult: AB::Expr::ONE,
336                idx_in_message: AB::Expr::ZERO,
337                value: local.value.map(Into::into),
338            },
339            local.is_first_in_message * local.has_interactions,
340        );
341        self.interaction_bus.receive(
342            builder,
343            local.proof_idx,
344            InteractionsFoldingMessage {
345                air_idx: local.air_idx.into(),
346                interaction_idx: local.interaction_idx.into(),
347                is_mult: AB::Expr::ZERO,
348                idx_in_message: AB::Expr::NEG_ONE,
349                value: local.value.map(Into::into),
350            },
351            local.is_bus_index,
352        );
353
354        self.transcript_bus.sample_ext(
355            builder,
356            local.proof_idx,
357            local.beta_tidx,
358            local.beta,
359            local.is_valid * local.is_first,
360        );
361
362        self.air_shape_bus.lookup_key(
363            builder,
364            local.proof_idx,
365            AirShapeBusMessage {
366                sort_idx: local.sort_idx.into(),
367                property_idx: AirShapeProperty::NumInteractions.to_field(),
368                value: (local.interaction_idx + AB::Expr::ONE) * local.has_interactions,
369            },
370            next_is_first_in_air_or_invalid * local.is_valid,
371        );
372        self.air_shape_bus.lookup_key(
373            builder,
374            local.proof_idx,
375            AirShapeBusMessage {
376                sort_idx: local.sort_idx.into(),
377                property_idx: AirShapeProperty::AirId.to_field(),
378                value: local.air_idx.into(),
379            },
380            local.is_first_in_air,
381        );
382        self.interactions_folding_input_bus.receive(
383            builder,
384            local.proof_idx,
385            InteractionsFoldingInputMessage {
386                tidx: local.beta_tidx,
387            },
388            local.is_first,
389        );
390
391        self.eq_3b_bus.receive(
392            builder,
393            local.proof_idx,
394            Eq3bMessage {
395                sort_idx: local.sort_idx,
396                interaction_idx: local.interaction_idx,
397                eq_3b: local.eq_3b,
398            },
399            local.has_interactions * local.is_first_in_message,
400        );
401    }
402}