1use binius_core::word::Word;
6use binius_field::BinaryField;
7use binius_ip::{
8 channel::{IPVerifierChannel, WordIPVerifierChannel},
9 sumcheck::{self, BatchSumcheckOutput},
10};
11use binius_math::{
12 line::extrapolate_line,
13 multilinear::eq::{eq_ind_partial_eval_scalars, eq_ind_zero},
14 univariate::evaluate_univariate,
15};
16use binius_utils::checked_arithmetics::log2_ceil_usize;
17use itertools::izip;
18
19use crate::{
20 basefold,
21 channel::{Error, IOPVerifierChannel, OracleSpec, TransparentEvalFn},
22 fri::FRIParams,
23 merkle_channel::MerkleIPVerifierChannel,
24};
25
26#[derive(Debug, Clone, Copy)]
28pub struct BaseFoldOracle {
29 index: usize,
30}
31
32struct QueuedRelation<Elem> {
34 transparent: TransparentEvalFn<Elem>,
36 claim: Elem,
38}
39
40pub struct BaseFoldVerifierChannel<'a, F, Channel>
51where
52 F: BinaryField,
53 Channel: MerkleIPVerifierChannel<F, Elem: From<F> + 'static>,
54{
55 channel: Channel,
58 oracle_specs: &'a [OracleSpec],
59 fri_params: &'a FRIParams<F>,
60 oracle_commitments: Vec<Channel::Commitment>,
61 queue: Vec<Vec<QueuedRelation<Channel::Elem>>>,
65}
66
67impl<'a, F, Channel> BaseFoldVerifierChannel<'a, F, Channel>
68where
69 F: BinaryField,
70 Channel: MerkleIPVerifierChannel<F, Elem: From<F> + 'static>,
71{
72 pub const fn new(
78 channel: Channel,
79 oracle_specs: &'a [OracleSpec],
80 fri_params: &'a FRIParams<F>,
81 ) -> Self {
82 Self {
83 channel,
84 oracle_specs,
85 fri_params,
86 oracle_commitments: Vec::new(),
87 queue: Vec::new(),
88 }
89 }
90
91 pub fn finish(self) -> Result<Channel, Error> {
104 let Self {
105 mut channel,
106 oracle_specs,
107 fri_params,
108 oracle_commitments,
109 queue,
110 } = self;
111
112 let n_remaining = oracle_specs.len() - queue.len();
113 assert!(n_remaining == 0, "finish called but {n_remaining} oracle specs remaining",);
114
115 if !queue.iter().all(Vec::is_empty) {
116 verify_batch_zk_basefold(
117 &mut channel,
118 oracle_specs,
119 fri_params,
120 &oracle_commitments,
121 queue,
122 )?;
123 }
124
125 Ok(channel)
126 }
127}
128
129fn verify_batch_zk_basefold<F, Channel>(
146 channel: &mut Channel,
147 oracle_specs: &[OracleSpec],
148 fri_params: &FRIParams<F>,
149 oracle_commitments: &[Channel::Commitment],
150 relations: Vec<Vec<QueuedRelation<Channel::Elem>>>,
151) -> Result<(), Error>
152where
153 F: BinaryField,
154 Channel: MerkleIPVerifierChannel<F, Elem: From<F> + 'static>,
155{
156 let n_committed = oracle_commitments.len();
157 assert_eq!(relations.len(), n_committed);
158
159 assert!(
161 relations.iter().all(|relations| !relations.is_empty()),
162 "expects at least one relation per committed oracle",
163 );
164
165 let max_n = oracle_specs
167 .iter()
168 .map(|spec| spec.log_msg_len)
169 .max()
170 .expect("at least one oracle");
171
172 let relations = batch_relations_per_oracle(channel, relations);
175
176 let n_zk = oracle_specs.iter().filter(|spec| spec.is_zk).count();
180 let sigmas = channel.recv_many(n_zk)?;
181 let gamma = (!sigmas.is_empty()).then(|| channel.sample());
182
183 let mut sigma_iter = sigmas.into_iter();
186 let sum_primes = izip!(&relations, oracle_specs)
187 .map(|(relation, spec)| {
188 if spec.is_zk {
189 let sigma = sigma_iter.next().expect("one σ per ZK oracle");
190 extrapolate_line(
191 relation.claim.clone(),
192 sigma,
193 gamma.clone().expect("γ sampled when ZK oracles present"),
194 )
195 } else {
196 relation.claim.clone()
197 }
198 })
199 .collect::<Vec<_>>();
200
201 let BatchSumcheckOutput {
203 batch_coeff: sumcheck_batch_coeff,
204 eval: sumcheck_reduced_eval,
205 challenges: sumcheck_challenges,
206 } = sumcheck::batch_verify::<F, _>(max_n, 2, &sum_primes, channel)?;
207
208 let alphas = channel.recv_many(n_committed)?;
210
211 let mut point = sumcheck_challenges;
213 point.reverse();
214
215 let contributions = izip!(relations, oracle_specs, &alphas)
217 .map(|(relation, spec, alpha_i)| {
218 let (eval_coords, padding_coords) = point.split_at(spec.log_msg_len);
219 let pad_eq = eq_ind_zero(padding_coords);
220 let transparent_eval = (relation.transparent)(eval_coords);
221 alpha_i.clone() * transparent_eval * pad_eq
222 })
223 .collect::<Vec<_>>();
224 let expected = evaluate_univariate(&contributions, &sumcheck_batch_coeff);
225 channel.assert_zero(sumcheck_reduced_eval - expected)?;
226
227 let log_n_oracles = log2_ceil_usize(n_committed);
232 let outer_challenges = channel.sample_many(log_n_oracles);
233 let eq_tensor = eq_ind_partial_eval_scalars(&outer_challenges);
234 let s_prime = izip!(fri_params.input_oracles(), oracle_specs, eq_tensor, alphas)
239 .map(|(fri_oracle, spec, eq_i, alpha_i)| {
240 let n_i = spec.log_msg_len;
241 let log_lift = fri_oracle.log_lift;
242 eq_i * alpha_i * eq_ind_zero(&point[n_i..][..log_lift])
243 })
244 .sum::<Channel::Elem>();
245
246 basefold::verify_mlecheck_basefold(
248 fri_params,
249 oracle_commitments,
250 s_prime,
251 &point,
252 gamma,
253 &outer_challenges,
254 channel,
255 )?;
256
257 Ok(())
258}
259
260fn batch_relations_per_oracle<F, Channel>(
278 channel: &mut Channel,
279 relations: Vec<Vec<QueuedRelation<Channel::Elem>>>,
280) -> Vec<QueuedRelation<Channel::Elem>>
281where
282 F: BinaryField,
283 Channel: MerkleIPVerifierChannel<F, Elem: From<F> + 'static>,
284{
285 let lambda = channel.sample();
286
287 relations
288 .into_iter()
289 .map(|mut relations| {
290 if relations.len() <= 1 {
292 return relations
293 .pop()
294 .expect("pre-condition: every committed oracle carries at least one relation");
295 }
296
297 let (transparents, claims): (Vec<_>, Vec<_>) = relations
298 .into_iter()
299 .map(|relation| (relation.transparent, relation.claim))
300 .unzip();
301 let claim = evaluate_univariate(&claims, &lambda);
302 let lambda = lambda.clone();
303
304 QueuedRelation {
305 transparent: Box::new(move |point: &[Channel::Elem]| {
306 let evals = transparents
307 .iter()
308 .map(|transparent| transparent(point))
309 .collect::<Vec<_>>();
310 evaluate_univariate(&evals, &lambda)
311 }),
312 claim,
313 }
314 })
315 .collect()
316}
317
318impl<F, Channel> IPVerifierChannel<F> for BaseFoldVerifierChannel<'_, F, Channel>
319where
320 F: BinaryField,
321 Channel: MerkleIPVerifierChannel<F, Elem: From<F> + 'static>,
322{
323 type Elem = Channel::Elem;
324
325 fn recv_one(&mut self) -> Result<Self::Elem, binius_ip::channel::Error> {
326 self.channel.recv_one()
327 }
328
329 fn recv_many(&mut self, n: usize) -> Result<Vec<Self::Elem>, binius_ip::channel::Error> {
330 self.channel.recv_many(n)
331 }
332
333 fn recv_array<const N: usize>(&mut self) -> Result<[Self::Elem; N], binius_ip::channel::Error> {
334 self.channel.recv_array()
335 }
336
337 fn recv_public_claim(&mut self) -> Result<Self::Elem, binius_ip::channel::Error> {
338 self.channel.recv_public_claim()
339 }
340
341 fn sample(&mut self) -> Self::Elem {
342 self.channel.sample()
343 }
344
345 fn observe_one(&mut self, val: F) -> Self::Elem {
346 self.channel.observe_one(val)
347 }
348
349 fn observe_many(&mut self, vals: &[F]) -> Vec<Self::Elem> {
350 self.channel.observe_many(vals)
351 }
352
353 fn assert_zero(&mut self, val: Self::Elem) -> Result<(), binius_ip::channel::Error> {
354 self.channel.assert_zero(val)
355 }
356}
357
358impl<F, Channel> WordIPVerifierChannel<F> for BaseFoldVerifierChannel<'_, F, Channel>
359where
360 F: BinaryField,
361 Channel: MerkleIPVerifierChannel<F, Elem: From<F> + 'static>,
362{
363 type Word = Channel::Word;
364
365 fn observe_words(&mut self, words: &[Word]) -> Vec<Self::Word> {
366 self.channel.observe_words(words)
367 }
368
369 fn subset_sum(&mut self, elems: &[Self::Elem], word: &Self::Word) -> Self::Elem {
370 self.channel.subset_sum(elems, word)
371 }
372
373 fn select(&mut self, elems: &[Self::Elem], word: &Self::Word) -> Self::Elem {
374 self.channel.select(elems, word)
375 }
376
377 fn sample_bits(&mut self, bits: usize) -> Self::Word {
378 self.channel.sample_bits(bits)
379 }
380
381 fn pack_words(&mut self, words: &[Self::Word]) -> Vec<Self::Elem> {
382 self.channel.pack_words(words)
383 }
384}
385
386impl<'a, F, Channel> IOPVerifierChannel<F> for BaseFoldVerifierChannel<'a, F, Channel>
387where
388 F: BinaryField,
389 Channel: MerkleIPVerifierChannel<F, Elem: From<F> + 'static>,
390{
391 type Oracle = BaseFoldOracle;
392
393 fn remaining_oracle_specs(&self) -> &[OracleSpec] {
394 &self.oracle_specs[self.queue.len()..]
395 }
396
397 fn recv_oracle(
398 &mut self,
399 _log_msg_len: usize,
400 _is_witness_dependent: bool,
401 ) -> Result<Self::Oracle, Error> {
402 assert!(
405 !self.remaining_oracle_specs().is_empty(),
406 "recv_oracle called but no remaining oracle specs"
407 );
408
409 let index = self.queue.len();
410
411 let fri_oracle = &self.fri_params.input_oracles()[index];
415 let depth = (self.fri_params.rs_code().log_dim() - fri_oracle.log_lift)
416 + self.fri_params.rs_code().log_inv_rate();
417
418 let spec = &self.oracle_specs[index];
422 assert_eq!(
423 fri_oracle.log_batch_size() + depth - self.fri_params.rs_code().log_inv_rate(),
424 spec.log_msg_len + usize::from(spec.is_zk),
425 "invariant: the FRI commitment shape must be consistent with the oracle spec's \
426 log_msg_len"
427 );
428
429 let commitment = self
430 .channel
431 .recv_merkle_commitment(1 << fri_oracle.log_batch_size(), depth)?;
432
433 self.oracle_commitments.push(commitment);
434 self.queue.push(Vec::new());
435
436 Ok(BaseFoldOracle { index })
437 }
438
439 fn verify_oracle_relation(
440 &mut self,
441 oracle: Self::Oracle,
442 transparent: TransparentEvalFn<Self::Elem>,
443 claim: Self::Elem,
444 ) -> Result<(), Error> {
445 let n_committed = self.queue.len();
448 self.queue
449 .get_mut(oracle.index)
450 .unwrap_or_else(|| {
451 panic!("oracle index {} out of bounds, expected < {n_committed}", oracle.index)
452 })
453 .push(QueuedRelation { transparent, claim });
454 Ok(())
455 }
456}