openvm_sdk/prover/deferral/
agg.rs1use 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 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 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 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 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 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 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
314struct 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}