openvm_sdk/prover/deferral/
single_circuit.rs

1use 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        // Generate deferral circuit proofs
74        let def_proofs = inputs
75            .byte_vec
76            .iter()
77            .map(|input| self.def_circuit_prover.prove(input))
78            .collect_vec();
79
80        // Extract leaf IO commits from the deferral circuit proofs
81        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        // Verify def-layer proofs and generate leaf-layer proofs
103        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        // Verify leaf-layer proofs and generate internal-for-leaf-layer proofs
121        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}