openvm_recursion_circuit/batch_constraint/expr_eval/symbolic_expression/
air.rs

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    // Args: (col_idx, is_next)
44    VarPreprocessed = 0,
45    // Args: (col_idx, is_next)
46    VarMain = 1,
47    // Args: (pv_idx,)
48    VarPublicValue = 2,
49    // Args: ()
50    SelIsFirst = 3,
51    // Args: ()
52    SelIsLast = 4,
53    // Args: ()
54    SelIsTransition = 5,
55    // Args: (val,)
56    Constant = 6,
57    // Args: (left_node_idx, right_node_idx)
58    Add = 7,
59    // Args: (left_node_idx, right_node_idx)
60    Sub = 8,
61    // Args: (node_idx,)
62    Neg = 9,
63    // Args: (left_node_idx, right_node_idx)
64    Mul = 10,
65    // Args: (node_idx,)
66    InteractionMult = 11,
67    // Args: (node_idx, idx_in_message)
68    InteractionMsgComp = 12,
69    // Args: (node_idx,)
70    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    /// Attributes that define this gate. For binary gates such as Add, Mul, etc.,
81    /// this contains the node_idx's. For InteractionMsgComp, it gives (node_idx, idx_in_message).
82    /// See [[NodeKind]].
83    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    // 0 = proof absent from this slot, 1 = proof present with absent air, 2 = proof present with
94    // present air
95    pub(in crate::batch_constraint) slot_state: T,
96    // Dynamic arguments. For Add/Mul/Sub, this splits into two extension-field elements.
97    // For selectors:
98    //   args[0..D_EF)   = sel_uni witness (base or rotated depending on selector type).
99    //   args[D_EF..2*D_EF) = eq-prefix witness (prod r_i or prod (1-r_i)).
100    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}
121// No columns provided: width is dynamic, depending on `cnt_proofs` and on whether
122// `dag_commit_subair` is present, and mixes several column structs.
123impl<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                // No cached trace: DagCommitCols come before the regular columns
191                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                // deg 2
276                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                // deg <= 4
299                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            // Selector
393            {
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                    // OK: cols.is_n_neg => air_present
424                    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            // is_interaction doesn't include bus index
441            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}