1use 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#[derive(Debug, Clone)]
25pub struct ProveZKOutput<F: Field> {
26 pub multilinear_evals: Vec<F>,
28 pub mask_eval: F,
30 pub challenges: Vec<F>,
32}
33
34pub 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
84pub struct Mask<P: PackedField, Data: Deref<Target = [P]> = Box<[P]>> {
99 n_vars: usize,
101 degree: usize,
103 buffer: FieldBuffer<P, Data>,
107}
108
109impl<F: Field, P: PackedField<Scalar = F>, Data: Deref<Target = [P]>> Mask<P, Data> {
110 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 pub const fn n_vars(&self) -> usize {
127 self.n_vars
128 }
129
130 pub const fn degree(&self) -> usize {
132 self.degree
133 }
134
135 const fn log_degree_plus_one(&self) -> usize {
137 (self.degree + 1).next_power_of_two().ilog2() as usize
138 }
139
140 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 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 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 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
192pub struct MleCheckMaskProver<F: Field, P: PackedField<Scalar = F>, Data: Deref<Target = [P]>> {
199 mask: Mask<P, Data>,
201 eval_point: Vec<F>,
203 n_vars_remaining: usize,
205 prefix_sum: F,
207 suffix_sums: Vec<F>,
209 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 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 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 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 let suffix_sum: F = self.suffix_sums[..var_idx].iter().copied().sum();
277
278 let constant_offset = self.prefix_sum + suffix_sum;
280
281 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 let new_claim = coeffs.evaluate(&challenge);
301
302 let var_idx = self.current_var_index();
303
304 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 vec![self.prefix_sum]
317 }
318
319 fn eval_point(&self) -> &[F] {
320 &self.eval_point[..self.n_vars_remaining]
323 }
324}
325
326pub 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 let mask_eval = mask.evaluate_mle(&eval_point);
371 channel.send_one(mask_eval);
372
373 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 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 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 let batched_round_coeffs = main_round_coeffs + &(mask_round_coeffs * batch_challenge);
413
414 channel.send_many(mlecheck::RoundProof::truncate(batched_round_coeffs).coeffs());
416
417 let challenge = channel.sample();
419 challenges.push(challenge);
420 main_prover.fold(challenge);
421 mask_prover.fold(challenge);
422 }
423
424 let main_evals = main_prover.finish();
426 let mask_evals = mask_prover.finish();
427 let mask_eval_out = mask_evals[0];
428
429 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 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 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 let eval_point: Vec<B128> = random_scalars(&mut rng, n_vars);
474
475 let mask = Mask::new(n_vars, degree, buffer.as_view());
477 let eval_claim = mask.evaluate_mle(&eval_point);
478
479 let prover = MleCheckMaskProver::new(mask, eval_point.clone(), eval_claim);
481
482 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
484 let output = prove_single_mlecheck(prover, &mut prover_transcript);
485
486 prover_transcript
488 .message()
489 .write_slice(&output.multilinear_evals);
490
491 let mut verifier_transcript = prover_transcript.into_verifier();
493 let sumcheck_output = mlecheck::verify(
494 &eval_point,
495 degree, eval_claim,
497 &mut verifier_transcript,
498 )
499 .unwrap();
500
501 let mask_eval_out: B128 = verifier_transcript.message().read().unwrap();
503
504 assert_eq!(mask_eval_out, sumcheck_output.eval);
507
508 let mut challenge_point = sumcheck_output.challenges;
510 challenge_point.reverse();
511
512 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 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 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 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 let eval_point: Vec<B128> = random_scalars(&mut rng, n_vars);
580
581 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 let main_prover = MleCheckMaskProver::new(main_mask, eval_point.clone(), main_eval_claim);
587
588 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 prover_transcript
595 .message()
596 .write_slice(&output.multilinear_evals);
597
598 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), main_eval_claim,
608 &mut verifier_transcript,
609 )
610 .unwrap();
611
612 let main_eval_out: B128 = verifier_transcript.message().read().unwrap();
614
615 assert_eq!(main_eval_out, eval);
617
618 let mut challenge_point = challenges;
620 challenge_point.reverse();
621
622 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 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 test_prove_with_degrees(2, 2);
637 }
638
639 #[test]
640 fn test_prove_mask_above_main_degree() {
641 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 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 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 let challenge_point: Vec<B128> = random_scalars(&mut rng, n_vars);
680
681 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 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 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 let challenge_point: Vec<B128> = random_scalars(&mut rng, n_vars);
724
725 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 let claimed_sum = inner_product_par(&mask_buffer, &libra_eval_tensor);
737
738 let prover =
740 bivariate_product_prover(&alloc, [mask_buffer.clone(), libra_eval_tensor], claimed_sum);
741
742 let mut prover_transcript = ProverTranscript::new(StdChallenger::default());
744 let output = prove_single(prover, &mut prover_transcript);
745
746 prover_transcript
748 .message()
749 .write_slice(&output.multilinear_evals);
750
751 let mut verifier_transcript = prover_transcript.into_verifier();
753 let sumcheck_output = verify(
754 log_size,
755 2, claimed_sum,
757 &mut verifier_transcript,
758 )
759 .unwrap();
760
761 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 assert_eq!(g_prime_eval * libra_eval_out, sumcheck_output.eval);
767
768 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 let expected_g_prime_eval = evaluate(&mask_buffer, &query_point);
785 assert_eq!(g_prime_eval, expected_g_prime_eval);
786 }
787}