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 pub is_first_in_message: T, 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 pub cur_sum: [T; D_EF],
58 pub beta: [T; D_EF],
59 pub eq_3b: [T; D_EF],
60
61 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 builder
149 .when(is_same_proof.clone())
150 .assert_bool(next.sort_idx - local.sort_idx);
151 builder
153 .when(is_same_air.clone())
154 .assert_bool(next.interaction_idx - local.interaction_idx);
155 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 builder
172 .when(is_same_air.clone())
173 .assert_eq(local.air_idx, next.air_idx);
174 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 builder
182 .when(local.has_interactions)
183 .assert_one(local.is_valid);
184 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 builder.when(local.has_interactions).assert_eq(
196 local.is_bus_index,
197 next_is_first_in_message_or_invalid.clone(),
198 );
199 builder
201 .when(local.has_interactions)
202 .assert_bool(local.is_bus_index + local.is_first_in_message);
203 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 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 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 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 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 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: 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: 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}