1use std::ops::Deref;
6
7use binius_compute::Allocator;
8use binius_field::{BinaryField, Field, PackedField};
9use binius_iop::{channel::OracleSpec, fri::FRIParams};
10use binius_ip_prover::{
11 channel::{IPProverChannel, WordIPProverChannel},
12 sumcheck::{
13 self, PaddedSumcheckDecorator, batch::BatchSumcheckOutput,
14 bivariate_product_evaluator::BivariateProductEvaluator, mle_store::MleStore,
15 round_evaluator::SharedSumcheckProver,
16 },
17};
18use binius_math::{
19 FieldBuffer, FieldSlice, FieldSliceMut, FieldVec, StructuredBuffer,
20 inner_product::inner_product_par,
21 line::extrapolate_line,
22 multilinear::eq::{eq_ind_partial_eval_scalars, eq_ind_zero},
23 ntt::AdditiveNTT,
24};
25use binius_utils::{
26 checked_arithmetics::log2_ceil_usize,
27 rayon::{
28 prelude::*,
29 task_size::{IndexedParallelIteratorExt, WorkPerItem},
30 },
31};
32use itertools::izip;
33use rand::{CryptoRng, SeedableRng, rngs::StdRng};
34
35use crate::{
36 basefold::prove_mlecheck_basefold,
37 channel::IOPProverChannel,
38 fri::{self, FRIFoldProver, MaskedCodeword},
39 merkle_channel::MerkleIPProverChannel,
40};
41
42#[derive(Debug, Clone, Copy)]
44pub struct BaseFoldOracle {
45 index: usize,
46}
47
48struct CommittedOracleData<P: PackedField, C, Data: Deref<Target = [P]>> {
50 mask: Option<FieldBuffer<P, Data>>,
53 codeword: FieldBuffer<P, Data>,
55 commitment: C,
57 message: Option<FieldBuffer<P, Data>>,
60}
61
62struct QueuedRelation<P: PackedField, Data: Deref<Target = [P]>> {
64 transparent: StructuredBuffer<P, Data>,
67 claim: P::Scalar,
69}
70
71struct BatchedRelation<P: PackedField, Data: Deref<Target = [P]>> {
73 transparent: FieldBuffer<P, Data>,
75 claim: P::Scalar,
77}
78
79pub struct BaseFoldProverChannel<'a, F, P, NTT, Channel, A>
97where
98 F: BinaryField,
99 P: PackedField<Scalar = F>,
100 NTT: AdditiveNTT<Field = F> + Sync,
101 Channel: MerkleIPProverChannel<F>,
102 A: Allocator,
103{
104 channel: Channel,
107 ntt: &'a NTT,
108 oracle_specs: Vec<OracleSpec>,
109 fri_params: FRIParams<F>,
111 committed_oracles: Vec<CommittedOracleData<P, Channel::Commitment, A::Vec<P>>>,
112 queue: Vec<Vec<QueuedRelation<P, A::Vec<P>>>>,
116 rng: StdRng,
117 alloc: A,
119}
120
121impl<'a, F, P, NTT, Channel, A> BaseFoldProverChannel<'a, F, P, NTT, Channel, A>
122where
123 F: BinaryField,
124 P: PackedField<Scalar = F>,
125 NTT: AdditiveNTT<Field = F> + Sync,
126 Channel: MerkleIPProverChannel<F>,
127 A: Allocator,
128{
129 pub fn new(
139 channel: Channel,
140 ntt: &'a NTT,
141 oracle_specs: Vec<OracleSpec>,
142 fri_params: FRIParams<F>,
143 mut rng: impl CryptoRng,
144 alloc: A,
145 ) -> Self {
146 Self {
147 channel,
148 ntt,
149 oracle_specs,
150 fri_params,
151 committed_oracles: Vec::new(),
152 queue: Vec::new(),
153 rng: StdRng::from_rng(&mut rng),
154 alloc,
155 }
156 }
157
158 pub fn finish(self) {
168 let Self {
169 mut channel,
170 ntt,
171 oracle_specs,
172 fri_params,
173 committed_oracles,
174 queue,
175 rng: _,
176 alloc,
177 } = self;
178
179 let n_remaining = oracle_specs.len() - queue.len();
180 assert!(n_remaining == 0, "finish called but {n_remaining} oracle specs remaining",);
181
182 if queue.iter().all(Vec::is_empty) {
183 return;
184 }
185
186 prove_batch_zk_basefold(
187 &mut channel,
188 ntt,
189 &oracle_specs,
190 &fri_params,
191 committed_oracles,
192 queue,
193 &alloc,
194 );
195 }
196}
197
198fn prove_batch_zk_basefold<A, F, P, NTT, Channel>(
210 channel: &mut Channel,
211 ntt: &NTT,
212 oracle_specs: &[OracleSpec],
213 fri_params: &FRIParams<F>,
214 mut committed_oracles: Vec<CommittedOracleData<P, Channel::Commitment, A::Vec<P>>>,
215 relations: Vec<Vec<QueuedRelation<P, A::Vec<P>>>>,
216 alloc: &A,
217) where
218 A: Allocator,
219 F: BinaryField,
220 P: PackedField<Scalar = F>,
221 NTT: AdditiveNTT<Field = F> + Sync,
222 Channel: MerkleIPProverChannel<F>,
223{
224 let n_committed = committed_oracles.len();
225 assert_eq!(oracle_specs.len(), n_committed);
226 assert_eq!(relations.len(), n_committed);
227
228 assert!(
232 relations.iter().all(|relations| !relations.is_empty()),
233 "expects at least one relation per committed oracle",
234 );
235
236 let mut messages = committed_oracles
239 .iter_mut()
240 .enumerate()
241 .map(|(index, oracle)| {
242 oracle
243 .message
244 .take()
245 .unwrap_or_else(|| panic!("oracle {index} was committed but never finalized"))
246 })
247 .collect::<Vec<_>>();
248
249 let relations = batch_relations_per_oracle(channel, relations, alloc);
252
253 let max_n = oracle_specs
255 .iter()
256 .map(|spec| spec.log_msg_len)
257 .max()
258 .expect("at least one oracle");
259
260 let any_zk_openings = oracle_specs.iter().any(|spec| spec.is_zk);
265 let (sigmas, gamma) = if any_zk_openings {
266 let _scope = tracing::debug_span!("Compute ZK mask opening values").entered();
267 let sigmas = izip!(&relations, oracle_specs, &committed_oracles)
268 .filter(|(_, spec, _)| spec.is_zk)
269 .map(|(relation, _, committed)| {
270 let mask = committed.mask.as_ref().expect("ZK oracle carries a mask");
271 inner_product_par(mask, &relation.transparent)
272 })
273 .collect::<Vec<_>>();
274 channel.send_many(&sigmas);
275
276 let gamma = channel.sample();
277
278 (sigmas, Some(gamma))
279 } else {
280 (Vec::new(), None)
281 };
282
283 for (message, spec, committed) in izip!(&mut messages, oracle_specs, &committed_oracles) {
286 let n_i = spec.log_msg_len;
287 assert_eq!(message.log_len(), n_i); if spec.is_zk {
290 let mask = committed.mask.as_ref().expect("ZK oracle carries a mask");
291 let gamma_broadcast = P::broadcast(gamma.expect("γ sampled when ZK oracles present"));
292
293 let _scope = tracing::debug_span!("Fold message and ZK mask", log_len = n_i).entered();
294 (message.as_mut(), mask.as_ref())
295 .into_par_iter()
296 .with_min_task(WorkPerItem::FieldMuls)
297 .for_each(|(message_i, &mask_i)| {
298 *message_i = extrapolate_line(*message_i, mask_i, gamma_broadcast);
299 });
300 }
301 }
302
303 let mut sigma_iter = sigmas.into_iter();
306 let provers = izip!(relations, &messages, oracle_specs)
307 .map(|(relation, message, spec)| {
308 let BatchedRelation { transparent, claim } = relation;
309 let n_i = spec.log_msg_len;
310 assert_eq!(transparent.log_len(), n_i); let sum_prime = if spec.is_zk {
315 let sigma = sigma_iter.next().expect("one σ per ZK oracle");
316 let gamma = gamma.expect("γ sampled when ZK oracles present");
317 extrapolate_line(claim, sigma, gamma)
318 } else {
319 claim
320 };
321
322 let mut store = MleStore::new(n_i, alloc);
323 let message_col = store.push(message.as_view());
324 let transparent_col = store.push_owned(transparent);
325 let inner = SharedSumcheckProver::new(
326 store,
327 [(sum_prime, BivariateProductEvaluator::new([message_col, transparent_col]))],
328 );
329 PaddedSumcheckDecorator::new(inner, max_n - n_i, vec![sum_prime], 2)
330 })
331 .collect::<Vec<_>>();
332
333 let BatchSumcheckOutput {
334 challenges,
335 multilinear_evals,
336 } = {
337 let _scope =
338 tracing::debug_span!("Reduce linear relations to committed openings").entered();
339 sumcheck::batch_prove(provers, channel)
340 };
341
342 let alphas = multilinear_evals
344 .iter()
345 .map(|evals| evals[0])
346 .collect::<Vec<_>>();
347 channel.send_many(&alphas);
348
349 let mut challenges = challenges;
356 challenges.reverse();
357 let point = &challenges;
358 let log_n_oracles = log2_ceil_usize(n_committed);
359 let outer_challenges = channel.sample_many(log_n_oracles);
360
361 let (combined, s_prime) = {
362 let _scope = tracing::debug_span!("Compute batched witness").entered();
363
364 let eq_tensor = eq_ind_partial_eval_scalars(&outer_challenges);
365
366 let mut combined = FieldBuffer::zeros_in(alloc, max_n);
367 let mut s_prime = F::ZERO;
368 for (fri_oracle, witness_prime, eq_i, alpha_i) in
369 izip!(fri_params.input_oracles(), messages, eq_tensor, alphas)
370 {
371 let n_i = witness_prime.log_len();
372 let log_lift = fri_oracle.log_lift;
376
377 place_repeated(combined.as_mut_view(), witness_prime.as_view(), eq_i, n_i + log_lift);
382
383 s_prime += eq_i * alpha_i * eq_ind_zero(&point[n_i..][..log_lift]);
385 }
386
387 (combined, s_prime)
388 };
389
390 let committed_codewords = committed_oracles
392 .into_iter()
393 .map(|committed| (committed.codeword, committed.commitment))
394 .collect();
395
396 let fri_folder = FRIFoldProver::new_batch(fri_params, ntt, committed_codewords);
397 prove_mlecheck_basefold(
398 combined,
399 point,
400 s_prime,
401 gamma,
402 &outer_challenges,
403 fri_folder,
404 channel,
405 alloc,
406 );
407}
408
409fn batch_relations_per_oracle<A, F, P, Channel>(
431 channel: &mut Channel,
432 relations: Vec<Vec<QueuedRelation<P, A::Vec<P>>>>,
433 alloc: &A,
434) -> Vec<BatchedRelation<P, A::Vec<P>>>
435where
436 A: Allocator,
437 F: BinaryField,
438 P: PackedField<Scalar = F>,
439 Channel: MerkleIPProverChannel<F>,
440{
441 let lambda = channel.sample();
442
443 relations
444 .into_iter()
445 .map(|relations| {
446 let mut relations = relations.into_iter();
447 let QueuedRelation {
448 transparent: first,
449 mut claim,
450 } = relations
451 .next()
452 .expect("pre-condition: every committed oracle carries at least one relation");
453
454 let mut transparent = first.materialize(alloc);
456
457 let mut coeff = lambda;
459 for relation in relations {
460 accumulate_scaled_structured(
461 transparent.as_mut_view(),
462 relation.transparent,
463 coeff,
464 );
465 claim += coeff * relation.claim;
466 coeff *= lambda;
467 }
468 BatchedRelation { transparent, claim }
469 })
470 .collect()
471}
472
473fn accumulate_scaled_structured<P: PackedField, Data: Deref<Target = [P]>>(
479 mut dst: FieldSliceMut<'_, P>,
480 src: StructuredBuffer<P, Data>,
481 scalar: P::Scalar,
482) {
483 match src {
484 StructuredBuffer::Buffer(buffer) => {
485 assert_eq!(buffer.log_len(), dst.log_len()); accumulate_scaled_buffer(dst, buffer.as_view(), P::broadcast(scalar));
487 }
488 StructuredBuffer::ZeroPadded {
489 inner,
490 log_n_blocks,
491 index,
492 } => {
493 let mut block = dst.chunk_mut(dst.log_len() - log_n_blocks, index);
494 accumulate_scaled_structured(block.chunk(), *inner, scalar);
495 }
496 }
497}
498
499fn place_repeated<P: PackedField>(
515 mut dst: FieldSliceMut<'_, P>,
516 src: FieldSlice<'_, P>,
517 scalar: P::Scalar,
518 log_block: usize,
519) {
520 assert!(src.log_len() <= log_block); assert!(log_block <= dst.log_len()); let scalar_broadcast = P::broadcast(scalar);
524 if log_block >= P::LOG_WIDTH {
525 let chunk_packed = 1usize << (log_block - P::LOG_WIDTH);
526 dst.as_mut().par_chunks_mut(chunk_packed).for_each(|chunk| {
527 let chunk_buf = FieldSliceMut::from_slice(log_block, chunk);
528 accumulate_scaled_buffer(chunk_buf, src.as_view(), scalar_broadcast);
529 });
530 } else {
531 let block_mask = (1usize << log_block) - 1;
535 let src_len = 1usize << src.log_len();
536 let lanes = P::WIDTH.min(1usize << dst.log_len());
537 let pattern = P::from_scalars((0..lanes).map(|lane| {
538 let position = lane & block_mask;
539 if position < src_len {
540 src.get(position)
541 } else {
542 P::Scalar::ZERO
543 }
544 }));
545 dst.as_mut()
546 .par_iter_mut()
547 .with_min_task(WorkPerItem::FieldMuls)
548 .for_each(|dst_i| *dst_i += scalar_broadcast * pattern);
549 }
550}
551
552fn accumulate_scaled_buffer<P: PackedField>(
553 mut dst: FieldSliceMut<'_, P>,
554 src: FieldSlice<'_, P>,
555 scalar_broadcast: P,
556) {
557 if src.log_len() >= P::LOG_WIDTH {
558 let src = src.as_ref();
559 dst.as_mut()
562 .par_iter_mut()
563 .zip(src.as_ref())
564 .with_min_task(WorkPerItem::FieldMuls)
565 .for_each(|(dst_i, src_i)| {
566 *dst_i += scalar_broadcast * *src_i;
567 });
568 } else {
569 let src = P::from_scalars(src.iter_scalars());
570 dst.as_mut()[0] += scalar_broadcast * src;
571 }
572}
573
574impl<'a, F, P, NTT, Channel, A> IPProverChannel<F>
575 for BaseFoldProverChannel<'a, F, P, NTT, Channel, A>
576where
577 F: BinaryField,
578 P: PackedField<Scalar = F>,
579 NTT: AdditiveNTT<Field = F> + Sync,
580 Channel: MerkleIPProverChannel<F>,
581 A: Allocator,
582{
583 fn send_one(&mut self, elem: F) {
584 self.channel.send_one(elem);
585 }
586
587 fn send_many(&mut self, elems: &[F]) {
588 self.channel.send_many(elems);
589 }
590
591 fn send_public_claim(&mut self, elem: F) {
592 self.channel.send_public_claim(elem);
593 }
594
595 fn observe_one(&mut self, val: F) {
596 self.channel.observe_one(val);
597 }
598
599 fn observe_many(&mut self, vals: &[F]) {
600 self.channel.observe_many(vals);
601 }
602
603 fn sample(&mut self) -> F {
604 self.channel.sample()
605 }
606}
607
608impl<F, P, NTT, Channel, A> WordIPProverChannel<F>
609 for BaseFoldProverChannel<'_, F, P, NTT, Channel, A>
610where
611 F: BinaryField,
612 P: PackedField<Scalar = F>,
613 NTT: AdditiveNTT<Field = F> + Sync,
614 Channel: MerkleIPProverChannel<F>,
615 A: Allocator,
616{
617 type Word = Channel::Word;
618
619 fn observe_words(&mut self, words: &[Self::Word]) {
620 self.channel.observe_words(words);
621 }
622
623 fn sample_bits(&mut self, bits: usize) -> Self::Word {
624 self.channel.sample_bits(bits)
625 }
626}
627
628impl<'a, F, P, NTT, Channel, A> IOPProverChannel<P, A>
629 for BaseFoldProverChannel<'a, F, P, NTT, Channel, A>
630where
631 F: BinaryField,
632 P: PackedField<Scalar = F>,
633 NTT: AdditiveNTT<Field = F> + Sync,
634 Channel: MerkleIPProverChannel<F>,
635 A: Allocator,
636{
637 type Oracle = BaseFoldOracle;
638
639 fn remaining_oracle_specs(&self) -> &[OracleSpec] {
640 &self.oracle_specs[self.queue.len()..]
641 }
642
643 fn send_oracle(&mut self, buffer: FieldSlice<'_, P>) -> Self::Oracle {
644 let remaining = self.remaining_oracle_specs();
645 assert!(!remaining.is_empty(), "send_oracle called but no remaining oracle specs");
646
647 let index = self.queue.len();
648 let spec = &remaining[0];
649
650 assert_eq!(
652 buffer.log_len(),
653 spec.log_msg_len,
654 "oracle buffer log_len mismatch: expected {}, got {}",
655 spec.log_msg_len,
656 buffer.log_len()
657 );
658
659 let (codeword, mask) = if spec.is_zk {
662 let MaskedCodeword { codeword, mask } = fri::encode_masked(
663 &self.fri_params,
664 index,
665 self.ntt,
666 buffer.as_view(),
667 &mut self.rng,
668 &self.alloc,
669 );
670 (codeword, Some(mask))
671 } else {
672 (
673 fri::encode_interleaved(
674 &self.fri_params,
675 index,
676 self.ntt,
677 buffer.as_view(),
678 &self.alloc,
679 ),
680 None,
681 )
682 };
683
684 let merkle_scope = tracing::debug_span!("Merkle commit").entered();
686 let leaf_size = 1 << self.fri_params.input_oracles()[index].log_batch_size();
687 let commitment = self
688 .channel
689 .send_merkle_commitment(codeword.as_view(), leaf_size);
690 drop(merkle_scope);
691
692 self.committed_oracles.push(CommittedOracleData {
693 mask,
694 codeword,
695 commitment,
696 message: None,
697 });
698 self.queue.push(Vec::new());
699
700 BaseFoldOracle { index }
701 }
702
703 fn prove_oracle_relation(
704 &mut self,
705 oracle: Self::Oracle,
706 transparent: StructuredBuffer<P, A::Vec<P>>,
707 claim: P::Scalar,
708 ) {
709 let n_committed = self.queue.len();
710 assert!(
711 oracle.index < n_committed,
712 "oracle index {} out of bounds, expected < {n_committed}",
713 oracle.index
714 );
715 let n_i = self.oracle_specs[oracle.index].log_msg_len;
716 assert_eq!(transparent.log_len(), n_i, "transparent log_len must match the oracle's");
717
718 self.queue[oracle.index].push(QueuedRelation { transparent, claim });
721 }
722
723 fn finalize_oracle(&mut self, oracle: Self::Oracle, buffer: FieldVec<P, A>) {
724 let committed = self
725 .committed_oracles
726 .get_mut(oracle.index)
727 .unwrap_or_else(|| panic!("oracle index {} out of bounds", oracle.index));
728 assert!(
729 committed.message.replace(buffer).is_none(),
730 "oracle {} finalized twice",
731 oracle.index
732 );
733 }
734}
735
736#[cfg(test)]
737mod tests {
738 use std::iter;
739
740 use binius_compute::GlobalAllocator;
741 use binius_field::{
742 BinaryField, Field, Ghash128b, Ghash128b as B128, PackedField, PackedGhash1x128b,
743 PackedGhash2x128b, PackedGhash4x128b, Random,
744 };
745 use binius_hash::{StdDigest, StdHashSuite};
746 use binius_iop::{
747 basefold::compiler::BaseFoldVerifierCompiler,
748 channel::{IOPVerifierChannel, OracleSpec},
749 fri::MinProofSizeStrategy,
750 merkle_tree::BinaryMerkleTreeScheme,
751 };
752 use binius_math::{
753 FieldBuffer,
754 inner_product::inner_product_buffers,
755 multilinear::eq::eq_ind_partial_eval,
756 ntt::{NeighborsLastSingleThread, domain_context::GaoMateerOnTheFly},
757 test_utils::{random_field_buffer, random_scalars},
758 };
759 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
760 use rand::{Rng, SeedableRng, rngs::StdRng};
761
762 use super::{IOPProverChannel, place_repeated};
763 use crate::basefold::compiler::BaseFoldProverCompiler;
764
765 type StdChallenger = HasherChallenger<StdDigest>;
766
767 const LOG_INV_RATE: usize = 1;
768 const SECURITY_BITS: usize = 32;
769
770 fn calculate_n_test_queries(security_bits: usize, log_inv_rate: usize) -> usize {
771 security_bits.div_ceil(log_inv_rate)
772 }
773
774 fn make_ntt(log_domain_size: usize) -> NeighborsLastSingleThread<GaoMateerOnTheFly<Ghash128b>> {
775 let domain_context = GaoMateerOnTheFly::generate(log_domain_size);
776 NeighborsLastSingleThread::new(domain_context)
777 }
778
779 fn make_merkle_scheme() -> BinaryMerkleTreeScheme<Ghash128b, StdHashSuite> {
780 BinaryMerkleTreeScheme::new()
781 }
782
783 fn generate_zk_oracle_data<F, P, R: Rng>(
784 rng: &mut R,
785 n_vars: usize,
786 ) -> (FieldBuffer<P>, FieldBuffer<P>, F)
787 where
788 F: BinaryField,
789 P: PackedField<Scalar = F>,
790 {
791 let buffer = random_field_buffer::<P>(&mut *rng, n_vars);
792 let evaluation_point = random_scalars::<F>(&mut *rng, n_vars);
793 let transparent_poly = eq_ind_partial_eval::<P>(&evaluation_point);
794 let evaluation_claim = inner_product_buffers(&buffer, &transparent_poly);
795 (buffer, transparent_poly, evaluation_claim)
796 }
797
798 #[test]
799 fn test_basefold_channel_single_oracle() {
800 type F = Ghash128b;
801 type P = PackedGhash1x128b;
802
803 let mut rng = StdRng::seed_from_u64(0);
804 let n_vars = 8;
805
806 let (buffer, transparent_poly, eval_claim) =
807 generate_zk_oracle_data::<F, P, _>(&mut rng, n_vars);
808
809 let n_test_queries = calculate_n_test_queries(SECURITY_BITS, LOG_INV_RATE);
810
811 let oracle_specs = vec![OracleSpec::new_zk(n_vars)];
812
813 let verifier_compiler = BaseFoldVerifierCompiler::new(
814 &make_merkle_scheme(),
815 oracle_specs,
816 LOG_INV_RATE,
817 n_test_queries,
818 &MinProofSizeStrategy,
819 );
820
821 let ntt = make_ntt(verifier_compiler.max_log_domain_size());
823 let prover_compiler =
824 BaseFoldProverCompiler::<P, _>::from_verifier_compiler(&verifier_compiler, ntt);
825
826 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
827 let prover_rng = StdRng::seed_from_u64(1);
828 let mut prover_channel = prover_compiler
829 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _, _>(
830 &mut prover_transcript,
831 prover_rng,
832 GlobalAllocator,
833 );
834
835 let oracle = prover_channel.send_oracle(buffer.as_view());
836 assert_eq!(oracle.index, 0);
837
838 prover_channel.prove_oracle_relation(oracle, transparent_poly.clone().into(), eval_claim);
839 prover_channel.finalize_oracle(oracle, buffer);
840 prover_channel.finish();
841
842 let mut verifier_transcript = prover_transcript.into_verifier();
844 let mut verifier_channel = verifier_compiler
845 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _>(
846 &mut verifier_transcript,
847 );
848
849 let v_oracle = verifier_channel.recv_oracle(n_vars, true).unwrap();
850
851 verifier_channel
852 .verify_oracle_relation(
853 v_oracle,
854 Box::new(move |point: &[F]| {
855 let eq = eq_ind_partial_eval::<P>(point);
856 inner_product_buffers(&transparent_poly, &eq)
857 }),
858 eval_claim,
859 )
860 .unwrap();
861 verifier_channel.finish().unwrap();
862 }
863
864 #[test]
865 fn test_basefold_channel_two_oracles() {
866 type F = Ghash128b;
867 type P = PackedGhash1x128b;
868
869 let mut rng = StdRng::seed_from_u64(0);
870 let n_vars_1 = 6;
871 let n_vars_2 = 8;
872
873 let (buffer_1, transparent_poly_1, eval_claim_1) =
874 generate_zk_oracle_data::<F, P, _>(&mut rng, n_vars_1);
875 let (buffer_2, transparent_poly_2, eval_claim_2) =
876 generate_zk_oracle_data::<F, P, _>(&mut rng, n_vars_2);
877
878 let n_test_queries = calculate_n_test_queries(SECURITY_BITS, LOG_INV_RATE);
879
880 let oracle_specs = vec![OracleSpec::new_zk(n_vars_1), OracleSpec::new_zk(n_vars_2)];
881
882 let verifier_compiler = BaseFoldVerifierCompiler::new(
883 &make_merkle_scheme(),
884 oracle_specs,
885 LOG_INV_RATE,
886 n_test_queries,
887 &MinProofSizeStrategy,
888 );
889
890 let ntt = make_ntt(verifier_compiler.max_log_domain_size());
892 let prover_compiler =
893 BaseFoldProverCompiler::<P, _>::from_verifier_compiler(&verifier_compiler, ntt);
894
895 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
896 let prover_rng = StdRng::seed_from_u64(1);
897 let mut prover_channel = prover_compiler
898 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _, _>(
899 &mut prover_transcript,
900 prover_rng,
901 GlobalAllocator,
902 );
903
904 let oracle_1 = prover_channel.send_oracle(buffer_1.as_view());
905 let oracle_2 = prover_channel.send_oracle(buffer_2.as_view());
906
907 prover_channel.prove_oracle_relation(
908 oracle_1,
909 transparent_poly_1.clone().into(),
910 eval_claim_1,
911 );
912 prover_channel.prove_oracle_relation(
913 oracle_2,
914 transparent_poly_2.clone().into(),
915 eval_claim_2,
916 );
917 prover_channel.finalize_oracle(oracle_1, buffer_1);
918 prover_channel.finalize_oracle(oracle_2, buffer_2);
919 prover_channel.finish();
920
921 let mut verifier_transcript = prover_transcript.into_verifier();
923 let mut verifier_channel = verifier_compiler
924 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _>(
925 &mut verifier_transcript,
926 );
927
928 let v_oracle_1 = verifier_channel.recv_oracle(n_vars_1, true).unwrap();
929 let v_oracle_2 = verifier_channel.recv_oracle(n_vars_2, true).unwrap();
930
931 let tp1 = transparent_poly_1;
932 let tp2 = transparent_poly_2;
933
934 verifier_channel
935 .verify_oracle_relation(
936 v_oracle_1,
937 Box::new(move |point: &[F]| {
938 let eq = eq_ind_partial_eval::<P>(point);
939 inner_product_buffers(&tp1, &eq)
940 }),
941 eval_claim_1,
942 )
943 .unwrap();
944 verifier_channel
945 .verify_oracle_relation(
946 v_oracle_2,
947 Box::new(move |point: &[F]| {
948 let eq = eq_ind_partial_eval::<P>(point);
949 inner_product_buffers(&tp2, &eq)
950 }),
951 eval_claim_2,
952 )
953 .unwrap();
954 verifier_channel.finish().unwrap();
955 }
956
957 fn run_zk_channel<P: PackedField<Scalar = Ghash128b>>(
961 n_vars_list: &[usize],
962 tamper: bool,
963 ) -> bool {
964 type F = Ghash128b;
965
966 let mut rng = StdRng::seed_from_u64(0);
967 let data: Vec<(FieldBuffer<P>, FieldBuffer<P>, F)> = n_vars_list
968 .iter()
969 .map(|&n| generate_zk_oracle_data::<F, P, _>(&mut rng, n))
970 .collect();
971
972 let n_test_queries = calculate_n_test_queries(SECURITY_BITS, LOG_INV_RATE);
973 let oracle_specs: Vec<OracleSpec> =
974 n_vars_list.iter().map(|&n| OracleSpec::new_zk(n)).collect();
975
976 let verifier_compiler = BaseFoldVerifierCompiler::new(
977 &make_merkle_scheme(),
978 oracle_specs,
979 LOG_INV_RATE,
980 n_test_queries,
981 &MinProofSizeStrategy,
982 );
983
984 let ntt = make_ntt(verifier_compiler.max_log_domain_size());
986 let prover_compiler =
987 BaseFoldProverCompiler::<P, _>::from_verifier_compiler(&verifier_compiler, ntt);
988
989 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
990 let prover_rng = StdRng::seed_from_u64(1);
991 let mut prover_channel = prover_compiler
992 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _, _>(
993 &mut prover_transcript,
994 prover_rng,
995 GlobalAllocator,
996 );
997
998 let oracles: Vec<_> = data
999 .iter()
1000 .map(|(buffer, _, _)| prover_channel.send_oracle(buffer.as_view()))
1001 .collect();
1002 for (oracle, (buffer, transparent, claim)) in iter::zip(oracles, &data) {
1003 prover_channel.prove_oracle_relation(oracle, transparent.clone().into(), *claim);
1004 prover_channel.finalize_oracle(oracle, buffer.clone());
1005 }
1006 prover_channel.finish();
1007
1008 let mut verifier_transcript = prover_transcript.into_verifier();
1010 let mut verifier_channel = verifier_compiler
1011 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _>(
1012 &mut verifier_transcript,
1013 );
1014
1015 let v_oracles: Vec<_> = n_vars_list
1016 .iter()
1017 .map(|&n| verifier_channel.recv_oracle(n, true).unwrap())
1018 .collect();
1019 for (i, (oracle, (_, transparent, claim))) in iter::zip(v_oracles, &data).enumerate() {
1020 let transparent = transparent.clone();
1021 let claim = if tamper && i == 0 {
1022 *claim + F::ONE
1023 } else {
1024 *claim
1025 };
1026 verifier_channel
1027 .verify_oracle_relation(
1028 oracle,
1029 Box::new(move |point: &[F]| {
1030 let eq = eq_ind_partial_eval::<P>(point);
1031 inner_product_buffers(&transparent, &eq)
1032 }),
1033 claim,
1034 )
1035 .expect("verify_oracle_relation only queues");
1036 }
1037 verifier_channel.finish().is_ok()
1038 }
1039
1040 fn run_mixed_channel<P: PackedField<Scalar = Ghash128b>>(
1043 specs: &[(usize, bool)],
1044 tamper: bool,
1045 ) -> bool {
1046 type F = Ghash128b;
1047
1048 let mut rng = StdRng::seed_from_u64(0);
1049 let data: Vec<(FieldBuffer<P>, FieldBuffer<P>, F)> = specs
1050 .iter()
1051 .map(|&(n, _)| generate_zk_oracle_data::<F, P, _>(&mut rng, n))
1052 .collect();
1053
1054 let n_test_queries = calculate_n_test_queries(SECURITY_BITS, LOG_INV_RATE);
1055 let oracle_specs: Vec<OracleSpec> = specs
1056 .iter()
1057 .map(|&(n, is_zk)| {
1058 if is_zk {
1059 OracleSpec::new_zk(n)
1060 } else {
1061 OracleSpec::new(n)
1062 }
1063 })
1064 .collect();
1065
1066 let verifier_compiler = BaseFoldVerifierCompiler::new(
1067 &make_merkle_scheme(),
1068 oracle_specs,
1069 LOG_INV_RATE,
1070 n_test_queries,
1071 &MinProofSizeStrategy,
1072 );
1073
1074 let ntt = make_ntt(verifier_compiler.max_log_domain_size());
1075 let prover_compiler =
1076 BaseFoldProverCompiler::<P, _>::from_verifier_compiler(&verifier_compiler, ntt);
1077
1078 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
1079 let prover_rng = StdRng::seed_from_u64(1);
1080 let mut prover_channel = prover_compiler
1081 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _, _>(
1082 &mut prover_transcript,
1083 prover_rng,
1084 GlobalAllocator,
1085 );
1086
1087 let oracles: Vec<_> = data
1088 .iter()
1089 .map(|(buffer, _, _)| prover_channel.send_oracle(buffer.as_view()))
1090 .collect();
1091 for (oracle, (buffer, transparent, claim)) in iter::zip(oracles, &data) {
1092 prover_channel.prove_oracle_relation(oracle, transparent.clone().into(), *claim);
1093 prover_channel.finalize_oracle(oracle, buffer.clone());
1094 }
1095 prover_channel.finish();
1096
1097 let mut verifier_transcript = prover_transcript.into_verifier();
1098 let mut verifier_channel = verifier_compiler
1099 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _>(
1100 &mut verifier_transcript,
1101 );
1102
1103 let v_oracles: Vec<_> = specs
1104 .iter()
1105 .map(|&(n, _)| verifier_channel.recv_oracle(n, true).unwrap())
1106 .collect();
1107 for (i, (oracle, (_, transparent, claim))) in iter::zip(v_oracles, &data).enumerate() {
1108 let transparent = transparent.clone();
1109 let claim = if tamper && i == 0 {
1110 *claim + F::ONE
1111 } else {
1112 *claim
1113 };
1114 verifier_channel
1115 .verify_oracle_relation(
1116 oracle,
1117 Box::new(move |point: &[F]| {
1118 let eq = eq_ind_partial_eval::<P>(point);
1119 inner_product_buffers(&transparent, &eq)
1120 }),
1121 claim,
1122 )
1123 .expect("verify_oracle_relation only queues");
1124 }
1125 verifier_channel.finish().is_ok()
1126 }
1127
1128 #[test]
1129 fn test_basefold_channel_three_oracles_non_power_of_two() {
1130 assert!(run_zk_channel::<PackedGhash1x128b>(&[5, 6, 8], false));
1133 }
1134
1135 #[test]
1158 fn place_repeated_matches_the_naive_placement() {
1159 fn check<P: PackedField<Scalar = B128>>(log_src: usize, log_block: usize, log_dst: usize) {
1160 let mut rng = StdRng::seed_from_u64(0);
1161 let src = random_field_buffer::<P>(&mut rng, log_src);
1162 let initial = random_field_buffer::<P>(&mut rng, log_dst);
1163 let scalar = B128::random(&mut rng);
1164
1165 let mut expected = initial.clone();
1168 for index in 0..1usize << log_dst {
1169 let position = index % (1usize << log_block);
1170 if position < 1usize << log_src {
1171 expected.set(index, expected.get(index) + scalar * src.get(position));
1172 }
1173 }
1174
1175 let mut actual = initial;
1176 place_repeated(actual.as_mut_view(), src.as_view(), scalar, log_block);
1177
1178 for index in 0..1usize << log_dst {
1179 assert_eq!(
1180 actual.get(index),
1181 expected.get(index),
1182 "P::LOG_WIDTH={}, log_src={log_src}, log_block={log_block}, log_dst={log_dst}, \
1183 index={index}",
1184 P::LOG_WIDTH,
1185 );
1186 }
1187 }
1188
1189 fn check_all_shapes<P: PackedField<Scalar = B128>>() {
1190 for log_dst in 0..=4 {
1191 for log_block in 0..=log_dst {
1192 for log_src in 0..=log_block {
1193 check::<P>(log_src, log_block, log_dst);
1194 }
1195 }
1196 }
1197 }
1198
1199 check_all_shapes::<PackedGhash1x128b>();
1200 check_all_shapes::<PackedGhash2x128b>();
1201 check_all_shapes::<PackedGhash4x128b>();
1202 }
1203
1204 #[test]
1205 fn batch_narrower_than_a_packed_element_proves() {
1206 const {
1207 assert!(
1208 PackedGhash4x128b::LOG_WIDTH > 1,
1209 "the fixture needs a packed element wider than the `[0, 1]` batch's lift block"
1210 );
1211 };
1212 for sizes in [[0, 1], [1, 2]] {
1213 assert!(
1214 run_zk_channel::<PackedGhash4x128b>(&sizes, false),
1215 "batch of {sizes:?}-variable oracles"
1216 );
1217 }
1218 }
1219
1220 #[test]
1225 fn test_basefold_channel_mixed_zk_non_zk() {
1226 assert!(run_mixed_channel::<PackedGhash1x128b>(&[(8, false), (6, true)], false));
1229 }
1230
1231 #[test]
1232 fn test_basefold_channel_zero_zk() {
1233 assert!(run_mixed_channel::<PackedGhash1x128b>(&[(6, false), (8, false)], false));
1235 }
1236
1237 #[test]
1238 fn test_basefold_channel_mixed_invalid_proof() {
1239 assert!(!run_mixed_channel::<PackedGhash1x128b>(&[(8, false), (6, true)], true));
1241 }
1242
1243 #[test]
1244 fn test_basefold_channel_invalid_proof() {
1245 assert!(!run_zk_channel::<PackedGhash1x128b>(&[6, 8], true));
1246 }
1247
1248 fn generate_oracle_relations<F, P, R: Rng>(
1251 rng: &mut R,
1252 n_vars: usize,
1253 n_relations: usize,
1254 ) -> (FieldBuffer<P>, Vec<(FieldBuffer<P>, F)>)
1255 where
1256 F: BinaryField,
1257 P: PackedField<Scalar = F>,
1258 {
1259 let buffer = random_field_buffer::<P>(&mut *rng, n_vars);
1260 let relations = (0..n_relations)
1261 .map(|_| {
1262 let point = random_scalars::<F>(&mut *rng, n_vars);
1263 let transparent = eq_ind_partial_eval::<P>(&point);
1264 let claim = inner_product_buffers(&buffer, &transparent);
1265 (transparent, claim)
1266 })
1267 .collect();
1268 (buffer, relations)
1269 }
1270
1271 fn run_multi_relation_channel(specs: &[(usize, bool, usize)], tamper: Option<usize>) -> bool {
1278 type F = Ghash128b;
1279 type P = PackedGhash1x128b;
1280
1281 let mut rng = StdRng::seed_from_u64(0);
1282 let data = specs
1283 .iter()
1284 .map(|&(n_vars, _, n_relations)| {
1285 generate_oracle_relations::<F, P, _>(&mut rng, n_vars, n_relations)
1286 })
1287 .collect::<Vec<_>>();
1288
1289 let max_relations = specs
1291 .iter()
1292 .map(|&(_, _, k)| k)
1293 .max()
1294 .expect("at least one oracle");
1295 let arrivals = (0..max_relations)
1296 .flat_map(|round| {
1297 specs
1298 .iter()
1299 .enumerate()
1300 .filter(move |&(_, &(_, _, k))| round < k)
1301 .map(move |(index, _)| (index, round))
1302 })
1303 .collect::<Vec<_>>();
1304
1305 let n_test_queries = calculate_n_test_queries(SECURITY_BITS, LOG_INV_RATE);
1306 let oracle_specs = specs
1307 .iter()
1308 .map(|&(n_vars, is_zk, _)| {
1309 if is_zk {
1310 OracleSpec::new_zk(n_vars)
1311 } else {
1312 OracleSpec::new(n_vars)
1313 }
1314 })
1315 .collect::<Vec<_>>();
1316
1317 let verifier_compiler = BaseFoldVerifierCompiler::new(
1318 &make_merkle_scheme(),
1319 oracle_specs,
1320 LOG_INV_RATE,
1321 n_test_queries,
1322 &MinProofSizeStrategy,
1323 );
1324
1325 let ntt = make_ntt(verifier_compiler.max_log_domain_size());
1327 let prover_compiler =
1328 BaseFoldProverCompiler::<P, _>::from_verifier_compiler(&verifier_compiler, ntt);
1329
1330 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
1331 let prover_rng = StdRng::seed_from_u64(1);
1332 let mut prover_channel = prover_compiler
1333 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _, _>(
1334 &mut prover_transcript,
1335 prover_rng,
1336 GlobalAllocator,
1337 );
1338
1339 let oracles = data
1340 .iter()
1341 .map(|(buffer, _)| prover_channel.send_oracle(buffer.as_view()))
1342 .collect::<Vec<_>>();
1343 for &(index, round) in &arrivals {
1344 let (transparent, claim) = &data[index].1[round];
1345 prover_channel.prove_oracle_relation(
1346 oracles[index],
1347 transparent.clone().into(),
1348 *claim,
1349 );
1350 }
1351 for (oracle, (buffer, _)) in iter::zip(&oracles, &data) {
1352 prover_channel.finalize_oracle(*oracle, buffer.clone());
1353 }
1354 prover_channel.finish();
1355
1356 let mut verifier_transcript = prover_transcript.into_verifier();
1358 let mut verifier_channel = verifier_compiler
1359 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _>(
1360 &mut verifier_transcript,
1361 );
1362
1363 let v_oracles = specs
1364 .iter()
1365 .map(|&(n_vars, _, _)| verifier_channel.recv_oracle(n_vars, true).unwrap())
1366 .collect::<Vec<_>>();
1367 for (position, &(index, round)) in arrivals.iter().enumerate() {
1368 let (transparent, claim) = &data[index].1[round];
1369 let transparent = transparent.clone();
1370 let claim = if tamper == Some(position) {
1371 *claim + F::ONE
1372 } else {
1373 *claim
1374 };
1375 verifier_channel
1376 .verify_oracle_relation(
1377 v_oracles[index],
1378 Box::new(move |point: &[F]| {
1379 let eq = eq_ind_partial_eval::<P>(point);
1380 inner_product_buffers(&transparent, &eq)
1381 }),
1382 claim,
1383 )
1384 .expect("verify_oracle_relation only queues");
1385 }
1386 verifier_channel.finish().is_ok()
1387 }
1388
1389 #[test]
1390 fn test_basefold_channel_two_relations_one_oracle() {
1391 assert!(run_multi_relation_channel(&[(6, true, 2)], None));
1393 }
1394
1395 #[test]
1396 fn test_basefold_channel_two_relations_one_oracle_invalid() {
1397 assert!(!run_multi_relation_channel(&[(6, true, 2)], Some(1)));
1399 }
1400
1401 #[test]
1402 fn test_basefold_channel_mixed_relation_counts() {
1403 assert!(run_multi_relation_channel(&[(8, false, 1), (6, true, 3)], None));
1406 }
1407
1408 #[test]
1409 fn test_basefold_channel_mixed_relation_counts_invalid() {
1410 assert!(!run_multi_relation_channel(&[(8, false, 2), (6, true, 2)], Some(2)));
1412 }
1413}