openvm_continuations/circuit/inner/verifier/
trace.rs

1use std::borrow::{Borrow, BorrowMut};
2
3use openvm_circuit::arch::POSEIDON2_WIDTH;
4use openvm_cpu_backend::CpuBackend;
5use openvm_stark_backend::{proof::Proof, prover::AirProvingContext};
6use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, F};
7use openvm_verify_stark_host::pvs::{
8    VerifierBasePvs, VerifierDefPvs, VkCommit, VERIFIER_PVS_AIR_ID,
9};
10use p3_field::{Field, PrimeCharacteristicRing, PrimeField32};
11use p3_matrix::dense::RowMajorMatrix;
12
13use crate::circuit::{
14    inner::{
15        verifier::air::{VerifierCombinedPvs, VerifierDeferralCols, VerifierPvsCols},
16        ProofsType,
17    },
18    subair::hash_slice_trace,
19    SingleAirTraceData,
20};
21
22#[derive(Copy, Clone)]
23pub enum VerifierChildLevel {
24    App,
25    Leaf,
26    InternalForLeaf,
27    InternalRecursive,
28}
29
30pub fn generate_proving_ctx(
31    proofs: &[Proof<BabyBearPoseidon2Config>],
32    proofs_type: ProofsType,
33    child_is_app: bool,
34    child_vk_commit: VkCommit<F>,
35    deferral_enabled: bool,
36) -> SingleAirTraceData<CpuBackend<BabyBearPoseidon2Config>> {
37    let num_proofs = proofs.len();
38    debug_assert!(num_proofs > 0);
39
40    if !deferral_enabled {
41        assert!(matches!(proofs_type, ProofsType::Vm))
42    }
43
44    let mut child_level = VerifierChildLevel::App;
45
46    if !child_is_app {
47        let proof = &proofs[0];
48        let child_pvs: &VerifierBasePvs<F> = proof.public_values[VERIFIER_PVS_AIR_ID].as_slice()
49            [0..VerifierBasePvs::<F>::width()]
50            .borrow();
51        child_level = match child_pvs.internal_flag {
52            F::ZERO => VerifierChildLevel::Leaf,
53            F::ONE => VerifierChildLevel::InternalForLeaf,
54            F::TWO => VerifierChildLevel::InternalRecursive,
55            _ => unreachable!(),
56        };
57    }
58
59    let height = num_proofs.next_power_of_two();
60    let base_width = VerifierPvsCols::<u8>::width();
61    let def_width = if deferral_enabled {
62        VerifierDeferralCols::<u8>::width()
63    } else {
64        0
65    };
66    let width = base_width + def_width;
67
68    let mut trace = vec![F::ZERO; height * width];
69    let mut chunks = trace.chunks_exact_mut(width);
70    let mut poseidon2_compress_inputs = vec![];
71    let mut poseidon2_permute_inputs = vec![];
72    let mut range_check_inputs = vec![];
73    let mut trailing_deferral_flag = F::ZERO;
74
75    for (proof_idx, proof) in proofs.iter().enumerate() {
76        let chunk = chunks.next().unwrap();
77        let (base_chunk, def_chunk) = chunk.split_at_mut(base_width);
78
79        let cols: &mut VerifierPvsCols<F> = base_chunk.borrow_mut();
80        cols.proof_idx = F::from_usize(proof_idx);
81        cols.is_valid = F::ONE;
82
83        if deferral_enabled {
84            let def_cols: &mut VerifierDeferralCols<_> = def_chunk.borrow_mut();
85            def_cols.is_last = F::from_bool(proof_idx + 1 == proofs.len());
86            if matches!(proofs_type, ProofsType::Deferral) {
87                def_cols.child_pvs.deferral_flag = F::ONE;
88                trailing_deferral_flag = def_cols.child_pvs.deferral_flag;
89            }
90        }
91
92        if !child_is_app {
93            let pv_chunk = proof.public_values[VERIFIER_PVS_AIR_ID].as_slice();
94            let (base_pv_chunk, def_pv_chunk) = pv_chunk.split_at(VerifierBasePvs::<u8>::width());
95
96            let base_pvs: &VerifierBasePvs<_> = base_pv_chunk.borrow();
97            cols.has_verifier_pvs = F::ONE;
98            cols.child_pvs = *base_pvs;
99
100            if deferral_enabled {
101                let def_cols: &mut VerifierDeferralCols<_> = def_chunk.borrow_mut();
102                let def_pvs: &VerifierDefPvs<_> = def_pv_chunk.borrow();
103                def_cols.child_pvs = *def_pvs;
104                trailing_deferral_flag = def_pvs.deferral_flag;
105            }
106        }
107
108        let depth = cols.child_pvs.recursion_depth.as_canonical_u32();
109        cols.recursion_flag = F::from_u32(depth.min(2));
110        cols.depth_inv = if depth >= 2 {
111            (cols.child_pvs.recursion_depth * (cols.child_pvs.recursion_depth - F::ONE)).inverse()
112        } else {
113            F::ZERO
114        };
115        range_check_inputs.push(depth as usize);
116    }
117
118    if deferral_enabled {
119        for chunk in chunks {
120            let (_, def_chunk) = chunk.split_at_mut(base_width);
121            let def_cols: &mut VerifierDeferralCols<_> = def_chunk.borrow_mut();
122            def_cols.child_pvs.deferral_flag = trailing_deferral_flag;
123        }
124    }
125
126    let first_row: &VerifierPvsCols<F> = trace[..base_width].borrow();
127    let mut base_pvs = first_row.child_pvs;
128
129    match child_level {
130        VerifierChildLevel::App => {
131            base_pvs.app_vk_commit = child_vk_commit;
132        }
133        VerifierChildLevel::Leaf => {
134            base_pvs.leaf_vk_commit = child_vk_commit;
135            base_pvs.internal_flag = F::ONE;
136        }
137        VerifierChildLevel::InternalForLeaf => {
138            base_pvs.internal_for_leaf_vk_commit = child_vk_commit;
139            base_pvs.internal_flag = F::TWO;
140            base_pvs.recursion_depth = F::ONE;
141        }
142        VerifierChildLevel::InternalRecursive => {
143            base_pvs.internal_recursive_vk_commit = child_vk_commit;
144            base_pvs.internal_flag = F::TWO;
145            base_pvs.recursion_depth = first_row.child_pvs.recursion_depth + F::ONE;
146        }
147    }
148
149    let deferral_flag_pv = match proofs_type {
150        ProofsType::Vm => F::ZERO,
151        ProofsType::Deferral => F::ONE,
152        ProofsType::Mix => {
153            assert_eq!(num_proofs, 2);
154            F::TWO
155        }
156        ProofsType::Combined => {
157            assert_eq!(num_proofs, 1);
158            F::TWO
159        }
160    };
161
162    let mut def_hook_commit = None;
163    if deferral_enabled && deferral_flag_pv == F::ONE && base_pvs.internal_flag == F::TWO {
164        let hash_elements = [
165            base_pvs.app_vk_commit.cached_commit,
166            base_pvs.app_vk_commit.vk_pre_hash,
167            base_pvs.leaf_vk_commit.cached_commit,
168            base_pvs.leaf_vk_commit.vk_pre_hash,
169            base_pvs.internal_for_leaf_vk_commit.cached_commit,
170            base_pvs.internal_for_leaf_vk_commit.vk_pre_hash,
171        ];
172
173        let mut row_compress_inputs = vec![];
174        let mut row_permute_inputs = vec![];
175        let (intermediate_states_vec, computed_def_hook_commit) = hash_slice_trace(
176            &hash_elements,
177            Some(&mut row_permute_inputs),
178            Some(&mut row_compress_inputs),
179        );
180        let intermediate_states: [[F; POSEIDON2_WIDTH]; 5] =
181            intermediate_states_vec.try_into().unwrap();
182
183        for chunk in trace.chunks_exact_mut(width) {
184            let (_, def_chunk) = chunk.split_at_mut(base_width);
185            let def_cols: &mut VerifierDeferralCols<_> = def_chunk.borrow_mut();
186            def_cols.intermediate_states = intermediate_states;
187        }
188
189        for &input in &row_compress_inputs {
190            poseidon2_compress_inputs.extend((0..height).map(|_| input));
191        }
192        for &input in &row_permute_inputs {
193            poseidon2_permute_inputs.extend((0..height).map(|_| input));
194        }
195        def_hook_commit = Some(computed_def_hook_commit);
196    }
197
198    let public_values = if deferral_enabled {
199        let last_row_def: &VerifierDeferralCols<F> =
200            trace[(num_proofs - 1) * width + base_width..num_proofs * width].borrow();
201        let mut def_pvs = last_row_def.child_pvs;
202        def_pvs.deferral_flag = deferral_flag_pv;
203
204        if let Some(def_hook_commit) = def_hook_commit {
205            def_pvs.def_hook_commit = def_hook_commit;
206        }
207
208        let mut combined = vec![F::ZERO; VerifierCombinedPvs::<u8>::width()];
209        let combined_pvs: &mut VerifierCombinedPvs<F> = combined.as_mut_slice().borrow_mut();
210        combined_pvs.base = base_pvs;
211        combined_pvs.def = def_pvs;
212        combined
213    } else {
214        base_pvs.to_vec()
215    };
216
217    SingleAirTraceData {
218        air_proving_ctx: AirProvingContext {
219            cached_mains: vec![],
220            common_main: RowMajorMatrix::new(trace, width),
221            public_values,
222        },
223        poseidon2_compress_inputs,
224        poseidon2_permute_inputs,
225        range_check_inputs,
226    }
227}