openvm_continuations/circuit/inner/def_pvs/
trace.rs

1use std::borrow::{Borrow, BorrowMut};
2
3use itertools::Itertools;
4use openvm_cpu_backend::CpuBackend;
5use openvm_poseidon2_air::POSEIDON2_WIDTH;
6use openvm_stark_backend::{proof::Proof, prover::AirProvingContext};
7use openvm_stark_sdk::config::baby_bear_poseidon2::{
8    poseidon2_compress_with_capacity, BabyBearPoseidon2Config, F,
9};
10use openvm_verify_stark_host::pvs::{DeferralPvs, DEF_PVS_AIR_ID};
11use p3_field::{PrimeCharacteristicRing, PrimeField32};
12use p3_matrix::dense::RowMajorMatrix;
13
14use crate::{
15    circuit::{
16        deferral::DEF_HOOK_PVS_AIR_ID,
17        inner::{def_pvs::air::DeferralPvsCols, ProofsType},
18    },
19    utils::digests_to_poseidon2_input,
20};
21
22pub fn generate_proving_ctx(
23    proofs: &[Proof<BabyBearPoseidon2Config>],
24    proofs_type: ProofsType,
25    child_is_app: bool,
26    absent_trace_pvs: Option<(DeferralPvs<F>, bool)>,
27) -> (
28    AirProvingContext<CpuBackend<BabyBearPoseidon2Config>>,
29    Vec<[F; POSEIDON2_WIDTH]>,
30    Vec<usize>,
31) {
32    assert!(
33        absent_trace_pvs.is_none()
34            || (matches!(proofs_type, ProofsType::Deferral) && proofs.len() == 1),
35        "absent_trace_pvs is only valid for single-proof deferral aggregation"
36    );
37    let mut proof_idxs = vec![];
38    let (num_rows, def_flag) = match proofs_type {
39        ProofsType::Vm => (1, 0),
40        ProofsType::Deferral => {
41            proof_idxs = (0..proofs.len()).collect_vec();
42            (proofs.len() + absent_trace_pvs.is_some() as usize, 1)
43        }
44        ProofsType::Mix => {
45            proof_idxs.push(1);
46            (1, 1)
47        }
48        ProofsType::Combined => {
49            proof_idxs.push(0);
50            (1, 2)
51        }
52    };
53
54    let width = DeferralPvsCols::<u8>::width();
55    let mut trace = vec![F::ZERO; num_rows * width];
56    let mut chunks = trace.chunks_exact_mut(width);
57
58    let mut child_pvs_vec = vec![];
59    let single_present_is_right = if let Some((_, is_right)) = absent_trace_pvs.as_ref() {
60        *is_right
61    } else {
62        false
63    };
64
65    for (row_idx, proof_idx) in proof_idxs.iter().enumerate() {
66        let proof = &proofs[*proof_idx];
67        let chunk = chunks.next().unwrap();
68        let cols: &mut DeferralPvsCols<F> = chunk.borrow_mut();
69        cols.row_idx = F::from_usize(row_idx);
70        cols.proof_idx = F::from_usize(*proof_idx);
71        cols.is_present = F::ONE;
72        cols.deferral_flag = F::from_usize(def_flag);
73        cols.has_verifier_pvs = F::from_bool(!child_is_app);
74        cols.single_present_is_right = F::from_bool(single_present_is_right);
75
76        let air_id = if child_is_app {
77            DEF_HOOK_PVS_AIR_ID
78        } else {
79            DEF_PVS_AIR_ID
80        };
81        let child_pvs: &DeferralPvs<_> = proof.public_values[air_id].as_slice().borrow();
82        cols.child_pvs = *child_pvs;
83        child_pvs_vec.push(cols.child_pvs);
84    }
85
86    if let Some((pvs, _)) = absent_trace_pvs {
87        let chunk = chunks.next().unwrap();
88        let cols: &mut DeferralPvsCols<F> = chunk.borrow_mut();
89        cols.row_idx = F::ONE;
90        cols.deferral_flag = F::from_usize(def_flag);
91        cols.has_verifier_pvs = F::from_bool(!child_is_app);
92        cols.single_present_is_right = F::from_bool(single_present_is_right);
93        cols.child_pvs = pvs;
94        child_pvs_vec.push(cols.child_pvs);
95    }
96
97    let mut poseidon2_inputs = vec![];
98    let mut range_check_inputs = vec![];
99    let mut public_values = vec![F::ZERO; DeferralPvs::<u8>::width()];
100    let pvs: &mut DeferralPvs<F> = public_values.as_mut_slice().borrow_mut();
101
102    if child_pvs_vec.len() == 1 {
103        *pvs = child_pvs_vec[0];
104    } else if child_pvs_vec.len() == 2 {
105        let first_child = child_pvs_vec[0];
106        let second_child = child_pvs_vec[1];
107        let (left_initial, right_initial, left_final, right_final) = if single_present_is_right {
108            (
109                second_child.initial_acc_hash,
110                first_child.initial_acc_hash,
111                second_child.final_acc_hash,
112                first_child.final_acc_hash,
113            )
114        } else {
115            (
116                first_child.initial_acc_hash,
117                second_child.initial_acc_hash,
118                first_child.final_acc_hash,
119                second_child.final_acc_hash,
120            )
121        };
122        pvs.initial_acc_hash = poseidon2_compress_with_capacity(left_initial, right_initial).0;
123        poseidon2_inputs.push(digests_to_poseidon2_input(left_initial, right_initial));
124        pvs.final_acc_hash = poseidon2_compress_with_capacity(left_final, right_final).0;
125        poseidon2_inputs.push(digests_to_poseidon2_input(left_final, right_final));
126        pvs.depth = first_child.depth + F::ONE;
127        pvs.node_idx = (first_child.node_idx - F::from_bool(single_present_is_right)).halve();
128        range_check_inputs.push(pvs.node_idx.as_canonical_u32() as usize);
129    }
130
131    (
132        AirProvingContext {
133            cached_mains: vec![],
134            common_main: RowMajorMatrix::new(trace, width),
135            public_values,
136        },
137        poseidon2_inputs,
138        range_check_inputs,
139    )
140}