Skip to main content

openvm_stark_backend/test_utils/dummy_airs/preprocessed_cached_air/
air.rs

1use p3_air::{Air, AirBuilder, BaseAir, BaseAirWithPublicValues, PairBuilder};
2use p3_field::{Field, PrimeCharacteristicRing};
3use p3_matrix::{dense::RowMajorMatrix, Matrix};
4
5use crate::{air_builders::PartitionedAirBuilder, PartitionedBaseAir};
6
7#[derive(Clone)]
8pub struct PreprocessedCachedAir {
9    sels: Vec<bool>,
10    num_cached_parts: usize,
11}
12
13impl PreprocessedCachedAir {
14    pub fn new(sels: Vec<bool>, num_cached_parts: usize) -> Self {
15        assert!(num_cached_parts > 0, "num_cached_parts must be at least 1");
16        Self {
17            sels,
18            num_cached_parts,
19        }
20    }
21}
22
23impl<F: Field> PartitionedBaseAir<F> for PreprocessedCachedAir {
24    fn cached_main_widths(&self) -> Vec<usize> {
25        vec![1; self.num_cached_parts]
26    }
27
28    fn common_main_width(&self) -> usize {
29        1
30    }
31}
32
33impl<F: Field> BaseAir<F> for PreprocessedCachedAir {
34    fn width(&self) -> usize {
35        1 + self.num_cached_parts
36    }
37
38    fn preprocessed_trace(&self) -> Option<RowMajorMatrix<F>> {
39        Some(RowMajorMatrix::new_col(
40            self.sels.iter().map(|&sel| F::from_bool(sel)).collect(),
41        ))
42    }
43}
44
45impl<F: Field> BaseAirWithPublicValues<F> for PreprocessedCachedAir {}
46
47impl<AB: PartitionedAirBuilder + PairBuilder> Air<AB> for PreprocessedCachedAir
48where
49    AB::F: Field,
50{
51    fn eval(&self, builder: &mut AB) {
52        let preprocessed = builder.preprocessed();
53        let preprocessed_local = preprocessed
54            .row_slice(0)
55            .expect("preprocessed window should have one row")[0]
56            .clone();
57        let preprocessed_next = preprocessed
58            .row_slice(1)
59            .expect("preprocessed window should have two rows")[0]
60            .clone();
61        let common_local = builder
62            .common_main()
63            .row_slice(0)
64            .expect("common main window should have one row")[0]
65            .clone();
66        let common_next = builder
67            .common_main()
68            .row_slice(1)
69            .expect("common main window should have two rows")[0]
70            .clone();
71        let cached_rows: Vec<_> = builder
72            .cached_mains()
73            .iter()
74            .map(|cached_main| {
75                let local = cached_main
76                    .row_slice(0)
77                    .expect("cached main window should have one row")[0]
78                    .clone();
79                let next = cached_main
80                    .row_slice(1)
81                    .expect("cached main window should have two rows")[0]
82                    .clone();
83                (local, next)
84            })
85            .collect();
86
87        debug_assert_eq!(cached_rows.len(), self.num_cached_parts);
88
89        builder
90            .assert_zero(preprocessed_local.clone() * (preprocessed_local.clone() - AB::Expr::ONE));
91        builder.when_first_row().assert_zero(common_local.clone());
92        builder
93            .when_transition()
94            .assert_one(common_next.clone() - common_local.clone());
95
96        // Each cached partition increments the previous partition by the selector bit.
97        let mut prev_local = common_local;
98        let mut prev_next = common_next;
99        for (cached_local, cached_next) in cached_rows {
100            builder.assert_eq(
101                prev_local.clone() + preprocessed_local.clone(),
102                cached_local.clone(),
103            );
104            builder.when_transition().assert_eq(
105                prev_next.clone() + preprocessed_next.clone(),
106                cached_next.clone(),
107            );
108            prev_local = cached_local;
109            prev_next = cached_next;
110        }
111    }
112}