1use std::ops::DerefMut;
9
10use binius_compute::Allocator;
11use binius_field::{Field, PackedField};
12use binius_iop::channel::{
13 OracleSchedule, OracleSpec,
14 merge::{Placement, place_oracles},
15};
16use binius_ip_prover::channel::{IPProverChannel, WordIPProverChannel};
17use binius_math::{FieldBuffer, FieldSlice, FieldVec, StructuredBuffer};
18
19use crate::channel::IOPProverChannel;
20
21#[derive(Debug, Clone, Copy)]
23pub struct MergeOracle {
24 index: usize,
25}
26
27pub struct MergeProverChannel<'a, P, A, C>
79where
80 P: PackedField,
81 A: Allocator,
82 C: IOPProverChannel<P, A>,
83{
84 inner: C,
86
87 schedule: &'a OracleSchedule,
89
90 placements: Vec<Placement>,
92
93 alloc: A,
95
96 open: Option<FieldVec<P, A>>,
100
101 n_sent: usize,
103
104 groups: Vec<C::Oracle>,
106}
107
108impl<'a, P, A, C> MergeProverChannel<'a, P, A, C>
109where
110 P: PackedField,
111 A: Allocator,
112 C: IOPProverChannel<P, A>,
113{
114 pub fn new(inner: C, schedule: &'a OracleSchedule, alloc: A) -> Self {
127 assert_eq!(
128 inner.remaining_oracle_specs(),
129 schedule.merged_specs(),
130 "inner channel must be configured with the schedule's merged specs"
131 );
132 Self {
133 inner,
134 schedule,
135 placements: place_oracles(schedule),
136 alloc,
137 open: None,
138 n_sent: 0,
139 groups: Vec::new(),
140 }
141 }
142
143 pub fn into_inner(self) -> C {
149 let n_remaining = self.placements.len() - self.n_sent;
150 assert!(n_remaining == 0, "into_inner called but {n_remaining} oracle specs remaining",);
151 self.inner
152 }
153}
154
155fn place_block<P, Data>(dst: &mut FieldBuffer<P, Data>, src: FieldSlice<'_, P>, block_index: usize)
163where
164 P: PackedField,
165 Data: DerefMut<Target = [P]>,
166{
167 let n = src.log_len();
171 let offset = block_index << n;
172 assert!(offset + (1 << n) <= dst.len(), "pre-condition: the block must fit in the destination");
173
174 if n >= P::LOG_WIDTH {
177 let n_words = 1 << (n - P::LOG_WIDTH);
179 let word_offset = block_index * n_words;
180 dst.as_mut()[word_offset..word_offset + n_words].copy_from_slice(src.as_ref());
181 } else {
182 for i in 0..1usize << n {
184 dst.set(offset + i, src.get(i));
185 }
186 }
187}
188
189impl<F, P, A, C> IPProverChannel<F> for MergeProverChannel<'_, P, A, C>
190where
191 F: Field,
192 P: PackedField<Scalar = F>,
193 A: Allocator,
194 C: IOPProverChannel<P, A>,
195{
196 fn send_one(&mut self, elem: F) {
197 self.inner.send_one(elem);
198 }
199
200 fn send_many(&mut self, elems: &[F]) {
201 self.inner.send_many(elems);
202 }
203
204 fn observe_one(&mut self, val: F) {
205 self.inner.observe_one(val);
206 }
207
208 fn observe_many(&mut self, vals: &[F]) {
209 self.inner.observe_many(vals);
210 }
211
212 fn sample(&mut self) -> F {
213 self.inner.sample()
214 }
215}
216
217impl<F, P, A, C> WordIPProverChannel<F> for MergeProverChannel<'_, P, A, C>
218where
219 F: Field,
220 P: PackedField<Scalar = F>,
221 A: Allocator,
222 C: IOPProverChannel<P, A> + WordIPProverChannel<F>,
223{
224 type Word = C::Word;
225
226 fn observe_words(&mut self, words: &[Self::Word]) {
227 self.inner.observe_words(words);
228 }
229
230 fn sample_bits(&mut self, bits: usize) -> Self::Word {
231 self.inner.sample_bits(bits)
232 }
233}
234
235impl<'a, F, P, A, C> IOPProverChannel<P, A> for MergeProverChannel<'a, P, A, C>
236where
237 F: Field,
238 P: PackedField<Scalar = F>,
239 A: Allocator,
240 C: IOPProverChannel<P, A>,
241{
242 type Oracle = MergeOracle;
243
244 fn remaining_oracle_specs(&self) -> &[OracleSpec] {
245 &self.schedule.specs()[self.n_sent..]
246 }
247
248 fn send_oracle(&mut self, buffer: FieldSlice<'_, P>) -> Self::Oracle {
249 let remaining = self.remaining_oracle_specs();
254 assert!(!remaining.is_empty(), "send_oracle called but no remaining oracle specs");
255 assert_eq!(buffer.log_len(), remaining[0].log_msg_len, "oracle size must match its spec");
256
257 let index = self.n_sent;
258 self.n_sent += 1;
259 let Placement {
260 round,
261 block_index,
262 combined_log_len,
263 } = self.placements[index];
264
265 let combined = self
270 .open
271 .get_or_insert_with(|| FieldBuffer::zeros_in(&self.alloc, combined_log_len));
272 place_block(combined, buffer, block_index);
273
274 let is_last_of_round = self
276 .placements
277 .get(index + 1)
278 .is_none_or(|next| next.round != round);
279 if is_last_of_round {
280 let combined = self
281 .open
282 .take()
283 .expect("open round buffer was just inserted");
284
285 let outer = self.inner.send_oracle(combined.as_view());
287 self.inner.finalize_oracle(outer.clone(), combined);
288 self.groups.push(outer);
289 }
290
291 MergeOracle { index }
292 }
293
294 fn prove_oracle_relation(
295 &mut self,
296 oracle: Self::Oracle,
297 transparent: StructuredBuffer<P, A::Vec<P>>,
298 claim: P::Scalar,
299 ) {
300 let n_i = self.schedule.specs()[oracle.index].log_msg_len;
301 assert_eq!(
302 transparent.log_len(),
303 n_i,
304 "transparent log_len mismatch: expected {n_i}, got {}",
305 transparent.log_len()
306 );
307 let Placement {
308 round,
309 block_index,
310 combined_log_len,
311 } = self.placements[oracle.index];
312
313 let padded = StructuredBuffer::ZeroPadded {
316 inner: Box::new(transparent),
317 log_n_blocks: combined_log_len - n_i,
318 index: block_index,
319 };
320
321 let outer = self.groups[round].clone();
322 self.inner.prove_oracle_relation(outer, padded, claim);
323 }
324
325 fn finalize_oracle(&mut self, _oracle: Self::Oracle, _buffer: FieldVec<P, A>) {
326 }
329}
330
331#[cfg(test)]
332mod tests {
333 use std::iter;
334
335 use binius_compute::GlobalAllocator;
336 use binius_field::{
337 BinaryField, Field, Ghash128b, PackedField, PackedGhash1x128b, PackedGhash4x128b,
338 };
339 use binius_hash::{StdDigest, StdHashSuite};
340 use binius_iop::{
341 basefold::compiler::BaseFoldVerifierCompiler,
342 channel::{
343 IOPVerifierChannel, OracleSchedule, OracleSpec, merge::MergeVerifierChannel,
344 naive::NaiveVerifierChannel,
345 },
346 fri::MinProofSizeStrategy,
347 merkle_tree::BinaryMerkleTreeScheme,
348 };
349 use binius_ip::channel::IPVerifierChannel;
350 use binius_ip_prover::channel::IPProverChannel;
351 use binius_math::{
352 FieldBuffer,
353 inner_product::inner_product_buffers,
354 multilinear::eq::eq_ind_partial_eval,
355 ntt::{NeighborsLastSingleThread, domain_context::GaoMateerOnTheFly},
356 test_utils::{random_field_buffer, random_scalars},
357 };
358 use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
359 use proptest::prelude::*;
360 use rand::{Rng, SeedableRng, rngs::StdRng};
361
362 use super::{IOPProverChannel, MergeProverChannel};
363 use crate::{basefold::compiler::BaseFoldProverCompiler, channel::naive::NaiveProverChannel};
364
365 type StdChallenger = HasherChallenger<StdDigest>;
366
367 fn generate_oracle_data<F, P, R: Rng>(
372 rng: &mut R,
373 n_vars: usize,
374 ) -> (FieldBuffer<P>, FieldBuffer<P>, F)
375 where
376 F: BinaryField,
377 P: PackedField<Scalar = F>,
378 {
379 let buffer = random_field_buffer::<P>(&mut *rng, n_vars);
380 let point = random_scalars::<F>(&mut *rng, n_vars);
381 let transparent = eq_ind_partial_eval::<P>(&point);
382 let claim = inner_product_buffers(&buffer, &transparent);
383 (buffer, transparent, claim)
384 }
385
386 fn run_merge_round_trip<P>(rounds: &[&[usize]], tamper: bool)
394 where
395 P: PackedField<Scalar = Ghash128b>,
396 {
397 type F = Ghash128b;
398
399 let mut rng = StdRng::seed_from_u64(0);
400
401 let fine_sizes: Vec<usize> = rounds
405 .iter()
406 .flat_map(|round| round.iter().copied())
407 .collect();
408 let data: Vec<(FieldBuffer<P>, FieldBuffer<P>, F)> = fine_sizes
409 .iter()
410 .map(|&n| generate_oracle_data::<F, P, _>(&mut rng, n))
411 .collect();
412
413 let mut schedule = OracleSchedule::new();
417 for sizes in rounds {
418 for &n in *sizes {
419 schedule.push(OracleSpec::new(n));
420 }
421 schedule.end_round();
422 }
423 let coarse_specs = schedule.merged_specs();
424
425 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
431 let naive_prover = NaiveProverChannel::new(&mut prover_transcript, coarse_specs.clone());
432 let mut merge_prover = MergeProverChannel::new(naive_prover, &schedule, GlobalAllocator);
433
434 let mut oracles = Vec::new();
435 let mut index = 0;
436 for sizes in rounds {
437 for _ in *sizes {
438 let (buffer, _, _) = &data[index];
439 oracles.push(merge_prover.send_oracle(buffer.as_view()));
440 index += 1;
441 }
442 IPProverChannel::sample(&mut merge_prover);
443 }
444 for (&oracle, (_, transparent, claim)) in iter::zip(&oracles, &data) {
448 merge_prover.prove_oracle_relation(oracle, transparent.clone().into(), *claim);
449 }
450 for (&oracle, (buffer, _, _)) in iter::zip(&oracles, &data) {
451 merge_prover.finalize_oracle(oracle, buffer.clone());
452 }
453 merge_prover.into_inner().finish();
454
455 let mut verifier_transcript = prover_transcript.into_verifier();
460 let naive_verifier = NaiveVerifierChannel::new(&mut verifier_transcript, &coarse_specs);
461 let mut merge_verifier = MergeVerifierChannel::new(naive_verifier, &schedule);
462
463 let mut v_oracles = Vec::new();
464 for sizes in rounds {
465 for &n in *sizes {
466 v_oracles.push(merge_verifier.recv_oracle(n, true).unwrap());
467 }
468 IPVerifierChannel::sample(&mut merge_verifier);
469 }
470 for (position, (&oracle, (_, transparent, claim))) in
471 iter::zip(&v_oracles, &data).enumerate()
472 {
473 let transparent = transparent.clone();
474 let claim = if tamper && position == 0 {
478 *claim + F::ONE
479 } else {
480 *claim
481 };
482 merge_verifier
483 .verify_oracle_relation(
484 oracle,
485 Box::new(move |point: &[F]| {
486 let eq = eq_ind_partial_eval::<P>(point);
487 inner_product_buffers(&transparent, &eq)
488 }),
489 claim,
490 )
491 .expect("verification only ever queues a relation, it does not check it here");
492 }
493 merge_verifier.into_inner().finish();
494 }
495
496 #[test]
497 fn single_oracle_round_trip() {
498 run_merge_round_trip::<PackedGhash1x128b>(&[&[6]], false);
501 }
502
503 #[test]
504 fn multi_round_round_trip() {
505 run_merge_round_trip::<PackedGhash1x128b>(&[&[3, 3], &[4, 2, 2], &[1]], false);
515 }
516
517 #[test]
518 fn multi_round_round_trip_narrow_packing() {
519 const {
524 assert!(
525 PackedGhash4x128b::LOG_WIDTH > 0,
526 "the fixture needs sub-packed-width oracle sizes to appear"
527 );
528 };
529 run_merge_round_trip::<PackedGhash4x128b>(&[&[3, 3], &[4, 2, 2], &[1]], false);
530 }
531
532 #[test]
533 fn zero_variable_oracle_round_trip() {
534 run_merge_round_trip::<PackedGhash1x128b>(&[&[0, 0, 3]], false);
537 }
538
539 #[test]
540 #[should_panic(expected = "NaiveVerifierChannel: inner product verification failed")]
541 fn tampered_claim_is_rejected() {
542 run_merge_round_trip::<PackedGhash1x128b>(&[&[3, 3], &[4, 2, 2]], true);
545 }
546
547 #[test]
548 fn multiple_relations_on_merged_oracle() {
549 type F = Ghash128b;
550 type P = PackedGhash1x128b;
551
552 let mut rng = StdRng::seed_from_u64(0);
555 let mut schedule = OracleSchedule::new();
556 schedule.push(OracleSpec::new(4));
557 schedule.push(OracleSpec::new(3));
558 schedule.end_round();
559 let (buffer_1, _, _) = generate_oracle_data::<F, P, _>(&mut rng, 4);
560 let (buffer_2, _, _) = generate_oracle_data::<F, P, _>(&mut rng, 3);
561
562 let relations_1: Vec<(FieldBuffer<P>, F)> = (0..2)
563 .map(|_| {
564 let point = random_scalars::<F>(&mut rng, 4);
565 let transparent = eq_ind_partial_eval::<P>(&point);
566 let claim = inner_product_buffers(&buffer_1, &transparent);
567 (transparent, claim)
568 })
569 .collect();
570 let relations_2: Vec<(FieldBuffer<P>, F)> = (0..2)
571 .map(|_| {
572 let point = random_scalars::<F>(&mut rng, 3);
573 let transparent = eq_ind_partial_eval::<P>(&point);
574 let claim = inner_product_buffers(&buffer_2, &transparent);
575 (transparent, claim)
576 })
577 .collect();
578
579 let coarse_specs = schedule.merged_specs();
582 assert_eq!(coarse_specs, [OracleSpec::new(5)]);
583
584 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
589 let naive_prover = NaiveProverChannel::new(&mut prover_transcript, coarse_specs.clone());
590 let mut merge_prover = MergeProverChannel::new(naive_prover, &schedule, GlobalAllocator);
591
592 let oracle_1 = merge_prover.send_oracle(buffer_1.as_view());
593 let oracle_2 = merge_prover.send_oracle(buffer_2.as_view());
594 for (transparent, claim) in &relations_1 {
595 merge_prover.prove_oracle_relation(oracle_1, transparent.clone().into(), *claim);
596 }
597 for (transparent, claim) in &relations_2 {
598 merge_prover.prove_oracle_relation(oracle_2, transparent.clone().into(), *claim);
599 }
600 merge_prover.finalize_oracle(oracle_1, buffer_1);
601 merge_prover.finalize_oracle(oracle_2, buffer_2);
602 merge_prover.into_inner().finish();
603
604 let mut verifier_transcript = prover_transcript.into_verifier();
609 let naive_verifier = NaiveVerifierChannel::new(&mut verifier_transcript, &coarse_specs);
610 let mut merge_verifier = MergeVerifierChannel::new(naive_verifier, &schedule);
611
612 let v_oracle_1 = merge_verifier.recv_oracle(4, true).unwrap();
613 let v_oracle_2 = merge_verifier.recv_oracle(3, true).unwrap();
614 for (transparent, claim) in relations_1 {
615 merge_verifier
616 .verify_oracle_relation(
617 v_oracle_1,
618 Box::new(move |point: &[F]| {
619 let eq = eq_ind_partial_eval::<P>(point);
620 inner_product_buffers(&transparent, &eq)
621 }),
622 claim,
623 )
624 .unwrap();
625 }
626 for (transparent, claim) in relations_2 {
627 merge_verifier
628 .verify_oracle_relation(
629 v_oracle_2,
630 Box::new(move |point: &[F]| {
631 let eq = eq_ind_partial_eval::<P>(point);
632 inner_product_buffers(&transparent, &eq)
633 }),
634 claim,
635 )
636 .unwrap();
637 }
638 merge_verifier.into_inner().finish();
639 }
640
641 fn run_merge_over_basefold<P>(rounds: &[&[usize]]) -> bool
648 where
649 P: PackedField<Scalar = Ghash128b>,
650 {
651 type F = Ghash128b;
652 const LOG_INV_RATE: usize = 1;
653 const N_TEST_QUERIES: usize = 32;
654
655 let mut rng = StdRng::seed_from_u64(0);
656 let mut schedule = OracleSchedule::new();
657 for sizes in rounds {
658 for &n in *sizes {
659 schedule.push(OracleSpec::new_zk(n));
660 }
661 schedule.end_round();
662 }
663 let data = schedule
664 .specs()
665 .iter()
666 .map(|spec| {
667 let buffer = random_field_buffer::<P>(&mut rng, spec.log_msg_len);
668 let relations = (0..2)
669 .map(|_| {
670 let point = random_scalars::<F>(&mut rng, spec.log_msg_len);
671 let transparent = eq_ind_partial_eval::<P>(&point);
672 let claim = inner_product_buffers(&buffer, &transparent);
673 (transparent, claim)
674 })
675 .collect::<Vec<_>>();
676 (buffer, relations)
677 })
678 .collect::<Vec<_>>();
679
680 let verifier_compiler = BaseFoldVerifierCompiler::new(
681 &BinaryMerkleTreeScheme::<F, StdHashSuite>::new(),
682 schedule.merged_specs(),
683 LOG_INV_RATE,
684 N_TEST_QUERIES,
685 &MinProofSizeStrategy,
686 );
687 let ntt = NeighborsLastSingleThread::new(GaoMateerOnTheFly::generate(
688 verifier_compiler.max_log_domain_size(),
689 ));
690 let prover_compiler =
691 BaseFoldProverCompiler::<P, _>::from_verifier_compiler(&verifier_compiler, ntt);
692
693 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
695 let basefold_prover = prover_compiler
696 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _, _>(
697 &mut prover_transcript,
698 StdRng::seed_from_u64(1),
699 GlobalAllocator,
700 );
701 let mut merge_prover = MergeProverChannel::new(basefold_prover, &schedule, GlobalAllocator);
702 let mut oracles = Vec::new();
703 let mut data_iter = data.iter();
704 for sizes in rounds {
705 for (buffer, _) in data_iter.by_ref().take(sizes.len()) {
706 oracles.push(merge_prover.send_oracle(buffer.as_view()));
707 }
708 IPProverChannel::sample(&mut merge_prover);
709 }
710 for (&oracle, (_, relations)) in iter::zip(&oracles, &data) {
711 for (transparent, claim) in relations {
712 merge_prover.prove_oracle_relation(oracle, transparent.clone().into(), *claim);
713 }
714 }
715 for (&oracle, (buffer, _)) in iter::zip(&oracles, &data) {
716 merge_prover.finalize_oracle(oracle, buffer.clone());
717 }
718 merge_prover.into_inner().finish();
719
720 let mut verifier_transcript = prover_transcript.into_verifier();
722 let basefold_verifier = verifier_compiler
723 .create_channel_from_transcript::<StdHashSuite, StdChallenger, _>(
724 &mut verifier_transcript,
725 );
726 let mut merge_verifier = MergeVerifierChannel::new(basefold_verifier, &schedule);
727 let mut v_oracles = Vec::new();
728 for sizes in rounds {
729 for &n in *sizes {
730 v_oracles.push(merge_verifier.recv_oracle(n, true).unwrap());
731 }
732 IPVerifierChannel::sample(&mut merge_verifier);
733 }
734 for (&oracle, (_, relations)) in iter::zip(&v_oracles, data) {
735 for (transparent, claim) in relations {
736 merge_verifier
737 .verify_oracle_relation(
738 oracle,
739 Box::new(move |point: &[F]| {
740 let eq = eq_ind_partial_eval::<P>(point);
741 inner_product_buffers(&transparent, &eq)
742 }),
743 claim,
744 )
745 .expect("verification only ever queues a relation, it does not check it here");
746 }
747 }
748 merge_verifier.into_inner().finish().is_ok()
749 }
750
751 #[test]
752 fn merge_over_basefold_round_trip() {
753 let rounds: &[&[usize]] = &[&[4, 2, 2], &[5], &[3, 3, 0]];
754 assert!(run_merge_over_basefold::<PackedGhash1x128b>(rounds));
755
756 assert!(run_merge_over_basefold::<PackedGhash4x128b>(rounds));
758 }
759
760 proptest! {
761 #[test]
762 fn round_trip_proptest(
763 rounds in prop::collection::vec(prop::collection::vec(0usize..5, 1..5), 1..5),
764 ) {
765 let round_refs: Vec<&[usize]> = rounds.iter().map(Vec::as_slice).collect();
772 run_merge_round_trip::<PackedGhash1x128b>(&round_refs, false);
773
774 run_merge_round_trip::<PackedGhash4x128b>(&round_refs, false);
777 }
778 }
779}