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

1use 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}