openvm_continuations/circuit/deferral/hook/onion/
trace.rs

1use std::borrow::BorrowMut;
2
3use openvm_circuit::arch::POSEIDON2_WIDTH;
4use openvm_cpu_backend::CpuBackend;
5use openvm_stark_backend::prover::AirProvingContext;
6use openvm_stark_sdk::config::baby_bear_poseidon2::{
7    poseidon2_compress_with_capacity, BabyBearPoseidon2Config, DIGEST_SIZE, F,
8};
9use p3_field::PrimeCharacteristicRing;
10use p3_matrix::dense::RowMajorMatrix;
11
12use crate::{
13    circuit::deferral::hook::onion::air::OnionHashCols, utils::digests_to_poseidon2_input,
14};
15
16pub type IoCommit = ([F; DIGEST_SIZE], [F; DIGEST_SIZE]);
17
18pub struct OnionTraceCtx {
19    pub proving_ctx: AirProvingContext<CpuBackend<BabyBearPoseidon2Config>>,
20    pub poseidon2_inputs: Vec<[F; POSEIDON2_WIDTH]>,
21    pub input_onion: [F; DIGEST_SIZE],
22    pub output_onion: [F; DIGEST_SIZE],
23}
24
25pub fn generate_proving_ctx(
26    def_circuit_commit: [F; DIGEST_SIZE],
27    io_commits: Vec<IoCommit>,
28) -> OnionTraceCtx {
29    let num_commits = io_commits.len();
30    debug_assert!(num_commits > 0);
31
32    let width = OnionHashCols::<u8>::width();
33    let height = (num_commits + 1).next_power_of_two();
34    let mut trace = vec![F::ZERO; height * width];
35    let mut poseidon2_inputs = Vec::with_capacity(2 * num_commits);
36
37    let mut current_input_onion = def_circuit_commit;
38    let mut current_output_onion = [F::ZERO; DIGEST_SIZE];
39
40    for (row_idx, (input_commit, output_commit)) in io_commits.iter().copied().enumerate() {
41        let cols: &mut OnionHashCols<F> =
42            trace[row_idx * width..(row_idx + 1) * width].borrow_mut();
43        cols.row_idx = F::from_usize(row_idx);
44        cols.is_valid = F::ONE;
45        cols.is_first = F::from_bool(row_idx == 0);
46        cols.input_commit = input_commit;
47        cols.output_commit = output_commit;
48        cols.input_onion = current_input_onion;
49        cols.output_onion = current_output_onion;
50
51        poseidon2_inputs.push(digests_to_poseidon2_input(
52            current_input_onion,
53            input_commit,
54        ));
55        poseidon2_inputs.push(digests_to_poseidon2_input(
56            current_output_onion,
57            output_commit,
58        ));
59
60        current_input_onion = poseidon2_compress_with_capacity(current_input_onion, input_commit).0;
61        current_output_onion =
62            poseidon2_compress_with_capacity(current_output_onion, output_commit).0;
63    }
64
65    // First invalid row stores the final onions so the final valid row can constrain its
66    // transition.
67    let first_invalid_row = num_commits;
68    let cols: &mut OnionHashCols<F> =
69        trace[first_invalid_row * width..(first_invalid_row + 1) * width].borrow_mut();
70    cols.row_idx = F::from_usize(first_invalid_row);
71    cols.input_onion = current_input_onion;
72    cols.output_onion = current_output_onion;
73
74    for row_idx in (first_invalid_row + 1)..height {
75        let cols: &mut OnionHashCols<F> =
76            trace[row_idx * width..(row_idx + 1) * width].borrow_mut();
77        cols.row_idx = F::from_usize(row_idx);
78    }
79
80    OnionTraceCtx {
81        proving_ctx: AirProvingContext::simple_no_pis(RowMajorMatrix::new(trace, width)),
82        poseidon2_inputs,
83        input_onion: current_input_onion,
84        output_onion: current_output_onion,
85    }
86}