openvm_continuations/circuit/deferral/inner/input/
trace.rs

1use std::borrow::{Borrow, BorrowMut};
2
3use openvm_cpu_backend::CpuBackend;
4use openvm_stark_backend::{proof::Proof, prover::AirProvingContext};
5use openvm_stark_sdk::config::baby_bear_poseidon2::{
6    poseidon2_compress_with_capacity, BabyBearPoseidon2Config, DIGEST_SIZE, F,
7};
8use p3_field::PrimeCharacteristicRing;
9use p3_matrix::dense::RowMajorMatrix;
10
11use crate::circuit::deferral::{
12    inner::input::air::InputCommitCols, DeferralAggregationPvs, DeferralCircuitPvs,
13    DEF_AGG_PVS_AIR_ID, DEF_CIRCUIT_PVS_AIR_ID,
14};
15
16type CachedCommitRow<F> = (usize, usize, [F; DIGEST_SIZE]);
17
18pub fn generate_proving_ctx(
19    proofs: &[Proof<BabyBearPoseidon2Config>],
20    child_is_agg: bool,
21) -> AirProvingContext<CpuBackend<BabyBearPoseidon2Config>> {
22    let num_proofs = proofs.len();
23    debug_assert!((1..=2).contains(&num_proofs));
24
25    let rows_per_proof = proofs
26        .iter()
27        .map(|proof| {
28            if child_is_agg {
29                1
30            } else {
31                count_cached_commitments(proof) + 1
32            }
33        })
34        .collect::<Vec<_>>();
35    let num_rows = rows_per_proof.iter().sum::<usize>();
36    let height = num_rows.next_power_of_two();
37    let width = InputCommitCols::<u8>::width();
38    let mut trace = vec![F::ZERO; height * width];
39    let mut row_idx = 0usize;
40
41    for (proof_idx, proof) in proofs.iter().enumerate() {
42        let initial_commit = if child_is_agg {
43            let child_pvs: &DeferralAggregationPvs<F> =
44                proof.public_values[DEF_AGG_PVS_AIR_ID].as_slice().borrow();
45            child_pvs.merkle_commit
46        } else {
47            let child_pvs: &DeferralCircuitPvs<F> = proof.public_values[DEF_CIRCUIT_PVS_AIR_ID]
48                .as_slice()
49                .borrow();
50            child_pvs.input_commit
51        };
52        let cached_rows = if child_is_agg {
53            Vec::new()
54        } else {
55            collect_cached_rows(proof)
56        };
57
58        let mut capacity = [F::ZERO; DIGEST_SIZE];
59
60        for (row_in_proof, (air_idx, cached_idx, current_commit)) in
61            std::iter::once((0usize, 0usize, initial_commit))
62                .chain(cached_rows.into_iter())
63                .enumerate()
64        {
65            let cols: &mut InputCommitCols<F> =
66                trace[row_idx * width..(row_idx + 1) * width].borrow_mut();
67            cols.is_valid = F::ONE;
68            cols.is_first = F::from_bool(row_in_proof == 0);
69            cols.proof_idx = F::from_usize(proof_idx);
70            cols.row_in_proof_idx = F::from_usize(row_in_proof);
71            cols.has_verifier_pvs = F::from_bool(child_is_agg);
72            cols.air_idx = F::from_usize(air_idx);
73            cols.cached_idx = F::from_usize(cached_idx);
74            cols.current_commit = current_commit;
75
76            if child_is_agg {
77                cols.res_left = [F::ZERO; DIGEST_SIZE];
78                cols.res_right = [F::ZERO; DIGEST_SIZE];
79            } else {
80                let (res_left, res_right) =
81                    poseidon2_compress_with_capacity(cols.current_commit, capacity);
82                cols.res_left = res_left;
83                cols.res_right = res_right;
84                capacity = res_right;
85            }
86            row_idx += 1;
87        }
88    }
89
90    AirProvingContext::simple_no_pis(RowMajorMatrix::new(trace, width))
91}
92
93fn count_cached_commitments(proof: &Proof<BabyBearPoseidon2Config>) -> usize {
94    proof
95        .trace_vdata
96        .iter()
97        .flatten()
98        .map(|vd| vd.cached_commitments.len())
99        .sum()
100}
101
102fn collect_cached_rows(proof: &Proof<BabyBearPoseidon2Config>) -> Vec<CachedCommitRow<F>> {
103    proof
104        .trace_vdata
105        .iter()
106        .enumerate()
107        .flat_map(|(air_idx, vdata)| {
108            vdata.iter().flat_map(move |vd| {
109                vd.cached_commitments
110                    .iter()
111                    .copied()
112                    .enumerate()
113                    .map(move |(cached_idx, cached_commit)| (air_idx, cached_idx, cached_commit))
114            })
115        })
116        .collect()
117}