openvm_continuations/prover/deferral/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::VkCommit;
16use tracing::instrument;
17
18use crate::{
19    circuit::{
20        deferral::inner::{DeferralInnerCircuit, DeferralInnerTraceGen},
21        Circuit,
22    },
23    prover::trace_heights_tracing_info,
24    SC,
25};
26
27mod trace;
28
29pub enum DeferralChildVkKind {
30    /// Child proofs are deferral verify proofs (consume DeferralCircuitPvs at air 0).
31    DeferralCircuit,
32    /// Child proofs are deferral aggregation inner proofs (consume DeferralAggregationPvs at air
33    /// 1).
34    DeferralAggregation,
35    /// Same as DeferralAggregation but uses this prover's own vk as child vk.
36    RecursiveSelf,
37}
38
39pub struct DeferralInnerProver<
40    PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
41    S: AggregationSubCircuit,
42    T,
43> {
44    pk: Arc<MultiStarkProvingKey<SC>>,
45    d_pk: DeviceMultiStarkProvingKey<PB>,
46    vk: Arc<MultiStarkVerifyingKey<SC>>,
47
48    agg_node_tracegen: T,
49
50    child_vk: Arc<MultiStarkVerifyingKey<SC>>,
51    child_vk_pcs_data: CommittedTraceData<PB>,
52    circuit: Arc<DeferralInnerCircuit<S>>,
53
54    self_vk_pcs_data: Option<CommittedTraceData<PB>>,
55}
56
57impl<PB, S, T> DeferralInnerProver<PB, S, T>
58where
59    PB: ProverBackend<Val = F, Challenge = EF, Commitment = Digest>,
60    S: AggregationSubCircuit,
61    PB::Matrix: Clone,
62{
63    #[instrument(name = "total_proof", skip_all)]
64    pub fn agg_prove<E: StarkEngine<SC = SC, PB = PB>>(
65        &self,
66        proofs: &[Proof<SC>],
67        child_vk_kind: DeferralChildVkKind,
68        child_merkle_depth: Option<usize>,
69    ) -> Result<Proof<SC>>
70    where
71        S: VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
72        T: DeferralInnerTraceGen<PB, EngineDeviceCtx<E>>,
73    {
74        let engine = E::new(self.pk.params.clone());
75        let ctx = self.generate_proving_ctx(
76            proofs,
77            child_vk_kind,
78            child_merkle_depth,
79            engine.device().device_ctx(),
80        );
81        if tracing::enabled!(tracing::Level::DEBUG) {
82            trace_heights_tracing_info::<_, SC>(&ctx.per_trace, &self.circuit.airs());
83        }
84        #[cfg(debug_assertions)]
85        if crate::prover::debug_checks_enabled() {
86            crate::prover::debug_constraints(&self.circuit, &ctx, &engine);
87        }
88        let proof = engine.prove(&self.d_pk, ctx)?;
89        #[cfg(debug_assertions)]
90        if crate::prover::debug_checks_enabled() {
91            engine.verify(&self.vk, &proof)?;
92        }
93        Ok(proof)
94    }
95    pub fn new<E: StarkEngine<SC = SC, PB = PB>>(
96        child_vk: Arc<MultiStarkVerifyingKey<SC>>,
97        system_params: SystemParams,
98        is_self_recursive: bool,
99    ) -> Self
100    where
101        S: VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
102        T: DeferralInnerTraceGen<PB, EngineDeviceCtx<E>>,
103    {
104        let verifier_circuit = S::new(
105            child_vk.clone(),
106            VerifierConfig {
107                continuations_enabled: true,
108                ..Default::default()
109            },
110        );
111        let engine = E::new(system_params);
112        let child_vk_pcs_data = verifier_circuit.commit_child_vk(&engine, &child_vk);
113        let circuit = Arc::new(DeferralInnerCircuit::new(Arc::new(verifier_circuit)));
114        let (pk, vk) = engine.keygen(&circuit.airs());
115        let d_pk = engine.device().transport_pk_to_device(&pk);
116        let self_vk_pcs_data = if is_self_recursive {
117            Some(circuit.verifier_circuit.commit_child_vk(&engine, &vk))
118        } else {
119            None
120        };
121        Self {
122            pk: Arc::new(pk),
123            d_pk,
124            vk: Arc::new(vk),
125            agg_node_tracegen: DeferralInnerTraceGen::new(),
126            child_vk,
127            child_vk_pcs_data,
128            circuit,
129            self_vk_pcs_data,
130        }
131    }
132
133    pub fn from_pk<E: StarkEngine<SC = SC, PB = PB>>(
134        child_vk: Arc<MultiStarkVerifyingKey<SC>>,
135        pk: Arc<MultiStarkProvingKey<SC>>,
136        is_self_recursive: bool,
137    ) -> Self
138    where
139        S: VerifierTraceGen<PB, SC, EngineDeviceCtx<E>>,
140        T: DeferralInnerTraceGen<PB, EngineDeviceCtx<E>>,
141    {
142        let verifier_circuit = S::new(
143            child_vk.clone(),
144            VerifierConfig {
145                continuations_enabled: true,
146                ..Default::default()
147            },
148        );
149        let engine = E::new(pk.params.clone());
150        let child_vk_pcs_data = verifier_circuit.commit_child_vk(&engine, &child_vk);
151        let circuit = Arc::new(DeferralInnerCircuit::new(Arc::new(verifier_circuit)));
152        let vk = Arc::new(pk.get_vk());
153        let d_pk = engine.device().transport_pk_to_device(&pk);
154        let self_vk_pcs_data = if is_self_recursive {
155            Some(circuit.verifier_circuit.commit_child_vk(&engine, &vk))
156        } else {
157            None
158        };
159        Self {
160            pk,
161            d_pk,
162            vk,
163            agg_node_tracegen: DeferralInnerTraceGen::new(),
164            child_vk,
165            child_vk_pcs_data,
166            circuit,
167            self_vk_pcs_data,
168        }
169    }
170
171    pub fn get_circuit(&self) -> Arc<DeferralInnerCircuit<S>> {
172        self.circuit.clone()
173    }
174
175    pub fn get_pk(&self) -> Arc<MultiStarkProvingKey<SC>> {
176        self.pk.clone()
177    }
178
179    pub fn get_vk(&self) -> Arc<MultiStarkVerifyingKey<SC>> {
180        self.vk.clone()
181    }
182
183    pub fn get_vk_commit(&self, is_self_recursive: bool) -> VkCommit<PB::Val> {
184        if is_self_recursive {
185            VkCommit {
186                cached_commit: self.self_vk_pcs_data.as_ref().unwrap().commitment,
187                vk_pre_hash: self.vk.pre_hash,
188            }
189        } else {
190            VkCommit {
191                cached_commit: self.child_vk_pcs_data.commitment,
192                vk_pre_hash: self.child_vk.pre_hash,
193            }
194        }
195    }
196}