1use std::{
2 marker::PhantomData,
3 sync::{Arc, OnceLock},
4};
5
6use eyre::eyre;
7use openvm_circuit::arch::{VmBuilder, VmExecutionConfig, VmExecutor};
8use openvm_sdk_config::{SdkVmConfig, TranspilerConfig};
9use openvm_stark_backend::StarkEngine;
10use openvm_transpiler::transpiler::Transpiler;
11#[cfg(feature = "evm-prove")]
12use {
13 crate::{
14 config::Halo2Config, halo2_params::CacheHalo2ParamsReader, keygen::Halo2ProvingKey,
15 prover::Halo2Prover,
16 },
17 openvm_static_verifier::StaticVerifierShape,
18 std::path::Path,
19};
20#[cfg(feature = "root-prover")]
21use {
22 crate::{keygen::RootProvingKey, prover::RootProver},
23 openvm_stark_backend::SystemParams,
24 openvm_stark_sdk::config::root_params_with_100_bits_security,
25};
26
27use crate::{
28 config::{AggregationConfig, AggregationSystemParams, AggregationTreeConfig, AppConfig},
29 keygen::{AggProvingKey, AppProvingKey, SdkCachedProvingKey},
30 prover::{AggProver, DeferralAggProver, DeferralHookCommits, MultiDeferralCircuitProver},
31 DeferralSetup, GenericSdk, SdkError, F, SC,
32};
33
34enum AppSource<VC> {
35 Config(AppConfig<VC>),
36 Pk(AppProvingKey<VC>),
37}
38
39#[allow(clippy::large_enum_variant)]
40enum AggSource {
41 Params(AggregationSystemParams),
42 Pk(AggProvingKey),
43}
44
45#[allow(clippy::large_enum_variant)]
46enum DeferralSource {
47 HookCommits(DeferralHookCommits),
48 MultiCircuitProver(MultiDeferralCircuitProver),
49 AggProver(DeferralAggProver),
50}
51
52#[cfg(feature = "root-prover")]
53enum RootSource {
54 Params(SystemParams),
55 Pk(RootProvingKey),
56}
57
58#[cfg(feature = "evm-prove")]
59enum Halo2Source {
60 Config {
61 shape: StaticVerifierShape,
62 config: Halo2Config,
63 },
64 Pk(Halo2ProvingKey),
65}
66
67pub struct GenericSdkBuilder<E, VB>
72where
73 E: StarkEngine<SC = SC>,
74 VB: VmBuilder<E>,
75 VB::VmConfig: VmExecutionConfig<F>,
76{
77 app_source: Option<AppSource<VB::VmConfig>>,
78 agg_source: Option<AggSource>,
79 #[cfg(feature = "root-prover")]
80 root_source: Option<RootSource>,
81 agg_tree_config: Option<AggregationTreeConfig>,
82 transpiler: Option<Transpiler<F>>,
83 deferral_source: Option<DeferralSource>,
84 #[cfg(feature = "evm-prove")]
85 halo2_source: Option<Halo2Source>,
86 #[cfg(feature = "evm-prove")]
87 halo2_params_reader: Option<CacheHalo2ParamsReader>,
88 _phantom: PhantomData<E>,
89}
90
91impl<E, VB> GenericSdkBuilder<E, VB>
92where
93 E: StarkEngine<SC = SC>,
94 VB: VmBuilder<E>,
95 VB::VmConfig: VmExecutionConfig<F>,
96{
97 pub fn new() -> Self {
99 Self::default()
100 }
101
102 fn set_once<T>(slot: &mut Option<T>, field_name: &str, value: T) {
103 assert!(slot.is_none(), "{field_name} already set");
104 *slot = Some(value);
105 }
106
107 fn init_once_lock<T>(value: Option<T>, field_name: &str) -> OnceLock<T> {
108 let lock = OnceLock::new();
109 if let Some(value) = value {
110 assert!(
111 lock.set(value).is_ok(),
112 "{field_name} should only be initialized once"
113 );
114 }
115 lock
116 }
117
118 fn agg_config_from_pk(agg_pk: &AggProvingKey) -> AggregationConfig {
119 AggregationConfig {
120 params: AggregationSystemParams {
121 leaf: agg_pk.prefix.leaf.params.clone(),
122 internal: agg_pk.internal_recursive.params.clone(),
123 },
124 }
125 }
126
127 #[cfg(feature = "evm-prove")]
128 fn halo2_shape_from_pk(halo2_pk: &Halo2ProvingKey) -> StaticVerifierShape {
129 halo2_pk.verifier.shape
130 }
131
132 #[cfg(feature = "evm-prove")]
133 fn halo2_config_from_pk(halo2_pk: &Halo2ProvingKey) -> Halo2Config {
134 Halo2Config {
135 wrapper_k: Some(halo2_pk.wrapper.pinning.metadata.config_params.k),
136 profiling: halo2_pk.profiling,
137 }
138 }
139
140 fn build_deferral_agg_prover(
141 agg_config: &AggregationConfig,
142 multi_deferral_circuit_prover: MultiDeferralCircuitProver,
143 ) -> Arc<DeferralAggProver> {
144 let agg_prover = AggProver::new(
145 multi_deferral_circuit_prover.def_hook_prover.get_vk(),
146 agg_config.clone(),
147 AggregationTreeConfig::deferral(),
148 Some(
149 multi_deferral_circuit_prover
150 .def_hook_prover
151 .get_cached_commit(),
152 ),
153 );
154 Arc::new(DeferralAggProver {
155 multi_deferral_circuit_prover: Arc::new(multi_deferral_circuit_prover),
156 agg_prover: Arc::new(agg_prover),
157 })
158 }
159
160 fn require_dependency(
161 has_value: bool,
162 value_name: &str,
163 has_dependency: bool,
164 dependency_name: &str,
165 ) -> Result<(), SdkError> {
166 if has_value && !has_dependency {
167 return Err(SdkError::Other(eyre!(
168 "`{value_name}` requires `{dependency_name}` to also be set"
169 )));
170 }
171 Ok(())
172 }
173
174 fn normalize_app_source(
175 app_source: AppSource<VB::VmConfig>,
176 ) -> (AppConfig<VB::VmConfig>, Option<AppProvingKey<VB::VmConfig>>) {
177 match app_source {
178 AppSource::Config(app_config) => (app_config, None),
179 AppSource::Pk(app_pk) => {
180 let app_config = app_pk.app_config();
181 (app_config, Some(app_pk))
182 }
183 }
184 }
185
186 fn normalize_agg_source(agg_source: AggSource) -> (AggregationConfig, Option<AggProvingKey>) {
187 match agg_source {
188 AggSource::Params(agg_params) => (AggregationConfig { params: agg_params }, None),
189 AggSource::Pk(agg_pk) => {
190 let agg_config = Self::agg_config_from_pk(&agg_pk);
191 (agg_config, Some(agg_pk))
192 }
193 }
194 }
195
196 #[cfg(feature = "root-prover")]
197 fn normalize_root_source(root_source: RootSource) -> (SystemParams, Option<RootProvingKey>) {
198 match root_source {
199 RootSource::Params(root_params) => (root_params, None),
200 RootSource::Pk(root_pk) => (root_pk.root_pk.params.clone(), Some(root_pk)),
201 }
202 }
203
204 #[cfg(feature = "evm-prove")]
205 fn normalize_halo2_source(
206 halo2_source: Halo2Source,
207 ) -> (StaticVerifierShape, Halo2Config, Option<Halo2ProvingKey>) {
208 match halo2_source {
209 Halo2Source::Config { shape, config } => (shape, config, None),
210 Halo2Source::Pk(halo2_pk) => (
211 Self::halo2_shape_from_pk(&halo2_pk),
212 Self::halo2_config_from_pk(&halo2_pk),
213 Some(halo2_pk),
214 ),
215 }
216 }
217
218 pub fn app_config(mut self, app_config: AppConfig<VB::VmConfig>) -> Self {
220 Self::set_once(
221 &mut self.app_source,
222 "app_source",
223 AppSource::Config(app_config),
224 );
225 self
226 }
227
228 pub fn app_pk(mut self, app_pk: AppProvingKey<VB::VmConfig>) -> Self {
230 Self::set_once(&mut self.app_source, "app_source", AppSource::Pk(app_pk));
231 self
232 }
233
234 pub fn agg_params(mut self, agg_params: AggregationSystemParams) -> Self {
236 Self::set_once(
237 &mut self.agg_source,
238 "agg_source",
239 AggSource::Params(agg_params),
240 );
241 self
242 }
243
244 pub fn agg_pk(mut self, agg_pk: AggProvingKey) -> Self {
246 Self::set_once(&mut self.agg_source, "agg_source", AggSource::Pk(agg_pk));
247 self
248 }
249
250 #[cfg(feature = "root-prover")]
251 pub fn root_params(mut self, root_params: SystemParams) -> Self {
253 Self::set_once(
254 &mut self.root_source,
255 "root_source",
256 RootSource::Params(root_params),
257 );
258 self
259 }
260
261 #[cfg(feature = "root-prover")]
262 pub fn root_pk(mut self, root_pk: RootProvingKey) -> Self {
264 Self::set_once(
265 &mut self.root_source,
266 "root_source",
267 RootSource::Pk(root_pk),
268 );
269 self
270 }
271
272 pub fn agg_tree_config(mut self, agg_tree_config: AggregationTreeConfig) -> Self {
274 Self::set_once(
275 &mut self.agg_tree_config,
276 "agg_tree_config",
277 agg_tree_config,
278 );
279 self
280 }
281
282 pub fn transpiler(mut self, transpiler: Transpiler<F>) -> Self {
285 Self::set_once(&mut self.transpiler, "transpiler", transpiler);
286 self
287 }
288
289 pub fn build_without_transpiler(self) -> Result<GenericSdk<E, VB>, SdkError>
296 where
297 VB: Default,
298 {
299 let app_source = self
300 .app_source
301 .ok_or_else(|| SdkError::Other(eyre!("`app_config` or `app_pk` must be set")))?;
302 let agg_source = self
303 .agg_source
304 .ok_or_else(|| SdkError::Other(eyre!("`agg_params` or `agg_pk` must be set")))?;
305 #[cfg(feature = "root-prover")]
306 let root_source = self
307 .root_source
308 .unwrap_or_else(|| RootSource::Params(root_params_with_100_bits_security()));
309 #[cfg(feature = "evm-prove")]
310 let halo2_source = self.halo2_source.unwrap_or(Halo2Source::Config {
311 shape: StaticVerifierShape::default(),
312 config: Halo2Config {
313 wrapper_k: None,
314 profiling: false,
315 },
316 });
317
318 #[cfg(feature = "evm-prove")]
319 Self::require_dependency(
320 matches!(halo2_source, Halo2Source::Pk(_)),
321 "halo2_pk",
322 matches!(root_source, RootSource::Pk(_)),
323 "root_pk",
324 )?;
325 #[cfg(feature = "root-prover")]
326 Self::require_dependency(
327 matches!(root_source, RootSource::Pk(_)),
328 "root_pk",
329 matches!(agg_source, AggSource::Pk(_)),
330 "agg_pk",
331 )?;
332 Self::require_dependency(
333 matches!(agg_source, AggSource::Pk(_)),
334 "agg_pk",
335 matches!(app_source, AppSource::Pk(_)),
336 "app_pk",
337 )?;
338
339 let Self {
340 app_source: _,
341 agg_source: _,
342 #[cfg(feature = "root-prover")]
343 root_source: _,
344 agg_tree_config,
345 transpiler,
346 deferral_source,
347 #[cfg(feature = "evm-prove")]
348 halo2_source: _,
349 #[cfg(feature = "evm-prove")]
350 halo2_params_reader,
351 _phantom: _,
352 } = self;
353
354 let (app_config, app_pk_seed) = Self::normalize_app_source(app_source);
355 let (agg_config, agg_pk_seed) = Self::normalize_agg_source(agg_source);
356 #[cfg(feature = "root-prover")]
357 let (root_params, root_pk_seed) = Self::normalize_root_source(root_source);
358
359 let executor = VmExecutor::new(app_config.app_vm_config.clone())
360 .map_err(|e| SdkError::Vm(e.into()))?;
361 let agg_tree_config = agg_tree_config.unwrap_or_default();
362
363 let deferral_setup = match deferral_source {
364 Some(DeferralSource::HookCommits(commits)) => DeferralSetup::Aware(commits),
365 Some(DeferralSource::MultiCircuitProver(multi_deferral_circuit_prover)) => {
366 DeferralSetup::Active(Self::build_deferral_agg_prover(
367 &agg_config,
368 multi_deferral_circuit_prover,
369 ))
370 }
371 Some(DeferralSource::AggProver(deferral_agg_prover)) => {
372 DeferralSetup::Active(Arc::new(deferral_agg_prover))
373 }
374 None => DeferralSetup::Disabled,
375 };
376 let def_hook_cached_commit = deferral_setup.hook_cached_commit();
377 #[cfg(feature = "root-prover")]
378 let def_hook_commit = deferral_setup.hook_commit();
379
380 let app_vm_vk = app_pk_seed
381 .as_ref()
382 .map(|app_pk| Arc::new(app_pk.app_vm_pk.vm_pk.get_vk()));
383 let app_pk = Self::init_once_lock(app_pk_seed, "app_pk");
384
385 let agg_prover_seed = agg_pk_seed.map(|agg_pk| {
386 let app_vm_vk = app_vm_vk.expect("validated `agg_pk` dependency on `app_pk`");
387 Arc::new(AggProver::from_pk(
388 app_vm_vk,
389 agg_pk,
390 agg_tree_config,
391 def_hook_cached_commit,
392 ))
393 });
394
395 #[cfg(feature = "root-prover")]
396 let root_prover_seed = root_pk_seed.map(|root_pk| {
397 let agg_prover = agg_prover_seed
398 .as_ref()
399 .expect("validated `root_pk` dependency on `agg_pk`");
400 let system_config = app_config.app_vm_config.as_ref();
401 let memory_dimensions = system_config.memory_config.memory_dimensions();
402 let num_user_pvs = system_config.num_public_values;
403 let internal_recursive_vk_commit = agg_prover
404 .internal_recursive_prover
405 .get_self_vk_pcs_data()
406 .unwrap()
407 .commitment
408 .into();
409
410 Arc::new(RootProver::from_pk(
411 agg_prover.internal_recursive_prover.get_vk(),
412 internal_recursive_vk_commit,
413 root_pk.root_pk,
414 memory_dimensions,
415 num_user_pvs,
416 def_hook_commit,
417 Some(root_pk.trace_heights),
418 ))
419 });
420
421 #[cfg(feature = "evm-prove")]
422 let halo2_params_reader =
423 halo2_params_reader.unwrap_or_else(CacheHalo2ParamsReader::new_with_default_params_dir);
424 #[cfg(feature = "evm-prove")]
425 let (halo2_shape, halo2_config, halo2_pk_seed) = Self::normalize_halo2_source(halo2_source);
426 #[cfg(feature = "evm-prove")]
427 let halo2_prover_seed =
428 halo2_pk_seed.map(|halo2_pk| Halo2Prover::new(&halo2_params_reader, halo2_pk));
429
430 Ok(GenericSdk {
431 app_config,
432 agg_config,
433 agg_tree_config,
434 #[cfg(feature = "root-prover")]
435 root_params,
436 #[cfg(feature = "evm-prove")]
437 halo2_shape,
438 #[cfg(feature = "evm-prove")]
439 halo2_config,
440 app_vm_builder: Default::default(),
441 transpiler,
442 executor,
443 app_pk,
444 agg_prover: Self::init_once_lock(agg_prover_seed, "agg_prover"),
445 #[cfg(feature = "root-prover")]
446 root_prover: Self::init_once_lock(root_prover_seed, "root_prover"),
447 deferral_setup,
448 #[cfg(feature = "evm-prove")]
449 halo2_params_reader,
450 #[cfg(feature = "evm-prove")]
451 halo2_prover: Self::init_once_lock(halo2_prover_seed, "halo2_prover"),
452 _phantom: PhantomData,
453 })
454 }
455
456 pub fn build(mut self) -> Result<GenericSdk<E, VB>, SdkError>
459 where
460 VB: Default,
461 VB::VmConfig: TranspilerConfig<F>,
462 {
463 if self.transpiler.is_none() {
464 self.transpiler = self.app_source.as_ref().map(|app_source| match app_source {
465 AppSource::Config(app_config) => app_config.app_vm_config.transpiler(),
466 AppSource::Pk(app_pk) => app_pk.app_vm_pk.vm_config.transpiler(),
467 });
468 }
469 self.build_without_transpiler()
470 }
471
472 pub fn deferral_hook_commits(mut self, commits: DeferralHookCommits) -> Self {
478 Self::set_once(
479 &mut self.deferral_source,
480 "deferral_source",
481 DeferralSource::HookCommits(commits),
482 );
483 self
484 }
485
486 pub fn multi_deferral_circuit_prover(
494 mut self,
495 multi_deferral_circuit_prover: MultiDeferralCircuitProver,
496 ) -> Self {
497 Self::set_once(
498 &mut self.deferral_source,
499 "deferral_source",
500 DeferralSource::MultiCircuitProver(multi_deferral_circuit_prover),
501 );
502 self
503 }
504
505 pub fn deferral_agg_prover(mut self, deferral_agg_prover: DeferralAggProver) -> Self {
511 Self::set_once(
512 &mut self.deferral_source,
513 "deferral_source",
514 DeferralSource::AggProver(deferral_agg_prover),
515 );
516 self
517 }
518
519 #[cfg(feature = "evm-prove")]
520 pub fn halo2_config(mut self, shape: StaticVerifierShape, config: Halo2Config) -> Self {
522 Self::set_once(
523 &mut self.halo2_source,
524 "halo2_source",
525 Halo2Source::Config { shape, config },
526 );
527 self
528 }
529
530 #[cfg(feature = "evm-prove")]
531 pub fn halo2_pk(mut self, halo2_pk: Halo2ProvingKey) -> Self {
533 Self::set_once(
534 &mut self.halo2_source,
535 "halo2_source",
536 Halo2Source::Pk(halo2_pk),
537 );
538 self
539 }
540
541 #[cfg(feature = "evm-prove")]
542 pub fn halo2_params_dir(mut self, params_dir: impl AsRef<Path>) -> Self {
544 Self::set_once(
545 &mut self.halo2_params_reader,
546 "halo2_params_reader",
547 CacheHalo2ParamsReader::new(params_dir),
548 );
549 self
550 }
551}
552
553impl<E, VB> Default for GenericSdkBuilder<E, VB>
554where
555 E: StarkEngine<SC = SC>,
556 VB: VmBuilder<E>,
557 VB::VmConfig: VmExecutionConfig<F>,
558{
559 fn default() -> Self {
560 Self {
561 app_source: None,
562 agg_source: None,
563 #[cfg(feature = "root-prover")]
564 root_source: None,
565 agg_tree_config: None,
566 transpiler: None,
567 deferral_source: None,
568 #[cfg(feature = "evm-prove")]
569 halo2_source: None,
570 #[cfg(feature = "evm-prove")]
571 halo2_params_reader: None,
572 _phantom: PhantomData,
573 }
574 }
575}
576
577impl<E, VB> GenericSdk<E, VB>
578where
579 E: StarkEngine<SC = SC>,
580 VB: VmBuilder<E>,
581 VB::VmConfig: VmExecutionConfig<F>,
582{
583 pub fn builder() -> GenericSdkBuilder<E, VB> {
585 GenericSdkBuilder::new()
586 }
587
588 pub fn from_cached_proving_key(
594 cached_pk: SdkCachedProvingKey<VB::VmConfig>,
595 ) -> Result<Self, SdkError>
596 where
597 VB: Default,
598 VB::VmConfig: TranspilerConfig<F>,
599 {
600 let SdkCachedProvingKey {
601 app_pk,
602 agg_pk,
603 deferral_pk,
604 deferral_agg_pk,
605 #[cfg(feature = "root-prover")]
606 root_pk,
607 } = cached_pk;
608
609 if deferral_pk.is_some() || deferral_agg_pk.is_some() {
610 return Err(SdkError::Other(eyre!(
611 "generic cached SDK cannot include deferral keys"
612 )));
613 }
614
615 let builder = Self::builder().app_pk(app_pk).agg_pk(agg_pk);
616 #[cfg(feature = "root-prover")]
617 let builder = if let Some(root_pk) = root_pk {
618 builder.root_pk(root_pk)
619 } else {
620 builder
621 };
622
623 builder.build()
624 }
625}
626
627impl<E, VB> GenericSdk<E, VB>
628where
629 E: StarkEngine<SC = SC>,
630 VB: VmBuilder<E, VmConfig = SdkVmConfig>,
631{
632 pub fn from_deferral_cached_proving_key(
638 cached_pk: SdkCachedProvingKey<SdkVmConfig>,
639 ) -> Result<Self, SdkError>
640 where
641 VB: Default,
642 {
643 let SdkCachedProvingKey {
644 app_pk,
645 agg_pk,
646 deferral_pk,
647 deferral_agg_pk,
648 #[cfg(feature = "root-prover")]
649 root_pk,
650 } = cached_pk;
651
652 let deferral_config = app_pk.app_vm_pk.vm_config.deferral.clone();
653
654 let builder = Self::builder().app_pk(app_pk).agg_pk(agg_pk);
655 #[cfg(feature = "root-prover")]
656 let builder = if let Some(root_pk) = root_pk {
657 builder.root_pk(root_pk)
658 } else {
659 builder
660 };
661
662 let builder = match (deferral_pk, deferral_agg_pk) {
663 (Some(deferral_pk), Some(deferral_agg_pk)) => {
664 let deferral_config = deferral_config.ok_or_else(|| {
665 SdkError::Other(eyre!("cached deferral keys require a deferral config"))
666 })?;
667 builder.deferral_agg_prover(DeferralAggProver::from_supported_deferral_pks(
668 &deferral_config,
669 deferral_pk,
670 deferral_agg_pk,
671 )?)
672 }
673 (None, None) => builder,
674 _ => {
675 return Err(SdkError::Other(eyre!(
676 "cached deferral keys are incomplete"
677 )));
678 }
679 };
680
681 builder.build()
682 }
683}