openvm_continuations/circuit/root/memory/
trace.rs

1use std::borrow::BorrowMut;
2
3use openvm_circuit::{
4    arch::POSEIDON2_WIDTH,
5    system::memory::{dimensions::MemoryDimensions, merkle::public_values::PUBLIC_VALUES_AS},
6};
7use openvm_cpu_backend::CpuBackend;
8use openvm_stark_backend::{
9    p3_util::log2_strict_usize, prover::AirProvingContext, StarkProtocolConfig,
10};
11use openvm_stark_sdk::config::baby_bear_poseidon2::{
12    poseidon2_compress_with_capacity, DIGEST_SIZE, F,
13};
14use p3_field::PrimeCharacteristicRing;
15use p3_matrix::dense::RowMajorMatrix;
16
17use crate::{circuit::root::memory::UserPvsInMemoryCols, utils::digests_to_poseidon2_input};
18
19pub fn generate_proving_input<SC: StarkProtocolConfig<F = F>>(
20    user_pv_commit: [F; DIGEST_SIZE],
21    merkle_proof: &[[F; DIGEST_SIZE]],
22    memory_dimensions: MemoryDimensions,
23    num_user_pvs: usize,
24) -> (AirProvingContext<CpuBackend<SC>>, Vec<[F; POSEIDON2_WIDTH]>) {
25    let merkle_proof_len = merkle_proof.len();
26    let num_layers = merkle_proof_len + 1;
27    let height = num_layers.next_power_of_two();
28    let width = UserPvsInMemoryCols::<u8>::width();
29
30    let mut trace = vec![F::ZERO; height * width];
31    let mut chunks = trace.chunks_exact_mut(width);
32    let mut current = user_pv_commit;
33
34    /*
35     * We can determine the public values' location in memory (and thus location in
36     * the memory merkle tree) from PUBLIC_VALUES_AS, the memory dimensions, and the
37     * number of user public values.
38     */
39    let pv_start_idx = memory_dimensions.label_to_index((PUBLIC_VALUES_AS, 0));
40    let pv_height = log2_strict_usize(num_user_pvs / DIGEST_SIZE);
41    let merkle_path_branch_bits = pv_start_idx >> pv_height;
42    let mut current_branch_bits = 0;
43
44    let mut poseidon2_compress_inputs = Vec::with_capacity(merkle_proof_len);
45
46    for (i, &sibling) in merkle_proof.iter().enumerate() {
47        let chunk = chunks.next().unwrap();
48        let cols: &mut UserPvsInMemoryCols<F> = chunk.borrow_mut();
49        let is_right_child = merkle_path_branch_bits & (1 << i) != 0;
50        current_branch_bits += (is_right_child as usize) << i;
51
52        cols.is_valid = if i == 0 { F::TWO } else { F::ONE };
53        cols.is_right_child = F::from_bool(is_right_child);
54        cols.node_commit = current;
55        cols.sibling = sibling;
56        cols.row_idx_exp_2 = F::from_usize(1 << i);
57        cols.merkle_path_branch_bits = F::from_usize(current_branch_bits);
58
59        let left = if is_right_child { sibling } else { current };
60        let right = if is_right_child { current } else { sibling };
61        current = poseidon2_compress_with_capacity(left, right).0;
62        poseidon2_compress_inputs.push(digests_to_poseidon2_input(left, right));
63    }
64
65    let last_chunk = chunks.next().unwrap();
66    let last_row: &mut UserPvsInMemoryCols<F> = last_chunk.borrow_mut();
67    last_row.is_valid = F::ONE;
68    last_row.node_commit = current;
69    last_row.row_idx_exp_2 = F::from_usize(1 << merkle_proof_len);
70    last_row.merkle_path_branch_bits = F::from_usize(current_branch_bits);
71
72    (
73        AirProvingContext::simple_no_pis(RowMajorMatrix::new(trace, width)),
74        poseidon2_compress_inputs,
75    )
76}