openvm_continuations/circuit/inner/vm_pvs/
air.rs

1use std::borrow::Borrow;
2
3use openvm_circuit::system::connector::DEFAULT_SUSPEND_EXIT_CODE;
4use openvm_circuit_primitives::{
5    utils::{and, assert_array_eq, not},
6    ColumnsAir, StructReflection, StructReflectionHelper,
7};
8use openvm_recursion_circuit::bus::{
9    CachedCommitBus, CachedCommitBusMessage, PublicValuesBus, PublicValuesBusMessage,
10};
11use openvm_recursion_circuit_derive::AlignedBorrow;
12use openvm_stark_backend::{
13    interaction::InteractionBuilder, BaseAirWithPublicValues, PartitionedBaseAir,
14};
15use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
16use openvm_verify_stark_host::pvs::{VmPvs, VM_PVS_AIR_ID};
17use p3_air::{Air, AirBuilder, AirBuilderWithPublicValues, BaseAir};
18use p3_field::PrimeCharacteristicRing;
19use p3_matrix::Matrix;
20
21use crate::circuit::inner::{
22    app::*,
23    bus::{PvsAirConsistencyBus, PvsAirConsistencyMessage},
24};
25
26#[repr(C)]
27#[derive(AlignedBorrow, StructReflection)]
28pub struct VmPvsCols<F> {
29    pub proof_idx: F,
30    pub is_valid: F,
31    pub is_last: F,
32    pub has_verifier_pvs: F,
33    pub child_pvs: VmPvs<F>,
34}
35
36#[derive(ColumnsAir)]
37#[columns_via(VmPvsCols<u8>)]
38pub struct VmPvsAir {
39    pub public_values_bus: PublicValuesBus,
40    pub cached_commit_bus: CachedCommitBus,
41    pub pvs_air_consistency_bus: PvsAirConsistencyBus,
42    pub deferral_enabled: bool,
43}
44
45impl<F> BaseAir<F> for VmPvsAir {
46    fn width(&self) -> usize {
47        VmPvsCols::<u8>::width() + (self.deferral_enabled as usize)
48    }
49}
50impl<F> BaseAirWithPublicValues<F> for VmPvsAir {
51    fn num_public_values(&self) -> usize {
52        VmPvs::<u8>::width()
53    }
54}
55impl<F> PartitionedBaseAir<F> for VmPvsAir {}
56
57impl<AB: AirBuilder + InteractionBuilder + AirBuilderWithPublicValues> Air<AB> for VmPvsAir {
58    fn eval(&self, builder: &mut AB) {
59        let main = builder.main();
60        let (local, next) = (
61            main.row_slice(0).expect("window should have two elements"),
62            main.row_slice(1).expect("window should have two elements"),
63        );
64
65        let base_cols_width = VmPvsCols::<AB::Var>::width();
66        let (base_local, def_local) = local.split_at(base_cols_width);
67        let (base_next, next_local) = next.split_at(base_cols_width);
68
69        let local: &VmPvsCols<AB::Var> = (*base_local).borrow();
70        let next: &VmPvsCols<AB::Var> = (*base_next).borrow();
71
72        /*
73         * If deferrals are enabled, this AIR expects an additional deferral_flag column. It
74         * can be either 0 or 2 here, and in the latter case there can only be one row.
75         */
76        let (deferral_flag, has_vm_pvs) = if self.deferral_enabled {
77            debug_assert_eq!(def_local.len(), 1);
78            debug_assert_eq!(next_local.len(), 1);
79            self.eval_deferrals(builder, local, def_local[0], next_local[0])
80        } else {
81            debug_assert_eq!(def_local.len(), 0);
82            debug_assert_eq!(next_local.len(), 0);
83            (AB::Expr::ZERO, AB::Expr::ONE)
84        };
85
86        /*
87         * Basic constraints for non-public value columns.
88         */
89        // constrain all valid rows are at the beginning
90        builder.assert_bool(local.is_valid);
91        builder
92            .when_first_row()
93            .assert_eq(local.is_valid, has_vm_pvs);
94        builder
95            .when_transition()
96            .assert_bool(local.is_valid - next.is_valid);
97
98        // constrain increasing proof_idx
99        builder.when_first_row().assert_zero(local.proof_idx);
100        builder
101            .when_transition()
102            .when(and(local.is_valid, next.is_valid))
103            .assert_eq(local.proof_idx + AB::Expr::ONE, next.proof_idx);
104
105        // constrain is_last, note proof_idx on the first row is 0
106        builder.assert_bool(local.is_last);
107        builder.when(local.is_last).assert_one(local.is_valid);
108        builder
109            .when(and(local.is_valid, not(local.is_last)))
110            .assert_one(next.is_valid);
111        builder
112            .when(local.is_last)
113            .assert_zero(next.is_valid * next.proof_idx);
114        builder
115            .when_last_row()
116            .when(local.is_valid)
117            .assert_one(local.is_last);
118
119        // constrain has_verifier_pvs, which will be compared with the other pv AIRs
120        builder.assert_bool(local.has_verifier_pvs);
121        builder
122            .when(local.has_verifier_pvs)
123            .assert_one(local.is_valid);
124        builder
125            .when(and(local.is_valid, next.is_valid))
126            .assert_eq(local.has_verifier_pvs, next.has_verifier_pvs);
127
128        /*
129         * We first constrain segment adjacency, i.e. that rows in the trace are such that the
130         * first row is the (chronologically) first segment, and adjacent rows correspond to
131         * adjacent segments.
132         */
133        // constrain that is_terminate is the last valid proof
134        builder.assert_bool(local.child_pvs.is_terminate);
135        builder
136            .when(local.child_pvs.is_terminate)
137            .assert_one(local.is_last);
138        builder
139            .when(local.child_pvs.is_terminate)
140            .assert_zero(local.child_pvs.exit_code);
141
142        // constrain that non-terminal segments exited successfully
143        builder
144            .when(and(local.is_valid, not(local.child_pvs.is_terminate)))
145            .assert_eq(
146                local.child_pvs.exit_code,
147                AB::F::from_u32(DEFAULT_SUSPEND_EXIT_CODE),
148            );
149
150        // when local and next are valid, constrain increasing proof_idx and adjacency
151        let mut when_both_valid = builder.when(and(local.is_valid, not(local.is_last)));
152        when_both_valid.assert_eq(local.child_pvs.final_pc, next.child_pvs.initial_pc);
153        assert_array_eq(
154            &mut when_both_valid,
155            local.child_pvs.final_root,
156            next.child_pvs.initial_root,
157        );
158
159        /*
160         * We receive public values from ProofShapeModule to ensure the values being read here
161         * are correct. The leaf verifier reads public values from PROGRAM_AIR_ID,
162         * CONNECTOR_AIR_ID, and MERKLE_AID_ID while the internal verifier reads the full
163         * VmPvs from VM_PVS_AIR_ID.
164         */
165        let is_leaf = not(local.has_verifier_pvs);
166        let is_internal = local.has_verifier_pvs;
167
168        let mut internal_pv_idx = 0u8;
169        let internal_air_id = is_internal * AB::Expr::from_usize(VM_PVS_AIR_ID);
170        let mut internal_pp = || {
171            let ret = is_internal * AB::Expr::from_u8(internal_pv_idx);
172            internal_pv_idx += 1;
173            ret
174        };
175
176        // receive program_commit
177        let cond_program_air_id =
178            is_leaf.clone() * AB::Expr::from_usize(PROGRAM_AIR_ID) + internal_air_id.clone();
179
180        for (didx, value) in local.child_pvs.program_commit.iter().enumerate() {
181            self.public_values_bus.receive(
182                builder,
183                local.proof_idx,
184                PublicValuesBusMessage {
185                    air_idx: cond_program_air_id.clone(),
186                    pv_idx: is_leaf.clone() * AB::Expr::from_usize(didx) + internal_pp(),
187                    value: (*value).into(),
188                },
189                local.is_valid * is_internal,
190            );
191        }
192
193        // receive connector public values
194        let cond_connector_air_id =
195            is_leaf.clone() * AB::Expr::from_usize(CONNECTOR_AIR_ID) + internal_air_id.clone();
196
197        self.public_values_bus.receive(
198            builder,
199            local.proof_idx,
200            PublicValuesBusMessage {
201                air_idx: cond_connector_air_id.clone(),
202                pv_idx: internal_pp(),
203                value: local.child_pvs.initial_pc.into(),
204            },
205            local.is_valid,
206        );
207
208        self.public_values_bus.receive(
209            builder,
210            local.proof_idx,
211            PublicValuesBusMessage {
212                air_idx: cond_connector_air_id.clone(),
213                pv_idx: is_leaf.clone() + internal_pp(),
214                value: local.child_pvs.final_pc.into(),
215            },
216            local.is_valid,
217        );
218
219        self.public_values_bus.receive(
220            builder,
221            local.proof_idx,
222            PublicValuesBusMessage {
223                air_idx: cond_connector_air_id.clone(),
224                pv_idx: is_leaf.clone() * AB::Expr::TWO + internal_pp(),
225                value: local.child_pvs.exit_code.into(),
226            },
227            local.is_valid,
228        );
229
230        self.public_values_bus.receive(
231            builder,
232            local.proof_idx,
233            PublicValuesBusMessage {
234                air_idx: cond_connector_air_id.clone(),
235                pv_idx: is_leaf.clone() * AB::Expr::from_u8(3) + internal_pp(),
236                value: local.child_pvs.is_terminate.into(),
237            },
238            local.is_valid,
239        );
240
241        // receive memory Merkle public values
242        let cond_merkle_air_id =
243            is_leaf.clone() * AB::Expr::from_usize(MERKLE_AIR_ID) + internal_air_id.clone();
244
245        for (didx, value) in local.child_pvs.initial_root.iter().enumerate() {
246            self.public_values_bus.receive(
247                builder,
248                local.proof_idx,
249                PublicValuesBusMessage {
250                    air_idx: cond_merkle_air_id.clone(),
251                    pv_idx: is_leaf.clone() * AB::Expr::from_usize(didx) + internal_pp(),
252                    value: (*value).into(),
253                },
254                local.is_valid,
255            );
256        }
257
258        for (didx, value) in local.child_pvs.final_root.iter().enumerate() {
259            self.public_values_bus.receive(
260                builder,
261                local.proof_idx,
262                PublicValuesBusMessage {
263                    air_idx: cond_merkle_air_id.clone(),
264                    pv_idx: is_leaf.clone() * AB::Expr::from_usize(didx + DIGEST_SIZE)
265                        + internal_pp(),
266                    value: (*value).into(),
267                },
268                local.is_valid,
269            );
270        }
271
272        /*
273         * At the leaf level, this AIR is responsible for receiving the cached trace commit
274         * program_commit.
275         */
276        self.cached_commit_bus.receive(
277            builder,
278            local.proof_idx,
279            CachedCommitBusMessage {
280                air_idx: AB::Expr::from_usize(PROGRAM_AIR_ID),
281                cached_idx: AB::Expr::from_usize(PROGRAM_CACHED_TRACE_INDEX),
282                global_cached_idx: AB::Expr::ZERO,
283                cached_commit: local.child_pvs.program_commit.map(Into::into),
284            },
285            local.is_valid * is_leaf,
286        );
287
288        /*
289         * We look up proof metadata from VerifierPvsAir here to ensure consistency on each row.
290         */
291        self.pvs_air_consistency_bus.lookup_key(
292            builder,
293            local.proof_idx,
294            PvsAirConsistencyMessage {
295                deferral_flag,
296                has_verifier_pvs: local.has_verifier_pvs.into(),
297            },
298            local.is_valid,
299        );
300
301        /*
302         * Finally, we need to constrain that the public values this AIR produces are consistent
303         * with the child's. Initial output pvs must match the first row, and final output pvs
304         * must match the last.
305         */
306        let &VmPvs::<_> {
307            program_commit,
308            initial_pc,
309            final_pc,
310            exit_code,
311            is_terminate,
312            initial_root,
313            final_root,
314        } = builder.public_values().borrow();
315
316        // constrain first proof pvs
317        builder
318            .when_first_row()
319            .assert_eq(local.child_pvs.initial_pc, initial_pc);
320        assert_array_eq(
321            &mut builder.when_first_row(),
322            local.child_pvs.initial_root,
323            initial_root,
324        );
325
326        // constrain last proof pvs
327        builder
328            .when(local.is_last)
329            .assert_eq(local.child_pvs.final_pc, final_pc);
330        builder
331            .when(local.is_last)
332            .assert_eq(local.child_pvs.exit_code, exit_code);
333        builder
334            .when(local.is_last)
335            .assert_eq(local.child_pvs.is_terminate, is_terminate);
336        assert_array_eq(
337            &mut builder.when(local.is_last),
338            local.child_pvs.final_root,
339            final_root,
340        );
341
342        // constrain program_commit
343        assert_array_eq(
344            &mut builder.when(local.is_valid),
345            local.child_pvs.program_commit,
346            program_commit,
347        );
348    }
349}
350
351impl VmPvsAir {
352    fn eval_deferrals<AB>(
353        &self,
354        builder: &mut AB,
355        local: &VmPvsCols<AB::Var>,
356        local_def_flag: AB::Var,
357        next_def_flag: AB::Var,
358    ) -> (AB::Expr, AB::Expr)
359    where
360        AB: AirBuilder + InteractionBuilder + AirBuilderWithPublicValues,
361    {
362        /*
363         * Constrain that deferral_flag must be in {0, 1, 2}. If:
364         * - deferral_flag == 0: all proofs have VmPvs only, ignore deferral-related constraints
365         * - deferral_flag == 1: all proofs have DeferralPvs only, there should be no valid rows
366         *   and output public values should all be 0
367         * - deferral_flag == 2: there is a single child proof with both sets of pvs
368         */
369        builder.assert_tern(local_def_flag);
370        builder.assert_eq(local_def_flag, next_def_flag);
371
372        let mut when_deferral_flag = builder.when(local_def_flag);
373        when_deferral_flag.assert_zero(local.proof_idx);
374
375        let mut when_deferral_flag_two = when_deferral_flag.when_ne(local_def_flag, AB::Expr::ONE);
376        when_deferral_flag_two.assert_one(local.is_valid);
377        when_deferral_flag_two.assert_one(local.is_last);
378
379        let mut when_deferral_flag_one = when_deferral_flag.when_ne(local_def_flag, AB::Expr::TWO);
380        when_deferral_flag_one.assert_zero(local.is_valid);
381
382        let vm_pvs: &VmPvs<_> = builder.public_values().borrow();
383        let vm_pvs = vm_pvs.as_slice().to_vec();
384
385        for value in vm_pvs {
386            builder
387                .when(local_def_flag)
388                .when_ne(local_def_flag, AB::Expr::TWO)
389                .assert_zero(value);
390        }
391
392        (
393            local_def_flag.into(),
394            (local_def_flag - AB::Expr::ONE).square(),
395        )
396    }
397}