openvm_recursion_circuit/batch_constraint/expr_eval/constraints_folding/
air.rs1use 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 builder
113 .when(is_same_air.clone())
114 .assert_one(next.constraint_idx - local.constraint_idx);
115 builder
117 .when(is_same_air.clone())
118 .assert_eq(local.air_idx, next.air_idx);
119 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 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 assert_array_eq(
144 &mut builder.when(AB::Expr::ONE - is_same_air.clone()),
145 local.cur_sum,
146 local.value,
147 );
148 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}