openvm_continuations/circuit/deferral/inner/input/
trace.rs1use 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}