openvm_deferral_circuit/count/
air.rs

1use std::borrow::Borrow;
2
3use openvm_circuit_primitives::{utils::not, ColumnsAir, StructReflection, StructReflectionHelper};
4use openvm_circuit_primitives_derive::AlignedBorrow;
5use openvm_stark_backend::{
6    interaction::InteractionBuilder,
7    p3_air::{Air, AirBuilder, BaseAir},
8    p3_field::PrimeCharacteristicRing,
9    p3_matrix::Matrix,
10    BaseAirWithPublicValues, PartitionedBaseAir,
11};
12
13use super::DeferralCircuitCountBus;
14
15#[repr(C)]
16#[derive(AlignedBorrow, StructReflection)]
17pub struct DeferralCircuitCountCols<T> {
18    pub is_valid: T,
19    pub row_idx: T,
20    pub mult: T,
21}
22
23#[derive(Clone, Copy, Debug, derive_new::new, ColumnsAir)]
24#[columns_via(DeferralCircuitCountCols<u8>)]
25pub struct DeferralCircuitCountAir {
26    pub lookup_bus: DeferralCircuitCountBus,
27    pub num_deferral_circuits: usize,
28}
29
30impl<F> BaseAir<F> for DeferralCircuitCountAir {
31    fn width(&self) -> usize {
32        DeferralCircuitCountCols::<F>::width()
33    }
34}
35impl<F> BaseAirWithPublicValues<F> for DeferralCircuitCountAir {}
36impl<F> PartitionedBaseAir<F> for DeferralCircuitCountAir {}
37
38impl<AB> Air<AB> for DeferralCircuitCountAir
39where
40    AB: InteractionBuilder,
41{
42    fn eval(&self, builder: &mut AB) {
43        let main = builder.main();
44        let local = main.row_slice(0).expect("row 0 present");
45        let next = main.row_slice(1).expect("row 1 present");
46
47        let local: &DeferralCircuitCountCols<AB::Var> = (*local).borrow();
48        let next: &DeferralCircuitCountCols<AB::Var> = (*next).borrow();
49
50        // Base constraints to ensure all valid rows are at the beginning of the trace,
51        // that row_idx increments properly, and that mult is 0 on invalid rows.
52        if self.num_deferral_circuits > 0 {
53            builder.when_first_row().assert_one(local.is_valid);
54            builder.assert_bool(local.is_valid);
55            builder
56                .when_transition()
57                .assert_bool(local.is_valid - next.is_valid);
58        } else {
59            builder.assert_zero(local.is_valid);
60        }
61
62        builder.when_first_row().assert_zero(local.row_idx);
63        builder
64            .when_transition()
65            .assert_one(next.row_idx - local.row_idx);
66
67        builder.when(not(local.is_valid)).assert_zero(local.mult);
68
69        // Constrain that there are exactly n = num_deferral_circuits valid rows. We do
70        // this by constraining that local.is_valid equals next.is_valid on rows such
71        // that row_idx + 1 != n, and then either that the last row is invalid (meaning
72        // the last valid row was n - 1) or that n is the trace height.
73        let num_valid = AB::F::from_usize(self.num_deferral_circuits);
74
75        builder
76            .when_transition()
77            .when_ne(next.row_idx, num_valid)
78            .assert_eq(local.is_valid, next.is_valid);
79        builder
80            .when_last_row()
81            .when(local.is_valid)
82            .assert_eq(local.row_idx + AB::Expr::ONE, num_valid);
83
84        // Provide the lookup for valid deferral indices.
85        self.lookup_bus
86            .receive(local.row_idx)
87            .eval(builder, local.mult);
88    }
89}