openvm_verify_stark_circuit/prover/
mod.rs

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        // WARNING: def_idx must match the original def_idx used when generating the pk,
201        // or else the generated proof will be incorrect.
202        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}