openvm_continuations/circuit/inner/vm_pvs/
trace.rs

1use std::borrow::{Borrow, BorrowMut};
2
3use openvm_circuit::system::{connector::VmConnectorPvs, memory::merkle::MemoryMerklePvs};
4use openvm_cpu_backend::CpuBackend;
5use openvm_stark_backend::{proof::Proof, prover::AirProvingContext};
6use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, DIGEST_SIZE, F};
7use openvm_verify_stark_host::pvs::{VmPvs, VM_PVS_AIR_ID};
8use p3_field::PrimeCharacteristicRing;
9use p3_matrix::dense::RowMajorMatrix;
10
11use crate::circuit::inner::{app::*, vm_pvs::air::VmPvsCols, ProofsType};
12
13pub fn generate_proving_ctx(
14    proofs: &[Proof<BabyBearPoseidon2Config>],
15    proofs_type: ProofsType,
16    child_is_app: bool,
17    deferral_enabled: bool,
18) -> AirProvingContext<CpuBackend<BabyBearPoseidon2Config>> {
19    debug_assert!(!proofs.is_empty());
20
21    let num_vm_proofs = match proofs_type {
22        ProofsType::Vm => proofs.len(),
23        ProofsType::Deferral => 0,
24        ProofsType::Mix | ProofsType::Combined => 1,
25    };
26
27    let height = num_vm_proofs.next_power_of_two();
28    let base_width = VmPvsCols::<u8>::width();
29    let width = base_width + deferral_enabled as usize;
30
31    let mut trace = vec![F::ZERO; height * width];
32    for (proof_idx, (proof, chunk)) in proofs[0..num_vm_proofs.max(1)]
33        .iter()
34        .zip(trace.chunks_exact_mut(width))
35        .enumerate()
36    {
37        let (base_chunk, def_chunk) = chunk.split_at_mut(base_width);
38        let cols: &mut VmPvsCols<F> = base_chunk.borrow_mut();
39        cols.proof_idx = F::from_usize(proof_idx);
40
41        if deferral_enabled {
42            def_chunk[0] = match proofs_type {
43                ProofsType::Vm | ProofsType::Mix => F::ZERO,
44                ProofsType::Deferral => F::ONE,
45                ProofsType::Combined => F::TWO,
46            };
47            if def_chunk[0] == F::ONE {
48                continue;
49            }
50        }
51
52        cols.is_valid = F::ONE;
53        cols.is_last = F::from_bool(proof_idx + 1 == num_vm_proofs);
54
55        if child_is_app {
56            cols.child_pvs.program_commit = proof.trace_vdata[PROGRAM_AIR_ID]
57                .as_ref()
58                .expect("program trace vdata must be present for app children")
59                .cached_commitments[PROGRAM_CACHED_TRACE_INDEX];
60
61            let &VmConnectorPvs {
62                initial_pc,
63                final_pc,
64                exit_code,
65                is_terminate,
66            } = proof.public_values[CONNECTOR_AIR_ID].as_slice().borrow();
67            cols.child_pvs.initial_pc = initial_pc;
68            cols.child_pvs.final_pc = final_pc;
69            cols.child_pvs.exit_code = exit_code;
70            cols.child_pvs.is_terminate = is_terminate;
71
72            let &MemoryMerklePvs::<_, DIGEST_SIZE> {
73                initial_root,
74                final_root,
75            } = proof.public_values[MERKLE_AIR_ID].as_slice().borrow();
76            cols.child_pvs.initial_root = initial_root;
77            cols.child_pvs.final_root = final_root;
78        } else {
79            cols.has_verifier_pvs = F::ONE;
80            let child_pvs: &VmPvs<F> = proof.public_values[VM_PVS_AIR_ID].as_slice().borrow();
81            cols.child_pvs = *child_pvs;
82        }
83    }
84
85    let mut public_values = vec![F::ZERO; VmPvs::<u8>::width()];
86    let pvs: &mut VmPvs<F> = public_values.as_mut_slice().borrow_mut();
87
88    if num_vm_proofs > 0 {
89        let first_row: &VmPvsCols<F> = trace[..base_width].borrow();
90        let last_row: &VmPvsCols<F> =
91            trace[(num_vm_proofs - 1) * width..(num_vm_proofs - 1) * width + base_width].borrow();
92
93        pvs.program_commit = first_row.child_pvs.program_commit;
94        pvs.initial_pc = first_row.child_pvs.initial_pc;
95        pvs.initial_root = first_row.child_pvs.initial_root;
96
97        pvs.final_pc = last_row.child_pvs.final_pc;
98        pvs.exit_code = last_row.child_pvs.exit_code;
99        pvs.is_terminate = last_row.child_pvs.is_terminate;
100        pvs.final_root = last_row.child_pvs.final_root;
101    }
102
103    AirProvingContext {
104        cached_mains: vec![],
105        common_main: RowMajorMatrix::new(trace, width),
106        public_values,
107    }
108}