openvm_recursion_circuit/whir/folding/
air.rs

1use 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    /// Distance from the leaf layer in the folding tree.
31    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        // value = left_value + term / (2x)
89        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}