openvm_sdk/prover/
root.rs1use std::sync::Arc;
2
3use eyre::Result;
4use openvm_circuit::system::memory::dimensions::MemoryDimensions;
5use openvm_continuations::{prover::engine_device_ctx, CommitBytes, RootSC, SC};
6use openvm_stark_backend::{
7 keygen::types::{MultiStarkProvingKey, MultiStarkVerifyingKey},
8 proof::Proof,
9 prover::ProvingContext,
10 StarkEngine, SystemParams,
11};
12use openvm_stark_sdk::config::baby_bear_poseidon2::Digest;
13use openvm_verify_stark_host::VmStarkProof;
14use tracing::info_span;
15
16cfg_if::cfg_if! {
17 if #[cfg(feature = "cuda")] {
18 use openvm_continuations::prover::RootGpuProver as RootInnerProver;
19 type E = openvm_cuda_backend::BabyBearBn254Poseidon2GpuEngine;
20 } else {
21 use openvm_continuations::prover::RootCpuProver as RootInnerProver;
22 type E = openvm_stark_sdk::config::baby_bear_bn254_poseidon2::BabyBearBn254Poseidon2CpuEngine;
23 }
24}
25
26pub struct RootProver(pub RootInnerProver);
27
28impl RootProver {
29 pub fn new(
30 internal_recursive_vk: Arc<MultiStarkVerifyingKey<SC>>,
31 internal_recursive_vk_commit: CommitBytes,
32 system_params: SystemParams,
33 memory_dimensions: MemoryDimensions,
34 num_user_pvs: usize,
35 def_hook_commit: Option<Digest>,
36 trace_heights: Option<Vec<usize>>,
37 ) -> Self {
38 let inner = RootInnerProver::new::<E>(
39 internal_recursive_vk,
40 internal_recursive_vk_commit,
41 system_params,
42 memory_dimensions,
43 num_user_pvs,
44 def_hook_commit.map(Into::into),
45 trace_heights,
46 );
47 Self(inner)
48 }
49
50 pub fn from_pk(
51 internal_recursive_vk: Arc<MultiStarkVerifyingKey<SC>>,
52 internal_recursive_vk_commit: CommitBytes,
53 pk: Arc<MultiStarkProvingKey<RootSC>>,
54 memory_dimensions: MemoryDimensions,
55 num_user_pvs: usize,
56 def_hook_commit: Option<Digest>,
57 trace_heights: Option<Vec<usize>>,
58 ) -> Self {
59 let inner = RootInnerProver::from_pk::<E>(
60 internal_recursive_vk,
61 internal_recursive_vk_commit,
62 pk,
63 memory_dimensions,
64 num_user_pvs,
65 def_hook_commit.map(Into::into),
66 trace_heights,
67 );
68 Self(inner)
69 }
70
71 pub fn create_engine(&self) -> E {
72 self.0.create_engine::<E>()
73 }
74
75 pub fn generate_proving_ctx(
76 &self,
77 input: VmStarkProof,
78 engine: &E,
79 ) -> Option<ProvingContext<<E as StarkEngine>::PB>> {
80 let ctx = info_span!("tracegen_attempt", group = format!("root")).in_scope(|| {
81 self.0.generate_proving_ctx(
82 input.inner,
83 &input.user_pvs_proof,
84 input.deferral_merkle_proofs.as_ref(),
85 engine_device_ctx(engine),
86 )
87 });
88 ctx
89 }
90
91 pub fn prove_from_ctx(
92 &self,
93 ctx: ProvingContext<<E as StarkEngine>::PB>,
94 engine: &E,
95 ) -> Result<Proof<RootSC>> {
96 let proof = info_span!("agg_layer", group = format!("root")).in_scope(|| {
97 info_span!("root").in_scope(|| self.0.root_prove_from_ctx::<E>(ctx, engine))
98 })?;
99 Ok(proof)
100 }
101
102 pub fn prove(
103 &self,
104 mut stark_proof: VmStarkProof,
105 engine: &E,
106 max_retries: usize,
107 mut wrap: impl FnMut(VmStarkProof) -> Result<VmStarkProof>,
108 ) -> Result<Proof<RootSC>> {
109 let mut attempt = 0usize;
110 let ctx = loop {
111 if let Some(ctx) = self.generate_proving_ctx(stark_proof.clone(), engine) {
112 break ctx;
113 }
114 if attempt >= max_retries {
115 return Err(eyre::eyre!(
116 "root tracegen returned None after {max_retries} retries"
117 ));
118 }
119 stark_proof = wrap(stark_proof)?;
120 attempt += 1;
121 };
122
123 #[cfg(test)]
126 for ((air_idx, air_ctx), expected_height) in ctx
127 .per_trace
128 .iter()
129 .zip(self.0.get_trace_heights().unwrap())
130 {
131 assert_eq!(
132 air_ctx.height(),
133 expected_height,
134 "height mismatch at {air_idx}"
135 );
136 }
137
138 self.prove_from_ctx(ctx, engine)
139 }
140}