openvm_continuations/prover/inner/
mod.rs

1use std::sync::Arc;
2
3use eyre::Result;
4use openvm_recursion_circuit::system::{AggregationSubCircuit, VerifierConfig, VerifierTraceGen};
5use openvm_stark_backend::{
6    keygen::types::{MultiStarkProvingKey, MultiStarkVerifyingKey},
7    proof::Proof,
8    prover::{
9        CommittedTraceData, DeviceDataTransporter, DeviceMultiStarkProvingKey, ProverBackend,
10        ProverDevice,
11    },
12    EngineDeviceCtx, StarkEngine, SystemParams,
13};
14use openvm_stark_sdk::config::baby_bear_poseidon2::{Digest, EF, F};
15use openvm_verify_stark_host::pvs::{DeferralPvs, VkCommit};
16use tracing::instrument;
17
18use crate::{
19    circuit::{
20        inner::{InnerCircuit, InnerTraceGen, ProofsType},
21        Circuit,
22    },
23    prover::trace_heights_tracing_info,
24    SC,
25};
26
27mod trace;
28
29/// Generates an aggregation proof for inner layers (leaf and internal).
30pub struct InnerAggregationProver<
31    PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
32    S: AggregationSubCircuit,
33    T,
34> {
35    pk: Arc<MultiStarkProvingKey<SC>>,
36    d_pk: DeviceMultiStarkProvingKey<PB>,
37    vk: Arc<MultiStarkVerifyingKey<SC>>,
38
39    agg_node_tracegen: T,
40
41    // TODO: tracegen currently requires storing these, we should revisit this
42    child_vk: Arc<MultiStarkVerifyingKey<SC>>,
43    child_vk_pcs_data: CommittedTraceData<PB>,
44    circuit: Arc<InnerCircuit<S>>,
45
46    self_vk_pcs_data: Option<CommittedTraceData<PB>>,
47}
48
49/// Struct to determine if InnerAggregationProver is proving a special case,
50/// i.e. if the child_vk is the app_vk or if it should use its own vk as child.
51#[derive(Clone, Copy)]
52pub enum ChildVkKind {
53    Standard,
54    App,
55    RecursiveSelf,
56}
57
58impl<PB, S, T> InnerAggregationProver<PB, S, T>
59where
60    PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
61    S: AggregationSubCircuit,
62    PB::Matrix: Clone,
63{
64    #[instrument(name = "total_proof", skip_all)]
65    pub fn agg_prove<E: StarkEngine<SC = SC, PB = PB>>(
66        &self,
67        proofs: &[Proof<SC>],
68        child_vk_kind: ChildVkKind,
69        proofs_type: ProofsType,
70        absent_trace_pvs: Option<(DeferralPvs<F>, bool)>,
71    ) -> Result<Proof<SC>>
72    where
73        S: VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
74        T: InnerTraceGen<PB, EngineDeviceCtx<E>>,
75    {
76        let engine = E::new(self.pk.params.clone());
77        let ctx = self.generate_proving_ctx(
78            proofs,
79            child_vk_kind,
80            proofs_type,
81            absent_trace_pvs,
82            engine.device().device_ctx(),
83        );
84        if tracing::enabled!(tracing::Level::DEBUG) {
85            trace_heights_tracing_info::<_, SC>(&ctx.per_trace, &self.circuit.airs());
86        }
87        #[cfg(debug_assertions)]
88        if crate::prover::debug_checks_enabled() {
89            crate::prover::debug_constraints(&self.circuit, &ctx, &engine);
90        }
91        let proof = engine.prove(&self.d_pk, ctx)?;
92        #[cfg(debug_assertions)]
93        if crate::prover::debug_checks_enabled() {
94            engine.verify(&self.vk, &proof)?;
95        }
96        Ok(proof)
97    }
98
99    pub fn agg_prove_no_def<E: StarkEngine<SC = SC, PB = PB>>(
100        &self,
101        proofs: &[Proof<SC>],
102        child_vk_kind: ChildVkKind,
103    ) -> Result<Proof<SC>>
104    where
105        S: VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
106        T: InnerTraceGen<PB, EngineDeviceCtx<E>>,
107    {
108        self.agg_prove::<E>(proofs, child_vk_kind, ProofsType::Vm, None)
109    }
110    pub fn new<E: StarkEngine<SC = SC, PB = PB>>(
111        child_vk: Arc<MultiStarkVerifyingKey<SC>>,
112        system_params: SystemParams,
113        is_self_recursive: bool,
114        def_hook_cached_commit: Option<Digest>,
115    ) -> Self
116    where
117        S: VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
118        T: InnerTraceGen<PB, EngineDeviceCtx<E>>,
119    {
120        let verifier_circuit = S::new(
121            child_vk.clone(),
122            VerifierConfig {
123                continuations_enabled: true,
124                ..Default::default()
125            },
126        );
127        let engine = E::new(system_params);
128        let child_vk_pcs_data = verifier_circuit.commit_child_vk(&engine, &child_vk);
129        let circuit = Arc::new(InnerCircuit::new(
130            Arc::new(verifier_circuit),
131            def_hook_cached_commit.map(|d| d.into()),
132        ));
133        let (pk, vk) = engine.keygen(&circuit.airs());
134        let d_pk = engine.device().transport_pk_to_device(&pk);
135        let self_vk_pcs_data = if is_self_recursive {
136            Some(circuit.verifier_circuit.commit_child_vk(&engine, &vk))
137        } else {
138            None
139        };
140        let agg_node_tracegen = InnerTraceGen::new(def_hook_cached_commit.is_some());
141        Self {
142            pk: Arc::new(pk),
143            d_pk,
144            vk: Arc::new(vk),
145            agg_node_tracegen,
146            child_vk,
147            child_vk_pcs_data,
148            circuit,
149            self_vk_pcs_data,
150        }
151    }
152
153    pub fn from_pk<E: StarkEngine<SC = SC, PB = PB>>(
154        child_vk: Arc<MultiStarkVerifyingKey<SC>>,
155        pk: Arc<MultiStarkProvingKey<SC>>,
156        is_self_recursive: bool,
157        def_hook_cached_commit: Option<Digest>,
158    ) -> Self
159    where
160        S: VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
161        T: InnerTraceGen<PB, EngineDeviceCtx<E>>,
162    {
163        let verifier_circuit = S::new(
164            child_vk.clone(),
165            VerifierConfig {
166                continuations_enabled: true,
167                ..Default::default()
168            },
169        );
170        let engine = E::new(pk.params.clone());
171        let child_vk_pcs_data = verifier_circuit.commit_child_vk(&engine, &child_vk);
172        let circuit = Arc::new(InnerCircuit::new(
173            Arc::new(verifier_circuit),
174            def_hook_cached_commit.map(|d| d.into()),
175        ));
176        let vk = Arc::new(pk.get_vk());
177        let d_pk = engine.device().transport_pk_to_device(&pk);
178        let self_vk_pcs_data = if is_self_recursive {
179            Some(circuit.verifier_circuit.commit_child_vk(&engine, &vk))
180        } else {
181            None
182        };
183        let agg_node_tracegen = InnerTraceGen::new(def_hook_cached_commit.is_some());
184        Self {
185            pk,
186            d_pk,
187            vk,
188            agg_node_tracegen,
189            child_vk,
190            child_vk_pcs_data,
191            circuit,
192            self_vk_pcs_data,
193        }
194    }
195
196    pub fn get_circuit(&self) -> Arc<InnerCircuit<S>> {
197        self.circuit.clone()
198    }
199
200    pub fn get_pk(&self) -> Arc<MultiStarkProvingKey<SC>> {
201        self.pk.clone()
202    }
203
204    pub fn get_vk(&self) -> Arc<MultiStarkVerifyingKey<SC>> {
205        self.vk.clone()
206    }
207
208    pub fn deferral_enabled(&self) -> bool {
209        self.circuit.def_hook_cached_commit.is_some()
210    }
211
212    pub fn get_vk_commit(&self, is_self_recursive: bool) -> VkCommit<PB::Val> {
213        if is_self_recursive {
214            VkCommit {
215                cached_commit: self.self_vk_pcs_data.as_ref().unwrap().commitment,
216                vk_pre_hash: self.vk.pre_hash,
217            }
218        } else {
219            VkCommit {
220                cached_commit: self.child_vk_pcs_data.commitment,
221                vk_pre_hash: self.child_vk.pre_hash,
222            }
223        }
224    }
225
226    pub fn get_self_vk_pcs_data(&self) -> Option<CommittedTraceData<PB>>
227    where
228        CommittedTraceData<PB>: Clone,
229    {
230        self.self_vk_pcs_data.clone()
231    }
232}