openvm_recursion_circuit/proof_shape/pvs/
trace.rs

1use std::borrow::BorrowMut;
2
3use openvm_stark_backend::proof::Proof;
4use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, F};
5use p3_field::PrimeCharacteristicRing;
6use p3_matrix::dense::RowMajorMatrix;
7
8use crate::{proof_shape::pvs::air::PublicValuesCols, system::Preflight, tracegen::RowMajorChip};
9
10pub struct PublicValuesTraceGenerator;
11
12impl RowMajorChip<F> for PublicValuesTraceGenerator {
13    type Ctx<'a> = (&'a [Proof<BabyBearPoseidon2Config>], &'a [Preflight]);
14
15    #[tracing::instrument(level = "trace", skip_all)]
16    fn generate_trace(
17        &self,
18        ctx: &Self::Ctx<'_>,
19        required_height: Option<usize>,
20    ) -> Option<RowMajorMatrix<F>> {
21        let (proofs, preflights) = ctx;
22        let num_valid_rows = proofs
23            .iter()
24            .map(|proof| {
25                proof
26                    .public_values
27                    .iter()
28                    .fold(0usize, |acc, per_air| acc + per_air.len())
29            })
30            .sum();
31        let height = if let Some(height) = required_height {
32            if height < num_valid_rows {
33                return None;
34            }
35            height
36        } else {
37            num_valid_rows.next_power_of_two()
38        };
39        let width = PublicValuesCols::<u8>::width();
40
41        debug_assert_eq!(proofs.len(), preflights.len());
42
43        let mut trace = vec![F::ZERO; height * width];
44        let mut chunks = trace.chunks_exact_mut(width);
45
46        for (proof_idx, (proof, preflight)) in proofs.iter().zip(preflights.iter()).enumerate() {
47            let mut row_idx = 0usize;
48
49            for ((air_idx, pvs), &starting_tidx) in proof
50                .public_values
51                .iter()
52                .enumerate()
53                .filter(|(_, per_air)| !per_air.is_empty())
54                .zip(&preflight.proof_shape.pvs_tidx)
55            {
56                let mut tidx = starting_tidx;
57
58                for (pv_idx, pv) in pvs.iter().enumerate() {
59                    let chunk = chunks.next().unwrap();
60                    let cols: &mut PublicValuesCols<F> = chunk.borrow_mut();
61
62                    cols.is_valid = F::ONE;
63
64                    cols.proof_idx = F::from_usize(proof_idx);
65                    cols.air_idx = F::from_usize(air_idx);
66                    cols.pv_idx = F::from_usize(pv_idx);
67
68                    cols.is_first_in_air = F::from_bool(pv_idx == 0);
69                    cols.is_first_in_proof = F::from_bool(row_idx == 0);
70
71                    cols.tidx = F::from_usize(tidx);
72                    cols.value = *pv;
73
74                    row_idx += 1;
75                    tidx += 1;
76                }
77            }
78        }
79
80        Some(RowMajorMatrix::new(trace, width))
81    }
82}