openvm_continuations/circuit/inner/verifier/
trace.rs1use 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}