openvm_verify_stark_circuit/output/
air.rs

1use std::{array::from_fn, borrow::Borrow};
2
3use itertools::{fold, Itertools};
4use openvm_circuit_primitives::{
5    utils::assert_array_eq, AlignedBorrow, ColumnsAir, StructReflection, StructReflectionHelper,
6};
7use openvm_continuations::utils::digests_to_poseidon2_input;
8use openvm_deferral_circuit::canonicity::{CanonicityAuxCols, CanonicitySubAir};
9use openvm_recursion_circuit::{
10    bus::{Poseidon2PermuteBus, Poseidon2PermuteMessage},
11    prelude::DIGEST_SIZE,
12    primitives::bus::{RangeCheckerBus, RangeCheckerBusMessage},
13};
14use openvm_stark_backend::{interaction::InteractionBuilder, PartitionedBaseAir};
15use p3_air::{Air, AirBuilder, BaseAir, BaseAirWithPublicValues};
16use p3_field::{PrimeCharacteristicRing, PrimeField32};
17use p3_matrix::Matrix;
18
19use crate::bus::{OutputCommitBus, OutputCommitMessage, OutputValBus, OutputValMessage};
20
21pub(crate) const F_NUM_BYTES: usize = 4;
22pub(crate) const VALS_IN_DIGEST: usize = exact_div_or_panic(DIGEST_SIZE, F_NUM_BYTES);
23
24const fn exact_div_or_panic(a: usize, b: usize) -> usize {
25    assert!(b != 0 && a.is_multiple_of(b), "non-exact division");
26    a / b
27}
28
29#[repr(C)]
30#[derive(AlignedBorrow, StructReflection)]
31pub struct DeferralOutputCommitCols<F> {
32    pub is_valid: F,
33    pub is_first: F,
34    pub row_idx: F,
35    pub output_len: F,
36
37    pub input_vals: [F; DIGEST_SIZE],
38    pub res_left: [F; DIGEST_SIZE],
39    pub res_right: [F; DIGEST_SIZE],
40
41    pub canonicity_aux: [CanonicityAuxCols<F>; VALS_IN_DIGEST],
42}
43
44#[derive(Debug, ColumnsAir)]
45#[columns_via(DeferralOutputCommitCols<u8>)]
46pub struct DeferralOutputCommitAir {
47    pub poseidon2_bus: Poseidon2PermuteBus,
48    pub range_bus: RangeCheckerBus,
49    pub output_val_bus: OutputValBus,
50    pub output_commit_bus: OutputCommitBus,
51
52    pub def_idx: usize,
53}
54
55impl<F> BaseAir<F> for DeferralOutputCommitAir {
56    fn width(&self) -> usize {
57        DeferralOutputCommitCols::<u8>::width()
58    }
59}
60impl<F> BaseAirWithPublicValues<F> for DeferralOutputCommitAir {}
61impl<F> PartitionedBaseAir<F> for DeferralOutputCommitAir {}
62
63impl<AB: AirBuilder + InteractionBuilder> Air<AB> for DeferralOutputCommitAir
64where
65    AB::F: PrimeField32,
66{
67    fn eval(&self, builder: &mut AB) {
68        let main = builder.main();
69        let local = main.row_slice(0).expect("row 0 present");
70        let next = main.row_slice(1).expect("row 1 present");
71
72        let local: &DeferralOutputCommitCols<AB::Var> = (*local).borrow();
73        let next: &DeferralOutputCommitCols<AB::Var> = (*next).borrow();
74
75        let is_transition = next.is_valid - next.is_first;
76        let is_last = local.is_valid - is_transition.clone();
77
78        /*
79         * Base constraints to ensure that all the valid rows are at the beginning,
80         * and that there is at least one valid row. Additionally, row_idx starts at
81         * 0 and increments.
82         */
83        builder.assert_bool(local.is_valid);
84        builder.when_first_row().assert_one(local.is_valid);
85        builder
86            .when_transition()
87            .assert_bool(local.is_valid - next.is_valid);
88
89        builder.when_first_row().assert_one(local.is_first);
90        builder.when_transition().assert_zero(next.is_first);
91
92        builder.when_first_row().assert_zero(local.row_idx);
93        builder
94            .when_transition()
95            .assert_eq(local.row_idx + AB::Expr::ONE, next.row_idx);
96
97        /*
98         * On the first row, input_vals should be [def_idx, output_len, 0, ...].
99         * We constrain def_idx against a constant, and output_len against the
100         * last valid row_idx.
101         */
102        let mut initial_state = [AB::Expr::ZERO; DIGEST_SIZE];
103        initial_state[0] = AB::Expr::from_usize(self.def_idx);
104        initial_state[1] = local.output_len.into();
105
106        assert_array_eq(
107            &mut builder.when_first_row(),
108            local.input_vals,
109            initial_state,
110        );
111
112        builder
113            .when(is_transition.clone())
114            .assert_eq(local.output_len, next.output_len);
115        builder.when(is_last.clone()).assert_eq(
116            local.output_len,
117            local.row_idx * AB::Expr::from_usize(DIGEST_SIZE),
118        );
119
120        /*
121         * On valid rows non-first we want to receive the next VALS_IN_DIGEST values
122         * and constrain that input_vals is their byte decomposition.
123         */
124        let next_f: [_; VALS_IN_DIGEST] = local
125            .input_vals
126            .chunks(F_NUM_BYTES)
127            .map(|c| {
128                fold(c.iter().enumerate(), AB::Expr::ZERO, |acc, (i, byte)| {
129                    acc + (AB::Expr::from_usize(1 << (i * 8)) * (*byte).into())
130                })
131            })
132            .collect_array()
133            .unwrap();
134
135        for byte in local.input_vals {
136            self.range_bus.lookup_key(
137                builder,
138                RangeCheckerBusMessage {
139                    value: byte.into(),
140                    max_bits: AB::Expr::from_u8(8),
141                },
142                local.is_valid - local.is_first,
143            );
144        }
145
146        self.output_val_bus.receive(
147            builder,
148            OutputValMessage {
149                values: next_f,
150                idx: local.row_idx - AB::Expr::ONE,
151            },
152            local.is_valid - local.is_first,
153        );
154
155        /*
156         * For each output value we need to constraint the canonicity of the byte
157         * decomposition.
158         */
159        let rcs = local
160            .input_vals
161            .chunks(F_NUM_BYTES)
162            .zip(local.canonicity_aux)
163            .map(|(x, aux)| {
164                CanonicitySubAir.assert_canonicity(
165                    builder,
166                    x,
167                    &aux,
168                    local.is_valid - local.is_first,
169                )
170            })
171            .collect_vec();
172
173        for rc in rcs {
174            self.range_bus.lookup_key(
175                builder,
176                RangeCheckerBusMessage {
177                    value: rc,
178                    max_bits: AB::Expr::from_u8(8),
179                },
180                local.is_valid - local.is_first,
181            );
182        }
183
184        /*
185         * Compute the output commit and send it on the last row. We sponge hash each
186         * valid input_vals and take the left child of the last row.
187         */
188        let next_capacity = from_fn(|i| is_transition.clone() * local.res_right[i]);
189        self.poseidon2_bus.lookup_key(
190            builder,
191            Poseidon2PermuteMessage {
192                input: digests_to_poseidon2_input(next.input_vals.map(Into::into), next_capacity),
193                output: digests_to_poseidon2_input(next.res_left, next.res_right).map(Into::into),
194            },
195            next.is_valid,
196        );
197
198        self.output_commit_bus.send(
199            builder,
200            OutputCommitMessage {
201                commit: local.res_left,
202            },
203            is_last,
204        );
205    }
206}