openvm_continuations/circuit/root/commit/
air.rs1use 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
19pub struct UserPvsCommitAir {
28 pub subair: MerkleTreeSubAir,
29 encoder: Encoder,
30 num_user_pvs: usize,
31}
32impl 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 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 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}