openvm_continuations/circuit/root/
trace.rs

1use std::borrow::Borrow;
2
3use itertools::Itertools;
4use openvm_circuit::{
5    arch::POSEIDON2_WIDTH,
6    system::memory::{dimensions::MemoryDimensions, merkle::public_values::UserPublicValuesProof},
7};
8#[cfg(feature = "cuda")]
9use openvm_circuit_primitives::hybrid_chip::cpu_proving_ctx_to_gpu;
10use openvm_cpu_backend::CpuBackend;
11#[cfg(feature = "cuda")]
12use openvm_cuda_backend::{BabyBearBn254Poseidon2HashScheme, GenericGpuBackend};
13#[cfg(feature = "cuda")]
14use openvm_cuda_common::stream::GpuDeviceCtx;
15use openvm_stark_backend::{
16    proof::Proof,
17    prover::{AirProvingContext, ProverBackend},
18    StarkProtocolConfig,
19};
20use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, DIGEST_SIZE, F};
21use openvm_verify_stark_host::pvs::{DeferralPvs, DEF_PVS_AIR_ID};
22use p3_field::PrimeField32;
23
24use crate::circuit::{
25    deferral::DeferralMerkleProofs,
26    root::{commit, memory},
27    SingleAirTraceData, SubCircuitTraceData,
28};
29
30// Trait that root provers use to remain generic in PB. Tracegen returns the AIR proving
31// contexts, Poseidon2 compress inputs, and Poseidon2 permute inputs to be fed to Poseidon2Air.
32pub trait RootTraceGen<PB: ProverBackend, DC: Clone + Send + Sync> {
33    fn new(deferral_enabled: bool) -> Self;
34    fn generate_pre_verifier_subcircuit_ctx(
35        &self,
36        proof: &Proof<BabyBearPoseidon2Config>,
37        user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, PB::Val>,
38        memory_dimensions: MemoryDimensions,
39        device_ctx: &DC,
40    ) -> SubCircuitTraceData<PB>;
41    fn generate_other_proving_ctxs(
42        &self,
43        proof: &Proof<BabyBearPoseidon2Config>,
44        memory_dimensions: MemoryDimensions,
45        deferral_merkle_proofs: Option<&DeferralMerkleProofs<PB::Val>>,
46        device_ctx: &DC,
47    ) -> (Vec<AirProvingContext<PB>>, Vec<[PB::Val; POSEIDON2_WIDTH]>);
48}
49
50pub struct RootTraceGenImpl {
51    pub deferral_enabled: bool,
52}
53
54impl<SC: StarkProtocolConfig<F = F>> RootTraceGen<CpuBackend<SC>, ()> for RootTraceGenImpl {
55    fn new(deferral_enabled: bool) -> Self {
56        Self { deferral_enabled }
57    }
58
59    fn generate_pre_verifier_subcircuit_ctx(
60        &self,
61        proof: &Proof<BabyBearPoseidon2Config>,
62        user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, F>,
63        memory_dimensions: MemoryDimensions,
64        _device_ctx: &(),
65    ) -> SubCircuitTraceData<CpuBackend<SC>> {
66        let SingleAirTraceData {
67            air_proving_ctx: verifier_ctx,
68            poseidon2_compress_inputs: verifier_compress_inputs,
69            poseidon2_permute_inputs: verifier_permute_inputs,
70            range_check_inputs,
71        } = super::verifier::generate_proving_ctx(proof, self.deferral_enabled);
72        let (commit_ctx, commit_inputs) =
73            commit::generate_proving_ctx(user_pvs_proof.public_values.clone());
74        let (memory_ctx, memory_inputs) = memory::generate_proving_input(
75            user_pvs_proof.public_values_commit,
76            &user_pvs_proof.proof,
77            memory_dimensions,
78            user_pvs_proof.public_values.len(),
79        );
80        SubCircuitTraceData {
81            air_proving_ctxs: vec![verifier_ctx, commit_ctx, memory_ctx],
82            poseidon2_compress_inputs: verifier_compress_inputs
83                .into_iter()
84                .chain(commit_inputs)
85                .chain(memory_inputs)
86                .collect_vec(),
87            poseidon2_permute_inputs: verifier_permute_inputs,
88            range_check_inputs,
89        }
90    }
91
92    fn generate_other_proving_ctxs(
93        &self,
94        proof: &Proof<BabyBearPoseidon2Config>,
95        memory_dimensions: MemoryDimensions,
96        deferral_merkle_proofs: Option<&DeferralMerkleProofs<F>>,
97        _device_ctx: &(),
98    ) -> (
99        Vec<AirProvingContext<CpuBackend<SC>>>,
100        Vec<[F; POSEIDON2_WIDTH]>,
101    ) {
102        let (paths_ctx, paths_inputs) = if let Some(deferral_merkle_proofs) = deferral_merkle_proofs
103        {
104            assert!(self.deferral_enabled);
105            let def_pvs: &DeferralPvs<F> = proof.public_values[DEF_PVS_AIR_ID].as_slice().borrow();
106            let depth = def_pvs.depth.as_canonical_u32() as usize;
107            let (ctx, inputs) = super::def_paths::generate_proving_input(
108                def_pvs.initial_acc_hash,
109                def_pvs.final_acc_hash,
110                &deferral_merkle_proofs.initial_merkle_proof,
111                &deferral_merkle_proofs.final_merkle_proof,
112                memory_dimensions,
113                depth,
114                depth == 0,
115            );
116            (Some(ctx), inputs)
117        } else {
118            assert!(!self.deferral_enabled);
119            (None, vec![])
120        };
121        (paths_ctx.into_iter().collect_vec(), paths_inputs)
122    }
123}
124
125#[cfg(feature = "cuda")]
126impl RootTraceGen<GenericGpuBackend<BabyBearBn254Poseidon2HashScheme>, GpuDeviceCtx>
127    for RootTraceGenImpl
128{
129    fn new(deferral_enabled: bool) -> Self {
130        Self { deferral_enabled }
131    }
132
133    fn generate_pre_verifier_subcircuit_ctx(
134        &self,
135        proof: &Proof<BabyBearPoseidon2Config>,
136        user_pvs_proof: &UserPublicValuesProof<DIGEST_SIZE, F>,
137        memory_dimensions: MemoryDimensions,
138        device_ctx: &GpuDeviceCtx,
139    ) -> SubCircuitTraceData<GenericGpuBackend<BabyBearBn254Poseidon2HashScheme>> {
140        let data: SubCircuitTraceData<CpuBackend<BabyBearPoseidon2Config>> =
141            <Self as RootTraceGen<CpuBackend<BabyBearPoseidon2Config>, ()>>::generate_pre_verifier_subcircuit_ctx(
142                self,
143                proof,
144                user_pvs_proof,
145                memory_dimensions,
146                &(),
147            );
148        SubCircuitTraceData {
149            air_proving_ctxs: data
150                .air_proving_ctxs
151                .into_iter()
152                .map(|c| cpu_proving_ctx_to_gpu::<BabyBearBn254Poseidon2HashScheme>(c, device_ctx))
153                .collect_vec(),
154            poseidon2_compress_inputs: data.poseidon2_compress_inputs,
155            poseidon2_permute_inputs: data.poseidon2_permute_inputs,
156            range_check_inputs: data.range_check_inputs,
157        }
158    }
159
160    fn generate_other_proving_ctxs(
161        &self,
162        proof: &Proof<BabyBearPoseidon2Config>,
163        memory_dimensions: MemoryDimensions,
164        deferral_merkle_proofs: Option<&DeferralMerkleProofs<F>>,
165        device_ctx: &GpuDeviceCtx,
166    ) -> (
167        Vec<AirProvingContext<GenericGpuBackend<BabyBearBn254Poseidon2HashScheme>>>,
168        Vec<[F; POSEIDON2_WIDTH]>,
169    ) {
170        let (cpu_ctxs, inputs) =
171            <Self as RootTraceGen<CpuBackend<BabyBearPoseidon2Config>, ()>>::generate_other_proving_ctxs(
172                self,
173                proof,
174                memory_dimensions,
175                deferral_merkle_proofs,
176                &(),
177            );
178        let gpu_proving_ctxs = cpu_ctxs
179            .into_iter()
180            .map(|c| cpu_proving_ctx_to_gpu::<BabyBearBn254Poseidon2HashScheme>(c, device_ctx))
181            .collect_vec();
182        (gpu_proving_ctxs, inputs)
183    }
184}