openvm_sdk/prover/deferral/
agg.rs

1use std::sync::Arc;
2
3use eyre::{eyre, Result};
4use openvm_circuit::system::memory::dimensions::MemoryDimensions;
5use openvm_continuations::{
6    circuit::deferral::dummy::dummy_deferral_circuit_vk,
7    prover::{DeferralCircuitProver, DeferralCircuitProverKey},
8    SC,
9};
10use openvm_sdk_config::deferral::{DeferralConfig, SupportedDeferral};
11use openvm_stark_backend::{keygen::types::MultiStarkVerifyingKey, proof::Proof, SystemParams};
12use openvm_stark_sdk::config::baby_bear_poseidon2::Digest;
13use serde::{Deserialize, Serialize};
14
15use crate::{
16    config::{AggregationConfig, AggregationSystemParams, AggregationTreeConfig},
17    keygen::{AggPrefixProvingKey, AggProvingKey, DeferralCircuitProvingKey, DeferralProvingKey},
18    prover::{AggProver, MultiDeferralCircuitProver, SingleDeferralCircuitProver},
19};
20
21cfg_if::cfg_if! {
22    if #[cfg(feature = "cuda")] {
23        use openvm_verify_stark_circuit::prover::DeferredVerifyGpuProver as VerifyProver;
24        use openvm_verify_stark_circuit::prover::DeferredVerifyGpuCircuitProver as VerifyCircuitProver;
25        type E = openvm_cuda_backend::BabyBearPoseidon2GpuEngine;
26    } else {
27        use openvm_verify_stark_circuit::prover::DeferredVerifyCpuProver as VerifyProver;
28        use openvm_verify_stark_circuit::prover::DeferredVerifyCpuCircuitProver as VerifyCircuitProver;
29        type E = openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2CpuEngine;
30    }
31}
32
33#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
34pub struct DeferralHookCommits {
35    pub hook_cached_commit: Digest,
36    pub hook_commit: Digest,
37}
38
39impl DeferralHookCommits {
40    pub fn from_system_params(
41        agg_params: &AggregationSystemParams,
42        hook_params: SystemParams,
43    ) -> Self {
44        deferral_hook_artifacts_from_system_params(agg_params, hook_params).commits
45    }
46}
47
48pub struct DeferralAggProver {
49    pub multi_deferral_circuit_prover: Arc<MultiDeferralCircuitProver>,
50    pub agg_prover: Arc<AggProver>,
51}
52
53impl DeferralAggProver {
54    pub fn get_pk(&self) -> AggProvingKey {
55        AggProvingKey {
56            prefix: AggPrefixProvingKey {
57                leaf: self.agg_prover.leaf_prover.get_pk(),
58                internal_for_leaf: self.agg_prover.internal_for_leaf_prover.get_pk(),
59            },
60            internal_recursive: self.agg_prover.internal_recursive_prover.get_pk(),
61        }
62    }
63
64    pub fn new(
65        agg_config: AggregationConfig,
66        multi_deferral_circuit_prover: Arc<MultiDeferralCircuitProver>,
67    ) -> Self {
68        let agg_prover = AggProver::new(
69            multi_deferral_circuit_prover.def_hook_prover.get_vk(),
70            agg_config,
71            AggregationTreeConfig::deferral(),
72            Some(
73                multi_deferral_circuit_prover
74                    .def_hook_prover
75                    .get_cached_commit(),
76            ),
77        );
78        DeferralAggProver {
79            multi_deferral_circuit_prover,
80            agg_prover: Arc::new(agg_prover),
81        }
82    }
83
84    pub fn from_pk(
85        pk: AggProvingKey,
86        multi_deferral_circuit_prover: Arc<MultiDeferralCircuitProver>,
87    ) -> DeferralAggProver {
88        let agg_prover = AggProver::from_pk(
89            multi_deferral_circuit_prover.def_hook_prover.get_vk(),
90            pk,
91            AggregationTreeConfig::deferral(),
92            Some(
93                multi_deferral_circuit_prover
94                    .def_hook_prover
95                    .get_cached_commit(),
96            ),
97        );
98        DeferralAggProver {
99            multi_deferral_circuit_prover,
100            agg_prover: Arc::new(agg_prover),
101        }
102    }
103
104    pub fn def_hook_cached_commit(&self) -> Digest {
105        self.multi_deferral_circuit_prover
106            .def_hook_prover
107            .get_cached_commit()
108    }
109
110    pub fn def_hook_commit(&self) -> Digest {
111        self.agg_prover.vm_or_hook_commit()
112    }
113
114    /// Reconstructs a [`DeferralAggProver`] from cached SDK proving keys for deferral circuits
115    /// whose prover type is known to the SDK.
116    ///
117    /// The cached app VM config supplies the ordered [`DeferralConfig`], including each supported
118    /// deferral kind and the commit that the guest VM expects. The cached deferral proving key
119    /// supplies the matching circuit proving material, plus the shared internal-recursive and hook
120    /// proving keys. This constructor stitches those pieces back into a
121    /// [`MultiDeferralCircuitProver`] and then restores the deferral aggregation prover from
122    /// `deferral_agg_pk`.
123    ///
124    /// Custom deferral circuits cannot be reconstructed here because the SDK does not know their
125    /// concrete prover types; callers should manually build a [`MultiDeferralCircuitProver`] or
126    /// [`DeferralAggProver`] for those cases.
127    pub(crate) fn from_supported_deferral_pks(
128        deferral_config: &DeferralConfig,
129        deferral_pk: DeferralProvingKey,
130        deferral_agg_pk: AggProvingKey,
131    ) -> Result<Self> {
132        if deferral_config.circuits.len() != deferral_pk.circuits.len() {
133            return Err(eyre!(
134                "cached deferral proving key circuit count does not match app VM deferral config"
135            ));
136        }
137
138        let mut deferral_circuit_pks = deferral_pk.circuits.into_iter();
139        let first_config = deferral_config
140            .circuits
141            .first()
142            .ok_or_else(|| eyre!("app VM deferral config has no circuits"))?;
143        let first_pk = deferral_circuit_pks
144            .next()
145            .ok_or_else(|| eyre!("cached deferral proving key has no circuits"))?;
146        let mut multi_deferral_circuit_prover =
147            MultiDeferralCircuitProver::from_single_circuit_prover(
148                supported_deferral_circuit_prover_from_pk(&first_config.def_type, first_pk)?,
149                deferral_pk.def_internal_recursive_pk,
150                deferral_pk.def_hook_pk,
151            );
152
153        for (config, circuit_pk) in deferral_config
154            .circuits
155            .iter()
156            .skip(1)
157            .zip(deferral_circuit_pks)
158        {
159            multi_deferral_circuit_prover =
160                multi_deferral_circuit_prover.with_single_circuit_prover(
161                    supported_deferral_circuit_prover_from_pk(&config.def_type, circuit_pk)?,
162                );
163        }
164
165        // Validate that the reconstructed prover's circuit commits match the cached app VM deferral
166        // config. Done unconditionally (not just under debug_assertions): a mismatch means the
167        // config and keys disagree, which would otherwise surface as an opaque proof/verify
168        // failure.
169        let reconstructed_config = multi_deferral_circuit_prover.make_config(
170            deferral_config
171                .circuits
172                .iter()
173                .map(|circuit| circuit.def_type.clone())
174                .collect(),
175        );
176        if reconstructed_config.circuits != deferral_config.circuits {
177            return Err(eyre!(
178                "cached deferral proving key does not match app VM deferral config \
179                 (def_type/commit mismatch)"
180            ));
181        }
182
183        Ok(DeferralAggProver::from_pk(
184            deferral_agg_pk,
185            Arc::new(multi_deferral_circuit_prover),
186        ))
187    }
188
189    /// Builds a [`DeferralAggProver`] backed by the verify-stark circuit, configured so an SDK
190    /// with the given params can recursively verify the VM STARK proofs it produces, including its
191    /// own deferral-carrying proofs.
192    ///
193    /// The deferral-enabled internal-recursive vk and the self-referential `def_hook_commit` are
194    /// derived internally from a dummy deferral circuit.
195    pub fn verify_stark(
196        agg_params: &AggregationSystemParams,
197        hook_params: SystemParams,
198        memory_dimensions: MemoryDimensions,
199        num_user_pvs: usize,
200    ) -> Self {
201        let hook_artifacts =
202            deferral_hook_artifacts_from_system_params(agg_params, hook_params.clone());
203        let agg_config = AggregationConfig {
204            params: agg_params.clone(),
205        };
206        let agg_prover = hook_artifacts.agg_prover;
207
208        // The deferral-path aggregation tree's internal-recursive vk is a universal copy of the VM
209        // internal-recursive vk that a verify-stark circuit verifies.
210        let ir_vk = agg_prover.internal_recursive_prover.get_vk();
211        let ir_cached_commit = agg_prover
212            .internal_recursive_prover
213            .get_self_vk_pcs_data()
214            .expect("internal-recursive prover must expose its self vk pcs data")
215            .commitment
216            .into();
217
218        // Construct the verify-stark MultiDeferralCircuitProver, which should have the same hook vk
219        // and cached commit as the dummy one.
220        let deferred_verify_prover = VerifyProver::new::<E>(
221            ir_vk,
222            ir_cached_commit,
223            agg_params.internal.clone(),
224            memory_dimensions,
225            num_user_pvs,
226            Some(hook_artifacts.commits.hook_commit.into()),
227            0,
228        );
229        let verify_stark_prover = VerifyCircuitProver::new(deferred_verify_prover);
230        let multi_deferral_circuit_prover =
231            MultiDeferralCircuitProver::new(verify_stark_prover, agg_config, hook_params);
232
233        assert_eq!(
234            multi_deferral_circuit_prover
235                .def_hook_prover
236                .get_vk()
237                .pre_hash,
238            hook_artifacts.hook_vk.pre_hash
239        );
240        assert_eq!(
241            multi_deferral_circuit_prover
242                .def_hook_prover
243                .get_cached_commit(),
244            hook_artifacts.commits.hook_cached_commit
245        );
246
247        // Return the deferral-enabled verify-stark DeferralAggProver.
248        Self {
249            multi_deferral_circuit_prover: Arc::new(multi_deferral_circuit_prover),
250            agg_prover,
251        }
252    }
253}
254
255fn supported_deferral_circuit_prover_from_pk(
256    def_type: &SupportedDeferral,
257    circuit_pk: DeferralCircuitProvingKey,
258) -> Result<SingleDeferralCircuitProver> {
259    match def_type {
260        SupportedDeferral::VerifyStark => {
261            let verify_circuit_pk = circuit_pk.def_circuit_pk.as_ref().clone();
262            let verify_circuit_prover =
263                <VerifyCircuitProver as DeferralCircuitProver<SC>>::from_pk(verify_circuit_pk);
264            Ok(SingleDeferralCircuitProver::from_pks(
265                verify_circuit_prover,
266                circuit_pk.agg_prefix_pk.leaf,
267                circuit_pk.agg_prefix_pk.internal_for_leaf,
268            ))
269        }
270        SupportedDeferral::Other(name) => Err(eyre!(
271            "custom deferral circuit provers need to be manually created for deferral `{name}`"
272        )),
273    }
274}
275
276struct DeferralHookArtifacts {
277    commits: DeferralHookCommits,
278    hook_vk: Arc<MultiStarkVerifyingKey<SC>>,
279    agg_prover: Arc<AggProver>,
280}
281
282fn deferral_hook_artifacts_from_system_params(
283    agg_params: &AggregationSystemParams,
284    hook_params: SystemParams,
285) -> DeferralHookArtifacts {
286    let dummy = DummyDefCircuitProver {
287        vk: dummy_deferral_circuit_vk::<E>(agg_params.internal.clone()),
288    };
289    let agg_config = AggregationConfig {
290        params: agg_params.clone(),
291    };
292    let dummy_multi_deferral_circuit_prover =
293        MultiDeferralCircuitProver::new(dummy, agg_config.clone(), hook_params);
294    let hook_vk = dummy_multi_deferral_circuit_prover.def_hook_prover.get_vk();
295    let hook_cached_commit = dummy_multi_deferral_circuit_prover
296        .def_hook_prover
297        .get_cached_commit();
298    let agg_prover = Arc::new(AggProver::new(
299        hook_vk.clone(),
300        agg_config,
301        AggregationTreeConfig::deferral(),
302        Some(hook_cached_commit),
303    ));
304    DeferralHookArtifacts {
305        commits: DeferralHookCommits {
306            hook_cached_commit,
307            hook_commit: agg_prover.vm_or_hook_commit(),
308        },
309        hook_vk,
310        agg_prover,
311    }
312}
313
314/// A dummy [`DeferralCircuitProver`] that only exposes a trivial verifying key. It exists solely to
315/// seed the deferral aggregation chain when deriving the deferral path fixed point; its `prove`
316/// method is never called.
317struct DummyDefCircuitProver {
318    vk: Arc<MultiStarkVerifyingKey<SC>>,
319}
320
321impl DeferralCircuitProver<SC> for DummyDefCircuitProver {
322    fn get_pk(&self) -> Arc<DeferralCircuitProverKey<SC>> {
323        unreachable!("DummyDefCircuitProver does not have proving material")
324    }
325
326    fn from_pk(_encoded_pk: DeferralCircuitProverKey<SC>) -> Self
327    where
328        Self: Sized,
329    {
330        unreachable!("DummyDefCircuitProver does not have proving material")
331    }
332
333    fn get_vk(&self) -> Arc<MultiStarkVerifyingKey<SC>> {
334        self.vk.clone()
335    }
336
337    fn prove(&self, _input_bytes: &[u8]) -> Proof<SC> {
338        unreachable!("DummyDefCircuitProver is only used to derive deferral path artifacts")
339    }
340
341    fn get_def_idx(&self) -> usize {
342        0
343    }
344}