openvm_sdk/prover/
agg.rs

1use std::sync::Arc;
2
3use eyre::Result;
4use itertools::Itertools;
5use openvm_circuit::arch::ContinuationVmProof;
6use openvm_continuations::{circuit::inner::ProofsType, prover::ChildVkKind};
7use openvm_recursion_circuit::{
8    prelude::Digest, system::check_param_compatibility, utils::poseidon2_hash_slice,
9};
10use openvm_stark_backend::{
11    codec::{Decode, Encode},
12    keygen::types::MultiStarkVerifyingKey,
13    p3_field::PrimeCharacteristicRing,
14    proof::Proof,
15};
16use openvm_stark_sdk::config::baby_bear_poseidon2::{poseidon2_compress_with_capacity, F};
17use openvm_verify_stark_host::{pvs::DeferralPvs, VmStarkProof};
18use tracing::info_span;
19
20use crate::{
21    config::{
22        AggregationConfig, AggregationTreeConfig, MAX_NUM_CHILDREN_INTERNAL, MAX_NUM_CHILDREN_LEAF,
23    },
24    keygen::{AggPrefixProvingKey, AggProvingKey},
25    prover::deferral::DeferralProof,
26    SC,
27};
28
29cfg_if::cfg_if! {
30    if #[cfg(feature = "cuda")] {
31        use openvm_continuations::prover::InnerGpuProver as InnerAggregationProver;
32        type E = openvm_cuda_backend::BabyBearPoseidon2GpuEngine;
33    } else {
34        use openvm_continuations::prover::InnerCpuProver as InnerAggregationProver;
35        type E = openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2CpuEngine;
36    }
37}
38
39pub struct AggProver {
40    pub leaf_prover: InnerAggregationProver<MAX_NUM_CHILDREN_LEAF>,
41    pub internal_for_leaf_prover: InnerAggregationProver<MAX_NUM_CHILDREN_INTERNAL>,
42    pub internal_recursive_prover: InnerAggregationProver<MAX_NUM_CHILDREN_INTERNAL>,
43    pub agg_tree_config: AggregationTreeConfig,
44}
45
46#[derive(Clone)]
47pub struct InternalLayerMetadata {
48    pub internal_recursive_layer: u32,
49    pub internal_node_idx: u32,
50    pub proofs_type: ProofsType,
51}
52
53impl AggProver {
54    pub fn keygen_prefix(
55        app_or_def_vk: Arc<MultiStarkVerifyingKey<SC>>,
56        agg_config: AggregationConfig,
57        def_hook_cached_commit: Option<Digest>,
58    ) -> AggPrefixProvingKey {
59        check_param_compatibility(
60            &app_or_def_vk.inner.params,
61            &agg_config.params.leaf,
62            &agg_config.params.internal,
63        );
64
65        let leaf_prover = InnerAggregationProver::<MAX_NUM_CHILDREN_LEAF>::new::<E>(
66            app_or_def_vk,
67            agg_config.params.leaf.clone(),
68            false,
69            def_hook_cached_commit,
70        );
71        let internal_for_leaf_prover = InnerAggregationProver::<MAX_NUM_CHILDREN_INTERNAL>::new::<E>(
72            leaf_prover.get_vk(),
73            agg_config.params.internal,
74            false,
75            def_hook_cached_commit,
76        );
77        AggPrefixProvingKey {
78            leaf: leaf_prover.get_pk(),
79            internal_for_leaf: internal_for_leaf_prover.get_pk(),
80        }
81    }
82
83    #[tracing::instrument(level = "info", fields(group = "agg_keygen"), skip_all)]
84    pub fn new(
85        app_or_def_vk: Arc<MultiStarkVerifyingKey<SC>>,
86        agg_config: AggregationConfig,
87        agg_tree_config: AggregationTreeConfig,
88        def_hook_cached_commit: Option<Digest>,
89    ) -> Self {
90        assert!(agg_tree_config.num_children_leaf <= MAX_NUM_CHILDREN_LEAF);
91        assert!(agg_tree_config.num_children_internal <= MAX_NUM_CHILDREN_INTERNAL);
92        check_param_compatibility(
93            &app_or_def_vk.inner.params,
94            &agg_config.params.leaf,
95            &agg_config.params.internal,
96        );
97
98        let leaf_prover = InnerAggregationProver::new::<E>(
99            app_or_def_vk,
100            agg_config.params.leaf.clone(),
101            false,
102            def_hook_cached_commit,
103        );
104        let internal_for_leaf_prover = InnerAggregationProver::new::<E>(
105            leaf_prover.get_vk(),
106            agg_config.params.internal.clone(),
107            false,
108            def_hook_cached_commit,
109        );
110        let internal_recursive_prover = InnerAggregationProver::new::<E>(
111            internal_for_leaf_prover.get_vk(),
112            agg_config.params.internal.clone(),
113            true,
114            def_hook_cached_commit,
115        );
116        Self {
117            leaf_prover,
118            internal_for_leaf_prover,
119            internal_recursive_prover,
120            agg_tree_config,
121        }
122    }
123
124    pub fn from_pk(
125        app_or_def_vk: Arc<MultiStarkVerifyingKey<SC>>,
126        agg_pk: AggProvingKey,
127        agg_tree_config: AggregationTreeConfig,
128        def_hook_cached_commit: Option<Digest>,
129    ) -> Self {
130        assert!(agg_tree_config.num_children_leaf <= MAX_NUM_CHILDREN_LEAF);
131        assert!(agg_tree_config.num_children_internal <= MAX_NUM_CHILDREN_INTERNAL);
132        check_param_compatibility(
133            &app_or_def_vk.inner.params,
134            &agg_pk.prefix.leaf.params,
135            &agg_pk.prefix.internal_for_leaf.params,
136        );
137
138        let leaf_prover = InnerAggregationProver::from_pk::<E>(
139            app_or_def_vk,
140            agg_pk.prefix.leaf,
141            false,
142            def_hook_cached_commit,
143        );
144        let internal_for_leaf_prover = InnerAggregationProver::from_pk::<E>(
145            leaf_prover.get_vk(),
146            agg_pk.prefix.internal_for_leaf,
147            false,
148            def_hook_cached_commit,
149        );
150        let internal_recursive_prover = InnerAggregationProver::from_pk::<E>(
151            internal_for_leaf_prover.get_vk(),
152            agg_pk.internal_recursive,
153            true,
154            def_hook_cached_commit,
155        );
156        Self {
157            leaf_prover,
158            internal_for_leaf_prover,
159            internal_recursive_prover,
160            agg_tree_config,
161        }
162    }
163
164    pub fn vm_or_hook_commit(&self) -> Digest {
165        let app_or_def_vk_commit = self.leaf_prover.get_vk_commit(false);
166        let leaf_vk_commit = self.internal_for_leaf_prover.get_vk_commit(false);
167        let internal_for_leaf_vk_commit = self.internal_recursive_prover.get_vk_commit(false);
168        let components = vec![
169            app_or_def_vk_commit.cached_commit,
170            app_or_def_vk_commit.vk_pre_hash,
171            leaf_vk_commit.cached_commit,
172            leaf_vk_commit.vk_pre_hash,
173            internal_for_leaf_vk_commit.cached_commit,
174            internal_for_leaf_vk_commit.vk_pre_hash,
175        ]
176        .into_flattened();
177        poseidon2_hash_slice(&components).0
178    }
179
180    pub fn prove_vm(
181        &self,
182        continuation_proof: ContinuationVmProof<SC>,
183    ) -> Result<(VmStarkProof, InternalLayerMetadata)> {
184        // Verify app-layer proofs and generate leaf-layer proofs
185        let leaf_proofs = info_span!("agg_layer", group = "leaf").in_scope(|| {
186            continuation_proof
187                .per_segment
188                .chunks(self.agg_tree_config.num_children_leaf)
189                .enumerate()
190                .map(|(leaf_node_idx, proofs)| {
191                    info_span!("single_leaf_agg", idx = leaf_node_idx).in_scope(|| {
192                        self.leaf_prover
193                            .agg_prove_no_def::<E>(proofs, ChildVkKind::App)
194                    })
195                })
196                .collect::<Result<Vec<_>>>()
197        })?;
198
199        // Verify leaf-layer proofs and generate internal-for-leaf-layer proofs
200        let mut internal_node_idx = -1;
201        let mut internal_proofs =
202            info_span!("agg_layer", group = "internal_for_leaf").in_scope(|| {
203                leaf_proofs
204                    .chunks(self.agg_tree_config.num_children_internal)
205                    .map(|proofs| {
206                        internal_node_idx += 1;
207                        info_span!("single_internal_agg", idx = internal_node_idx).in_scope(|| {
208                            self.internal_for_leaf_prover
209                                .agg_prove_no_def::<E>(proofs, ChildVkKind::Standard)
210                        })
211                    })
212                    .collect::<Result<Vec<_>>>()
213            })?;
214
215        // Verify internal-for-leaf-layer proofs and generate internal-recursive-layer proofs
216        internal_proofs =
217            info_span!("agg_layer", group = "internal_recursive.0").in_scope(|| {
218                internal_proofs
219                    .chunks(self.agg_tree_config.num_children_internal)
220                    .map(|proofs| {
221                        internal_node_idx += 1;
222                        info_span!("single_internal_agg", idx = internal_node_idx).in_scope(|| {
223                            self.internal_recursive_prover
224                                .agg_prove_no_def::<E>(proofs, ChildVkKind::Standard)
225                        })
226                    })
227                    .collect::<Result<Vec<_>>>()
228            })?;
229
230        // Recursively verify internal-layer proofs until only 1 remains
231        let mut internal_recursive_layer = 1;
232        while internal_proofs.len() > 1 {
233            internal_proofs = info_span!(
234                "agg_layer",
235                group = format!("internal_recursive.{internal_recursive_layer}")
236            )
237            .in_scope(|| {
238                internal_proofs
239                    .chunks(self.agg_tree_config.num_children_internal)
240                    .map(|proofs| {
241                        internal_node_idx += 1;
242                        info_span!("single_internal_agg", idx = internal_node_idx).in_scope(|| {
243                            self.internal_recursive_prover
244                                .agg_prove_no_def::<E>(proofs, ChildVkKind::RecursiveSelf)
245                        })
246                    })
247                    .collect::<Result<Vec<_>>>()
248            })?;
249            internal_recursive_layer += 1;
250        }
251
252        Ok((
253            VmStarkProof {
254                inner: internal_proofs.pop().unwrap(),
255                user_pvs_proof: continuation_proof.user_public_values,
256                deferral_merkle_proofs: None,
257            },
258            InternalLayerMetadata {
259                internal_recursive_layer: internal_recursive_layer as u32,
260                internal_node_idx: internal_node_idx as u32,
261                proofs_type: ProofsType::Vm,
262            },
263        ))
264    }
265
266    pub fn prove_def(&self, input: Vec<DeferralProof>) -> Result<(DeferralProof, u32)> {
267        assert!(!input.is_empty());
268        assert!(input.len().is_power_of_two());
269
270        // Leaf round: hook-level → leaf-level
271        let mut proofs = info_span!("agg_layer", group = "def_leaf")
272            .in_scope(|| reduce_def_round(input, ChildVkKind::App, &self.leaf_prover))?;
273
274        // Internal-for-leaf round: leaf-level → i4l-level
275        proofs = info_span!("agg_layer", group = "def_internal_for_leaf").in_scope(|| {
276            reduce_def_round(
277                proofs,
278                ChildVkKind::Standard,
279                &self.internal_for_leaf_prover,
280            )
281        })?;
282
283        // Internal-recursive round 0: i4l-level → ir-level
284        proofs = info_span!("agg_layer", group = "def_internal_recursive.0").in_scope(|| {
285            reduce_def_round(
286                proofs,
287                ChildVkKind::Standard,
288                &self.internal_recursive_prover,
289            )
290        })?;
291
292        // Internal-recursive rounds: ir-level → ir-level until single proof remains
293        let mut layer = 1;
294        while proofs.len() > 1 {
295            proofs = info_span!(
296                "agg_layer",
297                group = format!("def_internal_recursive.{layer}")
298            )
299            .in_scope(|| {
300                reduce_def_round(
301                    proofs,
302                    ChildVkKind::RecursiveSelf,
303                    &self.internal_recursive_prover,
304                )
305            })?;
306            layer += 1;
307        }
308
309        Ok((proofs.pop().unwrap(), layer))
310    }
311
312    pub fn prove_mixed(
313        &self,
314        mut vm_proof: VmStarkProof,
315        def_proof: DeferralProof,
316        metadata: &mut InternalLayerMetadata,
317        mut def_internal_recursive_layer: u32,
318    ) -> Result<VmStarkProof> {
319        let DeferralProof::Present(mut def_inner) = def_proof else {
320            return Ok(vm_proof);
321        };
322
323        // Aggregation requires equal child recursion_depth values, so we wrap the
324        // shallower side until both proofs are at the same internal-recursive depth.
325        while metadata.internal_recursive_layer < def_internal_recursive_layer {
326            vm_proof = self.wrap_proof(vm_proof, metadata)?;
327        }
328        while def_internal_recursive_layer < metadata.internal_recursive_layer {
329            def_inner = self.wrap_def_inner(def_inner, def_internal_recursive_layer)?;
330            def_internal_recursive_layer += 1;
331        }
332
333        vm_proof.inner = info_span!(
334            "agg_layer",
335            group = format!("internal_recursive.{}", metadata.internal_recursive_layer)
336        )
337        .in_scope(|| {
338            metadata.internal_recursive_layer += 1;
339            info_span!("single_internal_agg", idx = metadata.internal_node_idx).in_scope(|| {
340                metadata.internal_node_idx += 1;
341                self.internal_recursive_prover.agg_prove::<E>(
342                    &[vm_proof.inner, def_inner],
343                    ChildVkKind::RecursiveSelf,
344                    ProofsType::Mix,
345                    None,
346                )
347            })
348        })?;
349
350        metadata.proofs_type = ProofsType::Combined;
351        Ok(vm_proof)
352    }
353
354    pub fn wrap_proof(
355        &self,
356        mut proof: VmStarkProof,
357        metadata: &mut InternalLayerMetadata,
358    ) -> Result<VmStarkProof> {
359        proof.inner = info_span!(
360            "agg_layer",
361            group = format!("internal_recursive.{}", metadata.internal_recursive_layer)
362        )
363        .in_scope(|| {
364            metadata.internal_recursive_layer += 1;
365            info_span!("single_internal_agg", idx = metadata.internal_node_idx).in_scope(|| {
366                metadata.internal_node_idx += 1;
367                self.internal_recursive_prover.agg_prove::<E>(
368                    &[proof.inner],
369                    ChildVkKind::RecursiveSelf,
370                    metadata.proofs_type,
371                    None,
372                )
373            })
374        })?;
375        Ok(proof)
376    }
377
378    pub(crate) fn wrap_def_inner(
379        &self,
380        mut proof: Proof<SC>,
381        def_internal_recursive_layer: u32,
382    ) -> Result<Proof<SC>> {
383        proof = info_span!(
384            "agg_layer",
385            group = format!("def_internal_recursive.{def_internal_recursive_layer}")
386        )
387        .in_scope(|| {
388            self.internal_recursive_prover.agg_prove::<E>(
389                &[proof],
390                ChildVkKind::RecursiveSelf,
391                ProofsType::Deferral,
392                None,
393            )
394        })?;
395        Ok(proof)
396    }
397}
398
399fn reduce_def_round<const N: usize>(
400    proofs: Vec<DeferralProof>,
401    kind: ChildVkKind,
402    prover: &InnerAggregationProver<N>,
403) -> Result<Vec<DeferralProof>> {
404    if proofs.len() == 1 {
405        // A singleton round can only happen when the entire round has one present input proof.
406        let DeferralProof::Present(p) = proofs.into_iter().next().unwrap() else {
407            panic!("singleton deferral round must contain a present proof");
408        };
409        return Ok(vec![DeferralProof::Present(prover.agg_prove::<E>(
410            &[p],
411            kind,
412            ProofsType::Deferral,
413            None,
414        )?)]);
415    }
416
417    assert!(
418        proofs.len().is_multiple_of(2),
419        "non-singleton deferral round must have an even number of proofs"
420    );
421
422    let mut next = Vec::with_capacity(proofs.len() / 2);
423    for (a, b) in proofs.into_iter().tuples() {
424        let combined = match (a, b) {
425            (DeferralProof::Present(p0), DeferralProof::Present(p1)) => DeferralProof::Present(
426                prover.agg_prove::<E>(&[p0, p1], kind, ProofsType::Deferral, None)?,
427            ),
428            (DeferralProof::Present(p), DeferralProof::Absent(pvs)) => {
429                // Absent is the right child (present is left, is_right = false)
430                DeferralProof::Present(prover.agg_prove::<E>(
431                    &[p],
432                    kind,
433                    ProofsType::Deferral,
434                    Some((pvs, false)),
435                )?)
436            }
437            (DeferralProof::Absent(pvs), DeferralProof::Present(p)) => {
438                // Absent is the left child (present is right, is_right = true)
439                DeferralProof::Present(prover.agg_prove::<E>(
440                    &[p],
441                    kind,
442                    ProofsType::Deferral,
443                    Some((pvs, true)),
444                )?)
445            }
446            (DeferralProof::Absent(pvs0), DeferralProof::Absent(pvs1)) => {
447                debug_assert_eq!(pvs0.depth, pvs1.depth);
448                debug_assert_eq!(pvs0.node_idx + F::ONE, pvs1.node_idx);
449                DeferralProof::Absent(DeferralPvs {
450                    initial_acc_hash: poseidon2_compress_with_capacity(
451                        pvs0.initial_acc_hash,
452                        pvs1.initial_acc_hash,
453                    )
454                    .0,
455                    final_acc_hash: poseidon2_compress_with_capacity(
456                        pvs0.final_acc_hash,
457                        pvs1.final_acc_hash,
458                    )
459                    .0,
460                    depth: pvs0.depth + F::ONE,
461                    node_idx: pvs0.node_idx.halve(),
462                })
463            }
464        };
465        next.push(combined);
466    }
467    Ok(next)
468}
469
470impl Encode for InternalLayerMetadata {
471    fn encode<W: std::io::Write>(&self, writer: &mut W) -> std::io::Result<()> {
472        self.internal_recursive_layer.encode(writer)?;
473        self.internal_node_idx.encode(writer)?;
474        let proofs_type_byte: u8 = match self.proofs_type {
475            ProofsType::Vm => 0,
476            ProofsType::Deferral => 1,
477            ProofsType::Mix => 2,
478            ProofsType::Combined => 3,
479        };
480        proofs_type_byte.encode(writer)
481    }
482}
483
484impl Decode for InternalLayerMetadata {
485    fn decode<R: std::io::Read>(reader: &mut R) -> std::io::Result<Self> {
486        let internal_recursive_layer = u32::decode(reader)?;
487        let internal_node_idx = u32::decode(reader)?;
488        let proofs_type = match u8::decode(reader)? {
489            0 => ProofsType::Vm,
490            1 => ProofsType::Deferral,
491            2 => ProofsType::Mix,
492            3 => ProofsType::Combined,
493            b => {
494                return Err(std::io::Error::new(
495                    std::io::ErrorKind::InvalidData,
496                    format!("invalid ProofsType byte: {b}"),
497                ))
498            }
499        };
500        Ok(Self {
501            internal_recursive_layer,
502            internal_node_idx,
503            proofs_type,
504        })
505    }
506}