binius_ip/sumcheck/common.rs
1// Copyright 2023-2025 Irreducible Inc.
2
3use std::ops::{Add, AddAssign, Index, Mul, MulAssign};
4
5use binius_field::{Field, field::FieldOps};
6use binius_math::univariate::evaluate_univariate;
7
8/// A univariate polynomial in monomial basis.
9///
10/// The coefficient at position `i` in the inner vector corresponds to the term $X^i$.
11#[derive(Debug, Clone, PartialEq, Eq)]
12pub struct RoundCoeffs<F>(pub Vec<F>);
13
14// The empty coefficient vector is the zero polynomial, a valid default for any element type.
15// Deriving the default would instead demand that every element type have its own default value.
16impl<F> Default for RoundCoeffs<F> {
17 fn default() -> Self {
18 // The zero polynomial has no coefficients.
19 Self(Vec::new())
20 }
21}
22
23impl<F> RoundCoeffs<F> {
24 /// Truncate the highest-degree coefficient to produce a more compact round proof.
25 ///
26 /// # Pre-conditions
27 ///
28 /// - The coefficient vector must be non-empty.
29 /// - A round polynomial always has degree at least one, so an empty vector signals a bug.
30 pub fn truncate(mut self) -> RoundProof<F> {
31 // Drop the highest-degree coefficient; the verifier reconstructs it from the claimed sum.
32 self.0.pop();
33 RoundProof(self)
34 }
35
36 /// The coefficients ordered from the constant term to the highest-degree term.
37 pub fn as_slice(&self) -> &[F] {
38 // Position `i` holds the coefficient of X^i, so the constant term comes first.
39 &self.0
40 }
41}
42
43impl<F: FieldOps> RoundCoeffs<F> {
44 /// Evaluate the polynomial at a point.
45 pub fn evaluate(&self, x: &F) -> F {
46 // Horner's method over the monomial coefficients.
47 evaluate_univariate(&self.0, x)
48 }
49
50 /// Batches round polynomials into one, weighting polynomial `i` by `batch_coeff^i`.
51 ///
52 /// The verifier weights the matching claims with the same coefficient.
53 /// Each claim therefore stays tied to its own round polynomial.
54 ///
55 /// An empty input is the zero polynomial.
56 pub fn batch(polys: Vec<Self>, batch_coeff: &F) -> Self {
57 // Horner from the highest weight down: acc <- acc * batch_coeff + poly.
58 polys
59 .into_iter()
60 .rfold(Self::default(), |acc, poly| acc * batch_coeff.clone() + &poly)
61 }
62
63 /// The endpoint values $(R(0), R(1))$ of the round polynomial.
64 ///
65 /// $R(0)$ is the constant coefficient $a_0$.
66 /// $R(1) = \sum_j a_j$ is the sum of all coefficients.
67 /// An empty coefficient vector is the zero polynomial, whose endpoints are both zero.
68 fn endpoints(&self) -> (F, F) {
69 // R(0) is the constant term, or zero when there are no coefficients.
70 let r_0 = self.0.first().cloned().unwrap_or_else(F::zero);
71 // R(1) is the polynomial at one, which sums every coefficient.
72 let r_1 = self.0.iter().cloned().sum();
73 (r_0, r_1)
74 }
75
76 /// The claimed sum $R(0) + R(1)$ that this round polynomial encodes.
77 ///
78 /// For a sumcheck round polynomial, this is the round's claimed sum.
79 /// The verifier expects the identity $s = R(0) + R(1)$ (see [`RoundProof::recover`]).
80 pub fn sum_over_endpoints(&self) -> F {
81 // The sumcheck claim is s = R(0) + R(1).
82 let (r_0, r_1) = self.endpoints();
83 r_0 + r_1
84 }
85
86 /// The claimed value $(1 - \alpha) R(0) + \alpha R(1)$ that this round polynomial encodes in an
87 /// MLE-check.
88 ///
89 /// This is the MLE-check analogue of [`Self::sum_over_endpoints`].
90 /// An MLE-check round polynomial satisfies $s = (1 - \alpha) R(0) + \alpha R(1)$.
91 /// Here $\alpha$ is the round's evaluation-point coordinate (see
92 /// [`crate::mlecheck::RoundProof::recover`]). Equivalently, this is the linear extrapolation
93 /// of $R$ from the endpoints $0$ and $1$ to $\alpha$.
94 pub fn lerp_over_endpoints(&self, alpha: F) -> F {
95 let (r_0, r_1) = self.endpoints();
96 // Line through the endpoints, sampled at alpha: R(0) + alpha * (R(1) - R(0)).
97 r_0.clone() + alpha * (r_1 - r_0)
98 }
99}
100
101impl<F: Field> RoundCoeffs<F> {
102 /// Multiplies this polynomial by the equality factor $\text{eq}(X, \alpha)$.
103 ///
104 /// $$
105 /// \text{eq}(X, \alpha) = (1 - \alpha) + (2 \alpha - 1) X
106 /// $$
107 ///
108 /// An MLE-check prover interpolates the prime polynomial, which carries no equality factor.
109 /// This multiplies the factor back in.
110 ///
111 /// Monomial form makes that one scaling and one shift.
112 /// Sampling the factored polynomial instead would cost an extra evaluation point.
113 ///
114 /// The factor is linear, so the result has one more coefficient than `self`.
115 pub fn mul_by_eq(&self, alpha: F) -> Self {
116 // NB: In characteristic 2, eq(X, alpha) simplifies to (1 + alpha) + X.
117 let (by_constant_term, mut by_linear_term) = if F::CHARACTERISTIC == 2 {
118 (self.clone() * (F::ONE + alpha), self.clone())
119 } else {
120 (self.clone() * (F::ONE - alpha), self.clone() * (alpha.double() - F::ONE))
121 };
122
123 // Prepending a zero coefficient multiplies the polynomial by X.
124 by_linear_term.0.insert(0, F::ZERO);
125 by_constant_term + &by_linear_term
126 }
127}
128
129impl<F: FieldOps> Add<&Self> for RoundCoeffs<F> {
130 type Output = Self;
131
132 fn add(mut self, rhs: &Self) -> Self::Output {
133 // Reuse the in-place addition and return the grown accumulator.
134 self += rhs;
135 self
136 }
137}
138
139impl<F: FieldOps> AddAssign<&Self> for RoundCoeffs<F> {
140 fn add_assign(&mut self, rhs: &Self) {
141 // The two polynomials may have different degrees, hence different coefficient counts.
142 // Pad the shorter accumulator with zeros so every addend coefficient has a partner.
143 if self.0.len() < rhs.0.len() {
144 self.0.resize(rhs.0.len(), F::zero());
145 }
146
147 // Add coefficient by coefficient at matching degrees.
148 for (lhs_i, rhs_i) in self.0.iter_mut().zip(rhs.0.iter()) {
149 *lhs_i += rhs_i;
150 }
151 }
152}
153
154impl<F: FieldOps> Mul<F> for RoundCoeffs<F> {
155 type Output = Self;
156
157 fn mul(mut self, rhs: F) -> Self::Output {
158 // Reuse the in-place scaling.
159 self *= rhs;
160 self
161 }
162}
163
164impl<F: FieldOps> MulAssign<F> for RoundCoeffs<F> {
165 fn mul_assign(&mut self, rhs: F) {
166 // Scaling a polynomial by a constant scales every coefficient.
167 for coeff in &mut self.0 {
168 *coeff *= &rhs;
169 }
170 }
171}
172
173impl<F: FieldOps> std::iter::Sum for RoundCoeffs<F> {
174 fn sum<I: Iterator<Item = Self>>(iter: I) -> Self {
175 // Start from the zero polynomial and accumulate by polynomial addition.
176 iter.fold(Self::default(), |acc, x| acc + &x)
177 }
178}
179
180impl<F> Index<usize> for RoundCoeffs<F> {
181 type Output = F;
182
183 fn index(&self, index: usize) -> &F {
184 // Position `i` selects the coefficient of X^i.
185 &self.0[index]
186 }
187}
188
189/// A sumcheck round proof is a univariate polynomial in monomial basis with the coefficient of the
190/// highest-degree term truncated off.
191///
192/// Since the verifier knows the claimed sum of the polynomial values at the points 0 and 1, the
193/// high-degree term coefficient can be easily recovered. Truncating the coefficient off saves a
194/// small amount of proof data.
195#[derive(Debug, Clone, PartialEq, Eq)]
196pub struct RoundProof<F>(pub RoundCoeffs<F>);
197
198// The empty proof is a valid default for any element type.
199// Deriving the default would instead demand that every element type have its own default value.
200impl<F> Default for RoundProof<F> {
201 fn default() -> Self {
202 // Mirror the empty (zero) polynomial default.
203 Self(RoundCoeffs::default())
204 }
205}
206
207impl<F> RoundProof<F> {
208 /// The truncated polynomial coefficients.
209 pub fn coeffs(&self) -> &[F] {
210 // Every coefficient except the truncated highest-degree term.
211 self.0.as_slice()
212 }
213}
214
215impl<F: FieldOps> RoundProof<F> {
216 /// Recovers all univariate polynomial coefficients from the compressed round proof.
217 ///
218 /// The prover has sent coefficients for the purported ith round polynomial
219 /// $r_i(X) = \sum_{j=0}^d a_j * X^j$.
220 /// However, the prover has not sent the highest degree coefficient $a_d$.
221 /// The verifier will need to recover this missing coefficient.
222 ///
223 /// Let $s$ denote the current round's claimed sum.
224 /// The verifier expects the round polynomial $r_i$ to satisfy the identity
225 /// $s = r_i(0) + r_i(1)$.
226 /// Using
227 /// $r_i(0) = a_0$
228 /// $r_i(1) = \sum_{j=0}^d a_j$
229 /// There is a unique $a_d$ that allows $r_i$ to satisfy the above identity.
230 /// Specifically
231 /// $a_d = s - a_0 - \sum_{j=0}^{d-1} a_j$
232 ///
233 /// Not sending the whole round polynomial is an optimization.
234 /// In the unoptimized version of the protocol, the verifier will halt and reject
235 /// if given a round polynomial that does not satisfy the above identity.
236 pub fn recover(self, sum: F) -> RoundCoeffs<F> {
237 let Self(RoundCoeffs(mut coeffs)) = self;
238 // The received coefficients are a_0 through a_{d-1}; the top term a_d was truncated.
239 // The claimed sum expands to s = R(0) + R(1) = a_0 + (a_0 + a_1 + ... + a_d).
240 // Solving for the missing term gives a_d = s - a_0 - (a_0 + a_1 + ... + a_{d-1}).
241 //
242 // The constant term a_0 is subtracted twice on purpose.
243 // The two copies cancel in characteristic 2 yet keep the identity correct over any field.
244
245 // a_0, or zero for the degenerate empty proof.
246 let first_coeff = coeffs.first().cloned().unwrap_or_else(F::zero);
247 // a_d = s - a_0 - sum_{j=0}^{d-1} a_j.
248 let last_coeff = sum - first_coeff - coeffs.iter().cloned().sum::<F>();
249 // Append the recovered top coefficient to rebuild the full polynomial.
250 coeffs.push(last_coeff);
251 RoundCoeffs(coeffs)
252 }
253}
254
255#[cfg(test)]
256mod tests {
257 use binius_field::{Field, Random, arch::OptimalB128 as B128};
258 use binius_math::{multilinear::eq::eq_one_var, test_utils::random_scalars};
259 use rand::prelude::*;
260
261 use super::*;
262
263 // Sumcheck round polynomials have small degree.
264 // Degrees 1 through 4 cover the range the round provers in this workspace produce.
265 const DEGREES: [usize; 4] = [1, 2, 3, 4];
266
267 // Deterministic RNG seeded to a fixed value so any failure reproduces exactly.
268 fn rng() -> StdRng {
269 StdRng::seed_from_u64(0)
270 }
271
272 #[test]
273 fn recover_round_trips_with_the_claimed_sum() {
274 let mut rng = rng();
275 // Invariant: truncating the top coefficient then recovering it rebuilds the polynomial,
276 // as long as recovery is given the true claimed sum s = R(0) + R(1).
277 for degree in DEGREES {
278 // A random round polynomial with degree + 1 coefficients.
279 let coeffs = RoundCoeffs(random_scalars::<B128>(&mut rng, degree + 1));
280
281 // The verifier only learns the claimed sum s = R(0) + R(1).
282 let sum = coeffs.sum_over_endpoints();
283 // The proof drops the highest-degree coefficient to save transcript space.
284 let proof = coeffs.clone().truncate();
285
286 // Stripping one coefficient leaves exactly `degree` of them.
287 assert_eq!(proof.coeffs().len(), degree);
288 // The claimed sum uniquely determines the missing coefficient, so recovery is exact.
289 assert_eq!(proof.recover(sum), coeffs);
290 }
291 }
292
293 #[test]
294 fn recovered_polynomial_satisfies_the_sumcheck_identity() {
295 let mut rng = rng();
296 // Invariant: the recovered polynomial must satisfy the verifier's check s = R(0) + R(1).
297 for degree in DEGREES {
298 let coeffs = RoundCoeffs(random_scalars::<B128>(&mut rng, degree + 1));
299 // Claimed sum taken from the original polynomial.
300 let sum = coeffs.sum_over_endpoints();
301 // Round-trip through truncation and recovery.
302 let recovered = coeffs.truncate().recover(sum);
303
304 // Evaluate the recovered polynomial at both endpoints and confirm they sum to s.
305 assert_eq!(recovered.evaluate(&B128::ZERO) + recovered.evaluate(&B128::ONE), sum);
306 }
307 }
308
309 #[test]
310 fn batch_commutes_with_evaluation() {
311 let mut rng = rng();
312 // Invariant: the prover's batched polynomial agrees with the verifier's batched claim.
313 //
314 // batch(R_0, .., R_{n-1})(x) == sum_i batch_coeff^i * R_i(x)
315 //
316 // The prover folds the round polynomials, the verifier folds the claim scalars.
317 // A mismatch would send a batched round proof the verifier cannot reproduce.
318 for degree in DEGREES {
319 // Mixed lengths, since batched provers need not all reach the same degree.
320 let polys = (1..=degree)
321 .map(|len| RoundCoeffs(random_scalars::<B128>(&mut rng, len + 1)))
322 .collect::<Vec<_>>();
323 let batch_coeff = B128::random(&mut rng);
324
325 let batched = RoundCoeffs::batch(polys.clone(), &batch_coeff);
326
327 // The batched polynomial reaches the degree of the longest input.
328 assert_eq!(batched.0.len(), degree + 1);
329 for x in random_scalars::<B128>(&mut rng, 4) {
330 // The verifier's side: Horner-fold the per-claim evaluations.
331 let evals = polys
332 .iter()
333 .map(|poly| poly.evaluate(&x))
334 .collect::<Vec<_>>();
335 assert_eq!(batched.evaluate(&x), evaluate_univariate(&evals, &batch_coeff));
336 }
337 }
338 }
339
340 #[test]
341 fn batching_nothing_gives_the_zero_polynomial() {
342 // A batched group with no provers must contribute nothing to the round proof.
343 let batched = RoundCoeffs::batch(Vec::new(), &B128::random(&mut rng()));
344 assert_eq!(batched, RoundCoeffs::<B128>::default());
345 }
346
347 #[test]
348 fn mul_by_eq_multiplies_pointwise() {
349 let mut rng = rng();
350 // Invariant: scaling by the equality factor is exact at every point.
351 //
352 // (R * eq)(x) == R(x) * eq(x, alpha)
353 //
354 // The prover interpolates the prime polynomial and multiplies the factor back in here,
355 // so a discrepancy would send a round polynomial the verifier cannot reproduce.
356 for degree in DEGREES {
357 let coeffs = RoundCoeffs(random_scalars::<B128>(&mut rng, degree + 1));
358 let alpha = B128::random(&mut rng);
359
360 let scaled = coeffs.mul_by_eq(alpha);
361
362 // The equality factor is linear, so the product gains exactly one degree.
363 assert_eq!(scaled.0.len(), coeffs.0.len() + 1);
364 for x in random_scalars::<B128>(&mut rng, 4) {
365 let eq_at_x = eq_one_var(x, alpha);
366 assert_eq!(scaled.evaluate(&x), coeffs.evaluate(&x) * eq_at_x);
367 }
368 }
369 }
370
371 #[test]
372 fn sum_over_endpoints_equals_evaluation_at_zero_and_one() {
373 let mut rng = rng();
374 // Invariant: the endpoint-sum shortcut equals R(0) + R(1) computed by direct evaluation.
375 for degree in DEGREES {
376 let coeffs = RoundCoeffs(random_scalars::<B128>(&mut rng, degree + 1));
377 // Direct evaluation at the two endpoints.
378 let expected = coeffs.evaluate(&B128::ZERO) + coeffs.evaluate(&B128::ONE);
379 // The shortcut must agree with direct evaluation.
380 assert_eq!(coeffs.sum_over_endpoints(), expected);
381 }
382 }
383
384 #[test]
385 fn lerp_over_endpoints_is_the_line_through_the_endpoints() {
386 let mut rng = rng();
387 // Invariant: the extrapolation is the straight line joining R(0) and R(1),
388 // sampled at the round coordinate alpha.
389 for degree in DEGREES {
390 let coeffs = RoundCoeffs(random_scalars::<B128>(&mut rng, degree + 1));
391 let alpha = B128::random(&mut rng);
392
393 // The two endpoint values that define the line.
394 let r_0 = coeffs.evaluate(&B128::ZERO);
395 let r_1 = coeffs.evaluate(&B128::ONE);
396 // The point on that line at alpha: R(0) + alpha * (R(1) - R(0)).
397 let expected = r_0 + alpha * (r_1 - r_0);
398
399 // The helper must reproduce that point.
400 assert_eq!(coeffs.lerp_over_endpoints(alpha), expected);
401 // alpha = 0 lands exactly on R(0).
402 assert_eq!(coeffs.lerp_over_endpoints(B128::ZERO), r_0);
403 // alpha = 1 lands exactly on R(1).
404 assert_eq!(coeffs.lerp_over_endpoints(B128::ONE), r_1);
405 }
406 }
407
408 #[test]
409 fn addition_is_pointwise_across_ragged_lengths() {
410 let mut rng = rng();
411 // Invariant: adding polynomials adds their evaluations at every point,
412 // even when the two operands have different degrees.
413 //
414 // Fixture: a degree-4 polynomial (5 coefficients) and a degree-1 polynomial (2).
415 //
416 // long : [a_0, a_1, a_2, a_3, a_4]
417 // short: [b_0, b_1]
418 // sum : [a_0+b_0, a_1+b_1, a_2, a_3, a_4]
419 let long = RoundCoeffs(random_scalars::<B128>(&mut rng, 5));
420 let short = RoundCoeffs(random_scalars::<B128>(&mut rng, 2));
421 let x = B128::random(&mut rng);
422
423 // Longer accumulator, shorter addend: the accumulator keeps its high-degree tail.
424 assert_eq!((long.clone() + &short).evaluate(&x), long.evaluate(&x) + short.evaluate(&x));
425 // Shorter accumulator, longer addend: the accumulator is padded up to the larger degree.
426 assert_eq!((short.clone() + &long).evaluate(&x), short.evaluate(&x) + long.evaluate(&x));
427 }
428
429 #[test]
430 fn scaling_multiplies_every_evaluation() {
431 let mut rng = rng();
432 // Invariant: scaling a polynomial by a constant scales its value at every point.
433 let coeffs = RoundCoeffs(random_scalars::<B128>(&mut rng, 4));
434 let scalar = B128::random(&mut rng);
435 let x = B128::random(&mut rng);
436
437 // (c * R)(x) must equal c * R(x).
438 assert_eq!((coeffs.clone() * scalar).evaluate(&x), coeffs.evaluate(&x) * scalar);
439 }
440
441 #[test]
442 fn sum_of_round_coeffs_adds_the_polynomials() {
443 let mut rng = rng();
444 // Invariant: summing polynomials sums their evaluations at every point.
445 // Fixture: three degree-2 polynomials (3 coefficients each).
446 let parts: Vec<_> = (0..3)
447 .map(|_| RoundCoeffs(random_scalars::<B128>(&mut rng, 3)))
448 .collect();
449 let x = B128::random(&mut rng);
450
451 // Fold the polynomials together with the iterator sum.
452 let total: RoundCoeffs<B128> = parts.iter().cloned().sum();
453 // Reference: add the individual evaluations at x.
454 let expected: B128 = parts.iter().map(|c| c.evaluate(&x)).sum();
455
456 assert_eq!(total.evaluate(&x), expected);
457 }
458
459 #[test]
460 fn empty_round_coeffs_is_the_zero_polynomial() {
461 // A polynomial with no coefficients is the zero polynomial.
462 let coeffs = RoundCoeffs::<B128>(vec![]);
463 // It evaluates to zero everywhere.
464 assert_eq!(coeffs.evaluate(&B128::random(rng())), B128::ZERO);
465 // Both endpoints are zero, so their sum is zero.
466 assert_eq!(coeffs.sum_over_endpoints(), B128::ZERO);
467 // The line through (0, 0) and (1, 0) is zero at every alpha.
468 assert_eq!(coeffs.lerp_over_endpoints(B128::random(rng())), B128::ZERO);
469 }
470}