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 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 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 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 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 let mut proofs = info_span!("agg_layer", group = "def_leaf")
272 .in_scope(|| reduce_def_round(input, ChildVkKind::App, &self.leaf_prover))?;
273
274 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 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 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 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 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 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 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}