openvm_continuations/circuit/deferral/inner/def_pvs/
trace.rs1use std::{
2 borrow::{Borrow, BorrowMut},
3 iter::once,
4};
5
6use itertools::Itertools;
7use openvm_cpu_backend::CpuBackend;
8use openvm_recursion_circuit::utils::poseidon2_hash_slice;
9use openvm_stark_backend::{proof::Proof, prover::AirProvingContext};
10use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, F};
11use p3_field::{PrimeCharacteristicRing, PrimeField32};
12use p3_matrix::dense::RowMajorMatrix;
13
14use crate::circuit::deferral::{
15 inner::def_pvs::air::{DeferralAggPvsAir, DeferralAggPvsCols},
16 utils::{def_internal_compress, def_leaf_compress, def_zero_hash},
17 DeferralAggregationPvs, DeferralCircuitPvs, DEF_AGG_PVS_AIR_ID, DEF_CIRCUIT_PVS_AIR_ID,
18 MAX_DEF_AGG_MERKLE_DEPTH,
19};
20
21pub struct DeferralAggPvsTraceCtx {
22 pub proving_ctx: AirProvingContext<CpuBackend<BabyBearPoseidon2Config>>,
23 pub range_check_inputs: Vec<usize>,
24}
25
26pub fn generate_proving_ctx(
27 proofs: &[Proof<BabyBearPoseidon2Config>],
28 child_is_agg: bool,
29 child_merkle_depth: Option<usize>,
30) -> DeferralAggPvsTraceCtx {
31 let num_proofs = proofs.len();
32 let is_wrapper = child_merkle_depth.is_none();
33 debug_assert!(!is_wrapper || num_proofs == 1);
34
35 let num_rows = if is_wrapper { 1usize } else { 2usize };
36 let width = DeferralAggPvsCols::<u8>::width();
37
38 debug_assert!((1..=2).contains(&num_proofs));
39
40 let mut trace = vec![F::ZERO; num_rows * width];
41 let mut num_def_circuit_proofs = F::ZERO;
42 let mut merkle_depth = F::ZERO;
43 let mut def_idx = F::ZERO;
44
45 for (proof_idx, (proof, chunk)) in proofs.iter().zip(trace.chunks_exact_mut(width)).enumerate()
46 {
47 let cols: &mut DeferralAggPvsCols<F> = chunk.borrow_mut();
48 cols.proof_idx = F::from_usize(proof_idx);
49 cols.is_present = F::ONE;
50 cols.has_verifier_pvs = F::from_bool(child_is_agg);
51
52 if child_is_agg {
53 let child_pvs: &DeferralAggregationPvs<F> =
54 proof.public_values[DEF_AGG_PVS_AIR_ID].as_slice().borrow();
55 cols.merkle_commit = child_pvs.merkle_commit;
56 cols.child_pvs.input_commit[0] = child_pvs.num_def_circuit_proofs;
57 cols.child_pvs.input_commit[1] = child_pvs.merkle_depth;
58 cols.child_pvs.def_idx = child_pvs.def_idx;
59 num_def_circuit_proofs += child_pvs.num_def_circuit_proofs;
60 if proof_idx == 0 {
61 merkle_depth = child_pvs.merkle_depth;
62 def_idx = child_pvs.def_idx;
63 } else {
64 debug_assert_eq!(merkle_depth, child_pvs.merkle_depth);
65 debug_assert_eq!(def_idx, child_pvs.def_idx);
66 }
67 } else {
68 let child_pvs: &DeferralCircuitPvs<F> = proof.public_values[DEF_CIRCUIT_PVS_AIR_ID]
69 .as_slice()
70 .borrow();
71 let commit_values = once(child_pvs.input_commit)
72 .chain(
73 proof
74 .trace_vdata
75 .iter()
76 .flatten()
77 .flat_map(|vdata| vdata.cached_commitments.iter().copied()),
78 )
79 .flatten()
80 .collect_vec();
81 let folded_input_commit = poseidon2_hash_slice(&commit_values).0;
82 cols.child_pvs = DeferralCircuitPvs {
83 input_commit: folded_input_commit,
84 output_commit: child_pvs.output_commit,
85 def_idx: child_pvs.def_idx,
86 };
87 let (tagged_input_commit, merkle_commit) =
88 def_leaf_compress(folded_input_commit, child_pvs.output_commit);
89 cols.tagged_input_commit = tagged_input_commit;
90 cols.merkle_commit = merkle_commit;
91 num_def_circuit_proofs += F::ONE;
92 if proof_idx == 0 {
93 def_idx = child_pvs.def_idx;
94 } else {
95 debug_assert_eq!(def_idx, child_pvs.def_idx);
96 }
97 }
98 }
99
100 if num_rows == 2 && num_proofs == 1 {
101 let cols: &mut DeferralAggPvsCols<F> = trace[width..2 * width].borrow_mut();
102 cols.proof_idx = F::ONE;
103 cols.has_verifier_pvs = F::from_bool(child_is_agg);
104 for (dst, value) in cols
105 .merkle_commit
106 .iter_mut()
107 .zip(DeferralAggPvsAir::depth_encoder().get_flag_pt(child_merkle_depth.unwrap()))
108 {
109 *dst = F::from_u32(value);
110 }
111 }
112
113 let mut public_values = vec![F::ZERO; DeferralAggregationPvs::<u8>::width()];
114 let pvs: &mut DeferralAggregationPvs<F> = public_values.as_mut_slice().borrow_mut();
115
116 if is_wrapper {
117 let first_row: &DeferralAggPvsCols<F> = trace[..width].borrow();
118 pvs.merkle_commit = first_row.merkle_commit;
119 pvs.merkle_depth = merkle_depth;
120 } else {
121 let right_child = if num_proofs == 1 {
122 def_zero_hash(child_merkle_depth.unwrap() + 1)
123 } else {
124 let second_row: &DeferralAggPvsCols<F> = trace[width..2 * width].borrow();
125 second_row.merkle_commit
126 };
127 let first_row: &mut DeferralAggPvsCols<F> = trace[..width].borrow_mut();
128 let (tagged_left_merkle, merkle_commit) =
129 def_internal_compress(first_row.merkle_commit, right_child);
130 first_row.tagged_left_merkle = tagged_left_merkle;
131 pvs.merkle_commit = merkle_commit;
132 pvs.merkle_depth = merkle_depth + F::ONE;
133 }
134 pvs.num_def_circuit_proofs = num_def_circuit_proofs;
135 pvs.def_idx = def_idx;
136
137 let merkle_depth = pvs.merkle_depth.as_canonical_u32() as usize;
138 let max_depth_minus_merkle_depth = MAX_DEF_AGG_MERKLE_DEPTH
139 .checked_sub(merkle_depth)
140 .expect("deferral aggregation merkle depth exceeds max depth");
141
142 DeferralAggPvsTraceCtx {
143 proving_ctx: AirProvingContext {
144 cached_mains: vec![],
145 common_main: RowMajorMatrix::new(trace, width),
146 public_values,
147 },
148 range_check_inputs: vec![max_depth_minus_merkle_depth],
149 }
150}