openvm_recursion_circuit/whir/folding/
air.rs1use core::borrow::Borrow;
2
3use openvm_circuit_primitives::{
4 utils::assert_array_eq, ColumnsAir, StructReflection, StructReflectionHelper,
5};
6use openvm_recursion_circuit_derive::AlignedBorrow;
7use openvm_stark_backend::{
8 interaction::InteractionBuilder, BaseAirWithPublicValues, PartitionedBaseAir,
9};
10use openvm_stark_sdk::config::baby_bear_poseidon2::{D_EF, F};
11use p3_air::{Air, AirBuilder, BaseAir};
12use p3_field::{extension::BinomiallyExtendable, PrimeCharacteristicRing};
13use p3_matrix::Matrix;
14
15use crate::{
16 utils::{base_to_ext, ext_field_multiply, ext_field_subtract},
17 whir::bus::{WhirAlphaBus, WhirAlphaMessage, WhirFoldingBus, WhirFoldingBusMessage},
18};
19
20#[repr(C)]
21#[derive(AlignedBorrow, StructReflection)]
22pub struct WhirFoldingCols<T> {
23 pub is_valid: T,
24 pub proof_idx: T,
25 pub whir_round: T,
26 pub query_idx: T,
27 pub is_root: T,
28 pub coset_shift: T,
29 pub coset_idx: T,
30 pub height: T,
32 pub twiddle: T,
33 pub coset_size: T,
34 pub z_final: T,
35 pub value: [T; 4],
36 pub left_value: [T; 4],
37 pub right_value: [T; 4],
38 pub y_final: [T; 4],
39 pub alpha: [T; 4],
40}
41
42#[derive(ColumnsAir)]
43#[columns_via(WhirFoldingCols<u8>)]
44pub struct WhirFoldingAir {
45 pub alpha_bus: WhirAlphaBus,
46 pub folding_bus: WhirFoldingBus,
47 pub k: usize,
48}
49
50impl BaseAirWithPublicValues<F> for WhirFoldingAir {}
51impl PartitionedBaseAir<F> for WhirFoldingAir {}
52
53impl BaseAir<F> for WhirFoldingAir {
54 fn width(&self) -> usize {
55 WhirFoldingCols::<F>::width()
56 }
57}
58
59impl<AB: AirBuilder<F = F> + InteractionBuilder> Air<AB> for WhirFoldingAir
60where
61 <AB::Expr as PrimeCharacteristicRing>::PrimeSubfield: BinomiallyExtendable<{ D_EF }>,
62{
63 fn eval(&self, builder: &mut AB) {
64 let main = builder.main();
65
66 let local = main.row_slice(0).expect("window should have two elements");
67 let local: &WhirFoldingCols<AB::Var> = (*local).borrow();
68
69 builder.assert_bool(local.is_valid);
70 builder.assert_bool(local.is_root);
71 builder.when(local.is_root).assert_one(local.is_valid);
72 builder.when(local.is_root).assert_one(local.twiddle);
73 builder.when(local.is_root).assert_zero(local.coset_idx);
74 builder
75 .when(local.is_root)
76 .assert_eq(local.height, AB::F::from_usize(self.k));
77 builder
78 .when(local.is_root)
79 .assert_eq(local.z_final, local.coset_shift * local.coset_shift);
80 assert_array_eq(&mut builder.when(local.is_root), local.value, local.y_final);
81
82 let x = local.twiddle * local.coset_shift;
83
84 let term = ext_field_multiply::<AB::Expr>(
85 ext_field_subtract::<AB::Expr>(local.alpha, base_to_ext::<AB::Expr>(x.clone())),
86 ext_field_subtract::<AB::Expr>(local.left_value, local.right_value),
87 );
88 assert_array_eq(
90 builder,
91 ext_field_multiply::<AB::Expr>(
92 ext_field_subtract::<AB::Expr>(local.value, local.left_value),
93 base_to_ext::<AB::Expr>(x * AB::Expr::TWO),
94 ),
95 term,
96 );
97
98 self.alpha_bus.lookup_key(
99 builder,
100 local.proof_idx,
101 WhirAlphaMessage {
102 idx: local.whir_round * AB::Expr::from_usize(self.k) + local.height - AB::Expr::ONE,
103 challenge: local.alpha.map(Into::into),
104 },
105 local.is_valid,
106 );
107 self.folding_bus.receive(
108 builder,
109 local.proof_idx,
110 WhirFoldingBusMessage {
111 whir_round: local.whir_round.into(),
112 query_idx: local.query_idx.into(),
113 height: local.height - AB::Expr::ONE,
114 coset_shift: local.coset_shift.into(),
115 coset_size: AB::Expr::TWO * local.coset_size,
116 coset_idx: local.coset_idx.into(),
117 twiddle: local.twiddle.into(),
118 value: local.left_value.map(Into::into),
119 z_final: local.z_final.into(),
120 y_final: local.y_final.map(Into::into),
121 },
122 local.is_valid,
123 );
124 self.folding_bus.receive(
125 builder,
126 local.proof_idx,
127 WhirFoldingBusMessage {
128 whir_round: local.whir_round.into(),
129 query_idx: local.query_idx.into(),
130 height: local.height - AB::Expr::ONE,
131 coset_shift: local.coset_shift.into(),
132 coset_size: AB::Expr::TWO * local.coset_size,
133 coset_idx: local.coset_idx + local.coset_size,
134 twiddle: -local.twiddle.into(),
135 value: local.right_value.map(Into::into),
136 z_final: local.z_final.into(),
137 y_final: local.y_final.map(Into::into),
138 },
139 local.is_valid,
140 );
141 self.folding_bus.send(
142 builder,
143 local.proof_idx,
144 WhirFoldingBusMessage {
145 whir_round: local.whir_round.into(),
146 query_idx: local.query_idx.into(),
147 height: local.height.into(),
148 coset_shift: local.coset_shift * local.coset_shift,
149 coset_size: local.coset_size.into(),
150 coset_idx: local.coset_idx.into(),
151 twiddle: local.twiddle * local.twiddle,
152 value: local.value.map(Into::into),
153 z_final: local.z_final.into(),
154 y_final: local.y_final.map(Into::into),
155 },
156 local.is_valid - local.is_root,
157 );
158 }
159}