1use std::ops::Deref;
5
6use binius_field::{Field, PackedField};
7use binius_math::FieldBuffer;
8use tracing::instrument;
9
10use crate::merkle_channel::MerkleIPProverChannel;
11
12pub trait ProxTestOracleProver<F: Field> {
15 type Commitment;
17
18 fn open_queries<Channel>(&self, indices: &[Channel::Word], channel: &mut Channel)
21 where
22 Channel: MerkleIPProverChannel<F, Commitment = Self::Commitment>;
23}
24
25pub struct BrakedownOracleProver<P, C, Data = Vec<P>>
29where
30 P: PackedField,
31 Data: Deref<Target = [P]>,
32{
33 codeword: FieldBuffer<P, Data>,
34 commitment: C,
35 log_lift: usize,
39}
40
41impl<P, C, Data> BrakedownOracleProver<P, C, Data>
42where
43 P: PackedField,
44 Data: Deref<Target = [P]>,
45{
46 pub const fn new(codeword: FieldBuffer<P, Data>, commitment: C, log_lift: usize) -> Self {
52 Self {
53 codeword,
54 commitment,
55 log_lift,
56 }
57 }
58}
59
60impl<F, P, C, Data> ProxTestOracleProver<F> for BrakedownOracleProver<P, C, Data>
61where
62 F: Field,
63 P: PackedField<Scalar = F>,
64 Data: Deref<Target = [P]>,
65{
66 type Commitment = C;
67
68 fn open_queries<Channel>(&self, indices: &[Channel::Word], channel: &mut Channel)
69 where
70 Channel: MerkleIPProverChannel<F, Commitment = Self::Commitment>,
71 {
72 let lifted_indices = indices
76 .iter()
77 .map(|index| index.clone() >> self.log_lift as u32)
78 .collect::<Vec<_>>();
79 channel.send_openings(&self.commitment, self.codeword.as_view(), &lifted_indices);
80 }
81}
82
83pub struct BatchBrakedownOracleProver<P, C, Data = Vec<P>>
90where
91 P: PackedField,
92 Data: Deref<Target = [P]>,
93{
94 oracles: Vec<BrakedownOracleProver<P, C, Data>>,
95}
96
97impl<P, C, Data> BatchBrakedownOracleProver<P, C, Data>
98where
99 P: PackedField,
100 Data: Deref<Target = [P]>,
101{
102 pub const fn new(oracles: Vec<BrakedownOracleProver<P, C, Data>>) -> Self {
104 Self { oracles }
105 }
106}
107
108impl<F, P, C, Data> ProxTestOracleProver<F> for BatchBrakedownOracleProver<P, C, Data>
109where
110 F: Field,
111 P: PackedField<Scalar = F>,
112 Data: Deref<Target = [P]>,
113{
114 type Commitment = C;
115
116 fn open_queries<Channel>(&self, indices: &[Channel::Word], channel: &mut Channel)
117 where
118 Channel: MerkleIPProverChannel<F, Commitment = Self::Commitment>,
119 {
120 for oracle in &self.oracles {
121 oracle.open_queries(indices, channel);
122 }
123 }
124}
125
126pub struct FRIOracleProver<F, C>
128where
129 F: Field,
130{
131 codeword: FieldBuffer<F>,
132 commitment: C,
133 coset_log_size: usize,
136}
137
138impl<F, C> FRIOracleProver<F, C>
139where
140 F: Field,
141{
142 pub const fn new(codeword: FieldBuffer<F>, commitment: C, coset_log_size: usize) -> Self {
147 Self {
148 codeword,
149 commitment,
150 coset_log_size,
151 }
152 }
153
154 const fn coset_log_size(&self) -> usize {
156 self.coset_log_size
157 }
158}
159
160impl<F, C> ProxTestOracleProver<F> for FRIOracleProver<F, C>
161where
162 F: Field,
163{
164 type Commitment = C;
165
166 fn open_queries<Channel>(&self, indices: &[Channel::Word], channel: &mut Channel)
167 where
168 Channel: MerkleIPProverChannel<F, Commitment = Self::Commitment>,
169 {
170 channel.send_openings(&self.commitment, self.codeword.as_view(), indices);
171 }
172}
173
174pub struct FRIQueryProver<F, P, C, Data = Vec<P>>
180where
181 F: Field,
182 P: PackedField<Scalar = F>,
183 Data: Deref<Target = [P]>,
184{
185 codeword_oracle: BatchBrakedownOracleProver<P, C, Data>,
186 fri_oracles: Vec<FRIOracleProver<F, C>>,
187}
188
189impl<F, P, C, Data> FRIQueryProver<F, P, C, Data>
190where
191 F: Field,
192 P: PackedField<Scalar = F>,
193 Data: Deref<Target = [P]>,
194{
195 pub const fn new(
201 codeword_oracle: BatchBrakedownOracleProver<P, C, Data>,
202 fri_oracles: Vec<FRIOracleProver<F, C>>,
203 ) -> Self {
204 Self {
205 codeword_oracle,
206 fri_oracles,
207 }
208 }
209
210 pub const fn n_oracles(&self) -> usize {
212 1 + self.fri_oracles.len()
213 }
214
215 #[instrument(skip_all, name = "fri::FRIQueryProver::prove_queries", level = "debug")]
226 pub fn prove_queries<Channel>(&self, indices: &[Channel::Word], channel: &mut Channel)
227 where
228 Channel: MerkleIPProverChannel<F, Commitment = C>,
229 {
230 self.codeword_oracle.open_queries(indices, channel);
231
232 let mut indices = indices.to_vec();
235 for fri_oracle in &self.fri_oracles {
236 indices = indices
237 .into_iter()
238 .map(|index| index >> fri_oracle.coset_log_size() as u32)
239 .collect();
240 fri_oracle.open_queries(&indices, channel);
241 }
242 }
243}