openvm_verify_stark_circuit/output/
trace.rs

1use std::{array::from_fn, borrow::BorrowMut};
2
3use openvm_circuit::arch::POSEIDON2_WIDTH;
4use openvm_continuations::utils::digests_to_poseidon2_input;
5use openvm_cpu_backend::CpuBackend;
6use openvm_deferral_circuit::canonicity::CanonicityTraceGen;
7use openvm_poseidon2_air::Permutation;
8use openvm_stark_backend::prover::{AirProvingContext, ProverBackend};
9use openvm_stark_sdk::config::baby_bear_poseidon2::{
10    poseidon2_perm, BabyBearPoseidon2Config, DIGEST_SIZE, F,
11};
12use p3_field::{PrimeCharacteristicRing, PrimeField32};
13use p3_matrix::dense::RowMajorMatrix;
14
15use crate::output::{DeferralOutputCommitCols, F_NUM_BYTES, VALS_IN_DIGEST};
16
17pub struct DeferralOutputCtx<PB: ProverBackend> {
18    pub proving_ctx: AirProvingContext<PB>,
19    pub poseidon2_inputs: Vec<[PB::Val; POSEIDON2_WIDTH]>,
20    pub range_inputs: Vec<usize>,
21    pub output_commit: [PB::Val; DIGEST_SIZE],
22}
23
24pub fn generate_proving_ctx(
25    app_exe_commit: [F; DIGEST_SIZE],
26    app_vm_commit: [F; DIGEST_SIZE],
27    user_pvs: Vec<F>,
28    def_idx: usize,
29) -> DeferralOutputCtx<CpuBackend<BabyBearPoseidon2Config>> {
30    debug_assert!(DIGEST_SIZE.is_multiple_of(F_NUM_BYTES));
31    debug_assert!(DIGEST_SIZE.is_multiple_of(VALS_IN_DIGEST));
32    debug_assert!(user_pvs.len().is_multiple_of(VALS_IN_DIGEST));
33
34    let mut input_val_rows = values_to_rows(&app_exe_commit);
35    input_val_rows.extend(values_to_rows(&app_vm_commit));
36    input_val_rows.extend(values_to_rows(&user_pvs));
37
38    let num_rows = input_val_rows.len() + 1;
39    let height = num_rows.next_power_of_two();
40    let width = DeferralOutputCommitCols::<u8>::width();
41    let mut trace = vec![F::ZERO; height * width];
42    let mut chunks = trace.chunks_exact_mut(width);
43
44    let mut poseidon2_permute_inputs = Vec::with_capacity(num_rows);
45    let mut range_inputs =
46        Vec::with_capacity(input_val_rows.len() * (DIGEST_SIZE + VALS_IN_DIGEST));
47    let output_len = input_val_rows.len() * DIGEST_SIZE;
48    let mut input_capacity = [F::ZERO; DIGEST_SIZE];
49    let mut output_commit = [F::ZERO; DIGEST_SIZE];
50    let perm = poseidon2_perm();
51
52    for row_idx in 0..height {
53        let row = chunks.next().unwrap();
54        let cols: &mut DeferralOutputCommitCols<F> = row.borrow_mut();
55        cols.row_idx = F::from_usize(row_idx);
56        if row_idx < num_rows {
57            cols.is_valid = F::ONE;
58            cols.is_first = F::from_bool(row_idx == 0);
59            cols.output_len = F::from_usize(output_len);
60
61            cols.input_vals = if row_idx == 0 {
62                let mut input = [F::ZERO; DIGEST_SIZE];
63                input[0] = F::from_usize(def_idx);
64                input[1] = F::from_usize(output_len);
65                input
66            } else {
67                let next_f = input_val_rows[row_idx - 1];
68                let input_vals = next_f_to_digest(next_f);
69                range_inputs.extend(input_vals.map(|b| b.as_canonical_u32() as usize));
70                for (bytes, aux) in input_vals
71                    .chunks_exact(F_NUM_BYTES)
72                    .zip(cols.canonicity_aux.iter_mut())
73                {
74                    let x_le = from_fn(|i| bytes[i]);
75                    let rc = CanonicityTraceGen::generate_subrow(&x_le, aux);
76                    range_inputs.push(rc as usize);
77                }
78                input_vals
79            };
80
81            let perm_input = digests_to_poseidon2_input(cols.input_vals, input_capacity);
82            poseidon2_permute_inputs.push(perm_input);
83
84            let perm_output = perm.permute(perm_input);
85            cols.res_left = perm_output[..DIGEST_SIZE].try_into().unwrap();
86            cols.res_right = perm_output[DIGEST_SIZE..].try_into().unwrap();
87
88            input_capacity = cols.res_right;
89            output_commit = cols.res_left;
90        }
91    }
92
93    DeferralOutputCtx {
94        proving_ctx: AirProvingContext::simple_no_pis(RowMajorMatrix::new(trace, width)),
95        poseidon2_inputs: poseidon2_permute_inputs,
96        range_inputs,
97        output_commit,
98    }
99}
100
101fn values_to_rows(values: &[F]) -> Vec<[F; VALS_IN_DIGEST]> {
102    values
103        .chunks_exact(VALS_IN_DIGEST)
104        .map(|chunk| chunk.try_into().unwrap())
105        .collect()
106}
107
108fn next_f_to_digest(next_f: [F; VALS_IN_DIGEST]) -> [F; DIGEST_SIZE] {
109    from_fn(|byte_idx| {
110        let f_idx = byte_idx / F_NUM_BYTES;
111        let byte_in_f = byte_idx % F_NUM_BYTES;
112        let f_u32 = next_f[f_idx].as_canonical_u32();
113        F::from_u8(f_u32.to_le_bytes()[byte_in_f])
114    })
115}