openvm_continuations/circuit/inner/vm_pvs/
air.rs1use 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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}