binius_ip/mlecheck.rs
1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use binius_field::{Field, field::FieldOps, util::powers};
5use binius_math::multilinear::eq::eq_ind_partial_eval_scalars;
6
7use crate::{
8 channel::IPVerifierChannel,
9 sumcheck::{self, RoundCoeffs, SumcheckOutput},
10};
11
12/// Returns `(m_n, m_d)` dimensions for a mask polynomial buffer.
13///
14/// The ZK MLE-check protocol uses a separable mask polynomial with `n_vars` univariate polynomials,
15/// each of degree `d`. The mask coefficients are stored in a `2^m_n × 2^m_d` matrix where:
16/// - `m_n`: log of number of rows (one per variable, padded to power of two)
17/// - `m_d`: log of row size (degree + 1 coefficients, padded to power of two)
18///
19/// The protocol imposes `n * d + 1` linear constraints on the coefficients. The `n_extra_dof`
20/// parameter specifies additional degrees of freedom needed (e.g., for FRI openings), which
21/// may increase `m_n` to ensure `2^(m_n + m_d) >= n * d + 1 + n_extra_dof`.
22///
23/// # Arguments
24///
25/// * `n_vars` - Number of variables (n).
26/// * `degree` - Degree of each univariate polynomial (d).
27/// * `n_extra_dof` - Number of additional degrees of freedom.
28///
29/// # Returns
30///
31/// A tuple `(m_n, m_d)` where the total mask buffer size is `2^(m_n + m_d)`.
32pub fn mask_buffer_dimensions(n_vars: usize, degree: usize, n_extra_dof: usize) -> (usize, usize) {
33 let min_buffer_size = n_vars * degree + 1 + n_extra_dof;
34 let m_d = (degree + 1).next_power_of_two().ilog2() as usize;
35 // m_n must be large enough to hold n_vars rows AND satisfy the DOF constraint
36 let m_n_for_vars = n_vars.next_power_of_two().ilog2() as usize;
37 let m_n_for_size = (min_buffer_size.next_power_of_two().ilog2() as usize).saturating_sub(m_d);
38 let m_n = m_n_for_vars.max(m_n_for_size);
39 (m_n, m_d)
40}
41
42/// Output of the zero-knowledge MLE-check verification.
43#[derive(Debug, Clone, PartialEq, Eq)]
44pub struct VerifyZKOutput<F> {
45 /// The reduced evaluation of the main polynomial at the challenge point.
46 pub eval: F,
47 /// The evaluation of the mask polynomial at the challenge point.
48 pub mask_eval: F,
49 /// The sequence of challenge values from each round.
50 pub challenges: Vec<F>,
51}
52
53/// An MLE-check protocol is an interactive protocol similar to sumcheck, but with modifications
54/// introduced in [Gruen24], Section 3.
55///
56/// The prover in an MLE-check argues a that for some $n$-variate polynomial
57/// $F(X_0, \ldots, X_{n-1})$ (which is not necessarily multilinear), for a given point
58/// $(z_0, \ldots, z_{n-1})$ and claimed value $s$, that
59///
60/// $$
61/// s = \sum_{v \in B_n} F(v) \cdot eq(v, z)
62/// $$
63///
64/// Unless $F$ is indeed multilinear, $s \ne F(z)$ necessarily. While the prover and verifier could
65/// engage in a standard sumcheck protocol to reduce this claim, it is concretely more efficient to
66/// use the optimized protocol from [Gruen24], which we call an "MLE-check".
67///
68/// [Gruen24]: <https://eprint.iacr.org/2024/108>
69///
70/// ## Arguments
71///
72/// * `point` - The evaluation point for the multilinear extension
73/// * `degree` - The degree of the univariate polynomial in each round
74/// * `eval` - The claimed multilinear-extension evaluation of the multivariate polynomial
75/// * `channel` - The channel for receiving prover messages and sampling challenges
76///
77/// ## Returns
78///
79/// Returns a `Result` containing the `SumcheckOutput` with the reduced evaluation and challenge
80/// point, or an error if verification fails.
81pub fn verify<F, C>(
82 point: &[C::Elem],
83 degree: usize,
84 mut eval: C::Elem,
85 channel: &mut C,
86) -> Result<SumcheckOutput<C::Elem>, sumcheck::Error>
87where
88 F: Field,
89 C: IPVerifierChannel<F>,
90{
91 let n_vars = point.len();
92
93 let mut challenges = Vec::with_capacity(n_vars);
94 for z_i in point.iter().rev() {
95 let round_proof = RoundProof(RoundCoeffs(channel.recv_many(degree)?));
96 let challenge = channel.sample();
97
98 let round_coeffs = round_proof.recover(eval, z_i.clone());
99 eval = round_coeffs.evaluate(&challenge);
100 challenges.push(challenge);
101 }
102
103 Ok(SumcheckOutput { eval, challenges })
104}
105
106/// Variation of the MLE-check protocol that provides the hiding property.
107///
108/// This protocol is based on the zero-knowledge sumcheck technique from [Libra], with a
109/// modification. When the field has characteristic 2, the Libra ZK-sumcheck protocol is not hiding.
110/// Instead, the mask polynomial $g$ is batched together with the multivariate polynomial whose MLE
111/// is being evaluated, whereas Libra would batch the mask polynomial together with the MLE itself.
112///
113/// [Libra]: <https://dl.acm.org/doi/10.1007/978-3-030-26954-8_24>
114pub fn verify_zk<F, C>(
115 point: &[C::Elem],
116 degree: usize,
117 eval: C::Elem,
118 channel: &mut C,
119) -> Result<VerifyZKOutput<C::Elem>, sumcheck::Error>
120where
121 F: Field,
122 C: IPVerifierChannel<F>,
123{
124 // Read the evaluation of the MLE of the mask polynomial (g).
125 let mask_eval = channel.recv_one()?;
126
127 // Randomly mix the evaluation claim with the mask evaluation claim.
128 let batch_challenge = channel.sample();
129 let batch_eval = eval + batch_challenge.clone() * mask_eval;
130
131 let SumcheckOutput {
132 eval: batch_eval_out,
133 challenges,
134 } = verify(point, degree, batch_eval, channel)?;
135
136 // Read the evaluation of the mask polynomial (g) at the sumcheck challenge point.
137 let mask_eval_out = channel.recv_one()?;
138
139 let eval_out = batch_eval_out - batch_challenge * mask_eval_out.clone();
140 Ok(VerifyZKOutput {
141 eval: eval_out,
142 mask_eval: mask_eval_out,
143 challenges,
144 })
145}
146
147/// An MLE-check round proof is a univariate polynomial in monomial basis with the coefficient of
148/// the lowest-degree term truncated off.
149///
150/// Since the verifier knows the claimed linear extrapolation of the polynomial values at the
151/// points 0 and 1, the low-degree term coefficient can be easily recovered. Truncating the
152/// coefficient off saves a small amount of proof data.
153///
154/// This is an analogous struct to [`sumcheck::RoundProof`], except that we truncate the low-degree
155/// coefficient instead of the high-degree coefficient.
156///
157/// In a sumcheck protocol, the verifier has a claimed sum $s$ and the round polynomial $R(X)$ must
158/// satisfy $R(0) + R(1) = s$. In an MLE-check protocol, the verifier has a claimed coordinate
159/// $\alpha$ and extrapolated value $s$ and the round polynomial must satisfy
160/// $(1 - \alpha) R(0) + \alpha R(1) = s$. This difference changes the recovery procedure and which
161/// polynomial coefficient is most convenient to truncate.
162#[derive(Debug, Default, Clone, PartialEq, Eq)]
163pub struct RoundProof<F>(pub RoundCoeffs<F>);
164
165impl<F> RoundProof<F> {
166 /// Truncates the polynomial coefficients to a round proof.
167 ///
168 /// Removes the first coefficient. See the struct documentation for more info.
169 ///
170 /// ## Pre-conditions
171 ///
172 /// * `coeffs` must not be empty
173 pub fn truncate(mut coeffs: RoundCoeffs<F>) -> Self {
174 coeffs.0.remove(0);
175 Self(coeffs)
176 }
177
178 /// Recovers all univariate polynomial coefficients from the compressed round proof.
179 ///
180 /// The prover has sent coefficients for the purported $i$'th round polynomial
181 /// $R(X) = \sum_{j=0}^d a_j * X^j$.
182 ///
183 /// However, the prover has not sent the lowest degree coefficient $a_0$. The verifier will
184 /// need to recover this missing coefficient.
185 ///
186 /// Let $s$ denote the current round's claimed sum and $\alpha_i$ be the $i$'th coordinate of
187 /// the evaluation point.
188 ///
189 /// The verifier expects the round polynomial $R_i$ to satisfy the identity
190 /// $s = (1 - \alpha) R(0) + \alpha R(1)$, or equivalently, $s = R(0) + (R(1) - R(0)) \alpha$.
191 ///
192 /// Using
193 /// $R(0) = a_0$
194 /// $R(1) = \sum_{j=0}^d a_j$
195 /// There is a unique $a_0$ that allows $R$ to satisfy the above identity. Specifically,
196 /// $a_0 = s - \alpha \sum_{j=1}^d a_j$.
197 pub fn recover(self, eval: F, alpha: F) -> RoundCoeffs<F>
198 where
199 F: FieldOps,
200 {
201 let Self(RoundCoeffs(mut coeffs)) = self;
202 let first_coeff = eval - alpha * coeffs.iter().cloned().sum::<F>();
203 coeffs.insert(0, first_coeff);
204 RoundCoeffs(coeffs)
205 }
206
207 /// The truncated polynomial coefficients.
208 pub fn coeffs(&self) -> &[F] {
209 &self.0.0
210 }
211}
212
213/// Evaluates the MLE of the libra_eval polynomial at a query point.
214///
215/// Computes the multilinear extension of `libra_eval_r` at the query point `(query_j, query_k)`:
216///
217/// ```text
218/// Σⱼ Σₖ eq(j, query_j) · eq(k, query_k) · r[j]^k
219/// ```
220///
221/// for `j < n_vars` and `k ≤ degree`, where `eq` is the equality indicator polynomial.
222///
223/// # Arguments
224///
225/// * `challenge_point` - The challenge point `r` from sumcheck (length `n_vars`)
226/// * `query_j` - Query point for the variable index (length `m_n`)
227/// * `query_k` - Query point for the power index (length `m_d`)
228/// * `n_vars` - Number of variables in the mask polynomial
229/// * `degree` - Degree of each univariate in the mask polynomial
230pub fn libra_eval<F: FieldOps>(
231 challenge_point: &[F],
232 query_j: &[F],
233 query_k: &[F],
234 n_vars: usize,
235 degree: usize,
236) -> F {
237 let eq_j = eq_ind_partial_eval_scalars(query_j);
238 let eq_k = eq_ind_partial_eval_scalars(query_k);
239
240 eq_j.iter()
241 .take(n_vars)
242 .zip(challenge_point)
243 .map(|(eq_j_val, r_j)| {
244 eq_k.iter()
245 .take(degree + 1)
246 .zip(powers(r_j.clone()))
247 .map(|(eq_k_val, r_j_power)| eq_j_val.clone() * eq_k_val.clone() * r_j_power)
248 .sum::<F>()
249 })
250 .sum()
251}
252
253#[cfg(test)]
254mod tests {
255 use binius_field::{Random, arch::OptimalB128 as B128};
256 use binius_math::{line::extrapolate_line, test_utils::random_scalars};
257 use rand::prelude::*;
258
259 use super::*;
260
261 fn test_recover_with_degree<F: Field>(mut rng: impl Rng, alpha: F, degree: usize) {
262 let coeffs = RoundCoeffs(random_scalars(&mut rng, degree + 1));
263
264 let v0 = coeffs.evaluate(&F::ZERO);
265 let v1 = coeffs.evaluate(&F::ONE);
266 let eval = extrapolate_line(v0, v1, alpha);
267
268 let proof = RoundProof::truncate(coeffs.clone());
269 assert_eq!(proof.recover(eval, alpha), coeffs);
270 }
271
272 #[test]
273 fn test_recover() {
274 let mut rng = StdRng::seed_from_u64(0);
275 let alpha = B128::random(&mut rng);
276
277 for degree in 0..4 {
278 // Test with random coordinate
279 test_recover_with_degree(&mut rng, alpha, degree);
280
281 // Test edge case coordinate values
282 test_recover_with_degree(&mut rng, B128::ZERO, degree);
283 test_recover_with_degree(&mut rng, B128::ONE, degree);
284 }
285 }
286}