openvm_continuations/circuit/deferral/hook/onion/
air.rs

1use std::borrow::Borrow;
2
3use openvm_circuit_primitives::{utils::not, ColumnsAir, StructReflection, StructReflectionHelper};
4use openvm_recursion_circuit::{
5    bus::{Poseidon2CompressBus, Poseidon2CompressMessage},
6    prelude::DIGEST_SIZE,
7    utils::assert_zeros,
8};
9use openvm_recursion_circuit_derive::AlignedBorrow;
10use openvm_stark_backend::{
11    interaction::InteractionBuilder, BaseAirWithPublicValues, PartitionedBaseAir,
12};
13use p3_air::{Air, AirBuilder, AirBuilderWithPublicValues, BaseAir};
14use p3_matrix::Matrix;
15
16use crate::{
17    circuit::deferral::hook::bus::{
18        DefCircuitCommitBus, DefCircuitCommitMessage, IoCommitBus, IoCommitMessage, OnionResultBus,
19        OnionResultMessage,
20    },
21    utils::digests_to_poseidon2_input,
22};
23
24#[repr(C)]
25#[derive(AlignedBorrow, StructReflection)]
26pub struct OnionHashCols<F> {
27    pub row_idx: F,
28    pub is_valid: F,
29    pub is_first: F,
30
31    pub input_commit: [F; DIGEST_SIZE],
32    pub output_commit: [F; DIGEST_SIZE],
33
34    pub input_onion: [F; DIGEST_SIZE],
35    pub output_onion: [F; DIGEST_SIZE],
36}
37
38#[derive(ColumnsAir)]
39#[columns_via(OnionHashCols<u8>)]
40pub struct OnionHashAir {
41    pub poseidon2_bus: Poseidon2CompressBus,
42    pub def_circuit_commit_bus: DefCircuitCommitBus,
43    pub io_commit_bus: IoCommitBus,
44    pub onion_res_bus: OnionResultBus,
45}
46
47impl<F> BaseAir<F> for OnionHashAir {
48    fn width(&self) -> usize {
49        OnionHashCols::<u8>::width()
50    }
51}
52impl<F> BaseAirWithPublicValues<F> for OnionHashAir {}
53impl<F> PartitionedBaseAir<F> for OnionHashAir {}
54
55impl<AB: AirBuilder + InteractionBuilder + AirBuilderWithPublicValues> Air<AB> for OnionHashAir {
56    fn eval(&self, builder: &mut AB) {
57        let main = builder.main();
58        let (local, next) = (
59            main.row_slice(0).expect("window should have two elements"),
60            main.row_slice(1).expect("window should have two elements"),
61        );
62        let local: &OnionHashCols<AB::Var> = (*local).borrow();
63        let next: &OnionHashCols<AB::Var> = (*next).borrow();
64
65        /*
66         * Base constraints to ensure that all the valid rows are at the beginning,
67         * and that there is at least one valid and one invalid row. The latter is
68         * important because the final onion values are read from the first invalid row.
69         */
70        builder.assert_bool(local.is_valid);
71        builder.when_first_row().assert_one(local.is_valid);
72        builder
73            .when_transition()
74            .assert_bool(local.is_valid - next.is_valid);
75        builder.when_last_row().assert_zero(local.is_valid);
76
77        builder.when_first_row().assert_one(local.is_first);
78        builder.when_transition().assert_zero(next.is_first);
79
80        builder.when_first_row().assert_zero(local.row_idx);
81        builder
82            .when_transition()
83            .assert_one(next.row_idx - local.row_idx);
84
85        /*
86         * On the first row we want input_onion to initially be def_circuit_commit and
87         * output_onion to be all zeroes.
88         */
89        assert_zeros(&mut builder.when(local.is_first), local.output_onion);
90        self.def_circuit_commit_bus.receive(
91            builder,
92            DefCircuitCommitMessage {
93                def_circuit_commit: local.input_onion,
94            },
95            local.is_first,
96        );
97
98        /*
99         * On valid rows we want to receive the input and output commit values and
100         * hash them with the current input and output onions. We send the final
101         * onion values on the transition from the last valid row to the first invalid row.
102         */
103        self.io_commit_bus.receive(
104            builder,
105            IoCommitMessage {
106                idx: local.row_idx,
107                input_commit: local.input_commit,
108                output_commit: local.output_commit,
109            },
110            local.is_valid,
111        );
112
113        self.poseidon2_bus.lookup_key(
114            builder,
115            Poseidon2CompressMessage {
116                input: digests_to_poseidon2_input(local.input_onion, local.input_commit),
117                output: next.input_onion,
118            },
119            local.is_valid,
120        );
121
122        self.poseidon2_bus.lookup_key(
123            builder,
124            Poseidon2CompressMessage {
125                input: digests_to_poseidon2_input(local.output_onion, local.output_commit),
126                output: next.output_onion,
127            },
128            local.is_valid,
129        );
130
131        self.onion_res_bus.send(
132            builder,
133            OnionResultMessage {
134                input_onion: next.input_onion,
135                output_onion: next.output_onion,
136                num_elements: next.row_idx,
137            },
138            local.is_valid * not(next.is_valid),
139        );
140    }
141}