Skip to main content

binius_ip_prover/sumcheck/
zk_mlecheck.rs

1// Copyright 2026 The Binius Developers
2
3//! Prover for the Libra mask polynomial in ZK MLE-check protocols.
4//!
5//! The Libra ZK-sumcheck protocol uses a masking polynomial g(X_0, ..., X_{n-1}) of the form:
6//! g = sum_{i=0}^{n-1} g_i(X_i)
7//!
8//! where each g_i(X) is a univariate polynomial of configurable degree. This separable structure
9//! allows efficient computation of round polynomials without iterating over the full hypercube.
10
11use std::{iter, ops::Deref};
12
13use binius_compute::Allocator;
14use binius_field::{Field, PackedField, util::powers};
15use binius_ip::{mlecheck, sumcheck::RoundCoeffs};
16use binius_math::{
17	FieldVec, field_buffer::FieldBuffer, line::extrapolate_line, univariate::evaluate_univariate,
18};
19
20use super::{common::MleCheckProver, round_state::RoundState};
21use crate::channel::IPProverChannel;
22
23/// Output of the ZK MLE-check proving protocol.
24#[derive(Debug, Clone)]
25pub struct ProveZKOutput<F: Field> {
26	/// Evaluations of the main multilinear polynomials at the challenge point.
27	pub multilinear_evals: Vec<F>,
28	/// Evaluation of the mask polynomial at the challenge point.
29	pub mask_eval: F,
30	/// Verifier challenges for each round of the sumcheck protocol.
31	pub challenges: Vec<F>,
32}
33
34/// Generates hypercube evaluations of the libra_eval polynomial.
35///
36/// For a mask polynomial with `n_vars` variables and `degree`, generates a `FieldBuffer`
37/// of size `2^(m_n + m_d)` containing:
38///
39/// ```text
40/// libra_eval_r(j, k) = r[j]^k  if j < n_vars and k ≤ degree
41///                    = 0       otherwise
42/// ```
43///
44/// where `j` and `k` are derived from the hypercube index: `idx = j * 2^m_d + k`.
45///
46/// # Arguments
47///
48/// * `alloc` - The allocator the expansion is drawn from
49/// * `challenge_point` - The challenge point `r` from sumcheck (length `n_vars`)
50/// * `n_vars` - Number of variables in the mask polynomial
51/// * `degree` - Degree of each univariate in the mask polynomial
52/// * `m_n` - Log of number of rows (must satisfy `2^m_n >= n_vars`)
53/// * `m_d` - Log of row size (must satisfy `2^m_d >= degree + 1`)
54///
55/// # Panics
56///
57/// Panics (in debug mode) if `n_vars > 2^m_n` or `degree + 1 > 2^m_d`.
58pub fn expand_libra_eval<A: Allocator, P: PackedField>(
59	alloc: &A,
60	challenge_point: &[P::Scalar],
61	n_vars: usize,
62	degree: usize,
63	m_n: usize,
64	m_d: usize,
65) -> FieldVec<P, A> {
66	debug_assert!(challenge_point.len() == n_vars);
67	debug_assert!(n_vars <= 1 << m_n);
68	debug_assert!(degree < 1 << m_d);
69
70	let log_size = m_n + m_d;
71	let mut buffer = FieldBuffer::zeros_in(alloc, log_size);
72	let row_stride = 1 << m_d;
73
74	for (j, &r_j) in challenge_point.iter().enumerate() {
75		let base_idx = j * row_stride;
76		for (k, power) in powers(r_j).take(degree + 1).enumerate() {
77			buffer.set(base_idx + k, power);
78		}
79	}
80
81	buffer
82}
83
84/// Libra mask polynomial for ZK MLE-check protocols.
85///
86/// Stores coefficients for a separable polynomial $g(X) = \sum_i g_i(X_i)$
87/// where each $g_i$ is a univariate polynomial of degree $d$.
88///
89/// The coefficients are stored in a `FieldBuffer` with `m_n + m_d` variables where:
90/// - `m_n = ceil(log2(n))` - log of number of variables
91/// - `m_d = ceil(log2(d + 1))` - log of degree + 1
92///
93/// The buffer is conceptually an `n × (d+1)` matrix padded to `2^m_n × 2^m_d`,
94/// with random values in the `n × (d+1)` submatrix and zeros elsewhere.
95///
96/// The type is generic over the buffer storage type `Data`, allowing it to work
97/// with both owned buffers (`Box<[P]>`) and borrowed slices.
98pub struct Mask<P: PackedField, Data: Deref<Target = [P]> = Box<[P]>> {
99	/// Number of variables (n)
100	n_vars: usize,
101	/// Degree of each univariate polynomial (d)
102	degree: usize,
103	/// Coefficients stored as a FieldBuffer with log_len = m_n + m_d.
104	/// Layout: row i contains the monomial coefficients [g_{i,0}, ..., g_{i,d}, 0, ..., 0].
105	/// Row i spans indices [i * 2^m_d, (i+1) * 2^m_d).
106	buffer: FieldBuffer<P, Data>,
107}
108
109impl<F: Field, P: PackedField<Scalar = F>, Data: Deref<Target = [P]>> Mask<P, Data> {
110	/// Creates a new mask polynomial from a pre-allocated buffer.
111	///
112	/// # Arguments
113	///
114	/// * `n_vars` - Number of variables (n).
115	/// * `degree` - Degree of each univariate polynomial (d).
116	/// * `buffer` - Buffer with log_len = m_n + m_d.
117	pub const fn new(n_vars: usize, degree: usize, buffer: FieldBuffer<P, Data>) -> Self {
118		Self {
119			n_vars,
120			degree,
121			buffer,
122		}
123	}
124
125	/// Returns the number of variables.
126	pub const fn n_vars(&self) -> usize {
127		self.n_vars
128	}
129
130	/// Returns the degree of each univariate polynomial.
131	pub const fn degree(&self) -> usize {
132		self.degree
133	}
134
135	/// Returns m_d = ceil(log2(degree + 1)).
136	const fn log_degree_plus_one(&self) -> usize {
137		(self.degree + 1).next_power_of_two().ilog2() as usize
138	}
139
140	/// Gets the coefficient $g_{i,j}$ (coefficient of $X^j$ in $g_i$).
141	pub fn get_coeff(&self, var_index: usize, coeff_index: usize) -> F {
142		debug_assert!(var_index < self.n_vars);
143		debug_assert!(coeff_index <= self.degree);
144		let row_stride = 1 << self.log_degree_plus_one();
145		self.buffer.get(var_index * row_stride + coeff_index)
146	}
147
148	/// Returns the monomial coefficients [g_{i,0}, g_{i,1}, ..., g_{i,d}] for variable i.
149	pub fn coeffs_for_var(&self, var_index: usize) -> impl Iterator<Item = F> + '_ {
150		debug_assert!(var_index < self.n_vars);
151		let m_d = self.log_degree_plus_one();
152		let row_stride = 1 << m_d;
153		let start = var_index * row_stride;
154		(0..=self.degree).map(move |j| self.buffer.get(start + j))
155	}
156
157	/// Evaluates g_i(x) for a specific variable using Horner's method.
158	pub fn evaluate_univariate(&self, var_index: usize, x: F) -> F {
159		let coeffs: Vec<_> = self.coeffs_for_var(var_index).collect();
160		evaluate_univariate(&coeffs, &x)
161	}
162
163	/// Computes the MLE of the mask polynomial at a point.
164	///
165	/// For a mask polynomial $g(X) = \sum_i g_i(X_i)$ where each $g_i$ is univariate,
166	/// the MLE at point $z$ is:
167	///
168	/// $$
169	/// \sum_{v \in \{0,1\}^n} g(v) \cdot eq(v, z) = \sum_i [(1-z_i) g_i(0) + z_i g_i(1)]
170	/// $$
171	///
172	/// This simplification arises because $\sum_{v_j \in \{0,1\}} eq_1(v_j, z_j) = 1$.
173	pub fn evaluate_mle(&self, eval_point: &[F]) -> F {
174		assert_eq!(eval_point.len(), self.n_vars);
175
176		iter::zip(0..self.n_vars, eval_point)
177			.map(|(i, &z_i)| {
178				let g_at_0 = self.get_coeff(i, 0);
179				let g_at_1 = self.evaluate_univariate(i, F::ONE);
180				extrapolate_line(g_at_0, g_at_1, z_i)
181			})
182			.sum()
183	}
184}
185
186impl<P: PackedField, Data: Deref<Target = [P]>> AsRef<FieldBuffer<P, Data>> for Mask<P, Data> {
187	fn as_ref(&self) -> &FieldBuffer<P, Data> {
188		&self.buffer
189	}
190}
191
192/// Prover for the Libra mask polynomial in ZK MLE-check.
193///
194/// The mask polynomial has the separable form $g(X_0, ..., X_{n-1}) = sum_{i} g_i(X_i)$,
195/// where each $g_i$ is a univariate polynomial of configurable degree.
196///
197/// This structure allows efficient round polynomial computation in O(degree) time per round.
198pub struct MleCheckMaskProver<F: Field, P: PackedField<Scalar = F>, Data: Deref<Target = [P]>> {
199	/// The mask polynomial (owned)
200	mask: Mask<P, Data>,
201	/// The evaluation point z (in high-to-low variable order)
202	eval_point: Vec<F>,
203	/// Number of variables remaining to process
204	n_vars_remaining: usize,
205	/// Accumulated sum of g_j(r_j) for already-folded variables
206	prefix_sum: F,
207	/// Precomputed (1-z_j)*g_j(0) + z_j*g_j(1) for each variable
208	suffix_sums: Vec<F>,
209	/// State: either last round coefficients (after execute) or current claim (after fold)
210	last_coeffs_or_claim: RoundState<RoundCoeffs<F>, F>,
211}
212
213impl<F: Field, P: PackedField<Scalar = F>, Data: Deref<Target = [P]>>
214	MleCheckMaskProver<F, P, Data>
215{
216	/// Creates a new prover for the Libra mask polynomial.
217	///
218	/// # Arguments
219	///
220	/// * `mask` - The mask polynomial (takes ownership).
221	/// * `eval_point` - The evaluation point z for the MLE-check claim, in high-to-low order.
222	/// * `eval_claim` - The claimed value of the MLE of g at the evaluation point.
223	///
224	/// # Panics
225	///
226	/// Panics if `mask.n_vars() != eval_point.len()`.
227	pub fn new(mask: Mask<P, Data>, eval_point: Vec<F>, eval_claim: F) -> Self {
228		assert_eq!(mask.n_vars(), eval_point.len(), "mask n_vars must match eval_point length");
229
230		let n_vars = eval_point.len();
231
232		// Precompute suffix_sums[j] = (1-z_j)*g_j(0) + z_j*g_j(1)
233		// That is the line through g_j(0) and g_j(1), read at z_j.
234		let suffix_sums: Vec<F> = iter::zip(0..n_vars, &eval_point)
235			.map(|(i, &z_j)| {
236				let g_at_0 = mask.get_coeff(i, 0);
237				let g_at_1 = mask.evaluate_univariate(i, F::ONE);
238				extrapolate_line(g_at_0, g_at_1, z_j)
239			})
240			.collect();
241
242		Self {
243			mask,
244			eval_point,
245			n_vars_remaining: n_vars,
246			prefix_sum: F::ZERO,
247			suffix_sums,
248			last_coeffs_or_claim: RoundState::Claim(eval_claim),
249		}
250	}
251
252	/// Returns the index of the current variable being processed.
253	/// Processing is high-to-low, so we start at n_vars-1 and decrease.
254	const fn current_var_index(&self) -> usize {
255		self.n_vars_remaining - 1
256	}
257}
258
259impl<F: Field, P: PackedField<Scalar = F>, Data: Deref<Target = [P]>> MleCheckProver<F>
260	for MleCheckMaskProver<F, P, Data>
261{
262	fn n_vars(&self) -> usize {
263		self.n_vars_remaining
264	}
265
266	fn execute(&mut self) -> Vec<RoundCoeffs<F>> {
267		self.last_coeffs_or_claim.claim();
268
269		assert_ne!(self.n_vars_remaining, 0, "execute called out of order; expected finish");
270
271		let var_idx = self.current_var_index();
272
273		// Compute suffix sum for variables that haven't been processed yet (indices 0 to var_idx-1)
274		// Since we process high-to-low (n-1, n-2, ..., 0), the suffix is the lower-indexed
275		// variables
276		let suffix_sum: F = self.suffix_sums[..var_idx].iter().copied().sum();
277
278		// Compute the constant offset: prefix_sum + suffix_sum
279		let constant_offset = self.prefix_sum + suffix_sum;
280
281		// Build the round polynomial R(X) = g_i(X) + constant_offset
282		// g_i(X) = sum_{k=0}^{d} a_{i,k} * X^k
283		// So coefficients are: [a_0 + offset, a_1, a_2, ..., a_d]
284		let mut round_coeffs_vec: Vec<F> = self.mask.coeffs_for_var(var_idx).collect();
285		if round_coeffs_vec.is_empty() {
286			round_coeffs_vec.push(constant_offset);
287		} else {
288			round_coeffs_vec[0] += constant_offset;
289		}
290
291		let round_coeffs = RoundCoeffs(round_coeffs_vec);
292		self.last_coeffs_or_claim = RoundState::Coeffs(round_coeffs.clone());
293		vec![round_coeffs]
294	}
295
296	fn fold(&mut self, challenge: F) {
297		let coeffs = self.last_coeffs_or_claim.coeffs();
298
299		// Evaluate round polynomial at challenge to get new claim
300		let new_claim = coeffs.evaluate(&challenge);
301
302		let var_idx = self.current_var_index();
303
304		// Update prefix_sum: add g_i(r_i)
305		self.prefix_sum += self.mask.evaluate_univariate(var_idx, challenge);
306
307		self.n_vars_remaining -= 1;
308		self.last_coeffs_or_claim = RoundState::Claim(new_claim);
309	}
310
311	fn finish(self) -> Vec<F> {
312		assert_eq!(self.n_vars_remaining, 0, "finish called out of order; sumcheck rounds remain");
313
314		// Final evaluation of g at the challenge point is prefix_sum
315		// (since g(r_0, ..., r_{n-1}) = sum_i g_i(r_i))
316		vec![self.prefix_sum]
317	}
318
319	fn eval_point(&self) -> &[F] {
320		// Return remaining coordinates (high-to-low means we return the first n_vars_remaining
321		// elements)
322		&self.eval_point[..self.n_vars_remaining]
323	}
324}
325
326/// Executes the zero-knowledge MLE-check proving protocol for a single multivariate polynomial.
327///
328/// This function proves a single MLE-check while batching with a Libra mask polynomial
329/// to achieve zero-knowledge. The mask polynomial has the separable form
330/// $g(X_0, \ldots, X_{n-1}) = \sum_i g_i(X_i)$ where each $g_i$ is a univariate polynomial.
331///
332/// # Protocol Flow
333///
334/// 1. Compute and write `mask_eval` (MLE of mask polynomial at the evaluation point)
335/// 2. Sample `batch_challenge` and batch evaluation claims
336/// 3. For each round, batch the main and mask round polynomials
337/// 4. Write `mask_eval_out` (mask evaluation at the challenge point)
338///
339/// # Arguments
340///
341/// * `main_prover` - The MLE-check prover for the main polynomial. Must carry exactly one claim.
342/// * `mask` - The mask polynomial.
343/// * `channel` - The channel for sending prover messages and sampling challenges
344///
345/// # Returns
346///
347/// Returns [`ProveZKOutput`] containing the main polynomial's multilinear evaluations
348/// and the round challenges.
349///
350/// # Pre-conditions
351///
352/// * The mask's univariate degree must be at least the degree of the main prover's round
353///   polynomials.
354/// * The two round polynomials are added coefficient by coefficient, so a shorter mask leaves the
355///   high-degree coefficients uncovered and the protocol is no longer hiding.
356///
357/// # Panics
358///
359/// Panics if the main prover emits more than one round polynomial.
360/// Panics if the mask's round polynomial is shorter than the main prover's.
361pub fn prove<F: Field, P: PackedField<Scalar = F>, Data: Deref<Target = [P]>>(
362	mut main_prover: impl MleCheckProver<F>,
363	mask: Mask<P, Data>,
364	channel: &mut impl IPProverChannel<F>,
365) -> ProveZKOutput<F> {
366	let n_vars = main_prover.n_vars();
367	let eval_point = main_prover.eval_point().to_vec();
368
369	// Compute and write mask_eval (MLE of mask polynomial at the evaluation point)
370	let mask_eval = mask.evaluate_mle(&eval_point);
371	channel.send_one(mask_eval);
372
373	// Sample batch challenge and construct mask prover
374	let batch_challenge: F = channel.sample();
375	let batched_mask_eval = batch_challenge * mask_eval;
376	let mut mask_prover = MleCheckMaskProver::new(mask, eval_point, batched_mask_eval);
377
378	let mut challenges = Vec::with_capacity(n_vars);
379
380	for _ in 0..n_vars {
381		// Execute both provers
382		let mut main_round_coeffs_vec = main_prover.execute();
383		assert_eq!(
384			main_round_coeffs_vec.len(),
385			1,
386			"prove requires a main prover with one claim, but it emitted {}",
387			main_round_coeffs_vec.len()
388		);
389		let main_round_coeffs = main_round_coeffs_vec.pop().expect("length checked above");
390
391		let mut mask_round_coeffs_vec = mask_prover.execute();
392		let mask_round_coeffs = mask_round_coeffs_vec
393			.pop()
394			.expect("mask prover has 1 claim");
395
396		// The batching below pads the shorter polynomial with zeros, so any coefficient beyond
397		// the mask's degree would reach the channel exactly as the witness produced it.
398		//
399		//     main   : [a_0, a_1, a_2]
400		//     mask   : [b_0, b_1]      -> padded to [b_0, b_1, 0]
401		//     batched: [a_0 + c*b_0, a_1 + c*b_1, a_2]
402		//                                         ^^^ unmasked
403		assert!(
404			mask_round_coeffs.0.len() >= main_round_coeffs.0.len(),
405			"the mask round polynomial has {} coefficients against the main round polynomial's \
406			 {}, so the excess would be sent unmasked",
407			mask_round_coeffs.0.len(),
408			main_round_coeffs.0.len()
409		);
410
411		// Batch the round coefficients: batched = main + batch_challenge * mask
412		let batched_round_coeffs = main_round_coeffs + &(mask_round_coeffs * batch_challenge);
413
414		// Write truncated coefficients to channel
415		channel.send_many(mlecheck::RoundProof::truncate(batched_round_coeffs).coeffs());
416
417		// Sample challenge and fold both provers
418		let challenge = channel.sample();
419		challenges.push(challenge);
420		main_prover.fold(challenge);
421		mask_prover.fold(challenge);
422	}
423
424	// Finish both provers
425	let main_evals = main_prover.finish();
426	let mask_evals = mask_prover.finish();
427	let mask_eval_out = mask_evals[0];
428
429	// Write final mask evaluation
430	channel.send_one(mask_eval_out);
431
432	ProveZKOutput {
433		multilinear_evals: main_evals,
434		mask_eval: mask_eval_out,
435		challenges,
436	}
437}
438
439#[cfg(test)]
440mod tests {
441	use binius_field::arch::OptimalB128;
442	use binius_ip::mlecheck::{self, mask_buffer_dimensions};
443	use binius_math::test_utils::{random_field_buffer, random_scalars};
444	use binius_transcript::{ProverTranscript, fiat_shamir::HasherChallenger};
445
446	type StdChallenger = HasherChallenger<sha2::Sha256>;
447	use rand::prelude::*;
448
449	use super::*;
450	use crate::sumcheck::prove_single_mlecheck;
451
452	type B128 = OptimalB128;
453
454	/// Evaluates the mask polynomial g(X) = sum_i g_i(X_i) at a point using the Mask struct.
455	fn evaluate_mask_polynomial<P: PackedField, Data: Deref<Target = [P]>>(
456		mask: &Mask<P, Data>,
457		point: &[P::Scalar],
458	) -> P::Scalar {
459		iter::zip(0..mask.n_vars(), point)
460			.map(|(i, &x)| mask.evaluate_univariate(i, x))
461			.sum()
462	}
463
464	fn test_mask_prover_with_degree(degree: usize) {
465		let n_vars = 6;
466		let mut rng = StdRng::seed_from_u64(0);
467
468		// Generate random mask buffer
469		let (m_n, m_d) = mask_buffer_dimensions(n_vars, degree, 0);
470		let buffer = random_field_buffer::<B128>(&mut rng, m_n + m_d);
471
472		// Generate random evaluation point
473		let eval_point: Vec<B128> = random_scalars(&mut rng, n_vars);
474
475		// Compute the MLE of the mask polynomial at eval_point
476		let mask = Mask::new(n_vars, degree, buffer.as_view());
477		let eval_claim = mask.evaluate_mle(&eval_point);
478
479		// Create the prover (takes ownership of a borrowed mask view)
480		let prover = MleCheckMaskProver::new(mask, eval_point.clone(), eval_claim);
481
482		// Run the proving protocol
483		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
484		let output = prove_single_mlecheck(prover, &mut prover_transcript);
485
486		// Write the multilinear evaluation to the transcript
487		prover_transcript
488			.message()
489			.write_slice(&output.multilinear_evals);
490
491		// Convert to verifier transcript and run verification
492		let mut verifier_transcript = prover_transcript.into_verifier();
493		let sumcheck_output = mlecheck::verify(
494			&eval_point,
495			degree, // round polynomial degree equals the univariate g_i degree
496			eval_claim,
497			&mut verifier_transcript,
498		)
499		.unwrap();
500
501		// Read the mask evaluation from the transcript
502		let mask_eval_out: B128 = verifier_transcript.message().read().unwrap();
503
504		// Verify the reduced evaluation equals the composition of the evaluations
505		// The mask polynomial is the single "multilinear" here, so its eval should match
506		assert_eq!(mask_eval_out, sumcheck_output.eval);
507
508		// Compute the challenge point (reverse for high-to-low order)
509		let mut challenge_point = sumcheck_output.challenges;
510		challenge_point.reverse();
511
512		// Check that the final evaluation matches direct computation
513		let mask = Mask::new(n_vars, degree, buffer.as_view());
514		let expected_eval = evaluate_mask_polynomial(&mask, &challenge_point);
515		assert_eq!(output.multilinear_evals[0], expected_eval);
516	}
517
518	#[test]
519	fn test_linear_mask() {
520		test_mask_prover_with_degree(1);
521	}
522
523	#[test]
524	fn test_quadratic_mask() {
525		test_mask_prover_with_degree(2);
526	}
527
528	#[test]
529	fn test_cubic_mask() {
530		test_mask_prover_with_degree(3);
531	}
532
533	#[test]
534	fn test_single_variable() {
535		let mut rng = StdRng::seed_from_u64(0);
536
537		// Single variable mask buffer
538		let (m_n, m_d) = mask_buffer_dimensions(1, 2, 0);
539		let buffer = random_field_buffer::<B128>(&mut rng, m_n + m_d);
540
541		let eval_point: Vec<B128> = random_scalars(&mut rng, 1);
542		let mask = Mask::new(1, 2, buffer.as_view());
543		let eval_claim = mask.evaluate_mle(&eval_point);
544
545		let prover = MleCheckMaskProver::new(mask, eval_point.clone(), eval_claim);
546
547		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
548		let output = prove_single_mlecheck(prover, &mut prover_transcript);
549
550		prover_transcript
551			.message()
552			.write_slice(&output.multilinear_evals);
553
554		let mut verifier_transcript = prover_transcript.into_verifier();
555		let sumcheck_output =
556			mlecheck::verify(&eval_point, 2, eval_claim, &mut verifier_transcript).unwrap();
557
558		let mut challenge_point = sumcheck_output.challenges;
559		challenge_point.reverse();
560
561		let mask = Mask::new(1, 2, buffer.as_view());
562		let expected_eval = evaluate_mask_polynomial(&mask, &challenge_point);
563		assert_eq!(output.multilinear_evals[0], expected_eval);
564	}
565
566	fn test_prove_with_degrees(main_degree: usize, mask_degree: usize) {
567		let n_vars = 6;
568		let mut rng = StdRng::seed_from_u64(0);
569
570		// Generate random main mask buffer (using Mask as a simple MleCheckProver for testing)
571		let (m_n, m_d) = mask_buffer_dimensions(n_vars, main_degree, 0);
572		let main_buffer = random_field_buffer::<B128>(&mut rng, m_n + m_d);
573
574		// Generate random ZK mask buffer
575		let (zk_m_n, zk_m_d) = mask_buffer_dimensions(n_vars, mask_degree, 0);
576		let zk_buffer = random_field_buffer::<B128>(&mut rng, zk_m_n + zk_m_d);
577
578		// Generate random evaluation point
579		let eval_point: Vec<B128> = random_scalars(&mut rng, n_vars);
580
581		// Compute the MLE of the main polynomial at eval_point
582		let main_mask = Mask::new(n_vars, main_degree, main_buffer.as_view());
583		let main_eval_claim = main_mask.evaluate_mle(&eval_point);
584
585		// Create the main prover (using MleCheckMaskProver as a simple MleCheckProver)
586		let main_prover = MleCheckMaskProver::new(main_mask, eval_point.clone(), main_eval_claim);
587
588		// Run the ZK proving protocol
589		let zk_mask = Mask::new(n_vars, mask_degree, zk_buffer.as_view());
590		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
591		let output = prove(main_prover, zk_mask, &mut prover_transcript);
592
593		// Write the main polynomial evaluation to the transcript
594		prover_transcript
595			.message()
596			.write_slice(&output.multilinear_evals);
597
598		// Convert to verifier transcript and run ZK verification
599		let mut verifier_transcript = prover_transcript.into_verifier();
600		let mlecheck::VerifyZKOutput {
601			eval,
602			mask_eval,
603			challenges,
604		} = mlecheck::verify_zk(
605			&eval_point,
606			main_degree.max(mask_degree), // batched polynomial degree
607			main_eval_claim,
608			&mut verifier_transcript,
609		)
610		.unwrap();
611
612		// Read the main polynomial evaluation from the transcript
613		let main_eval_out: B128 = verifier_transcript.message().read().unwrap();
614
615		// Verify the reduced evaluation matches
616		assert_eq!(main_eval_out, eval);
617
618		// Compute the challenge point (reverse for high-to-low order)
619		let mut challenge_point = challenges;
620		challenge_point.reverse();
621
622		// Check that the final main evaluation matches direct computation
623		let main_mask = Mask::new(n_vars, main_degree, main_buffer.as_view());
624		let expected_main_eval = evaluate_mask_polynomial(&main_mask, &challenge_point);
625		assert_eq!(output.multilinear_evals[0], expected_main_eval);
626
627		// Check that the mask evaluation matches the zk_mask evaluation at the challenge point
628		let zk_mask = Mask::new(n_vars, mask_degree, zk_buffer.as_view());
629		let expected_mask_eval = evaluate_mask_polynomial(&zk_mask, &challenge_point);
630		assert_eq!(mask_eval, expected_mask_eval);
631	}
632
633	#[test]
634	fn test_prove() {
635		// Both round polynomials carry three coefficients, so every one of them is covered.
636		test_prove_with_degrees(2, 2);
637	}
638
639	#[test]
640	fn test_prove_mask_above_main_degree() {
641		// A mask that overshoots is still hiding.
642		//
643		//     main   : [a_0, a_1]
644		//     mask   : [b_0, b_1, b_2]
645		//     batched: [a_0 + c*b_0, a_1 + c*b_1, c*b_2]
646		//
647		// The verifier is told the batched degree, so the proof still checks out.
648		test_prove_with_degrees(1, 2);
649	}
650
651	#[test]
652	#[should_panic(expected = "would be sent unmasked")]
653	fn test_prove_mask_below_main_degree() {
654		// A mask one degree short leaves the top coefficient of every round in the clear.
655		//
656		//     main   : [a_0, a_1, a_2]
657		//     mask   : [b_0, b_1]      -> zero-padded to [b_0, b_1, 0]
658		//     batched: [a_0 + c*b_0, a_1 + c*b_1, a_2]
659		//                                         ^^^ witness-dependent, unmasked
660		//
661		// The transcript truncates the constant term, so this coefficient is sent every round.
662		test_prove_with_degrees(2, 1);
663	}
664
665	#[test]
666	fn test_libra_eval_inner_product_equals_mask_eval() {
667		use binius_compute::GlobalAllocator;
668		use binius_math::inner_product::inner_product_buffers;
669
670		let mut rng = StdRng::seed_from_u64(0);
671		let n_vars = 6;
672		let degree = 2;
673
674		// Create random mask buffer
675		let (m_n, m_d) = mask_buffer_dimensions(n_vars, degree, 0);
676		let mask_buffer = random_field_buffer::<B128>(&mut rng, m_n + m_d);
677
678		// Create random challenge point
679		let challenge_point: Vec<B128> = random_scalars(&mut rng, n_vars);
680
681		// Compute g(r) using direct computation.evaluate()
682		let mask = Mask::new(n_vars, degree, mask_buffer.as_view());
683		let direct_eval: B128 = (0..n_vars)
684			.map(|i| mask.evaluate_univariate(i, challenge_point[i]))
685			.sum();
686
687		// Compute <g', libra_eval_r> using inner product
688		let libra_eval_tensor = expand_libra_eval::<_, B128>(
689			&GlobalAllocator,
690			&challenge_point,
691			n_vars,
692			degree,
693			m_n,
694			m_d,
695		);
696		let inner_product_eval = inner_product_buffers(&mask_buffer, &libra_eval_tensor);
697
698		assert_eq!(
699			direct_eval, inner_product_eval,
700			"Inner product <g', libra_eval_r> should equal g(r)"
701		);
702	}
703
704	#[test]
705	fn test_libra_eval_sumcheck() {
706		use binius_compute::GlobalAllocator;
707		use binius_ip::{mlecheck::libra_eval, sumcheck::verify};
708		use binius_math::{inner_product::inner_product_par, multilinear::evaluate::evaluate};
709
710		use crate::sumcheck::{bivariate_product_prover, prove_single};
711
712		let mut rng = StdRng::seed_from_u64(0);
713		let n_vars = 6;
714		let degree = 2;
715		let alloc = GlobalAllocator;
716
717		// Create random mask buffer (g')
718		let (m_n, m_d) = mask_buffer_dimensions(n_vars, degree, 0);
719		let log_size = m_n + m_d;
720		let mask_buffer = random_field_buffer::<B128>(&mut rng, log_size);
721
722		// Create random challenge point r
723		let challenge_point: Vec<B128> = random_scalars(&mut rng, n_vars);
724
725		// Generate libra_eval tensor
726		let libra_eval_tensor = expand_libra_eval::<_, B128>(
727			&GlobalAllocator,
728			&challenge_point,
729			n_vars,
730			degree,
731			m_n,
732			m_d,
733		);
734
735		// Compute the claimed sum: <g', libra_eval_r> = g(r)
736		let claimed_sum = inner_product_par(&mask_buffer, &libra_eval_tensor);
737
738		// Create the bivariate product sumcheck prover
739		let prover =
740			bivariate_product_prover(&alloc, [mask_buffer.clone(), libra_eval_tensor], claimed_sum);
741
742		// Run the proving protocol
743		let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
744		let output = prove_single(prover, &mut prover_transcript);
745
746		// Write the multilinear evaluations to the transcript
747		prover_transcript
748			.message()
749			.write_slice(&output.multilinear_evals);
750
751		// Verify with sumcheck verifier
752		let mut verifier_transcript = prover_transcript.into_verifier();
753		let sumcheck_output = verify(
754			log_size,
755			2, // degree 2 for bivariate product
756			claimed_sum,
757			&mut verifier_transcript,
758		)
759		.unwrap();
760
761		// Read the multilinear evaluations
762		let multilinear_evals: Vec<B128> = verifier_transcript.message().read_vec(2).unwrap();
763		let [g_prime_eval, libra_eval_out] = [multilinear_evals[0], multilinear_evals[1]];
764
765		// Check that the product equals the reduced evaluation
766		assert_eq!(g_prime_eval * libra_eval_out, sumcheck_output.eval);
767
768		// Verify libra_eval_out using libra_eval
769		// The sumcheck binds variables high-to-low, so we need to reverse to get low-to-high order
770		// In the buffer layout, the low-order bits encode k (first m_d variables) and
771		// high-order bits encode j (next m_n variables)
772		let mut query_point = sumcheck_output.challenges;
773		query_point.reverse();
774		let (query_k, query_j) = query_point.split_at(m_d);
775
776		let expected_libra_eval_out =
777			libra_eval::<B128>(&challenge_point, query_j, query_k, n_vars, degree);
778		assert_eq!(
779			libra_eval_out, expected_libra_eval_out,
780			"libra_eval should match the sumcheck-reduced evaluation"
781		);
782
783		// Also verify g' evaluation directly
784		let expected_g_prime_eval = evaluate(&mask_buffer, &query_point);
785		assert_eq!(g_prime_eval, expected_g_prime_eval);
786	}
787}