openvm_recursion_circuit/primitives/exp_bits_len/
air.rs1use core::borrow::Borrow;
2
3use openvm_circuit_primitives::{ColumnsAir, StructReflection, StructReflectionHelper};
4use openvm_recursion_circuit_derive::AlignedBorrow;
5use openvm_stark_backend::{
6 interaction::InteractionBuilder, BaseAirWithPublicValues, PartitionedBaseAir,
7};
8use p3_air::{Air, AirBuilder, BaseAir};
9use p3_baby_bear::BabyBear;
10use p3_field::{PrimeCharacteristicRing, PrimeField32};
11use p3_matrix::Matrix;
12
13use crate::primitives::{
14 bus::{ExpBitsLenBus, ExpBitsLenMessage, RightShiftBus, RightShiftMessage},
15 exp_bits_len::trace::{LOW_BITS_COUNT, NUM_BITS_MAX_PLUS_ONE},
16};
17
18#[repr(C)]
19#[derive(AlignedBorrow, StructReflection)]
20pub struct ExpBitsLenCols<T> {
21 pub is_valid: T,
23 pub is_first: T,
25 pub bit_idx: T,
27 pub base: T,
29 pub bit_src: T,
31 pub num_bits: T,
33 pub apply_bit: T,
35 pub low_bits_left: T,
37 pub in_low_region: T,
39 pub result: T,
41 pub result_multiplier: T,
43 pub bit_src_mod_2: T,
45 pub low_bits_are_zero: T,
47 pub high_bits_all_one: T,
49
50 pub bit_src_original: T,
52 pub shift_mult: T,
54}
55
56#[derive(Debug, derive_new::new, ColumnsAir)]
57#[columns_via(ExpBitsLenCols<u8>)]
58pub struct ExpBitsLenAir {
59 pub exp_bits_len_bus: ExpBitsLenBus,
60 pub right_shift_bus: RightShiftBus,
61}
62
63fn assert_babybear_field<F: PrimeField32>() {
64 assert_eq!(
65 F::ORDER_U32,
66 BabyBear::ORDER_U32,
67 "ExpBitsLenAir is hard-coded for the BabyBear modulus; canonicality constraints assume p = 15 * 2^27 + 1",
68 );
69}
70
71impl<F: PrimeField32> BaseAirWithPublicValues<F> for ExpBitsLenAir {}
72impl<F: PrimeField32> PartitionedBaseAir<F> for ExpBitsLenAir {}
73
74impl<F: PrimeField32> BaseAir<F> for ExpBitsLenAir {
75 fn width(&self) -> usize {
76 assert_babybear_field::<F>();
77 ExpBitsLenCols::<F>::width()
78 }
79}
80
81impl<AB: AirBuilder + InteractionBuilder> Air<AB> for ExpBitsLenAir
82where
83 AB::F: PrimeField32,
84{
85 fn eval(&self, builder: &mut AB) {
86 assert_babybear_field::<AB::F>();
87 let main = builder.main();
88
89 let (local, next) = (
90 main.row_slice(0).expect("window should have two elements"),
91 main.row_slice(1).expect("window should have two elements"),
92 );
93 let local: &ExpBitsLenCols<AB::Var> = (*local).borrow();
94 let next: &ExpBitsLenCols<AB::Var> = (*next).borrow();
95
96 let is_transition = next.is_valid.into() - next.is_first.into();
97 let local_is_last = local.is_valid.into() - is_transition.clone();
98 let local_in_high_region =
99 local.is_valid.into() - local.in_low_region.into() - local_is_last.clone();
100 let next_is_not_low_region = next.is_valid.into() - next.in_low_region.into();
101
102 builder.assert_bool(local.is_valid);
103 builder.assert_bool(local.is_first);
104 builder.assert_bool(is_transition.clone());
105 builder
106 .when(is_transition.clone())
107 .assert_one(local.is_valid);
108 builder.assert_bool(local.apply_bit);
109 builder.assert_bool(local.in_low_region);
110 builder.assert_bool(local.bit_src_mod_2);
111 builder.assert_bool(local.low_bits_are_zero);
112 builder.assert_bool(local.high_bits_all_one);
113 builder.assert_bool(local_in_high_region.clone());
114 builder.when(local.is_first).assert_one(local.is_valid);
115 builder.when(local.num_bits).assert_one(local.apply_bit);
116 builder
117 .when(local.low_bits_left)
118 .assert_one(local.in_low_region);
119
120 builder.assert_eq(
121 local.result_multiplier - AB::Expr::ONE,
122 local.apply_bit * local.bit_src_mod_2 * (local.base - AB::Expr::ONE),
123 );
124
125 builder
126 .when_first_row()
127 .assert_eq(local.is_valid, local.is_first);
128 builder.when(local.is_first).assert_one(local.in_low_region);
129 builder
130 .when(local.is_first)
131 .assert_zero(local_is_last.clone());
132 builder.when(local.is_first).assert_zero(local.bit_idx);
133 builder
134 .when(local.is_first)
135 .assert_eq(local.low_bits_left, AB::Expr::from_usize(LOW_BITS_COUNT));
136 builder
137 .when(local.is_first)
138 .assert_one(local.low_bits_are_zero);
139 builder
140 .when(local.is_first)
141 .assert_zero(local.high_bits_all_one);
142
143 builder
144 .when(local_is_last.clone())
145 .assert_zero(local.in_low_region);
146 builder.when(local_is_last.clone()).assert_eq(
147 local.bit_idx,
148 AB::Expr::from_usize(NUM_BITS_MAX_PLUS_ONE - 1),
149 );
150 builder
151 .when(local_is_last.clone())
152 .assert_zero(local.bit_src);
153 builder
154 .when(local_is_last.clone())
155 .assert_zero(local.num_bits);
156 builder
157 .when(local_is_last.clone())
158 .assert_zero(local.apply_bit);
159 builder
160 .when(local_is_last.clone())
161 .assert_zero(local.low_bits_left);
162 builder.when(local_is_last.clone()).assert_one(local.result);
163 builder
164 .when(local_is_last.clone())
165 .assert_one(local.result_multiplier);
166 builder
167 .when(local_is_last.clone())
168 .assert_zero(local.high_bits_all_one * (AB::Expr::ONE - local.low_bits_are_zero));
169
170 builder
171 .when(is_transition.clone())
172 .assert_eq(next.bit_idx, local.bit_idx + AB::Expr::ONE);
173 builder
174 .when(is_transition.clone())
175 .assert_eq(next.base, local.base * local.base);
176 builder.when(is_transition.clone()).assert_eq(
177 local.bit_src,
178 next.bit_src * AB::Expr::TWO + local.bit_src_mod_2,
179 );
180 builder
181 .when(is_transition.clone())
182 .assert_eq(next.num_bits, local.num_bits - local.apply_bit);
183 builder.when(is_transition.clone()).assert_eq(
184 next.low_bits_left,
185 local.low_bits_left - local.in_low_region,
186 );
187 builder
188 .when(is_transition.clone())
189 .assert_eq(local.result, next.result * local.result_multiplier);
190 builder.when(local.in_low_region).assert_eq(
191 next.low_bits_are_zero,
192 local.low_bits_are_zero * (AB::Expr::ONE - local.bit_src_mod_2),
193 );
194 builder
195 .when(local_in_high_region.clone())
196 .assert_eq(next.low_bits_are_zero, local.low_bits_are_zero);
197 builder
198 .when(local.in_low_region * next.in_low_region)
199 .assert_zero(next.high_bits_all_one);
200 builder
201 .when(local.in_low_region * next_is_not_low_region.clone())
202 .assert_one(next.high_bits_all_one);
203 builder.when(local_in_high_region).assert_eq(
204 next.high_bits_all_one,
205 local.high_bits_all_one * local.bit_src_mod_2,
206 );
207
208 self.exp_bits_len_bus.add_key_with_lookups(
209 builder,
210 ExpBitsLenMessage {
211 base: local.base,
212 bit_src: local.bit_src,
213 num_bits: local.num_bits,
214 result: local.result,
215 },
216 local.is_first,
217 );
218
219 builder.when(local.shift_mult).assert_one(local.is_valid);
220
221 builder
222 .when(local.is_first)
223 .assert_eq(local.bit_src, local.bit_src_original);
224 builder
225 .when(is_transition)
226 .assert_eq(local.bit_src_original, next.bit_src_original);
227
228 self.right_shift_bus.add_key_with_lookups(
229 builder,
230 RightShiftMessage {
231 input: local.bit_src_original,
232 shift_bits: local.bit_idx,
233 result: local.bit_src,
234 },
235 local.shift_mult,
236 );
237 }
238}