1use std::{marker::PhantomData, sync::Arc};
2
3use eyre::Result;
4use openvm_circuit::system::memory::{
5 dimensions::MemoryDimensions, merkle::public_values::UserPublicValuesProof,
6};
7#[cfg(debug_assertions)]
8use openvm_continuations::prover::debug_constraints;
9use openvm_continuations::{
10 circuit::{deferral::DeferralMerkleProofs, Circuit},
11 prover::{DeferralCircuitProver, DeferralCircuitProverKey},
12 CommitBytes, VkCommitBytes, SC,
13};
14use openvm_cpu_backend::CpuBackend;
15#[cfg(feature = "cuda")]
16use openvm_cuda_backend::{BabyBearPoseidon2GpuEngine, GpuBackend};
17use openvm_recursion_circuit::system::{
18 AggregationSubCircuit, VerifierConfig, VerifierSubCircuit, VerifierTraceGen,
19};
20use openvm_stark_backend::{
21 codec::{Decode, Encode},
22 keygen::types::{MultiStarkProvingKey, MultiStarkVerifyingKey},
23 proof::Proof,
24 prover::{CommittedTraceData, DeviceDataTransporter, ProverBackend, ProverDevice},
25 EngineDeviceCtx, StarkEngine, SystemParams,
26};
27use openvm_stark_sdk::config::baby_bear_poseidon2::{
28 BabyBearPoseidon2CpuEngine, Digest, DIGEST_SIZE, EF, F,
29};
30use openvm_verify_stark_host::VmStarkProof;
31use p3_field::{Field, PrimeField32};
32use serde::{Deserialize, Serialize};
33use tracing::instrument;
34
35use crate::{DeferredVerifyCircuit, DeferredVerifyTraceGen, DeferredVerifyTraceGenImpl};
36
37mod trace;
38
39pub type DeferredVerifyCpuProver =
40 DeferredVerifyProver<CpuBackend<SC>, VerifierSubCircuit<1>, DeferredVerifyTraceGenImpl>;
41pub type DeferredVerifyCpuCircuitProver = DeferredVerifyCircuitProver<
42 BabyBearPoseidon2CpuEngine,
43 VerifierSubCircuit<1>,
44 DeferredVerifyTraceGenImpl,
45>;
46
47#[cfg(feature = "cuda")]
48pub type DeferredVerifyGpuProver =
49 DeferredVerifyProver<GpuBackend, VerifierSubCircuit<1>, DeferredVerifyTraceGenImpl>;
50#[cfg(feature = "cuda")]
51pub type DeferredVerifyGpuCircuitProver = DeferredVerifyCircuitProver<
52 BabyBearPoseidon2GpuEngine,
53 VerifierSubCircuit<1>,
54 DeferredVerifyTraceGenImpl,
55>;
56
57pub struct DeferredVerifyProver<
58 PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
59 S: AggregationSubCircuit,
60 T,
61> {
62 pk: Arc<MultiStarkProvingKey<SC>>,
63 vk: Arc<MultiStarkVerifyingKey<SC>>,
64
65 agg_node_tracegen: T,
66
67 child_vk: Arc<MultiStarkVerifyingKey<SC>>,
68 child_vk_pcs_data: CommittedTraceData<PB>,
69 circuit: Arc<DeferredVerifyCircuit<S>>,
70}
71
72impl<PB, S, T> DeferredVerifyProver<PB, S, T>
73where
74 PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
75 S: AggregationSubCircuit,
76 PB::Matrix: Clone,
77{
78 #[instrument(name = "total_proof", skip_all)]
79 pub fn prove<E: StarkEngine<SC = SC, PB = PB>>(
80 &self,
81 proof: Proof<SC>,
82 user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, PB::Val>,
83 deferral_merkle_proofs: Option<&DeferralMerkleProofs<PB::Val>>,
84 ) -> Result<Proof<SC>>
85 where
86 S: AggregationSubCircuit + VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
87 T: DeferredVerifyTraceGen<PB, EngineDeviceCtx<E>>,
88 {
89 assert!(
90 deferral_merkle_proofs.is_none() || self.circuit.def_hook_commit.is_some(),
91 "def_hook_commit must be defined to verify child proof with deferrals"
92 );
93 let engine = E::new(self.pk.params.clone());
94 let ctx = self.generate_proving_ctx(
95 proof,
96 user_pvs_proof,
97 deferral_merkle_proofs,
98 engine.device().device_ctx(),
99 );
100 #[cfg(debug_assertions)]
101 debug_constraints(&self.circuit, &ctx, &engine);
102 let d_pk = engine.device().transport_pk_to_device(self.pk.as_ref());
103 let proof = engine.prove(&d_pk, ctx)?;
104 #[cfg(debug_assertions)]
105 engine.verify(&self.vk, &proof)?;
106 Ok(proof)
107 }
108
109 #[instrument(name = "total_proof", skip_all)]
110 pub fn prove_no_def<E: StarkEngine<SC = SC, PB = PB>>(
111 &self,
112 proof: Proof<SC>,
113 user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, PB::Val>,
114 ) -> Result<Proof<SC>>
115 where
116 S: AggregationSubCircuit + VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
117 T: DeferredVerifyTraceGen<PB, EngineDeviceCtx<E>>,
118 {
119 self.prove::<E>(proof, user_pvs_proof, None)
120 }
121
122 pub fn new<E: StarkEngine<SC = SC, PB = PB>>(
123 child_vk: Arc<MultiStarkVerifyingKey<SC>>,
124 internal_recursive_cached_commit: CommitBytes,
125 system_params: SystemParams,
126 memory_dimensions: MemoryDimensions,
127 num_user_pvs: usize,
128 def_hook_commit: Option<CommitBytes>,
129 def_idx: usize,
130 ) -> Self
131 where
132 S: AggregationSubCircuit + VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
133 T: DeferredVerifyTraceGen<PB, EngineDeviceCtx<E>>,
134 E::PD: DeviceDataTransporter<SC, PB> + Clone,
135 PB::Val: Field + PrimeField32,
136 PB::Matrix: Clone,
137 {
138 let verifier_circuit = S::new(
139 child_vk.clone(),
140 VerifierConfig {
141 continuations_enabled: true,
142 final_state_bus_enabled: true,
143 has_cached: true,
144 },
145 );
146 let engine = E::new(system_params);
147 let internal_recursive_vk_commit = VkCommitBytes {
148 cached_commit: internal_recursive_cached_commit,
149 vk_pre_hash: child_vk.pre_hash.into(),
150 };
151 let child_vk_pcs_data = verifier_circuit.commit_child_vk(&engine, &child_vk);
152 let circuit = Arc::new(DeferredVerifyCircuit::new(
153 Arc::new(verifier_circuit),
154 internal_recursive_vk_commit,
155 def_hook_commit,
156 memory_dimensions,
157 num_user_pvs,
158 def_idx,
159 ));
160 let (pk, vk) = engine.keygen(&circuit.airs());
161 Self {
162 pk: Arc::new(pk),
163 vk: Arc::new(vk),
164 agg_node_tracegen: T::new(def_hook_commit.is_some()),
165 child_vk,
166 child_vk_pcs_data,
167 circuit,
168 }
169 }
170
171 #[allow(clippy::too_many_arguments)]
172 pub fn from_pk<E: StarkEngine<SC = SC, PB = PB>>(
173 child_vk: Arc<MultiStarkVerifyingKey<SC>>,
174 internal_recursive_cached_commit: CommitBytes,
175 pk: Arc<MultiStarkProvingKey<SC>>,
176 memory_dimensions: MemoryDimensions,
177 num_user_pvs: usize,
178 def_hook_commit: Option<CommitBytes>,
179 def_idx: usize,
180 ) -> Self
181 where
182 S: AggregationSubCircuit + VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
183 T: DeferredVerifyTraceGen<PB, EngineDeviceCtx<E>>,
184 PB::Matrix: Clone,
185 {
186 let verifier_circuit = S::new(
187 child_vk.clone(),
188 VerifierConfig {
189 continuations_enabled: true,
190 final_state_bus_enabled: true,
191 has_cached: true,
192 },
193 );
194 let internal_recursive_vk_commit = VkCommitBytes {
195 cached_commit: internal_recursive_cached_commit,
196 vk_pre_hash: child_vk.pre_hash.into(),
197 };
198 let engine = E::new(pk.params.clone());
199 let child_vk_pcs_data = verifier_circuit.commit_child_vk(&engine, &child_vk);
200 let circuit = Arc::new(DeferredVerifyCircuit::new(
203 Arc::new(verifier_circuit),
204 internal_recursive_vk_commit,
205 def_hook_commit,
206 memory_dimensions,
207 num_user_pvs,
208 def_idx,
209 ));
210 let vk = Arc::new(pk.get_vk());
211 Self {
212 pk,
213 vk,
214 agg_node_tracegen: T::new(def_hook_commit.is_some()),
215 child_vk,
216 child_vk_pcs_data,
217 circuit,
218 }
219 }
220
221 pub fn get_circuit(&self) -> Arc<DeferredVerifyCircuit<S>> {
222 self.circuit.clone()
223 }
224
225 pub fn get_pk(&self) -> Arc<MultiStarkProvingKey<SC>> {
226 self.pk.clone()
227 }
228
229 pub fn get_vk(&self) -> Arc<MultiStarkVerifyingKey<SC>> {
230 self.vk.clone()
231 }
232
233 pub fn get_cached_commit(&self) -> <PB as ProverBackend>::Commitment {
234 self.child_vk_pcs_data.commitment
235 }
236}
237
238pub struct DeferredVerifyCircuitProver<
239 E: StarkEngine<SC = SC>,
240 S: AggregationSubCircuit + VerifierTraceGen<E::PB, SC, EngineDeviceCtx<E>>,
241 T: DeferredVerifyTraceGen<E::PB, EngineDeviceCtx<E>>,
242> {
243 prover: DeferredVerifyProver<E::PB, S, T>,
244 phantom: PhantomData<E>,
245}
246
247impl<E, S, T> DeferredVerifyCircuitProver<E, S, T>
248where
249 E: StarkEngine<SC = SC>,
250 E::PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
251 S: AggregationSubCircuit + VerifierTraceGen<E::PB, SC, EngineDeviceCtx<E>>,
252 T: DeferredVerifyTraceGen<E::PB, EngineDeviceCtx<E>>,
253{
254 pub fn new(prover: DeferredVerifyProver<E::PB, S, T>) -> Self {
255 Self {
256 prover,
257 phantom: PhantomData,
258 }
259 }
260}
261
262impl<PB, S, T, E> DeferralCircuitProver<SC> for DeferredVerifyCircuitProver<E, S, T>
263where
264 PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
265 S: AggregationSubCircuit + VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
266 T: DeferredVerifyTraceGen<PB, EngineDeviceCtx<E>>,
267 E: StarkEngine<PB = PB, SC = SC>,
268 PB::Matrix: Clone,
269{
270 fn from_pk(pk: DeferralCircuitProverKey<SC>) -> Self {
271 let aux = DeferralVerifyProvingAux::decode(&mut pk.aux.as_slice())
272 .expect("failed to decode verify-stark deferral proving aux");
273 Self::new(DeferredVerifyProver::from_pk::<E>(
274 aux.child_vk,
275 aux.internal_recursive_cached_commit,
276 pk.base_pk,
277 aux.memory_dimensions,
278 aux.num_user_pvs,
279 aux.def_hook_commit,
280 aux.def_idx,
281 ))
282 }
283
284 fn get_pk(&self) -> Arc<DeferralCircuitProverKey<SC>> {
285 let aux = DeferralVerifyProvingAux {
286 child_vk: self.prover.child_vk.clone(),
287 internal_recursive_cached_commit: self
288 .prover
289 .circuit
290 .internal_recursive_vk_commit
291 .cached_commit,
292 memory_dimensions: self.prover.circuit.memory_dimensions,
293 num_user_pvs: self.prover.circuit.num_user_pvs,
294 def_hook_commit: self.prover.circuit.def_hook_commit,
295 def_idx: self.prover.circuit.def_idx,
296 };
297 let mut encoded_aux = Vec::new();
298 aux.encode(&mut encoded_aux)
299 .expect("failed to encode verify-stark deferral proving aux");
300 Arc::new(DeferralCircuitProverKey {
301 base_pk: self.prover.get_pk(),
302 aux: encoded_aux,
303 })
304 }
305
306 fn get_vk(&self) -> Arc<MultiStarkVerifyingKey<SC>> {
307 self.prover.get_vk()
308 }
309
310 fn prove(&self, input_bytes: &[u8]) -> Proof<SC> {
311 let vm_proof = VmStarkProof::decode_from_bytes(input_bytes).unwrap();
312 self.prover
313 .prove::<E>(
314 vm_proof.inner,
315 &vm_proof.user_pvs_proof,
316 vm_proof.deferral_merkle_proofs.as_ref(),
317 )
318 .expect("DeferredVerifyProver::prove failed")
319 }
320
321 fn get_def_idx(&self) -> usize {
322 self.prover.circuit.def_idx
323 }
324
325 fn cached_commits(&self) -> Vec<CommitBytes> {
326 vec![self.prover.get_cached_commit().into()]
327 }
328}
329
330#[derive(Clone, Serialize, Deserialize)]
331struct DeferralVerifyProvingAux {
332 child_vk: Arc<MultiStarkVerifyingKey<SC>>,
333 internal_recursive_cached_commit: CommitBytes,
334 memory_dimensions: MemoryDimensions,
335 num_user_pvs: usize,
336 def_hook_commit: Option<CommitBytes>,
337 def_idx: usize,
338}
339
340impl Encode for DeferralVerifyProvingAux {
341 fn encode<W: std::io::Write>(&self, writer: &mut W) -> std::io::Result<()> {
342 let bytes = bitcode::serialize(self).map_err(std::io::Error::other)?;
343 std::io::Write::write_all(writer, &(bytes.len() as u64).to_le_bytes())?;
344 std::io::Write::write_all(writer, &bytes)
345 }
346}
347
348impl Decode for DeferralVerifyProvingAux {
349 fn decode<R: std::io::Read>(reader: &mut R) -> std::io::Result<Self> {
350 let mut len_bytes = [0u8; 8];
351 std::io::Read::read_exact(reader, &mut len_bytes)?;
352 let len = u64::from_le_bytes(len_bytes) as usize;
353 let mut bytes = vec![0u8; len];
354 std::io::Read::read_exact(reader, &mut bytes)?;
355 bitcode::deserialize(&bytes).map_err(std::io::Error::other)
356 }
357}