1use core::array;
2use std::{borrow::Borrow, sync::Arc};
3
4use openvm_circuit_primitives::{encoder::Encoder, utils::assert_array_eq, ColumnsAir, SubAir};
5use openvm_recursion_circuit_derive::AlignedBorrow;
6use openvm_stark_backend::{
7 air_builders::PartitionedAirBuilder, interaction::InteractionBuilder, BaseAirWithPublicValues,
8 PartitionedBaseAir,
9};
10use openvm_stark_sdk::config::baby_bear_poseidon2::D_EF;
11use p3_air::{Air, AirBuilder, AirBuilderWithPublicValues, BaseAir};
12use p3_field::{extension::BinomiallyExtendable, Field, PrimeCharacteristicRing};
13use p3_matrix::Matrix;
14use strum::{EnumCount, IntoEnumIterator};
15use strum_macros::EnumIter;
16
17use crate::{
18 batch_constraint::{
19 bus::{
20 ConstraintsFoldingBus, ConstraintsFoldingMessage, InteractionsFoldingBus,
21 InteractionsFoldingMessage, SymbolicExpressionBus, SymbolicExpressionMessage,
22 },
23 expr_eval::{dag_commit_cols_to_cached_cols, DagCommitCols, DagCommitPvs, DagCommitSubAir},
24 },
25 bus::{
26 AirPresenceBus, AirPresenceBusMessage, AirShapeBus, AirShapeBusMessage, AirShapeProperty,
27 ColumnClaimsBus, ColumnClaimsMessage, HyperdimBus, HyperdimBusMessage, PublicValuesBus,
28 PublicValuesBusMessage, SelHypercubeBus, SelHypercubeBusMessage, SelUniBus,
29 SelUniBusMessage,
30 },
31 utils::{
32 base_to_ext, ext_field_add, ext_field_multiply, ext_field_multiply_scalar,
33 ext_field_subtract, scalar_subtract_ext_field,
34 },
35};
36
37pub(in crate::batch_constraint) const NUM_FLAGS: usize = 4;
38pub(in crate::batch_constraint) const ENCODER_MAX_DEGREE: u32 = 2;
39pub(in crate::batch_constraint) const FLAG_MODULUS: u32 = ENCODER_MAX_DEGREE + 1;
40
41#[derive(Debug, Clone, Copy, EnumIter, EnumCount)]
42pub(crate) enum NodeKind {
43 VarPreprocessed = 0,
45 VarMain = 1,
47 VarPublicValue = 2,
49 SelIsFirst = 3,
51 SelIsLast = 4,
53 SelIsTransition = 5,
55 Constant = 6,
57 Add = 7,
59 Sub = 8,
61 Neg = 9,
63 Mul = 10,
65 InteractionMult = 11,
67 InteractionMsgComp = 12,
69 InteractionBusIndex = 13,
71}
72
73#[derive(AlignedBorrow, Copy, Clone)]
74#[repr(C)]
75pub struct CachedSymbolicExpressionColumns<T> {
76 pub(in crate::batch_constraint) flags: [T; NUM_FLAGS],
77
78 pub(in crate::batch_constraint) air_idx: T,
79 pub(in crate::batch_constraint) node_or_interaction_idx: T,
80 pub(in crate::batch_constraint) attrs: [T; 3],
84 pub(in crate::batch_constraint) fanout: T,
85
86 pub(in crate::batch_constraint) is_constraint: T,
87 pub(in crate::batch_constraint) constraint_idx: T,
88}
89
90#[derive(AlignedBorrow, Copy, Clone)]
91#[repr(C)]
92pub struct SingleMainSymbolicExpressionColumns<T> {
93 pub(in crate::batch_constraint) slot_state: T,
96 pub(in crate::batch_constraint) args: [T; 2 * D_EF],
101 pub(in crate::batch_constraint) sort_idx: T,
102 pub(in crate::batch_constraint) n_abs: T,
103 pub(in crate::batch_constraint) is_n_neg: T,
104}
105
106pub struct SymbolicExpressionAir<F: Field> {
107 pub expr_bus: SymbolicExpressionBus,
108 pub hyperdim_bus: HyperdimBus,
109 pub air_shape_bus: AirShapeBus,
110 pub air_presence_bus: AirPresenceBus,
111 pub column_claims_bus: ColumnClaimsBus,
112 pub interactions_folding_bus: InteractionsFoldingBus,
113 pub constraints_folding_bus: ConstraintsFoldingBus,
114 pub public_values_bus: PublicValuesBus,
115 pub sel_hypercube_bus: SelHypercubeBus,
116 pub sel_uni_bus: SelUniBus,
117
118 pub cnt_proofs: usize,
119 pub dag_commit_subair: Option<Arc<DagCommitSubAir<F>>>,
120}
121impl<F: Field> ColumnsAir for SymbolicExpressionAir<F> {}
124
125impl<F: Field> SymbolicExpressionAir<F> {
126 fn has_cached(&self) -> bool {
127 self.dag_commit_subair.is_none()
128 }
129}
130
131impl<F: Field> BaseAirWithPublicValues<F> for SymbolicExpressionAir<F> {
132 fn num_public_values(&self) -> usize {
133 if self.has_cached() {
134 0
135 } else {
136 DagCommitPvs::<F>::width()
137 }
138 }
139}
140
141impl<F: Field> PartitionedBaseAir<F> for SymbolicExpressionAir<F> {
142 fn cached_main_widths(&self) -> Vec<usize> {
143 if self.has_cached() {
144 vec![CachedSymbolicExpressionColumns::<F>::width()]
145 } else {
146 vec![]
147 }
148 }
149
150 fn common_main_width(&self) -> usize {
151 SingleMainSymbolicExpressionColumns::<F>::width() * self.cnt_proofs
152 + if self.has_cached() {
153 0
154 } else {
155 DagCommitCols::<F>::width()
156 }
157 }
158}
159
160impl<F: Field> BaseAir<F> for SymbolicExpressionAir<F> {
161 fn width(&self) -> usize {
162 let single_main_width = SingleMainSymbolicExpressionColumns::<F>::width();
163 if self.has_cached() {
164 CachedSymbolicExpressionColumns::<F>::width() + single_main_width * self.cnt_proofs
165 } else {
166 DagCommitCols::<F>::width() + single_main_width * self.cnt_proofs
167 }
168 }
169}
170
171impl<AB: PartitionedAirBuilder + InteractionBuilder + AirBuilderWithPublicValues> Air<AB>
172 for SymbolicExpressionAir<AB::F>
173where
174 <AB::Expr as PrimeCharacteristicRing>::PrimeSubfield: BinomiallyExtendable<{ D_EF }>,
175{
176 fn eval(&self, builder: &mut AB) {
177 let main_local = builder
178 .common_main()
179 .row_slice(0)
180 .expect("window should have at least one row")
181 .to_vec();
182 let main_next = builder
183 .common_main()
184 .row_slice(1)
185 .expect("window should have at least two rows")
186 .to_vec();
187
188 let (cached_local_vec, main_local_slice, main_next_slice) =
189 if let Some(subair) = self.dag_commit_subair.as_ref() {
190 debug_assert!(!self.has_cached());
192 let commit_width = DagCommitCols::<AB::Var>::width();
193 let (commit_local, rest_local) = main_local.as_slice().split_at(commit_width);
194 let (commit_next, rest_next) = main_next.as_slice().split_at(commit_width);
195 subair.eval(builder, (commit_local, commit_next));
196
197 let cached_local_vec = dag_commit_cols_to_cached_cols(commit_local).to_vec();
198 (cached_local_vec, rest_local, rest_next)
199 } else {
200 debug_assert!(self.has_cached());
201 let cached_local_vec = builder.cached_mains()[0]
202 .row_slice(0)
203 .expect("window should have at least one row")
204 .to_vec();
205 (
206 cached_local_vec,
207 main_local.as_slice(),
208 main_next.as_slice(),
209 )
210 };
211
212 let cached_cols: &CachedSymbolicExpressionColumns<AB::Var> =
213 cached_local_vec.as_slice().borrow();
214 let main_cols: Vec<&SingleMainSymbolicExpressionColumns<AB::Var>> = main_local_slice
215 .chunks(SingleMainSymbolicExpressionColumns::<AB::Var>::width())
216 .map(|chunk| chunk.borrow())
217 .collect();
218 let next_main_cols: Vec<&SingleMainSymbolicExpressionColumns<AB::Var>> = main_next_slice
219 .chunks(SingleMainSymbolicExpressionColumns::<AB::Var>::width())
220 .map(|chunk| chunk.borrow())
221 .collect();
222
223 let enc = Encoder::new(NodeKind::COUNT, ENCODER_MAX_DEGREE, true);
224 assert_eq!(enc.width(), NUM_FLAGS);
225 let flags = cached_cols.flags;
226 let is_valid_row = enc.is_valid::<AB>(&flags);
227
228 let is_arg0_node_idx = enc.contains_flag::<AB>(
229 &flags,
230 &[
231 NodeKind::Add,
232 NodeKind::Sub,
233 NodeKind::Mul,
234 NodeKind::Neg,
235 NodeKind::InteractionMult,
236 NodeKind::InteractionMsgComp,
237 ]
238 .map(|x| x as usize),
239 );
240 let is_arg1_node_idx = enc.contains_flag::<AB>(
241 &flags,
242 &[NodeKind::Add, NodeKind::Sub, NodeKind::Mul].map(|x| x as usize),
243 );
244
245 for (proof_idx, (&cols, &next_cols)) in main_cols.iter().zip(&next_main_cols).enumerate() {
246 let proof_idx = AB::F::from_usize(proof_idx);
247
248 let slot_state: AB::Expr = cols.slot_state.into();
249 let next_slot_state: AB::Expr = next_cols.slot_state.into();
250 let proof_present_in_slot = slot_state.clone()
251 * (AB::Expr::from_u8(3) - slot_state.clone())
252 * AB::F::TWO.inverse();
253 let next_proof_present_in_slot = next_slot_state.clone()
254 * (AB::Expr::from_u8(3) - next_slot_state)
255 * AB::F::TWO.inverse();
256 let air_present =
257 slot_state.clone() * (slot_state.clone() - AB::Expr::ONE) * AB::F::TWO.inverse();
258
259 let arg_ef0: [AB::Var; D_EF] = cols.args[..D_EF].try_into().unwrap();
260 let arg_ef1: [AB::Var; D_EF] = cols.args[D_EF..2 * D_EF].try_into().unwrap();
261
262 builder.assert_tern(cols.slot_state);
263 builder
264 .when(cols.is_n_neg)
265 .assert_eq(cols.slot_state, AB::Expr::TWO);
266 builder
267 .when(air_present.clone())
268 .assert_one(is_valid_row.clone());
269 builder
270 .when_transition()
271 .assert_eq(proof_present_in_slot.clone(), next_proof_present_in_slot);
272
273 let mut value = [AB::Expr::ZERO; D_EF];
274 for node_kind in NodeKind::iter() {
275 let sel = enc.get_flag_expr::<AB>(node_kind as usize, &flags);
277 let expr = match node_kind {
278 NodeKind::Add => ext_field_add::<AB::Expr>(arg_ef0, arg_ef1),
279 NodeKind::Sub => ext_field_subtract::<AB::Expr>(arg_ef0, arg_ef1),
280 NodeKind::Neg => scalar_subtract_ext_field::<AB::Expr>(AB::F::ZERO, arg_ef0),
281 NodeKind::Mul => ext_field_multiply::<AB::Expr>(arg_ef0, arg_ef1),
282 NodeKind::Constant => base_to_ext(cached_cols.attrs[0]),
283 NodeKind::VarPublicValue => base_to_ext(cols.args[0]),
284 NodeKind::SelIsFirst => ext_field_multiply(arg_ef0, arg_ef1),
285 NodeKind::SelIsLast => ext_field_multiply(arg_ef0, arg_ef1),
286 NodeKind::SelIsTransition => scalar_subtract_ext_field(
287 AB::Expr::ONE,
288 ext_field_multiply(arg_ef0, arg_ef1),
289 ),
290 NodeKind::VarPreprocessed
291 | NodeKind::VarMain
292 | NodeKind::InteractionMult
293 | NodeKind::InteractionMsgComp => arg_ef0.map(Into::into),
294 NodeKind::InteractionBusIndex => {
295 base_to_ext(cached_cols.attrs[0] + AB::Expr::ONE)
296 }
297 };
298 value = ext_field_add::<AB::Expr>(
300 value,
301 ext_field_multiply_scalar::<AB::Expr>(expr, sel),
302 );
303 }
304
305 self.expr_bus.add_key_with_lookups(
306 builder,
307 proof_idx,
308 SymbolicExpressionMessage {
309 air_idx: cached_cols.air_idx.into(),
310 node_idx: cached_cols.node_or_interaction_idx.into(),
311 value: value.clone(),
312 },
313 air_present.clone() * cached_cols.fanout,
314 );
315 self.expr_bus.lookup_key(
316 builder,
317 proof_idx,
318 SymbolicExpressionMessage {
319 air_idx: cached_cols.air_idx,
320 node_idx: cached_cols.attrs[0],
321 value: arg_ef0,
322 },
323 air_present.clone() * is_arg0_node_idx.clone(),
324 );
325 self.expr_bus.lookup_key(
326 builder,
327 proof_idx,
328 SymbolicExpressionMessage {
329 air_idx: cached_cols.air_idx,
330 node_idx: cached_cols.attrs[1],
331 value: arg_ef1,
332 },
333 air_present.clone() * is_arg1_node_idx.clone(),
334 );
335
336 let is_var = enc.contains_flag::<AB>(
337 &flags,
338 &[NodeKind::VarMain, NodeKind::VarPreprocessed].map(|x| x as usize),
339 );
340 self.column_claims_bus.receive(
341 builder,
342 proof_idx,
343 ColumnClaimsMessage {
344 sort_idx: cols.sort_idx.into(),
345 part_idx: cached_cols.attrs[1].into(),
346 col_idx: cached_cols.attrs[0].into(),
347 claim: array::from_fn(|i| cols.args[i].into()),
348 is_rot: cached_cols.attrs[2].into(),
349 },
350 is_var * air_present.clone(),
351 );
352 self.public_values_bus.receive(
353 builder,
354 proof_idx,
355 PublicValuesBusMessage {
356 air_idx: cached_cols.air_idx,
357 pv_idx: cached_cols.attrs[0],
358 value: cols.args[0],
359 },
360 enc.get_flag_expr::<AB>(NodeKind::VarPublicValue as usize, &flags)
361 * air_present.clone(),
362 );
363 self.air_shape_bus.lookup_key(
364 builder,
365 proof_idx,
366 AirShapeBusMessage {
367 sort_idx: cols.sort_idx.into(),
368 property_idx: AirShapeProperty::AirId.to_field(),
369 value: cached_cols.air_idx.into(),
370 },
371 air_present.clone(),
372 );
373 self.air_presence_bus.lookup_key(
374 builder,
375 proof_idx,
376 AirPresenceBusMessage {
377 air_idx: cached_cols.air_idx.into(),
378 is_present: air_present.clone(),
379 },
380 proof_present_in_slot * is_valid_row.clone(),
381 );
382 self.hyperdim_bus.lookup_key(
383 builder,
384 proof_idx,
385 HyperdimBusMessage {
386 sort_idx: cols.sort_idx,
387 n_abs: cols.n_abs,
388 n_sign_bit: cols.is_n_neg,
389 },
390 air_present.clone(),
391 );
392 {
394 let is_sel = enc.contains_flag::<AB>(
395 &flags,
396 &[
397 NodeKind::SelIsFirst,
398 NodeKind::SelIsLast,
399 NodeKind::SelIsTransition,
400 ]
401 .map(|x| x as usize),
402 );
403
404 let is_first = enc.get_flag_expr::<AB>(NodeKind::SelIsFirst as usize, &flags);
405 self.sel_uni_bus.lookup_key(
406 builder,
407 proof_idx,
408 SelUniBusMessage {
409 n: AB::Expr::NEG_ONE * cols.n_abs * cols.is_n_neg,
410 is_first: is_first.clone(),
411 value: arg_ef0.map(Into::into),
412 },
413 air_present.clone() * is_sel.clone(),
414 );
415 self.sel_hypercube_bus.lookup_key(
416 builder,
417 proof_idx,
418 SelHypercubeBusMessage {
419 n: cols.n_abs.into(),
420 is_first: is_first.clone(),
421 value: arg_ef1.map(Into::into),
422 },
423 is_sel.clone() * (air_present.clone() - cols.is_n_neg),
425 );
426 assert_array_eq(
427 &mut builder.when(is_sel.clone() * cols.is_n_neg),
428 arg_ef1,
429 [
430 AB::Expr::ONE,
431 AB::Expr::ZERO,
432 AB::Expr::ZERO,
433 AB::Expr::ZERO,
434 ],
435 );
436 }
437 let is_mult = enc.get_flag_expr::<AB>(NodeKind::InteractionMult as usize, &flags);
438 let is_bus_index =
439 enc.get_flag_expr::<AB>(NodeKind::InteractionBusIndex as usize, &flags);
440 let is_interaction = enc.contains_flag::<AB>(
442 &flags,
443 &[NodeKind::InteractionMult, NodeKind::InteractionMsgComp].map(|x| x as usize),
444 );
445 self.interactions_folding_bus.send(
446 builder,
447 proof_idx,
448 InteractionsFoldingMessage {
449 air_idx: cached_cols.air_idx.into(),
450 interaction_idx: cached_cols.node_or_interaction_idx.into(),
451 is_mult,
452 idx_in_message: cached_cols.attrs[1].into(),
453 value: value.clone(),
454 },
455 is_interaction * air_present.clone(),
456 );
457 self.interactions_folding_bus.send(
458 builder,
459 proof_idx,
460 InteractionsFoldingMessage {
461 air_idx: cached_cols.air_idx.into(),
462 interaction_idx: cached_cols.node_or_interaction_idx.into(),
463 is_mult: AB::Expr::ZERO,
464 idx_in_message: AB::Expr::NEG_ONE,
465 value: value.clone(),
466 },
467 is_bus_index * air_present.clone(),
468 );
469 self.constraints_folding_bus.send(
470 builder,
471 proof_idx,
472 ConstraintsFoldingMessage {
473 air_idx: cached_cols.air_idx.into(),
474 constraint_idx: cached_cols.constraint_idx.into(),
475 value: value.clone(),
476 },
477 cached_cols.is_constraint * air_present,
478 );
479 }
480 }
481}