openvm_sdk/prover/deferral/
single_circuit.rs1use std::{borrow::Borrow, iter::once, sync::Arc};
2
3use eyre::Result;
4use itertools::Itertools;
5use openvm_continuations::{
6 circuit::deferral::{hook::DeferralIoCommit, DeferralCircuitPvs, DEF_CIRCUIT_PVS_AIR_ID},
7 prover::{DeferralChildVkKind, DeferralCircuitProver},
8 SC,
9};
10use openvm_recursion_circuit::utils::poseidon2_hash_slice;
11use openvm_stark_backend::{keygen::types::MultiStarkProvingKey, proof::Proof, SystemParams};
12use openvm_stark_sdk::config::baby_bear_poseidon2::{Digest, F};
13use openvm_verify_stark_host::pvs::VkCommit;
14use tracing::info_span;
15
16use crate::DeferralInput;
17
18cfg_if::cfg_if! {
19 if #[cfg(feature = "cuda")] {
20 use openvm_continuations::prover::DeferralInnerGpuProver as DeferralInnerProver;
21 type E = openvm_cuda_backend::BabyBearPoseidon2GpuEngine;
22 } else {
23 use openvm_continuations::prover::DeferralInnerCpuProver as DeferralInnerProver;
24 type E = openvm_stark_sdk::config::baby_bear_poseidon2::BabyBearPoseidon2CpuEngine;
25 }
26}
27
28pub struct SingleDeferralCircuitProver {
29 pub def_circuit_prover: Box<dyn DeferralCircuitProver<SC> + Send + Sync>,
30 pub leaf_prover: DeferralInnerProver,
31 pub internal_for_leaf_prover: DeferralInnerProver,
32}
33
34pub struct SingleDeferralCircuitResult {
35 pub internal_for_leaf_proofs: Vec<Proof<SC>>,
36 pub leaf_io_commits: Vec<DeferralIoCommit<F>>,
37}
38
39impl SingleDeferralCircuitProver {
40 pub fn new<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
41 def_circuit_prover: DP,
42 leaf_params: SystemParams,
43 internal_params: SystemParams,
44 ) -> Self {
45 let leaf_prover =
46 DeferralInnerProver::new::<E>(def_circuit_prover.get_vk(), leaf_params, false);
47 let internal_for_leaf_prover =
48 DeferralInnerProver::new::<E>(leaf_prover.get_vk(), internal_params, false);
49 Self {
50 def_circuit_prover: Box::new(def_circuit_prover),
51 leaf_prover,
52 internal_for_leaf_prover,
53 }
54 }
55
56 pub fn from_pks<DP: DeferralCircuitProver<SC> + Send + Sync + 'static>(
57 def_circuit_prover: DP,
58 leaf_pk: Arc<MultiStarkProvingKey<SC>>,
59 internal_for_leaf_pk: Arc<MultiStarkProvingKey<SC>>,
60 ) -> Self {
61 let leaf_prover =
62 DeferralInnerProver::from_pk::<E>(def_circuit_prover.get_vk(), leaf_pk, false);
63 let internal_for_leaf_prover =
64 DeferralInnerProver::from_pk::<E>(leaf_prover.get_vk(), internal_for_leaf_pk, false);
65 Self {
66 def_circuit_prover: Box::new(def_circuit_prover),
67 leaf_prover,
68 internal_for_leaf_prover,
69 }
70 }
71
72 pub fn prove(&self, inputs: &DeferralInput) -> Result<SingleDeferralCircuitResult> {
73 let def_proofs = inputs
75 .byte_vec
76 .iter()
77 .map(|input| self.def_circuit_prover.prove(input))
78 .collect_vec();
79
80 let leaf_io_commits = def_proofs
82 .iter()
83 .map(|proof| {
84 let pvs: &DeferralCircuitPvs<F> = proof.public_values[DEF_CIRCUIT_PVS_AIR_ID]
85 .as_slice()
86 .borrow();
87 let commit_values = once(pvs.input_commit)
88 .chain(
89 proof
90 .trace_vdata
91 .iter()
92 .flatten()
93 .flat_map(|vdata| vdata.cached_commitments.iter().copied()),
94 )
95 .flatten()
96 .collect_vec();
97 let folded_input_commit = poseidon2_hash_slice(&commit_values).0;
98 (folded_input_commit, pvs.output_commit)
99 })
100 .collect();
101
102 let child_merkle_depth = (def_proofs.len() != 1).then_some(0);
104 let leaf_proofs = info_span!("agg_layer", group = "def_leaf").in_scope(|| {
105 def_proofs
106 .chunks(2)
107 .enumerate()
108 .map(|(leaf_node_idx, proofs)| {
109 info_span!("single_leaf_agg", idx = leaf_node_idx).in_scope(|| {
110 self.leaf_prover.agg_prove::<E>(
111 proofs,
112 DeferralChildVkKind::DeferralCircuit,
113 child_merkle_depth,
114 )
115 })
116 })
117 .collect::<Result<Vec<_>>>()
118 })?;
119
120 let mut internal_node_idx = 0u32;
122 let child_merkle_depth = (leaf_proofs.len() != 1).then_some(1);
123 let internal_for_leaf_proofs = info_span!("agg_layer", group = "internal_for_leaf")
124 .in_scope(|| {
125 leaf_proofs
126 .chunks(2)
127 .map(|proofs| {
128 let ret = info_span!("single_internal_agg", idx = internal_node_idx)
129 .in_scope(|| {
130 self.internal_for_leaf_prover.agg_prove::<E>(
131 proofs,
132 DeferralChildVkKind::DeferralAggregation,
133 child_merkle_depth,
134 )
135 });
136 internal_node_idx += 1;
137 ret
138 })
139 .collect::<Result<Vec<_>>>()
140 })?;
141
142 Ok(SingleDeferralCircuitResult {
143 internal_for_leaf_proofs,
144 leaf_io_commits,
145 })
146 }
147
148 pub fn circuit_commit(&self, internal_for_leaf_vk_commit: VkCommit<F>) -> Digest {
149 let def_vk_commit = self.leaf_prover.get_vk_commit(false);
150 let leaf_vk_commit = self.internal_for_leaf_prover.get_vk_commit(false);
151
152 let vk_commit_components = vec![
153 def_vk_commit.cached_commit,
154 def_vk_commit.vk_pre_hash,
155 leaf_vk_commit.cached_commit,
156 leaf_vk_commit.vk_pre_hash,
157 internal_for_leaf_vk_commit.cached_commit,
158 internal_for_leaf_vk_commit.vk_pre_hash,
159 ];
160 poseidon2_hash_slice(&vk_commit_components.into_flattened()).0
161 }
162}