openvm_recursion_circuit/batch_constraint/expr_eval/constraints_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        ConstraintsFoldingBus, ConstraintsFoldingMessage, EqNOuterBus, EqNOuterMessage,
19        ExpressionClaimBus, ExpressionClaimMessage,
20    },
21    bus::{
22        AirShapeBus, AirShapeBusMessage, AirShapeProperty, ConstraintsFoldingInputBus,
23        ConstraintsFoldingInputMessage, NLiftBus, NLiftMessage, TranscriptBus,
24    },
25    subairs::nested_for_loop::{NestedForLoopIoCols, NestedForLoopSubAir},
26    utils::{assert_zeros, ext_field_add, ext_field_multiply, ext_field_multiply_scalar},
27};
28
29#[derive(AlignedBorrow, Copy, Clone, StructReflection)]
30#[repr(C)]
31pub struct ConstraintsFoldingCols<T> {
32    pub is_valid: T,
33    pub is_first: T,
34    pub proof_idx: T,
35
36    pub air_idx: T,
37    pub sort_idx: T,
38    pub constraint_idx: T,
39    pub n_lift: T,
40
41    pub lambda_tidx: T,
42    pub lambda: [T; D_EF],
43
44    pub value: [T; D_EF],
45    pub cur_sum: [T; D_EF],
46    pub eq_n: [T; D_EF],
47
48    pub is_first_in_air: T,
49}
50
51#[derive(ColumnsAir)]
52#[columns_via(ConstraintsFoldingCols<u8>)]
53pub struct ConstraintsFoldingAir {
54    pub transcript_bus: TranscriptBus,
55    pub constraint_bus: ConstraintsFoldingBus,
56    pub expression_claim_bus: ExpressionClaimBus,
57    pub eq_n_outer_bus: EqNOuterBus,
58    pub n_lift_bus: NLiftBus,
59    pub air_shape_bus: AirShapeBus,
60    pub constraints_folding_input_bus: ConstraintsFoldingInputBus,
61}
62
63impl<F> BaseAirWithPublicValues<F> for ConstraintsFoldingAir {}
64impl<F> PartitionedBaseAir<F> for ConstraintsFoldingAir {}
65
66impl<F> BaseAir<F> for ConstraintsFoldingAir {
67    fn width(&self) -> usize {
68        ConstraintsFoldingCols::<F>::width()
69    }
70}
71
72impl<AB: AirBuilder + InteractionBuilder> Air<AB> for ConstraintsFoldingAir
73where
74    <AB::Expr as PrimeCharacteristicRing>::PrimeSubfield: BinomiallyExtendable<{ D_EF }>,
75{
76    fn eval(&self, builder: &mut AB) {
77        let main = builder.main();
78        let (local, next) = (
79            main.row_slice(0).expect("window should have two elements"),
80            main.row_slice(1).expect("window should have two elements"),
81        );
82
83        let local: &ConstraintsFoldingCols<AB::Var> = (*local).borrow();
84        let next: &ConstraintsFoldingCols<AB::Var> = (*next).borrow();
85
86        type LoopSubAir = NestedForLoopSubAir<2>;
87        LoopSubAir {}.eval(
88            builder,
89            (
90                NestedForLoopIoCols {
91                    is_enabled: local.is_valid,
92                    counter: [local.proof_idx, local.sort_idx],
93                    is_first: [local.is_first, local.is_first_in_air],
94                }
95                .map_into(),
96                NestedForLoopIoCols {
97                    is_enabled: next.is_valid,
98                    counter: [next.proof_idx, next.sort_idx],
99                    is_first: [next.is_first, next.is_first_in_air],
100                }
101                .map_into(),
102            ),
103        );
104        builder.when_first_row().assert_zero(local.proof_idx);
105        builder.when(local.is_first).assert_zero(local.sort_idx);
106
107        let is_same_proof = next.is_valid - next.is_first;
108        let is_same_air = next.is_valid - next.is_first_in_air;
109
110        // =========================== indices consistency ===============================
111        // When we are within one air, constraint_idx increases by 1
112        builder
113            .when(is_same_air.clone())
114            .assert_one(next.constraint_idx - local.constraint_idx);
115        // air_idx doesn't change within an air
116        builder
117            .when(is_same_air.clone())
118            .assert_eq(local.air_idx, next.air_idx);
119        // First constraint_idx within an air is zero
120        builder
121            .when(local.is_first_in_air)
122            .assert_zero(local.constraint_idx);
123        builder
124            .when(is_same_air.clone())
125            .assert_eq(local.n_lift, next.n_lift);
126
127        // ======================== lambda and cur sum consistency ============================
128        assert_array_eq(&mut builder.when(is_same_proof), local.lambda, next.lambda);
129        assert_array_eq(
130            &mut builder.when(is_same_air.clone()),
131            local.cur_sum,
132            ext_field_add(
133                local.value,
134                ext_field_multiply::<AB::Expr>(local.lambda, next.cur_sum),
135            ),
136        );
137        assert_array_eq(
138            &mut builder.when(is_same_air.clone()),
139            local.eq_n,
140            next.eq_n,
141        );
142        // numerator and the last element of the message are just the corresponding values
143        assert_array_eq(
144            &mut builder.when(AB::Expr::ONE - is_same_air.clone()),
145            local.cur_sum,
146            local.value,
147        );
148        // If we don't have constraints then `value` is zero
149        assert_zeros(
150            &mut builder
151                .when(local.is_first_in_air)
152                .when(not::<AB::Expr>(is_same_air.clone())),
153            local.value,
154        );
155
156        self.n_lift_bus.receive(
157            builder,
158            local.proof_idx,
159            NLiftMessage {
160                air_idx: local.air_idx,
161                n_lift: local.n_lift,
162            },
163            local.is_first_in_air * local.is_valid,
164        );
165        self.constraint_bus.receive(
166            builder,
167            local.proof_idx,
168            ConstraintsFoldingMessage {
169                air_idx: local.air_idx.into(),
170                constraint_idx: local.constraint_idx - AB::Expr::ONE,
171                value: local.value.map(Into::into),
172            },
173            local.is_valid * (AB::Expr::ONE - local.is_first_in_air),
174        );
175        let folded_sum: [AB::Expr; D_EF] = ext_field_add(
176            ext_field_multiply_scalar::<AB::Expr>(next.cur_sum, is_same_air.clone()),
177            ext_field_multiply_scalar::<AB::Expr>(local.cur_sum, AB::Expr::ONE - is_same_air),
178        );
179        self.expression_claim_bus.send(
180            builder,
181            local.proof_idx,
182            ExpressionClaimMessage {
183                is_interaction: AB::Expr::ZERO,
184                idx: local.sort_idx.into(),
185                value: ext_field_multiply(folded_sum, local.eq_n),
186            },
187            local.is_first_in_air * local.is_valid,
188        );
189        self.constraints_folding_input_bus.receive(
190            builder,
191            local.proof_idx,
192            ConstraintsFoldingInputMessage {
193                tidx: local.lambda_tidx,
194            },
195            local.is_first,
196        );
197        self.transcript_bus.sample_ext(
198            builder,
199            local.proof_idx,
200            local.lambda_tidx,
201            local.lambda,
202            local.is_valid * local.is_first,
203        );
204
205        self.eq_n_outer_bus.lookup_key(
206            builder,
207            local.proof_idx,
208            EqNOuterMessage {
209                is_sharp: AB::Expr::ZERO,
210                n: local.n_lift.into(),
211                value: local.eq_n.map(Into::into),
212            },
213            local.is_first_in_air * local.is_valid,
214        );
215
216        self.air_shape_bus.lookup_key(
217            builder,
218            local.proof_idx,
219            AirShapeBusMessage {
220                sort_idx: local.sort_idx.into(),
221                property_idx: AirShapeProperty::AirId.to_field(),
222                value: local.air_idx.into(),
223            },
224            local.is_first_in_air,
225        );
226    }
227}