openvm_continuations/circuit/root/commit/
air.rs

1use std::borrow::Borrow;
2
3use itertools::Itertools;
4use openvm_circuit_primitives::{encoder::Encoder, utils::assert_array_eq, ColumnsAir, SubAir};
5use openvm_recursion_circuit::{bus::Poseidon2CompressBus, utils::assert_zeros};
6use openvm_stark_backend::{
7    interaction::InteractionBuilder, BaseAirWithPublicValues, PartitionedBaseAir,
8};
9use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
10use p3_air::{Air, AirBuilder, AirBuilderWithPublicValues, BaseAir};
11use p3_field::PrimeCharacteristicRing;
12use p3_matrix::Matrix;
13
14pub use crate::circuit::subair::MerkleTreeCols;
15use crate::circuit::subair::{MerkleRootBus, MerkleTreeInternalBus, MerkleTreeSubAir};
16
17pub(super) const MAX_ENCODER_DEGREE: u32 = 3;
18
19/**
20 * Builds a binary Merkle tree to decommit and expose or emit the raw user public values.
21 * Constrains that:
22 * - leaf nodes read single digests from encoder-selected exposed public values, compress them
23 *   with zeros, and compute leaf hashes
24 * - internal nodes receive children from an internal permutation bus
25 * - root commitment is sent to `MerkleRootBus`
26 */
27pub struct UserPvsCommitAir {
28    pub subair: MerkleTreeSubAir,
29    encoder: Encoder,
30    num_user_pvs: usize,
31}
32// No columns provided: width is dynamic — `MerkleTreeCols` followed by encoder flags whose count
33// depends on `num_user_pvs`.
34impl ColumnsAir for UserPvsCommitAir {}
35
36impl UserPvsCommitAir {
37    pub fn new(
38        poseidon2_compress_bus: Poseidon2CompressBus,
39        merkle_root_bus: MerkleRootBus,
40        merkle_tree_internal_bus: MerkleTreeInternalBus,
41        num_user_pvs: usize,
42    ) -> Self {
43        // Each leaf consumes `DIGEST_SIZE` public values, which are compressed with zeros
44        // to compute the leaf hash. We require at least one leaf, and a full binary tree.
45        assert!(num_user_pvs >= DIGEST_SIZE);
46        assert!(num_user_pvs.is_multiple_of(DIGEST_SIZE));
47        assert!((num_user_pvs / DIGEST_SIZE).is_power_of_two());
48        let encoder = Encoder::new(num_user_pvs / DIGEST_SIZE, MAX_ENCODER_DEGREE, true);
49
50        UserPvsCommitAir {
51            subair: MerkleTreeSubAir::new(
52                poseidon2_compress_bus,
53                merkle_root_bus,
54                merkle_tree_internal_bus,
55                0,
56                false,
57            ),
58            encoder,
59            num_user_pvs,
60        }
61    }
62}
63
64impl<F> BaseAir<F> for UserPvsCommitAir {
65    fn width(&self) -> usize {
66        MerkleTreeCols::<u8>::width() + self.encoder.width()
67    }
68}
69impl<F> BaseAirWithPublicValues<F> for UserPvsCommitAir {
70    fn num_public_values(&self) -> usize {
71        self.num_user_pvs
72    }
73}
74impl<F> PartitionedBaseAir<F> for UserPvsCommitAir {}
75
76impl<AB: AirBuilder + InteractionBuilder + AirBuilderWithPublicValues> Air<AB>
77    for UserPvsCommitAir
78{
79    fn eval(&self, builder: &mut AB) {
80        let main = builder.main();
81        let (local, next) = (
82            main.row_slice(0).expect("window should have two elements"),
83            main.row_slice(1).expect("window should have two elements"),
84        );
85
86        let const_width = MerkleTreeCols::<u8>::width();
87        let row_idx_flags = &(*local)[const_width..];
88
89        let local: &MerkleTreeCols<AB::Var> = (*local)[..const_width].borrow();
90        let next: &MerkleTreeCols<AB::Var> = (*next)[..const_width].borrow();
91
92        let num_rows = AB::F::from_usize(2 * self.num_user_pvs / DIGEST_SIZE);
93        self.subair
94            .eval(builder, (local, next, num_rows.into(), None));
95
96        /*
97         * Constrain that the left_child of each leaf node at row_idx corresponds to this
98         * AIR's public values. Leaf nodes correspond to the raw user public values, and
99         * leaf rows should be in order of their position in the public values vector. A
100         * row is a leaf node if its receive_type == 1.
101         */
102        let is_leaf = local.receive_type * (AB::Expr::TWO - local.receive_type);
103        assert_zeros(&mut builder.when(is_leaf.clone()), local.right_child);
104
105        debug_assert_eq!(self.encoder.width(), row_idx_flags.len());
106        self.encoder.eval(builder, row_idx_flags);
107        builder.assert_eq(self.encoder.is_valid::<AB>(row_idx_flags), is_leaf.clone());
108
109        let pvs = builder.public_values().iter().copied().collect_vec();
110        let mut pvs_digest = [AB::Expr::ZERO; DIGEST_SIZE];
111        for (pv_chunk_idx, pvs_chunk) in pvs.chunks(DIGEST_SIZE).enumerate() {
112            let selected = self
113                .encoder
114                .get_flag_expr::<AB>(pv_chunk_idx, row_idx_flags);
115            builder
116                .when(selected.clone())
117                .assert_eq(AB::Expr::from_usize(pv_chunk_idx), local.row_idx);
118            for digest_idx in 0..DIGEST_SIZE {
119                pvs_digest[digest_idx] += selected.clone() * pvs_chunk[digest_idx].into();
120            }
121        }
122
123        assert_array_eq(
124            builder,
125            pvs_digest,
126            local.left_child.map(|x| x * is_leaf.clone()),
127        );
128    }
129}