openvm_verify_stark_circuit/output/
air.rs1use std::{array::from_fn, borrow::Borrow};
2
3use itertools::{fold, Itertools};
4use openvm_circuit_primitives::{
5 utils::assert_array_eq, AlignedBorrow, ColumnsAir, StructReflection, StructReflectionHelper,
6};
7use openvm_continuations::utils::digests_to_poseidon2_input;
8use openvm_deferral_circuit::canonicity::{CanonicityAuxCols, CanonicitySubAir};
9use openvm_recursion_circuit::{
10 bus::{Poseidon2PermuteBus, Poseidon2PermuteMessage},
11 prelude::DIGEST_SIZE,
12 primitives::bus::{RangeCheckerBus, RangeCheckerBusMessage},
13};
14use openvm_stark_backend::{interaction::InteractionBuilder, PartitionedBaseAir};
15use p3_air::{Air, AirBuilder, BaseAir, BaseAirWithPublicValues};
16use p3_field::{PrimeCharacteristicRing, PrimeField32};
17use p3_matrix::Matrix;
18
19use crate::bus::{OutputCommitBus, OutputCommitMessage, OutputValBus, OutputValMessage};
20
21pub(crate) const F_NUM_BYTES: usize = 4;
22pub(crate) const VALS_IN_DIGEST: usize = exact_div_or_panic(DIGEST_SIZE, F_NUM_BYTES);
23
24const fn exact_div_or_panic(a: usize, b: usize) -> usize {
25 assert!(b != 0 && a.is_multiple_of(b), "non-exact division");
26 a / b
27}
28
29#[repr(C)]
30#[derive(AlignedBorrow, StructReflection)]
31pub struct DeferralOutputCommitCols<F> {
32 pub is_valid: F,
33 pub is_first: F,
34 pub row_idx: F,
35 pub output_len: F,
36
37 pub input_vals: [F; DIGEST_SIZE],
38 pub res_left: [F; DIGEST_SIZE],
39 pub res_right: [F; DIGEST_SIZE],
40
41 pub canonicity_aux: [CanonicityAuxCols<F>; VALS_IN_DIGEST],
42}
43
44#[derive(Debug, ColumnsAir)]
45#[columns_via(DeferralOutputCommitCols<u8>)]
46pub struct DeferralOutputCommitAir {
47 pub poseidon2_bus: Poseidon2PermuteBus,
48 pub range_bus: RangeCheckerBus,
49 pub output_val_bus: OutputValBus,
50 pub output_commit_bus: OutputCommitBus,
51
52 pub def_idx: usize,
53}
54
55impl<F> BaseAir<F> for DeferralOutputCommitAir {
56 fn width(&self) -> usize {
57 DeferralOutputCommitCols::<u8>::width()
58 }
59}
60impl<F> BaseAirWithPublicValues<F> for DeferralOutputCommitAir {}
61impl<F> PartitionedBaseAir<F> for DeferralOutputCommitAir {}
62
63impl<AB: AirBuilder + InteractionBuilder> Air<AB> for DeferralOutputCommitAir
64where
65 AB::F: PrimeField32,
66{
67 fn eval(&self, builder: &mut AB) {
68 let main = builder.main();
69 let local = main.row_slice(0).expect("row 0 present");
70 let next = main.row_slice(1).expect("row 1 present");
71
72 let local: &DeferralOutputCommitCols<AB::Var> = (*local).borrow();
73 let next: &DeferralOutputCommitCols<AB::Var> = (*next).borrow();
74
75 let is_transition = next.is_valid - next.is_first;
76 let is_last = local.is_valid - is_transition.clone();
77
78 builder.assert_bool(local.is_valid);
84 builder.when_first_row().assert_one(local.is_valid);
85 builder
86 .when_transition()
87 .assert_bool(local.is_valid - next.is_valid);
88
89 builder.when_first_row().assert_one(local.is_first);
90 builder.when_transition().assert_zero(next.is_first);
91
92 builder.when_first_row().assert_zero(local.row_idx);
93 builder
94 .when_transition()
95 .assert_eq(local.row_idx + AB::Expr::ONE, next.row_idx);
96
97 let mut initial_state = [AB::Expr::ZERO; DIGEST_SIZE];
103 initial_state[0] = AB::Expr::from_usize(self.def_idx);
104 initial_state[1] = local.output_len.into();
105
106 assert_array_eq(
107 &mut builder.when_first_row(),
108 local.input_vals,
109 initial_state,
110 );
111
112 builder
113 .when(is_transition.clone())
114 .assert_eq(local.output_len, next.output_len);
115 builder.when(is_last.clone()).assert_eq(
116 local.output_len,
117 local.row_idx * AB::Expr::from_usize(DIGEST_SIZE),
118 );
119
120 let next_f: [_; VALS_IN_DIGEST] = local
125 .input_vals
126 .chunks(F_NUM_BYTES)
127 .map(|c| {
128 fold(c.iter().enumerate(), AB::Expr::ZERO, |acc, (i, byte)| {
129 acc + (AB::Expr::from_usize(1 << (i * 8)) * (*byte).into())
130 })
131 })
132 .collect_array()
133 .unwrap();
134
135 for byte in local.input_vals {
136 self.range_bus.lookup_key(
137 builder,
138 RangeCheckerBusMessage {
139 value: byte.into(),
140 max_bits: AB::Expr::from_u8(8),
141 },
142 local.is_valid - local.is_first,
143 );
144 }
145
146 self.output_val_bus.receive(
147 builder,
148 OutputValMessage {
149 values: next_f,
150 idx: local.row_idx - AB::Expr::ONE,
151 },
152 local.is_valid - local.is_first,
153 );
154
155 let rcs = local
160 .input_vals
161 .chunks(F_NUM_BYTES)
162 .zip(local.canonicity_aux)
163 .map(|(x, aux)| {
164 CanonicitySubAir.assert_canonicity(
165 builder,
166 x,
167 &aux,
168 local.is_valid - local.is_first,
169 )
170 })
171 .collect_vec();
172
173 for rc in rcs {
174 self.range_bus.lookup_key(
175 builder,
176 RangeCheckerBusMessage {
177 value: rc,
178 max_bits: AB::Expr::from_u8(8),
179 },
180 local.is_valid - local.is_first,
181 );
182 }
183
184 let next_capacity = from_fn(|i| is_transition.clone() * local.res_right[i]);
189 self.poseidon2_bus.lookup_key(
190 builder,
191 Poseidon2PermuteMessage {
192 input: digests_to_poseidon2_input(next.input_vals.map(Into::into), next_capacity),
193 output: digests_to_poseidon2_input(next.res_left, next.res_right).map(Into::into),
194 },
195 next.is_valid,
196 );
197
198 self.output_commit_bus.send(
199 builder,
200 OutputCommitMessage {
201 commit: local.res_left,
202 },
203 is_last,
204 );
205 }
206}