Skip to main content

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}