openvm_recursion_circuit/batch_constraint/eq_airs/eq_uni/
air.rs1use 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 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 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 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}