openvm_recursion_circuit/primitives/exp_bits_len/
air.rs

1use 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    /// Marks rows that belong to an `ExpBitsLen` request rather than trailing padding.
22    pub is_valid: T,
23    /// Marks the first row of a 32-row request block. Only these rows publish on the lookup bus.
24    pub is_first: T,
25    /// Bit position carried by this row. Rows `0..30` decompose bits, row `31` is terminal.
26    pub bit_idx: T,
27    /// `base^(2^bit_idx)`.
28    pub base: T,
29    /// Remaining suffix of the canonical 31-bit decomposition after shifting by `bit_idx`.
30    pub bit_src: T,
31    /// Remaining number of low bits that still affect `result`.
32    pub num_bits: T,
33    /// Boolean witness for `num_bits != 0`.
34    pub apply_bit: T,
35    /// Countdown for the low 27 bits used in the BabyBear `< p` canonicality check.
36    pub low_bits_left: T,
37    /// Boolean witness for `low_bits_left != 0`.
38    pub in_low_region: T,
39    /// Running product for the requested low-bit exponentiation.
40    pub result: T,
41    /// Multiplies `result` by either `1` or `base` on the next transition.
42    pub result_multiplier: T,
43    /// Current decomposition bit.
44    pub bit_src_mod_2: T,
45    /// Running flag: all low bits `b0..b26` seen so far are zero.
46    pub low_bits_are_zero: T,
47    /// Running flag: all high bits `b27..b30` seen so far are one.
48    pub high_bits_all_one: T,
49
50    /// Original bit_src value
51    pub bit_src_original: T,
52    /// Indicator for if this show should send a shift message
53    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}