openvm_continuations/circuit/root/memory/
air.rs

1use std::borrow::Borrow;
2
3use openvm_circuit::system::memory::{
4    dimensions::MemoryDimensions, merkle::public_values::PUBLIC_VALUES_AS,
5};
6use openvm_circuit_primitives::{ColumnsAir, StructReflection, StructReflectionHelper, SubAir};
7use openvm_recursion_circuit::bus::Poseidon2CompressBus;
8use openvm_recursion_circuit_derive::AlignedBorrow;
9use openvm_stark_backend::{
10    interaction::InteractionBuilder, p3_util::log2_strict_usize, BaseAirWithPublicValues,
11    PartitionedBaseAir,
12};
13use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
14use p3_air::{Air, AirBuilder, BaseAir};
15use p3_field::{Field, PrimeCharacteristicRing};
16use p3_matrix::Matrix;
17
18use crate::circuit::{
19    root::bus::{MemoryMerkleCommitBus, MemoryMerkleCommitMessage},
20    subair::{
21        MerklePathRowView, MerklePathSubAir, MerklePathSubAirContext, MerkleRootBus,
22        MerkleRootMessage,
23    },
24};
25
26#[repr(C)]
27#[derive(AlignedBorrow, StructReflection)]
28pub struct UserPvsInMemoryCols<F> {
29    // 0 for invalid, 1 for valid, 2 for valid first row
30    pub is_valid: F,
31    pub is_right_child: F,
32    pub node_commit: [F; DIGEST_SIZE],
33    pub sibling: [F; DIGEST_SIZE],
34
35    // 2^row_idx, used to (a) constrain merkle proof height and (b) accumulate
36    // merkle_path_branch_bits
37    pub row_idx_exp_2: F,
38    pub merkle_path_branch_bits: F,
39}
40
41#[derive(ColumnsAir)]
42#[columns_via(UserPvsInMemoryCols<u8>)]
43pub struct UserPvsInMemoryAir {
44    pub merkle_path_subair: MerklePathSubAir,
45    pub merkle_root_bus: MerkleRootBus,
46    pub memory_merkle_commit_bus: MemoryMerkleCommitBus,
47}
48
49impl UserPvsInMemoryAir {
50    pub fn new(
51        poseidon2_compress_bus: Poseidon2CompressBus,
52        merkle_root_bus: MerkleRootBus,
53        memory_merkle_commit_bus: MemoryMerkleCommitBus,
54        memory_dimensions: MemoryDimensions,
55        num_user_pvs: usize,
56    ) -> Self {
57        assert!(memory_dimensions.addr_space_height > 1);
58        let pv_start_idx = memory_dimensions.label_to_index((PUBLIC_VALUES_AS, 0));
59        let pv_height = log2_strict_usize(num_user_pvs / DIGEST_SIZE);
60        let merkle_path_branch_bits = u32::try_from(pv_start_idx >> pv_height)
61            .expect("merkle_path_branch_bits must fit in u32");
62        let expected_proof_len = memory_dimensions.overall_height() - pv_height;
63        Self {
64            merkle_path_subair: MerklePathSubAir::new(
65                poseidon2_compress_bus,
66                expected_proof_len,
67                merkle_path_branch_bits,
68            ),
69            merkle_root_bus,
70            memory_merkle_commit_bus,
71        }
72    }
73}
74
75impl<F> BaseAir<F> for UserPvsInMemoryAir {
76    fn width(&self) -> usize {
77        UserPvsInMemoryCols::<u8>::width()
78    }
79}
80impl<F> BaseAirWithPublicValues<F> for UserPvsInMemoryAir {}
81impl<F> PartitionedBaseAir<F> for UserPvsInMemoryAir {}
82
83impl<AB: AirBuilder + InteractionBuilder> Air<AB> for UserPvsInMemoryAir {
84    fn eval(&self, builder: &mut AB) {
85        let main = builder.main();
86        let (local, next) = (
87            main.row_slice(0).expect("window should have two elements"),
88            main.row_slice(1).expect("window should have two elements"),
89        );
90        let local: &UserPvsInMemoryCols<AB::Var> = (*local).borrow();
91        let next: &UserPvsInMemoryCols<AB::Var> = (*next).borrow();
92
93        /*
94         * Receive the user public values commit on the first row. The first DIGEST_SIZE
95         * elements of perm state are this merkle tree node's commit. We do not receive
96         * the number of rows here.
97         */
98        self.merkle_root_bus.receive(
99            builder,
100            MerkleRootMessage {
101                merkle_root: local.node_commit.map(Into::into),
102                idx: AB::Expr::ZERO,
103                num_rows_or_zero: AB::Expr::ZERO,
104            },
105            local.is_valid * (local.is_valid - AB::F::ONE) * AB::F::TWO.inverse(),
106        );
107
108        self.merkle_path_subair.eval(
109            builder,
110            (
111                MerklePathSubAirContext {
112                    local: MerklePathRowView {
113                        is_valid: &local.is_valid,
114                        is_right_child: &local.is_right_child,
115                        node_commit: &local.node_commit,
116                        sibling: &local.sibling,
117                        row_idx_exp_2: &local.row_idx_exp_2,
118                        merkle_path_branch_bits: &local.merkle_path_branch_bits,
119                    },
120                    next: MerklePathRowView {
121                        is_valid: &next.is_valid,
122                        is_right_child: &next.is_right_child,
123                        node_commit: &next.node_commit,
124                        sibling: &next.sibling,
125                        row_idx_exp_2: &next.row_idx_exp_2,
126                        merkle_path_branch_bits: &next.merkle_path_branch_bits,
127                    },
128                },
129                AB::Expr::ZERO,
130                AB::Expr::ZERO,
131                AB::Expr::ZERO,
132            ),
133        );
134
135        /*
136         * Receive the final memory merkle root on the last valid row.
137         */
138        self.memory_merkle_commit_bus.receive(
139            builder,
140            MemoryMerkleCommitMessage {
141                merkle_root: local.node_commit,
142            },
143            local.is_valid * (AB::Expr::ONE - next.is_valid).square(),
144        );
145    }
146}