Skip to main content

binius_iop_prover/basefold/
channel.rs

1// Copyright 2026 The Binius Developers
2
3//! BaseFold ZK implementation of the IOP prover channel.
4
5use 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/// Oracle handle returned by [`BaseFoldProverChannel::send_oracle`].
43#[derive(Debug, Clone, Copy)]
44pub struct BaseFoldOracle {
45	index: usize,
46}
47
48/// Committed oracle data stored internally.
49struct CommittedOracleData<P: PackedField, C, Data: Deref<Target = [P]>> {
50	/// The mask buffer generated during [`fri::encode_masked`] for a ZK oracle, held by the
51	/// channel because it is the only party that knows it. `None` for a non-ZK (unmasked) oracle.
52	mask: Option<FieldBuffer<P, Data>>,
53	/// RS-encoded codeword, drawn from the channel's allocator.
54	codeword: FieldBuffer<P, Data>,
55	/// The Merkle commitment handle for query proofs, owning the committed tree.
56	commitment: C,
57	/// The committed multilinear message `pi_i`, backed by the caller's allocator. Handed over by
58	/// [`IOPProverChannel::finalize_oracle`], and `None` until then.
59	message: Option<FieldBuffer<P, Data>>,
60}
61
62/// A committed-oracle relation queued for the single batched opening.
63struct QueuedRelation<P: PackedField, Data: Deref<Target = [P]>> {
64	/// The transparent multilinear `t` the message is opened against, backed by the caller's
65	/// allocator. Its zero padding is never written.
66	transparent: StructuredBuffer<P, Data>,
67	/// The claimed inner product `s = <pi, t>`.
68	claim: P::Scalar,
69}
70
71/// One oracle's queued relations, batched into a single relation.
72struct BatchedRelation<P: PackedField, Data: Deref<Target = [P]>> {
73	/// The batched transparent `T`, with every value written out.
74	transparent: FieldBuffer<P, Data>,
75	/// The batched claim `S = <pi, T>`.
76	claim: P::Scalar,
77}
78
79/// A prover channel that uses ZK BaseFold for all oracle commitments and openings.
80///
81/// This channel owns an [`StdRng`] and generates random masks internally during
82/// [`send_oracle`](IOPProverChannel::send_oracle). The caller provides only the raw witness
83/// buffer (not doubled). The channel handles:
84/// - Generating a random mask of equal length
85/// - Interleaving witness and mask for FRI commitment
86/// - Running ZK BaseFold proofs in [`Self::finish`]
87///
88/// # Type Parameters
89///
90/// - `F`: The binary field type
91/// - `P`: The packed field type with `Scalar = F`
92/// - `NTT`: The additive NTT for Reed-Solomon encoding
93/// - `Channel`: The Merkle channel carrying all prover interaction
94/// - `A`: The allocator the queued messages are drawn from, and the one [`Self::finish`] runs the
95///   opening with
96pub 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	/// The Merkle channel carrying all prover interaction: field elements, challenges,
105	/// commitments, and openings.
106	channel: Channel,
107	ntt: &'a NTT,
108	oracle_specs: Vec<OracleSpec>,
109	/// The combined FRI parameters over all committed oracles.
110	fri_params: FRIParams<F>,
111	committed_oracles: Vec<CommittedOracleData<P, Channel::Commitment, A::Vec<P>>>,
112	/// Oracle relations queued by [`IOPProverChannel::prove_oracle_relation`], indexed by oracle
113	/// index and opened together in [`Self::finish`]. One entry per committed oracle, so its
114	/// length is also the number of oracles committed so far.
115	queue: Vec<Vec<QueuedRelation<P, A::Vec<P>>>>,
116	rng: StdRng,
117	/// The allocator every codeword, mask and encode temporary is drawn from.
118	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	/// Creates a new BaseFold ZK prover channel over a Merkle channel from precomputed FRI
130	/// parameters.
131	///
132	/// The FRI parameters should already account for ZK (log_batch_size = 1, doubled message
133	/// length).
134	///
135	/// The RNG seeds the channel's own generator, whose only output is the ZK masks.
136	/// A mask is what hides a committed witness at the positions the verifier opens.
137	/// Hiding is therefore only as strong as this RNG, so it must be a cryptographic one.
138	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	/// Consumes the channel and proves the single combined opening over **all** committed oracles.
159	///
160	/// All oracle relations queued by
161	/// [`prove_oracle_relation`](IOPProverChannel::prove_oracle_relation) across every call are
162	/// processed here in one batch: masking, one batched sumcheck reducing the masked claims to a
163	/// shared point `r`, then one combined FRI opening over every committed oracle
164	/// (in oracle-index order). Mirrors [`BaseFoldVerifierChannel::finish`].
165	///
166	/// [`BaseFoldVerifierChannel::finish`]: binius_iop::basefold::channel::BaseFoldVerifierChannel::finish
167	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
198/// Proves the combined ZK BaseFold opening over all committed oracles.
199///
200/// This drives `channel` — the Merkle channel taken from the destructured
201/// [`BaseFoldProverChannel`] — through its [`MerkleIPProverChannel`] interface: it sends the
202/// masked inner products σ_i, runs one batched sumcheck reducing the masked claims to a shared
203/// point `r`, then opens all committed oracles together with a single combined FRI. Mirrors
204/// [`binius_iop::basefold::channel::BaseFoldVerifierChannel::finish`].
205///
206/// Everything runs in oracle-index order: `relations` arrives indexed by oracle, as do the
207/// per-oracle data (`oracle_specs`, `fri_params`, `committed_oracles`), the masking inner products
208/// σ_i, the sumcheck provers, the reduced evaluations α_i, and the FRI openings.
209fn 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	// TODO: Remove this limitation, it shouldn't be necessary. It is currently because of how the
229	// sumcheck reduces to the multilinear evaluations (alphas): an oracle with no relation gets no
230	// sumcheck prover, so its α would have to come from a plain multilinear evaluation.
231	assert!(
232		relations.iter().all(|relations| !relations.is_empty()),
233		"expects at least one relation per committed oracle",
234	);
235
236	// Take ownership of the committed messages π_i, leaving the masks and codewords in place. Every
237	// committed oracle must have been handed back with `finalize_oracle`.
238	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	// Batch each oracle's claims into one, so everything below runs exactly one relation per
250	// committed oracle.
251	let relations = batch_relations_per_oracle(channel, relations, alloc);
252
253	// `𝐧 = max_i log_msg_len_i`, the variable count of the combined opening / materialized buffer.
254	let max_n = oracle_specs
255		.iter()
256		.map(|spec| spec.log_msg_len)
257		.max()
258		.expect("at least one oracle");
259
260	// === Masking step (whitepaper 7.2) ===
261	// Only ZK oracles are masked. Send their σ_i = ⟨ω_i, T_i⟩ against the batched transparent (one
262	// per ZK oracle), then sample the single shared masking challenge γ — skipped entirely when no
263	// ZK oracle is present.
264	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	// Blind each ZK oracle's message in place: π_i' = (1-γ)π_i + γω_i. A non-ZK oracle's message
284	// passes through unmasked.
285	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); // pre-condition
288
289		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	// === Phase A: batched sumcheck on the masked claims ⟨π_i', T_i⟩ = s_i' ===
304	// One prover per committed oracle, in oracle-index order, each padded to `max_n`.
305	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); // pre-condition
311
312			// ZK oracle: mask the claim with σ_i.
313			// Non-ZK oracle: the claim passes through unmasked.
314			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	// Reduced oracle evaluations α_i = π_i'(ρ_i), one per committed oracle in oracle-index order.
343	let alphas = multilinear_evals
344		.iter()
345		.map(|evals| evals[0])
346		.collect::<Vec<_>>();
347	channel.send_many(&alphas);
348
349	// === Phase B: single combined-FRI MLE-check over the piecewise-concatenated oracle ===
350	// Collapse the oracle-index variables up front at sampled batching challenges `r'`: build the
351	// combined multilinear 𝛑(X) = Σ_i e[i]·π_i^↑(X) with e = eq(·, r') into one 2^𝐧 buffer, and the
352	// combined target s' = 𝛑(r) = Σ_i e[i]·α_i·∏_{j≥n_i}(1 - r_j).
353	// `batch_prove` returns binding-order challenges; reverse to variable-indexed (low-to-high),
354	// so that ρ_i is the first n_i coords.
355	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			// Each oracle occupies the low 2^{n_i} of every 2^{log_lift}·2^{n_i}-sized lift block,
373			// and that block is *repeated* across the 2^{log_repeat} high dims so the small
374			// oracle is constant along them (matching the FRI lift/repeat structure).
375			let log_lift = fri_oracle.log_lift;
376
377			// Repeat placement: add scalar · π_i' into the first 2^{n_i} entries of each of the
378			// 2^{log_repeat} chunks of size 2^{n_i + log_lift}.
379			// Borrow as a slice before the closure: the allocator's buffer type is only `Send`, so
380			// a closure capturing the owned buffer would not be `Sync` as `for_each` requires.
381			place_repeated(combined.as_mut_view(), witness_prime.as_view(), eq_i, n_i + log_lift);
382
383			// Repeat dims contribute 1; only the lift dims contribute an eq-to-zero factor.
384			s_prime += eq_i * alpha_i * eq_ind_zero(&point[n_i..][..log_lift]);
385		}
386
387		(combined, s_prime)
388	};
389
390	// Codeword commitments in oracle-index order, matching `open_fri_params.input_oracles()`.
391	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
409/// Batches each oracle's queued relations down to one, in oracle-index order.
410///
411/// An oracle carrying `k > 1` claims has them folded into a single claim against a single
412/// transparent, using a batching challenge λ shared by every oracle:
413///
414/// ```text
415/// T_i = Σ_j λ^j · t_ij     the combined transparent
416/// S_i = Σ_j λ^j · s_ij     the combined claim
417/// ```
418///
419/// The inner product is linear in the transparent, so `⟨π_i, T_i⟩ = S_i` holds exactly when every
420/// `⟨π_i, t_ij⟩ = s_ij` does, except with probability at most `Σ_i (k_i - 1) / |F|` over λ. λ is
421/// drawn after every claim it combines is already bound to the transcript, so no claim can be
422/// chosen as a function of it. The batched per-oracle claims are then combined again by the
423/// sumcheck's own outer batching coefficient.
424///
425/// Each zero-padded `t_ij` is accumulated into only its explicit block of `T_i`, so it costs only
426/// that block's size. An oracle whose first transparent is a plain buffer accumulates into that
427/// buffer, and an oracle carrying only that relation folds nothing.
428///
429/// Mirrors [`binius_iop::basefold::channel`]'s verifier-side batching.
430fn 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			// A plain first transparent becomes the accumulator with no copy.
455			let mut transparent = first.materialize(alloc);
456
457			// Powers λ, λ², … scale the oracle's remaining relations into their blocks.
458			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
473/// Adds `scalar · src` into `dst`, touching only the block `src` holds explicitly.
474///
475/// ## Preconditions
476///
477/// * `src.log_len() == dst.log_len()`
478fn 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()); // precondition
486			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
499/// Adds `scalar · src` into the low `2^src.log_len()` scalars of every `2^log_block`-sized block
500/// of `dst`.
501///
502/// This is the lift/repeat placement of one oracle into the combined buffer: the oracle occupies
503/// the low part of a lift block, and that block repeats across the high dims so the oracle is
504/// constant along them.
505///
506/// A block spanning at least one whole packed element gets a chunk of the buffer to itself. A
507/// narrower block does not: several of them then share one element, and no chunking can express
508/// the placement. The scalars of one element are the same repeating pattern for every element, so
509/// that pattern is built once and added to all of them.
510///
511/// ## Preconditions
512///
513/// * `src.log_len() <= log_block <= dst.log_len()`
514fn 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); // precondition
521	assert!(log_block <= dst.log_len()); // precondition
522
523	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		// Lane `k` of every element sits at position `k % 2^log_block` of its block, and carries
532		// the oracle only over the low `2^src.log_len()` of them. A buffer shorter than one
533		// element leaves its high lanes out of the pattern, so they stay zero.
534		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		// This accumulation already runs inside a parallel loop over chunks.
560		// One chunk is small, so a second split here would only add handoff cost.
561		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		// ZK channel expects raw witness buffer (NOT doubled).
651		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		// Encode oracle `index` of the combined FRI parameters. ZK oracles interleave a fresh mask
660		// (`encode_masked`); non-ZK oracles encode the message alone (`encode_interleaved`).
661		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		// Commit the codeword over the Merkle channel, with one interleaved coset per leaf.
685		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		// Queue the relation under its oracle; the actual opening (masking + sumcheck + combined
719		// FRI) happens once, over all committed oracles, in [`Self::finish`].
720		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		// === PROVER SIDE ===
822		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		// === VERIFIER SIDE ===
843		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		// === PROVER SIDE ===
891		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		// === VERIFIER SIDE ===
922		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	/// Runs a full prove/verify cycle of the Batched ZK BaseFold channel over oracles of the given
958	/// sizes. If `tamper`, the verifier's claim on the first oracle is corrupted; verification must
959	/// then fail. Returns whether verification accepted.
960	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		// === PROVER SIDE ===
985		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		// === VERIFIER SIDE ===
1009		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	/// Like `run_zk_channel` but with per-oracle `(n_vars, is_zk)` flags, exercising the mixed
1041	/// ZK/non-ZK opening. If `tamper`, the verifier's claim on the first oracle is corrupted.
1042	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		// 3 oracles (not a power of two) of unequal sizes: exercises oracle padding (Lifted FRI)
1131		// and the `⌈log 3⌉ = 2` outer oracle-combine rounds.
1132		assert!(run_zk_channel::<PackedGhash1x128b>(&[5, 6, 8], false));
1133	}
1134
1135	/// A batch whose lift blocks are narrower than one packed field element must still prove.
1136	///
1137	/// Placing an oracle into the combined buffer chunks that buffer by the lift block. Every
1138	/// block is lifted to the combined dimension, so a block narrower than a packed element means
1139	/// the whole buffer is — there is no whole-packed chunk to place into, and the placement has
1140	/// to write into part of a single element instead.
1141	///
1142	/// Whether that happens is a function of the packed width alone, so this pins the width rather
1143	/// than leaving it to `-Ctarget-cpu=native` and the host: under `PackedGhash4x128b` the
1144	/// `[0, 1]` batch reaches the case on every machine, while the 128-bit type the other tests use
1145	/// never does. That batch also lifts its first oracle (`log_lift = 1`), so the sub-packed
1146	/// placement is exercised on a lifted block rather than only on an unlifted one.
1147	///
1148	/// The `[1, 2]` batch sits just the other side of the boundary — its blocks are exactly one
1149	/// packed element — and covers the chunked path that the clamp must leave alone.
1150	/// [`place_repeated`] must match the definition it implements, at every shape.
1151	///
1152	/// The two regimes it splits on — a lift block spanning whole packed elements, and several
1153	/// blocks sharing one element — are selected by `log_block` against `P::LOG_WIDTH`, so the grid
1154	/// runs three packed widths against every `(log_src, log_block, log_dst)` they admit. That
1155	/// reaches shapes the FRI parameters do not currently produce, which is the point: the
1156	/// placement should not depend on which of them the optimizer happens to choose.
1157	#[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			// The definition: `scalar * src` lands in the low `2^log_src` scalars of each
1166			// `2^log_block`-sized block, and nowhere else.
1167			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	// Heterogeneous mixed/zero-ZK openings: each non-ZK oracle's batch fold is routed to the
1221	// *later* window of the first-fold challenge slice `[early ++ outer ++ later]`, so the
1222	// non-ZK oracles' batch-fold challenges come from the leading MLE rounds (which follow the
1223	// outer challenges in feed order) and land correctly in the FirstFold's later window.
1224	#[test]
1225	fn test_basefold_channel_mixed_zk_non_zk() {
1226		// One non-ZK oracle (8 vars) and one ZK oracle (6 vars): exercises conditional masking,
1227		// the heterogeneous combined-buffer lift/repeat placement, and the non-ZK unmasked commit.
1228		assert!(run_mixed_channel::<PackedGhash1x128b>(&[(8, false), (6, true)], false));
1229	}
1230
1231	#[test]
1232	fn test_basefold_channel_zero_zk() {
1233		// All non-ZK oracles: γ must never be sampled and the proof must still verify.
1234		assert!(run_mixed_channel::<PackedGhash1x128b>(&[(6, false), (8, false)], false));
1235	}
1236
1237	#[test]
1238	fn test_basefold_channel_mixed_invalid_proof() {
1239		// Tampering the claim on a mixed batch must be rejected.
1240		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	/// Generates a committed buffer of `n_vars` variables together with `n_relations` independent
1249	/// `(transparent, claim)` pairs, each opening the buffer at a different point.
1250	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	/// Runs a full prove/verify cycle over oracles described as `(n_vars, is_zk, n_relations)`.
1272	///
1273	/// The relations are queued round-robin over the oracles, so they arrive interleaved and the
1274	/// channel's grouping by oracle is exercised. If `tamper` is set, the verifier's claim at that
1275	/// arrival position is corrupted; verification must then fail. Returns whether verification
1276	/// accepted.
1277	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		// Arrival order of the relations, as `(oracle position, relation position)`.
1290		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		// === PROVER SIDE ===
1326		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		// === VERIFIER SIDE ===
1357		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		// Two claims on the same committed oracle: the channel batches them behind one λ.
1392		assert!(run_multi_relation_channel(&[(6, true, 2)], None));
1393	}
1394
1395	#[test]
1396	fn test_basefold_channel_two_relations_one_oracle_invalid() {
1397		// Tampering the second of the two claims must be rejected.
1398		assert!(!run_multi_relation_channel(&[(6, true, 2)], Some(1)));
1399	}
1400
1401	#[test]
1402	fn test_basefold_channel_mixed_relation_counts() {
1403		// A non-ZK oracle with one claim and a ZK oracle with three, arriving interleaved: only the
1404		// multi-claim oracle draws a λ, and the σ accounting must follow the batched relations.
1405		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		// Tampering a claim on the non-ZK oracle in a mixed batch must be rejected.
1411		assert!(!run_multi_relation_channel(&[(8, false, 2), (6, true, 2)], Some(2)));
1412	}
1413}