1use std::iter::{self, repeat_with};
5
6use binius_core::word::Word;
7use binius_field::{BinaryField, FieldOps};
8use binius_math::ntt::domain_context::GaoMateerOnTheFly;
9use binius_utils::checked_arithmetics::log2_ceil_usize;
10
11use super::{
12 batch::{BatchBrakedownOracle, BrakedownOracle, FRIOracle, ProxTestOracle, fold_coset},
13 common::FRIParams,
14 error::Error,
15};
16use crate::merkle_channel::MerkleIPVerifierChannel;
17
18pub struct FRIQueryVerifier<'a, F, E, C>
30where
31 F: BinaryField,
32{
33 params: &'a FRIParams<F>,
34 terminal_commitment: C,
36 final_challenges: &'a [E],
38 codeword_oracle: BatchBrakedownOracle<E, C>,
40 fri_oracles: Vec<FRIOracle<E, C, GaoMateerOnTheFly<F>>>,
42}
43
44impl<'a, F, E, C> FRIQueryVerifier<'a, F, E, C>
45where
46 F: BinaryField,
47 E: FieldOps<Scalar = F> + From<F>,
48 C: Clone,
49{
50 pub fn new(
51 params: &'a FRIParams<F>,
52 codeword_commitment: &C,
53 round_commitments: &[C],
54 challenges: &'a [E],
55 ) -> Self {
56 Self::new_batch(
57 params,
58 std::slice::from_ref(codeword_commitment),
59 round_commitments,
60 challenges,
61 )
62 }
63
64 pub fn new_batch(
78 params: &'a FRIParams<F>,
79 codeword_commitments: &[C],
80 round_commitments: &[C],
81 challenges: &'a [E],
82 ) -> Self {
83 assert_eq!(
84 codeword_commitments.len(),
85 params.input_oracles().len(),
86 "precondition: codeword_commitments.len() must equal params.input_oracles().len()"
87 );
88 assert_eq!(
89 round_commitments.len(),
90 params.n_oracles(),
91 "precondition: round_commitments.len() must equal params.n_oracles()"
92 );
93 assert_eq!(
94 challenges.len(),
95 params.n_fold_rounds(),
96 "precondition: challenges.len() must equal params.n_fold_rounds()"
97 );
98
99 let log_dim = params.rs_code().log_dim();
104 for spec in params.input_oracles() {
105 assert!(
106 spec.log_lift <= log_dim,
107 "precondition: input oracle dimension must not exceed the reduced code dimension"
108 );
109 }
110
111 let index_bits = params.index_bits();
114 let max_early = params
121 .input_oracles()
122 .iter()
123 .map(|spec| spec.log_early_batch_size)
124 .max()
125 .expect("input_oracles is non-empty as an invariant");
126 let max_later = params
127 .input_oracles()
128 .iter()
129 .map(|spec| spec.log_later_batch_size)
130 .max()
131 .expect("input_oracles is non-empty as an invariant");
132 let log_n_oracles = log2_ceil_usize(params.input_oracles().len());
133 let early_challenges = &challenges[..max_early];
134 let outer_challenges = challenges[max_early..max_early + log_n_oracles].to_vec();
135 let later_challenges = &challenges[max_early + log_n_oracles..params.log_batch_size()];
136 let codeword_sub_oracles = iter::zip(codeword_commitments, params.input_oracles())
137 .map(|(commitment, spec)| {
138 let early_window = &early_challenges[max_early - spec.log_early_batch_size..];
141 let later_window = &later_challenges[max_later - spec.log_later_batch_size..];
142 let fold_challenges: Vec<E> =
143 early_window.iter().chain(later_window).cloned().collect();
144 BrakedownOracle::new(fold_challenges, commitment.clone(), spec.log_lift)
145 })
146 .collect();
147 let codeword_oracle = BatchBrakedownOracle::new(codeword_sub_oracles, outer_challenges);
148
149 let domain_context = GaoMateerOnTheFly::generate(params.rs_code().log_len());
155 let mut fri_oracles = Vec::with_capacity(params.fold_arities().len());
156 let mut depth = index_bits;
157 let mut fold_round = params.log_batch_size();
158 for (round_commitment, &arity) in iter::zip(round_commitments, params.fold_arities()) {
159 depth -= arity;
160 fri_oracles.push(FRIOracle::new(
161 challenges[fold_round..fold_round + arity].to_vec(),
162 round_commitment.clone(),
163 depth,
164 domain_context.clone(),
165 ));
166 fold_round += arity;
167 }
168
169 let final_challenges = &challenges[fold_round..];
170 let terminal_commitment = round_commitments
171 .last()
172 .expect("round_commitments is non-empty as an invariant")
173 .clone();
174
175 Self {
176 params,
177 terminal_commitment,
178 final_challenges,
179 codeword_oracle,
180 fri_oracles,
181 }
182 }
183
184 pub const fn n_oracles(&self) -> usize {
186 self.params.n_oracles()
187 }
188
189 pub fn verify<Channel>(&self, channel: &mut Channel) -> Result<E, Error>
190 where
191 Channel: MerkleIPVerifierChannel<F, Commitment = C, Elem = E>,
192 {
193 let mut indices = repeat_with(|| channel.sample_bits(self.params.index_bits()))
195 .take(self.params.n_test_queries())
196 .collect::<Vec<_>>();
197
198 let mut claims = self.codeword_oracle.open_queries(&indices, channel)?;
201 for (oracle, &arity) in self.fri_oracles.iter().zip(self.params.fold_arities()) {
202 claims = oracle.reduce_queries(&indices, &claims, channel)?;
203 indices = indices
204 .into_iter()
205 .map(|index| index >> arity as u32)
206 .collect();
207 }
208
209 self.verify_terminal_queries(&claims, &indices, channel)
211 }
212
213 fn verify_terminal_queries<Channel>(
220 &self,
221 claims: &[E],
222 indices: &[Channel::Word],
223 channel: &mut Channel,
224 ) -> Result<E, Error>
225 where
226 Channel: MerkleIPVerifierChannel<F, Commitment = C, Elem = E>,
227 {
228 let n_final_challenges = self.params.n_final_challenges();
229 let log_inv_rate = self.params.rs_code().log_inv_rate();
230
231 let terminate_codeword = channel.recv_committed_vector(&self.terminal_commitment)?;
232
233 iter::zip(claims, indices).try_for_each(|(claim, index)| {
235 let entry = channel.select(&terminate_codeword, index);
236 channel.assert_zero(claim.clone() - entry)
237 })?;
238
239 let domain_context = GaoMateerOnTheFly::generate(self.params.rs_code().log_len());
242 let log_len = n_final_challenges + log_inv_rate;
243 let repetition_codeword = terminate_codeword
244 .chunks(1 << n_final_challenges)
245 .enumerate()
246 .map(|(coset_index, coset)| {
247 let coset_index = Channel::Word::from(Word::from_u64(coset_index as u64));
250 fold_coset(
251 &domain_context,
252 log_len,
253 &coset_index,
254 self.final_challenges,
255 coset.to_vec(),
256 channel,
257 )
258 })
259 .collect::<Vec<_>>();
260
261 let final_value = repetition_codeword[0].clone();
262
263 repetition_codeword[1..]
265 .iter()
266 .try_for_each(|entry| channel.assert_zero(entry.clone() - final_value.clone()))?;
267
268 Ok(final_value)
269 }
270}