openvm_recursion_circuit/proof_shape/pvs/
trace.rs1use 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}