openvm_continuations/prover/root/
mod.rs

1use std::sync::Arc;
2
3use eyre::Result;
4use openvm_circuit::system::memory::dimensions::MemoryDimensions;
5use openvm_recursion_circuit::{
6    batch_constraint::expr_eval::CachedTraceRecord,
7    system::{AggregationSubCircuit, VerifierConfig, VerifierTraceGen},
8};
9use openvm_stark_backend::{
10    keygen::types::{MultiStarkProvingKey, MultiStarkVerifyingKey},
11    proof::Proof,
12    prover::{DeviceDataTransporter, ProverBackend, ProvingContext},
13    EngineDeviceCtx, StarkEngine, SystemParams,
14};
15use openvm_stark_sdk::config::baby_bear_poseidon2::{EF, F};
16use p3_bn254::Bn254;
17use p3_field::{Field, PrimeField32};
18use tracing::instrument;
19
20use crate::{
21    circuit::{
22        root::{RootCircuit, RootTraceGen},
23        Circuit,
24    },
25    prover::trace_heights_tracing_info,
26    CommitBytes, RootSC, VkCommitBytes, SC,
27};
28
29mod trace;
30
31/// RootProver does not store a device context because it uses late binding:
32/// the engine (and its device context) is created at prove time, not at construction.
33/// This allows the prover to be device-agnostic until proving begins.
34pub struct RootProver<S: AggregationSubCircuit, T> {
35    pk: Arc<MultiStarkProvingKey<RootSC>>,
36    vk: Arc<MultiStarkVerifyingKey<RootSC>>,
37
38    agg_node_tracegen: T,
39
40    child_vk: Arc<MultiStarkVerifyingKey<SC>>,
41    cached_trace_record: CachedTraceRecord,
42    circuit: Arc<RootCircuit<S>>,
43    trace_heights: Option<Vec<usize>>,
44}
45
46impl<S: AggregationSubCircuit, T> RootProver<S, T> {
47    pub fn create_engine<E>(&self) -> E
48    where
49        E: StarkEngine<SC = RootSC>,
50    {
51        E::new(self.pk.params.clone())
52    }
53
54    #[instrument(name = "total_proof", skip_all)]
55    pub fn root_prove_from_ctx<E>(
56        &self,
57        ctx: ProvingContext<E::PB>,
58        engine: &E,
59    ) -> Result<Proof<RootSC>>
60    where
61        E: StarkEngine<SC = RootSC>,
62        E::PB: ProverBackend<Val = F, Challenge = EF, Commitment = [Bn254; 1]>,
63        <E::PB as ProverBackend>::Matrix: Clone,
64        S: VerifierTraceGen<E::PB, RootSC, EngineDeviceCtx<E>>,
65        T: RootTraceGen<E::PB, EngineDeviceCtx<E>>,
66    {
67        if tracing::enabled!(tracing::Level::DEBUG) {
68            trace_heights_tracing_info::<_, RootSC>(&ctx.per_trace, &self.circuit.airs());
69        }
70        #[cfg(debug_assertions)]
71        if crate::prover::debug_checks_enabled() {
72            crate::prover::debug_constraints(&self.circuit, &ctx, engine);
73        }
74        let d_pk = engine.device().transport_pk_to_device(self.pk.as_ref());
75        let proof = engine.prove(&d_pk, ctx)?;
76        #[cfg(debug_assertions)]
77        if crate::prover::debug_checks_enabled() {
78            engine.verify(&self.vk, &proof)?;
79        }
80        Ok(proof)
81    }
82}
83
84impl<S: AggregationSubCircuit, T> RootProver<S, T> {
85    pub fn new<E>(
86        child_vk: Arc<MultiStarkVerifyingKey<SC>>,
87        internal_recursive_cached_commit: CommitBytes,
88        system_params: SystemParams,
89        memory_dimensions: MemoryDimensions,
90        num_user_pvs: usize,
91        def_hook_commit: Option<CommitBytes>,
92        trace_heights: Option<Vec<usize>>,
93    ) -> Self
94    where
95        E: StarkEngine<SC = RootSC>,
96        E::PB: ProverBackend<Val = F, Challenge = EF, Commitment = [Bn254; 1]>,
97        S: VerifierTraceGen<E::PB, RootSC, EngineDeviceCtx<E>>,
98        T: RootTraceGen<E::PB, EngineDeviceCtx<E>>,
99        E::PD: DeviceDataTransporter<RootSC, E::PB> + Clone,
100        <E::PB as ProverBackend>::Val: Field + PrimeField32,
101        <E::PB as ProverBackend>::Matrix: Clone,
102    {
103        let verifier_circuit = S::new(
104            child_vk.clone(),
105            VerifierConfig {
106                continuations_enabled: true,
107                has_cached: false,
108                ..Default::default()
109            },
110        );
111        let cached_trace_record = verifier_circuit.cached_trace_record(&child_vk);
112        let engine = E::new(system_params);
113        let internal_recursive_vk_commit = VkCommitBytes {
114            cached_commit: internal_recursive_cached_commit,
115            vk_pre_hash: child_vk.pre_hash.into(),
116        };
117        let circuit = Arc::new(RootCircuit::new(
118            Arc::new(verifier_circuit),
119            internal_recursive_vk_commit,
120            def_hook_commit,
121            memory_dimensions,
122            num_user_pvs,
123        ));
124        let (pk, vk) = engine.keygen(&circuit.airs());
125        Self {
126            pk: Arc::new(pk),
127            vk: Arc::new(vk),
128            agg_node_tracegen: T::new(def_hook_commit.is_some()),
129            child_vk,
130            cached_trace_record,
131            circuit,
132            trace_heights,
133        }
134    }
135
136    pub fn from_pk<E>(
137        child_vk: Arc<MultiStarkVerifyingKey<SC>>,
138        internal_recursive_cached_commit: CommitBytes,
139        pk: Arc<MultiStarkProvingKey<RootSC>>,
140        memory_dimensions: MemoryDimensions,
141        num_user_pvs: usize,
142        def_hook_commit: Option<CommitBytes>,
143        trace_heights: Option<Vec<usize>>,
144    ) -> Self
145    where
146        E: StarkEngine<SC = RootSC>,
147        E::PB: ProverBackend<Val = F, Challenge = EF, Commitment = [Bn254; 1]>,
148        S: VerifierTraceGen<E::PB, RootSC, EngineDeviceCtx<E>>,
149        T: RootTraceGen<E::PB, EngineDeviceCtx<E>>,
150        <E::PB as ProverBackend>::Val: Field + PrimeField32,
151        <E::PB as ProverBackend>::Matrix: Clone,
152    {
153        let verifier_circuit = S::new(
154            child_vk.clone(),
155            VerifierConfig {
156                continuations_enabled: true,
157                has_cached: false,
158                ..Default::default()
159            },
160        );
161        let cached_trace_record = verifier_circuit.cached_trace_record(&child_vk);
162        let internal_recursive_vk_commit = VkCommitBytes {
163            cached_commit: internal_recursive_cached_commit,
164            vk_pre_hash: child_vk.pre_hash.into(),
165        };
166        let circuit = Arc::new(RootCircuit::new(
167            Arc::new(verifier_circuit),
168            internal_recursive_vk_commit,
169            def_hook_commit,
170            memory_dimensions,
171            num_user_pvs,
172        ));
173        let vk = Arc::new(pk.get_vk());
174        Self {
175            pk,
176            vk,
177            agg_node_tracegen: T::new(def_hook_commit.is_some()),
178            child_vk,
179            cached_trace_record,
180            circuit,
181            trace_heights,
182        }
183    }
184
185    pub fn get_circuit(&self) -> Arc<RootCircuit<S>> {
186        self.circuit.clone()
187    }
188
189    pub fn get_pk(&self) -> Arc<MultiStarkProvingKey<RootSC>> {
190        self.pk.clone()
191    }
192
193    pub fn get_vk(&self) -> Arc<MultiStarkVerifyingKey<RootSC>> {
194        self.vk.clone()
195    }
196
197    pub fn get_trace_heights(&self) -> Option<Vec<usize>> {
198        self.trace_heights.clone()
199    }
200}