openvm_recursion_circuit/stacking/univariate/
air.rs1use 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 pub proof_idx: F,
31 pub is_valid: F,
32 pub is_first: F,
33 pub is_last: F,
34
35 pub tidx: F,
37 pub u_0: [F; D_EF],
38 pub u_0_pow: [F; D_EF],
39
40 pub coeff: [F; D_EF],
42
43 pub coeff_idx: F,
45 pub coeff_is_d: F,
46 pub s_0_sum_over_d: [F; D_EF],
47
48 pub poly_rand_eval: [F; D_EF],
50}
51
52#[derive(ColumnsAir)]
53#[columns_via(UnivariateRoundCols<u8>)]
54pub struct UnivariateRoundAir {
55 pub transcript_bus: TranscriptBus,
57
58 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 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 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 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 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 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}