openvm_continuations/prover/root/
mod.rs1use 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
31pub 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}