openvm_sdk/prover/
stark.rs

1use std::{borrow::Borrow, sync::Arc};
2
3use eyre::Result;
4use openvm_circuit::{
5    arch::{
6        hasher::poseidon2::vm_poseidon2_hasher, instructions::exe::VmExe, Executor,
7        MeteredExecutor, PreflightExecutor, VmBuilder, VmExecutionConfig,
8    },
9    system::memory::merkle::MerkleTree,
10};
11use openvm_stark_backend::{p3_field::PrimeField32, StarkEngine, Val};
12use openvm_stark_sdk::config::baby_bear_poseidon2::{Digest, F};
13use openvm_verify_stark_host::{
14    pvs::{DeferralPvs, DEF_PVS_AIR_ID},
15    vk::VerificationBaseline,
16    VmStarkProof,
17};
18
19use crate::{
20    prover::{
21        deferral::compute_deferral_merkle_proofs, vm::types::VmProvingKey, AggProver, AppProver,
22        InternalLayerMetadata,
23    },
24    DeferralInput, DeferralSetup, StdIn, SC,
25};
26
27pub struct StarkProver<E, VB>
28where
29    E: StarkEngine,
30    VB: VmBuilder<E>,
31{
32    pub app_prover: AppProver<E, VB>,
33    pub agg_prover: Arc<AggProver>,
34    pub deferral_setup: DeferralSetup,
35}
36
37impl<E, VB> StarkProver<E, VB>
38where
39    E: StarkEngine<SC = SC>,
40    VB: VmBuilder<E>,
41    Val<SC>: PrimeField32,
42{
43    pub fn new(
44        vm_builder: VB,
45        app_vm_pk: &VmProvingKey<VB::VmConfig>,
46        app_exe: Arc<VmExe<Val<SC>>>,
47        agg_prover: Arc<AggProver>,
48        deferral_setup: DeferralSetup,
49    ) -> Result<Self> {
50        Ok(Self {
51            app_prover: AppProver::new(vm_builder, app_vm_pk, app_exe)?,
52            agg_prover,
53            deferral_setup,
54        })
55    }
56
57    pub fn set_program_name(&mut self, program_name: impl AsRef<str>) -> &mut Self {
58        self.app_prover.set_program_name(program_name);
59        self
60    }
61
62    pub fn with_program_name(mut self, program_name: impl AsRef<str>) -> Self {
63        self.set_program_name(program_name);
64        self
65    }
66
67    pub fn prove(
68        &mut self,
69        vm_input: StdIn<Val<SC>>,
70        def_inputs: &[DeferralInput],
71    ) -> Result<(VmStarkProof, InternalLayerMetadata)>
72    where
73        <VB::VmConfig as VmExecutionConfig<Val<SC>>>::Executor: Executor<Val<SC>>
74            + MeteredExecutor<Val<SC>>
75            + PreflightExecutor<Val<SC>, VB::RecordArena>,
76    {
77        let has_deferrals = self.deferral_setup.hook_commit().is_some();
78        let memory_dimensions = self.app_prover.memory_dimensions();
79
80        // Build the initial memory merkle tree before proving (needed for deferral proofs).
81        let initial_merkle_tree = if has_deferrals {
82            let hasher = vm_poseidon2_hasher();
83            let initial_memory = &self
84                .app_prover
85                .instance()
86                .state()
87                .as_ref()
88                .expect("initial state should exist before proving")
89                .memory
90                .memory;
91            Some(MerkleTree::from_memory(
92                initial_memory,
93                &memory_dimensions,
94                &hasher,
95            ))
96        } else {
97            None
98        };
99
100        let continuation_proof = self.app_prover.prove(vm_input)?;
101        let (mut stark_proof, mut internal_metadata) =
102            self.agg_prover.prove_vm(continuation_proof)?;
103
104        // Skip aggregation unless some circuit received a deferred call. Note that
105        // deferrals are also skipped if def_inputs is an empty slice.
106        if def_inputs.iter().any(|input| !input.is_empty()) {
107            let def_agg_prover = self.deferral_setup.prover().ok_or_else(|| {
108                eyre::eyre!("non-empty deferral inputs require a deferral aggregation prover")
109            })?;
110            let def_hook_proofs = def_agg_prover
111                .multi_deferral_circuit_prover
112                .prove(def_inputs)?;
113            let (def_proof, def_internal_recursive_layer) =
114                def_agg_prover.agg_prover.prove_def(def_hook_proofs)?;
115            stark_proof = self.agg_prover.prove_mixed(
116                stark_proof,
117                def_proof,
118                &mut internal_metadata,
119                def_internal_recursive_layer,
120            )?;
121        }
122
123        // We add one additional internal_recursive layer to reduce the proof size.
124        const ADDITIONAL_INTERNAL_RECURSIVE_LAYERS: usize = 1;
125        for _ in 0..ADDITIONAL_INTERNAL_RECURSIVE_LAYERS {
126            stark_proof = self
127                .agg_prover
128                .wrap_proof(stark_proof, &mut internal_metadata)?;
129        }
130
131        // Generate deferral merkle proofs if deferrals are enabled.
132        if has_deferrals {
133            let hasher = vm_poseidon2_hasher();
134            let final_memory = &self
135                .app_prover
136                .instance()
137                .state()
138                .as_ref()
139                .expect("final state should exist after proving")
140                .memory
141                .memory;
142            let final_merkle_tree =
143                MerkleTree::from_memory(final_memory, &memory_dimensions, &hasher);
144
145            let def_pvs: &DeferralPvs<F> = stark_proof.inner.public_values[DEF_PVS_AIR_ID]
146                .as_slice()
147                .borrow();
148            let depth = def_pvs.depth.as_canonical_u32() as usize;
149
150            stark_proof.deferral_merkle_proofs = Some(compute_deferral_merkle_proofs(
151                memory_dimensions,
152                initial_merkle_tree.as_ref().unwrap(),
153                &final_merkle_tree,
154                depth,
155            ));
156        }
157
158        Ok((stark_proof, internal_metadata))
159    }
160
161    pub fn generate_baseline(&self) -> VerificationBaseline {
162        VerificationBaseline {
163            app_exe_commit: self.app_prover.app_exe_commit(),
164            memory_dimensions: self.app_prover.memory_dimensions(),
165            num_user_pvs: self.app_prover.num_user_pvs(),
166            app_vk_commit: self.agg_prover.leaf_prover.get_vk_commit(false),
167            leaf_vk_commit: self
168                .agg_prover
169                .internal_for_leaf_prover
170                .get_vk_commit(false),
171            internal_for_leaf_vk_commit: self
172                .agg_prover
173                .internal_recursive_prover
174                .get_vk_commit(false),
175            internal_recursive_vk_commit: self
176                .agg_prover
177                .internal_recursive_prover
178                .get_vk_commit(true),
179            expected_def_hook_commit: self.deferral_setup.hook_commit(),
180        }
181    }
182
183    pub fn app_vm_commit(&self) -> Digest {
184        self.agg_prover.vm_or_hook_commit()
185    }
186}