openvm_sdk/prover/
root.rs

1use 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        // Internal sanity (SDK tests only): a successful tracegen must land at
124        // exactly the root verifier's fixed, expected trace heights.
125        #[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}