openvm_verify_stark_circuit/output/
trace.rs1use 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}