openvm_sdk/prover/deferral/
multi_circuit.rs

1use std::sync::Arc;
2
3use eyre::Result;
4use itertools::Itertools;
5use openvm_continuations::{
6    prover::{DeferralChildVkKind, DeferralCircuitProver},
7    SC,
8};
9use openvm_deferral_circuit::{DeferralExtension, DeferralFn};
10use openvm_recursion_circuit::utils::poseidon2_hash_slice;
11use openvm_sdk_config::deferral::{DeferralCircuitConfig, DeferralConfig, SupportedDeferral};
12use openvm_stark_backend::{
13    keygen::types::MultiStarkProvingKey,
14    p3_field::{PrimeCharacteristicRing, PrimeField32},
15    proof::Proof,
16    SystemParams,
17};
18use openvm_stark_sdk::config::baby_bear_poseidon2::{
19    poseidon2_compress_with_capacity, DIGEST_SIZE, F,
20};
21use openvm_verify_stark_host::pvs::DeferralPvs;
22use tracing::info_span;
23
24use crate::{
25    config::AggregationConfig,
26    keygen::{AggPrefixProvingKey, DeferralCircuitProvingKey, DeferralProvingKey},
27    prover::SingleDeferralCircuitProver,
28    DeferralInput,
29};
30
31cfg_if::cfg_if! {
32    if #[cfg(feature = "cuda")] {
33        use openvm_continuations::prover::DeferralInnerGpuProver as DeferralInnerProver;
34        use openvm_continuations::prover::DeferralHookGpuProver as DeferralHookProver;
35        type E = openvm_cuda_backend::BabyBearPoseidon2GpuEngine;
36    } else {
37        use openvm_continuations::prover::DeferralInnerCpuProver as DeferralInnerProver;
38        use openvm_continuations::prover::DeferralHookCpuProver as DeferralHookProver;
39        type E = openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2CpuEngine;
40    }
41}
42
43#[allow(clippy::large_enum_variant)]
44#[derive(Clone)]
45pub enum DeferralProof {
46    Present(Proof<SC>),
47    Absent(DeferralPvs<F>),
48}
49
50/// Proves the configured set of deferral circuits through their hook proofs.
51pub struct MultiDeferralCircuitProver {
52    pub single_circuit_provers: Vec<SingleDeferralCircuitProver>,
53    pub internal_recursive_prover: DeferralInnerProver,
54    pub def_hook_prover: DeferralHookProver,
55}
56
57impl MultiDeferralCircuitProver {
58    pub fn new<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
59        def_circuit_prover: DP,
60        agg_config: AggregationConfig,
61        hook_params: SystemParams,
62    ) -> Self {
63        assert_eq!(def_circuit_prover.get_def_idx(), 0);
64        let single_circuit_prover = SingleDeferralCircuitProver::new(
65            def_circuit_prover,
66            agg_config.params.leaf,
67            agg_config.params.internal.clone(),
68        );
69        let internal_recursive_prover = DeferralInnerProver::new::<E>(
70            single_circuit_prover.internal_for_leaf_prover.get_vk(),
71            agg_config.params.internal,
72            true,
73        );
74        let internal_recursive_cached_commit = internal_recursive_prover
75            .get_vk_commit(true)
76            .cached_commit
77            .into();
78        let def_hook_prover = DeferralHookProver::new::<E>(
79            internal_recursive_prover.get_vk(),
80            internal_recursive_cached_commit,
81            hook_params,
82        );
83        Self {
84            single_circuit_provers: vec![single_circuit_prover],
85            internal_recursive_prover,
86            def_hook_prover,
87        }
88    }
89
90    pub fn from_pks<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
91        def_circuit_prover: DP,
92        def_prefix_pk: AggPrefixProvingKey,
93        def_internal_recursive_pk: Arc<MultiStarkProvingKey<SC>>,
94        def_hook_pk: Arc<MultiStarkProvingKey<SC>>,
95    ) -> Self {
96        assert_eq!(def_circuit_prover.get_def_idx(), 0);
97        let single_circuit_prover = SingleDeferralCircuitProver::from_pks(
98            def_circuit_prover,
99            def_prefix_pk.leaf,
100            def_prefix_pk.internal_for_leaf,
101        );
102        Self::from_single_circuit_prover(
103            single_circuit_prover,
104            def_internal_recursive_pk,
105            def_hook_pk,
106        )
107    }
108
109    pub fn from_single_circuit_prover(
110        single_circuit_prover: SingleDeferralCircuitProver,
111        def_internal_recursive_pk: Arc<MultiStarkProvingKey<SC>>,
112        def_hook_pk: Arc<MultiStarkProvingKey<SC>>,
113    ) -> Self {
114        assert_eq!(single_circuit_prover.def_circuit_prover.get_def_idx(), 0);
115        let internal_recursive_prover = DeferralInnerProver::from_pk::<E>(
116            single_circuit_prover.internal_for_leaf_prover.get_vk(),
117            def_internal_recursive_pk,
118            true,
119        );
120        let internal_recursive_cached_commit = internal_recursive_prover
121            .get_vk_commit(true)
122            .cached_commit
123            .into();
124        let def_hook_prover = DeferralHookProver::from_pk::<E>(
125            internal_recursive_prover.get_vk(),
126            internal_recursive_cached_commit,
127            def_hook_pk,
128        );
129        Self {
130            single_circuit_provers: vec![single_circuit_prover],
131            internal_recursive_prover,
132            def_hook_prover,
133        }
134    }
135
136    pub fn with_prover<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
137        mut self,
138        def_circuit_prover: DP,
139    ) -> Self {
140        assert_eq!(
141            def_circuit_prover.get_def_idx(),
142            self.single_circuit_provers.len()
143        );
144        let leaf_params = self.single_circuit_provers[0]
145            .leaf_prover
146            .get_vk()
147            .inner
148            .params
149            .clone();
150        let internal_params = self.internal_recursive_prover.get_vk().inner.params.clone();
151        let single_circuit_prover =
152            SingleDeferralCircuitProver::new(def_circuit_prover, leaf_params, internal_params);
153        self.single_circuit_provers.push(single_circuit_prover);
154        self
155    }
156
157    pub fn with_prover_from_pk<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
158        self,
159        def_circuit_prover: DP,
160        def_prefix_pk: AggPrefixProvingKey,
161    ) -> Self {
162        assert_eq!(
163            def_circuit_prover.get_def_idx(),
164            self.single_circuit_provers.len()
165        );
166        let single_circuit_prover = SingleDeferralCircuitProver::from_pks(
167            def_circuit_prover,
168            def_prefix_pk.leaf,
169            def_prefix_pk.internal_for_leaf,
170        );
171        self.with_single_circuit_prover(single_circuit_prover)
172    }
173
174    pub fn with_single_circuit_prover(
175        mut self,
176        single_circuit_prover: SingleDeferralCircuitProver,
177    ) -> Self {
178        assert_eq!(
179            single_circuit_prover.def_circuit_prover.get_def_idx(),
180            self.single_circuit_provers.len()
181        );
182        self.single_circuit_provers.push(single_circuit_prover);
183        self
184    }
185
186    pub fn prove(&self, inputs: &[DeferralInput]) -> Result<Vec<DeferralProof>> {
187        // Generate internal-for-leaf proofs and leaf IO commits per circuit
188        let per_circuit = self
189            .single_circuit_provers
190            .iter()
191            .zip_eq(inputs)
192            .map(|(prover, inputs)| prover.prove(inputs))
193            .collect::<Result<Vec<_>>>()?;
194
195        // For each circuit: do internal recursive aggregation then generate the hook proof
196        let mut per_circuit = per_circuit
197            .into_iter()
198            .enumerate()
199            .map(|(def_idx, res)| {
200                let mut proofs = res.internal_for_leaf_proofs;
201                if proofs.is_empty() {
202                    let def_circuit_commit = self.single_circuit_provers[def_idx]
203                        .circuit_commit(self.internal_recursive_prover.get_vk_commit(false));
204                    Ok(DeferralProof::Absent(absent_deferral_pvs(
205                        def_idx,
206                        def_circuit_commit,
207                    )))
208                } else {
209                    let mut merkle_depth = 2usize;
210                    let mut layer = 0usize;
211
212                    // Aggregate internal-for-leaf proofs down to a single proof.
213                    // First pass uses DeferralAggregation (children are i4l proofs);
214                    // subsequent passes use RecursiveSelf (children are ir proofs).
215                    loop {
216                        let is_first = layer == 0;
217                        let child_merkle_depth = if proofs.len() > 1 {
218                            let d = merkle_depth;
219                            merkle_depth += 1;
220                            Some(d)
221                        } else {
222                            None
223                        };
224
225                        proofs = info_span!(
226                            "agg_layer",
227                            group = format!("internal_recursive.{layer}"),
228                            circuit = def_idx
229                        )
230                        .in_scope(|| {
231                            proofs
232                                .chunks(2)
233                                .enumerate()
234                                .map(|(idx, chunk)| {
235                                    let kind = if is_first {
236                                        DeferralChildVkKind::DeferralAggregation
237                                    } else {
238                                        DeferralChildVkKind::RecursiveSelf
239                                    };
240                                    info_span!("single_internal_agg", idx = idx).in_scope(|| {
241                                        self.internal_recursive_prover.agg_prove::<E>(
242                                            chunk,
243                                            kind,
244                                            child_merkle_depth,
245                                        )
246                                    })
247                                })
248                                .collect::<Result<Vec<_>>>()
249                        })?;
250
251                        layer += 1;
252                        if proofs.len() == 1 {
253                            break;
254                        }
255                    }
256
257                    // Generate the deferral hook proof
258                    Ok(DeferralProof::Present(
259                        self.def_hook_prover
260                            .prove::<E>(proofs.pop().unwrap(), res.leaf_io_commits)?,
261                    ))
262                }
263            })
264            .collect::<Result<Vec<_>>>()?;
265
266        // Pad returned vector up to a power of two length with absent deferral proofs
267        let target_length = per_circuit.len().next_power_of_two();
268        for def_idx in per_circuit.len()..target_length {
269            per_circuit.push(DeferralProof::Absent(absent_deferral_pvs(
270                def_idx,
271                [F::ZERO; DIGEST_SIZE],
272            )));
273        }
274
275        Ok(per_circuit)
276    }
277
278    pub fn make_config(&self, supported_deferrals: Vec<SupportedDeferral>) -> DeferralConfig {
279        let vk_commit = self.internal_recursive_prover.get_vk_commit(false);
280        let circuits = self
281            .single_circuit_provers
282            .iter()
283            .zip_eq(supported_deferrals)
284            .map(|(p, deferral)| {
285                let commit = p.circuit_commit(vk_commit).into();
286                DeferralCircuitConfig {
287                    def_type: deferral,
288                    commit,
289                }
290            })
291            .collect();
292        DeferralConfig::new(circuits)
293    }
294
295    pub fn make_extension(&self, fns: Vec<Arc<DeferralFn>>) -> DeferralExtension {
296        let vk_commit = self.internal_recursive_prover.get_vk_commit(false);
297        let def_circuit_commits = self
298            .single_circuit_provers
299            .iter()
300            .map(|p| {
301                p.circuit_commit(vk_commit)
302                    .iter()
303                    .flat_map(|f| f.to_unique_u32().to_le_bytes())
304                    .collect::<Vec<_>>()
305                    .try_into()
306                    .unwrap()
307            })
308            .collect();
309        DeferralExtension {
310            fns,
311            def_circuit_commits,
312        }
313    }
314
315    pub fn get_pk(&self) -> DeferralProvingKey {
316        let circuits = self
317            .single_circuit_provers
318            .iter()
319            .map(|single_circuit_prover| DeferralCircuitProvingKey {
320                def_circuit_pk: single_circuit_prover.def_circuit_prover.get_pk(),
321                agg_prefix_pk: AggPrefixProvingKey {
322                    leaf: single_circuit_prover.leaf_prover.get_pk(),
323                    internal_for_leaf: single_circuit_prover.internal_for_leaf_prover.get_pk(),
324                },
325            })
326            .collect();
327        DeferralProvingKey {
328            circuits,
329            def_internal_recursive_pk: self.internal_recursive_prover.get_pk(),
330            def_hook_pk: self.def_hook_prover.get_pk(),
331        }
332    }
333}
334
335fn absent_deferral_pvs(def_idx: usize, def_circuit_commit: [F; DIGEST_SIZE]) -> DeferralPvs<F> {
336    let input_acc_hash = poseidon2_hash_slice(&def_circuit_commit).0;
337    let output_acc_hash = poseidon2_hash_slice(&[F::ZERO]).0;
338    let combined_hash = poseidon2_compress_with_capacity(input_acc_hash, output_acc_hash).0;
339    DeferralPvs {
340        initial_acc_hash: combined_hash,
341        final_acc_hash: combined_hash,
342        depth: F::ONE,
343        node_idx: F::from_usize(def_idx),
344    }
345}