openvm_recursion_circuit/batch_constraint/eq_airs/eq_uni/
air.rs

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        BatchConstraintConductorBus, BatchConstraintConductorMessage,
19        BatchConstraintInnerMessageType, EqZeroNBus, EqZeroNMessage,
20    },
21    subairs::nested_for_loop::{NestedForLoopIoCols, NestedForLoopSubAir},
22    utils::{assert_one_ext, ext_field_add, ext_field_multiply, ext_field_one_minus},
23};
24
25#[derive(AlignedBorrow, Clone, Copy, StructReflection)]
26#[repr(C)]
27pub struct EqUniCols<T> {
28    pub is_valid: T,
29    pub is_first: T,
30    pub proof_idx: T,
31
32    pub x: [T; D_EF],
33    pub y: [T; D_EF],
34    pub res: [T; D_EF],
35    pub idx: T,
36}
37
38#[derive(ColumnsAir)]
39#[columns_via(EqUniCols<u8>)]
40pub struct EqUniAir {
41    pub zero_n_bus: EqZeroNBus,
42    pub r_xi_bus: BatchConstraintConductorBus,
43    pub l_skip: usize,
44}
45
46impl<F> BaseAirWithPublicValues<F> for EqUniAir {}
47impl<F> PartitionedBaseAir<F> for EqUniAir {}
48
49impl<F> BaseAir<F> for EqUniAir {
50    fn width(&self) -> usize {
51        EqUniCols::<F>::width()
52    }
53}
54
55impl<AB: AirBuilder + InteractionBuilder> Air<AB> for EqUniAir
56where
57    <AB::Expr as PrimeCharacteristicRing>::PrimeSubfield: BinomiallyExtendable<{ D_EF }>,
58{
59    fn eval(&self, builder: &mut AB) {
60        let main = builder.main();
61        let (local, next) = (
62            main.row_slice(0).expect("window should have two elements"),
63            main.row_slice(1).expect("window should have two elements"),
64        );
65
66        let local: &EqUniCols<AB::Var> = (*local).borrow();
67        let next: &EqUniCols<AB::Var> = (*next).borrow();
68
69        // Summary:
70        // - Proof loop: rely on the nested for-loop sub-AIR to enforce the standard
71        //   `is_valid`/`is_first`/`proof_idx` sequencing across proofs.
72        // - idx handling: start `idx` at zero, increment it on each transition within the proof,
73        //   keep it zero on invalid rows, and ensure the last valid row reaches `l_skip`.
74        // - Values recalculation: during transitions, square both `x` and `y`, and update `res` via
75        //   `(x + y) * res + (1 - x) * (1 - y)`.
76
77        // Enforce the standard proof loop flags (is_valid/is_first/proof_idx).
78        type LoopSubAir = NestedForLoopSubAir<1>;
79        LoopSubAir {}.eval(
80            builder,
81            (
82                NestedForLoopIoCols {
83                    is_enabled: local.is_valid,
84                    counter: [local.proof_idx],
85                    is_first: [local.is_first],
86                }
87                .map_into(),
88                NestedForLoopIoCols {
89                    is_enabled: next.is_valid,
90                    counter: [next.proof_idx],
91                    is_first: [next.is_first],
92                }
93                .map_into(),
94            ),
95        );
96
97        let local_is_last = next.is_first + not(next.is_valid);
98        builder.assert_bool(local_is_last.clone());
99
100        self.r_xi_bus.lookup_key(
101            builder,
102            local.proof_idx,
103            BatchConstraintConductorMessage {
104                msg_type: BatchConstraintInnerMessageType::Xi.to_field(),
105                idx: AB::Expr::ZERO,
106                value: local.x.map(|x| x.into()),
107            },
108            local.is_valid * local.is_first,
109        );
110        self.r_xi_bus.lookup_key(
111            builder,
112            local.proof_idx,
113            BatchConstraintConductorMessage {
114                msg_type: BatchConstraintInnerMessageType::R.to_field(),
115                idx: AB::Expr::ZERO,
116                value: local.y.map(|x| x.into()),
117            },
118            local.is_valid * local.is_first,
119        );
120
121        let inv_p2 = AB::F::ONE.halve().exp_u64(self.l_skip as u64);
122        self.zero_n_bus.send(
123            builder,
124            local.proof_idx,
125            EqZeroNMessage {
126                is_sharp: AB::Expr::ZERO,
127                value: local.res.map(|x| x * inv_p2),
128            },
129            local.is_valid * local_is_last.clone(),
130        );
131
132        let is_transition = next.is_valid * (AB::Expr::ONE - next.is_first);
133
134        // ======================== idx handling ==========================
135        builder.when(local.is_first).assert_zero(local.idx);
136        builder
137            .when(is_transition.clone())
138            .assert_eq(next.idx, local.idx + AB::Expr::ONE);
139        builder.when(not(local.is_valid)).assert_zero(local.idx);
140        builder
141            .when(local_is_last)
142            .when(local.is_valid)
143            .assert_eq(local.idx, AB::Expr::from_usize(self.l_skip));
144
145        // ======================== Values recalculation ==========================
146        assert_one_ext(&mut builder.when(local.is_first), local.res);
147
148        let mut when_transition = builder.when(is_transition);
149        assert_array_eq(
150            &mut when_transition,
151            next.x.map(Into::into),
152            ext_field_multiply::<AB::Expr>(local.x, local.x),
153        );
154        assert_array_eq(
155            &mut when_transition,
156            next.y.map(Into::into),
157            ext_field_multiply::<AB::Expr>(local.y, local.y),
158        );
159
160        let x_plus_y = ext_field_add::<AB::Expr>(local.x, local.y);
161        let one_minus_x = ext_field_one_minus::<AB::Expr>(local.x);
162        let one_minus_y = ext_field_one_minus::<AB::Expr>(local.y);
163        let next_res_expected = ext_field_add::<AB::Expr>(
164            ext_field_multiply::<AB::Expr>(x_plus_y, local.res),
165            ext_field_multiply::<AB::Expr>(one_minus_x, one_minus_y),
166        );
167        assert_array_eq(
168            &mut when_transition,
169            next.res.map(Into::into),
170            next_res_expected,
171        );
172    }
173}