Skip to main content

binius_prover/protocols/bitand/
prover.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::ops::Deref;
5
6use binius_compute::Allocator;
7use binius_core::word::Word;
8use binius_field::{BinaryField, PackedField, Rijndael8b as B8};
9use binius_ip_prover::sumcheck::{common::MleCheckProver, quadratic_mlecheck_prover};
10use binius_math::{BinarySubspace, univariate::EvaluationDomain};
11use binius_verifier::{
12	config::PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES, protocols::bitand::ROWS_PER_HYPERCUBE_VERTEX,
13};
14
15use super::sumcheck_round_messages;
16use crate::fold_word::BitAxisFolder;
17
18/// Prover for the univariate-skip round of the AND constraint reduction.
19///
20/// It computes the round message over the two operand columns, then folds the columns at the
21/// verifier's challenge into the MLE-check summand.
22///
23/// See [`binius_verifier::protocols::bitand`] for the protocol specification.
24///
25/// The columns are generic over their backing store `Data` (anything that dereferences to
26/// `[Word]`), so callers can supply pooled buffers ([`PoolVec`](binius_compute::PoolVec)) or plain
27/// `Vec<Word>` interchangeably.
28pub struct UnivariateRoundProver<F, Data>
29where
30	F: BinaryField,
31{
32	log_words: usize,
33	first_col: Data,
34	second_col: Data,
35	big_field_zerocheck_challenges: Vec<F>,
36	univariate_round_message: [F; ROWS_PER_HYPERCUBE_VERTEX],
37	univariate_round_message_domain: BinarySubspace<F>,
38}
39
40impl<F, Data> UnivariateRoundProver<F, Data>
41where
42	F: BinaryField + From<B8>,
43	Data: Deref<Target = [Word]>,
44{
45	/// Computes the univariate-skip message over the two operand columns.
46	///
47	/// The message holds the evaluations of the univariate round polynomial R₀(Z), which encodes
48	/// the AND constraint across the oblong dimension.
49	///
50	/// The C operand of the AND constraint `A & B ^ C = 0` is not an input.
51	/// The prover derives it word-by-word as `A & B`.
52	///
53	/// # Why deriving C is sound
54	///
55	/// - A satisfying witness makes `C = A & B` hold on every row.
56	/// - Folding is F2-linear on word bits.
57	/// - Equal words therefore fold to equal field elements.
58	/// - So an honest prover emits the exact same transcript as with an explicit C column.
59	/// - A cheating witness is still rejected.
60	/// - The shift reduction later checks the claimed C evaluation against the committed witness.
61	///
62	/// # Arguments
63	///
64	/// * `log_words` - Base-2 logarithm of the constraint axis's length
65	/// * `first_col` - The oblong multilinear polynomial A in the AND constraint A & B ^ C = 0
66	/// * `second_col` - The oblong multilinear polynomial B in the AND constraint
67	/// * `big_field_zerocheck_challenges` - Challenges Z_{k+1},...,Zₙ in the large field `F`
68	/// * `prover_message_domain` - The domain for evaluating the univariate polynomial
69	///
70	/// The two columns must have equal length, at most `1 << log_words`. A column shorter than the
71	/// axis has its remaining rows read as zero, in both the round-1 message and the fold; the
72	/// reduction skips them rather than working over them.
73	///
74	/// # Implementation Details
75	///
76	/// This function:
77	/// 1. Computes the equality indicator polynomial from the big field challenges
78	/// 2. Uses the NTT lookup to efficiently compute the univariate polynomial evaluations
79	/// 3. Caches these evaluations for later use in the [`round_message`](Self::round_message)
80	///    method
81	pub fn compute_message(
82		log_words: usize,
83		first_col: Data,
84		second_col: Data,
85		big_field_zerocheck_challenges: Vec<F>,
86		prover_message_domain: &BinarySubspace<B8>,
87	) -> Self {
88		let univariate_round_message = tracing::debug_span!("Compute univariate round message")
89			.in_scope(|| {
90				sumcheck_round_messages::univariate_round_message_extension_domain::<F>(
91					log_words,
92					&first_col,
93					&second_col,
94					&big_field_zerocheck_challenges,
95					prover_message_domain,
96				)
97			});
98
99		Self {
100			log_words,
101			first_col,
102			second_col,
103			univariate_round_message,
104			big_field_zerocheck_challenges,
105			univariate_round_message_domain: prover_message_domain.isomorphic(),
106		}
107	}
108
109	/// The message to send: R₀(Z) on the extension domain.
110	///
111	/// These are exactly `ROWS_PER_HYPERCUBE_VERTEX` field elements that represent R₀(Z) for Z in
112	/// the upper half of the univariate domain. [`compute_message`](Self::compute_message)
113	/// computes them; this method returns the cached result.
114	pub const fn round_message(&self) -> &[F; ROWS_PER_HYPERCUBE_VERTEX] {
115		&self.univariate_round_message
116	}
117
118	/// The univariate round polynomial at `challenge`: the claim the multilinear rounds prove.
119	///
120	/// The polynomial is zero on the base half of the domain, and the round message holds its
121	/// evaluations on the upper half.
122	pub fn univariate_claim(&self, challenge: F) -> F {
123		let mut coeffs = vec![F::ZERO; 2 * ROWS_PER_HYPERCUBE_VERTEX];
124		coeffs[ROWS_PER_HYPERCUBE_VERTEX..].copy_from_slice(&self.univariate_round_message);
125		self.univariate_round_message_domain
126			.extrapolate(&coeffs, &challenge)
127	}
128
129	/// Folds A, B and the derived C = A & B at the univariate challenge.
130	///
131	/// Returns the MLE-check summand over the ℓ_and constraint variables, at the zerocheck point.
132	/// Fixing Z to `challenge` reduces the oblong multilinears to standard multilinears over the
133	/// remaining variables, and the returned prover proves the sumcheck claim:
134	/// R₀(z) = ∑_{X₀,...,Xₙ₋₁ ∈ {0,1}} (A(z,X₀,...,Xₙ₋₁)·B(z,X₀,...,Xₙ₋₁) -
135	/// C(z,X₀,...,Xₙ₋₁))·eq(X₀,...,Xₙ₋₁; r₀,...,rₙ₋₁)
136	///
137	/// The folded columns are allocated in `alloc`, which the returned prover borrows.
138	///
139	/// # Process
140	///
141	/// 1. Creates a fold lookup table over the univariate domain of the round message
142	/// 2. Folds A, B, and the derived C = A & B at Z = challenge, in one fused pass
143	/// 3. Combines the zerocheck challenges (small field + big field)
144	/// 4. Evaluates the univariate polynomial at the challenge to get the sumcheck claim
145	/// 5. Constructs the AND reduction sumcheck prover with the folded multilinears
146	pub fn fold<'alloc, PChallenge: PackedField<Scalar = F>, A: Allocator>(
147		self,
148		alloc: &'alloc A,
149		challenge: F,
150	) -> impl MleCheckProver<F> + 'alloc {
151		let claim = self.univariate_claim(challenge);
152		let round_message_domain = &self.univariate_round_message_domain;
153		let univariate_domain = round_message_domain.reduce_dim(round_message_domain.dim() - 1);
154		let lagrange_evals = univariate_domain.lagrange_evals(&challenge);
155		let folder = BitAxisFolder::new(&lagrange_evals);
156
157		let proving_polys =
158			folder.fold_bitand_operands::<PChallenge, _>(alloc, &self.first_col, &self.second_col);
159
160		let upcasted_small_field_challenges = PROVER_SMALL_FIELD_ZEROCHECK_CHALLENGES
161			.iter()
162			.copied()
163			.take(self.log_words)
164			.map(F::from);
165
166		let verifier_field_zerocheck_challenges = upcasted_small_field_challenges
167			.chain(self.big_field_zerocheck_challenges)
168			.collect::<Vec<_>>();
169
170		quadratic_mlecheck_prover(
171			alloc,
172			proving_polys,
173			|[a, b, c]| a * b - c,
174			|[a, b, _]| a * b,
175			verifier_field_zerocheck_challenges,
176			claim,
177		)
178	}
179}
180
181#[cfg(test)]
182mod test {
183	use std::{iter, iter::repeat_with};
184
185	use binius_compute::GlobalAllocator;
186	use binius_core::word::Word;
187	use binius_field::{Rijndael8b, arch::OptimalPackedB128};
188	use binius_math::{
189		BinarySubspace, FieldBuffer, multilinear::evaluate::evaluate, univariate::EvaluationDomain,
190	};
191	use binius_transcript::ProverTranscript;
192	use binius_verifier::{
193		config::{B128, StdChallenger},
194		protocols::bitand::{AndCheckOutput, SKIPPED_VARS},
195		verify_bitand_reduction,
196	};
197	use rand::prelude::*;
198
199	use crate::{fold_word::BitAxisFolder, protocols::bitand::prove};
200
201	fn random_words(log_num_words: usize, mut rng: impl Rng) -> Vec<Word> {
202		repeat_with(|| Word(rng.random()))
203			.take(1 << log_num_words)
204			.collect()
205	}
206
207	#[test]
208	fn test_transcript_prover_verifies() {
209		let mut prover_challenger = ProverTranscript::new(StdChallenger::default());
210		let log_num_rows = 6;
211		let mut rng = StdRng::seed_from_u64(0);
212
213		let first_mlv = random_words(log_num_rows, &mut rng);
214		let second_mlv = random_words(log_num_rows, &mut rng);
215		// The prover receives only the A and B columns.
216		// This materialized C = A & B feeds only the verifier-side fold check at the end.
217		let third_mlv: Vec<Word> = iter::zip(&first_mlv, &second_mlv)
218			.map(|(&a, &b)| a & b)
219			.collect();
220
221		// Agreed-upon proof parameter
222		let prover_message_domain = BinarySubspace::<Rijndael8b>::with_dim(SKIPPED_VARS + 1);
223		let verifier_message_domain = prover_message_domain.isomorphic();
224
225		let prove_output = prove::<_, B128, OptimalPackedB128, _, _>(
226			[first_mlv.clone(), second_mlv.clone()],
227			&[],
228			&mut prover_challenger,
229			&GlobalAllocator,
230		);
231
232		let mut verifier_challenger = prover_challenger.into_verifier();
233		let verify_output = verify_bitand_reduction(
234			log_num_rows,
235			&verifier_message_domain,
236			&[],
237			&mut verifier_challenger,
238		)
239		.unwrap();
240
241		assert_eq!(prove_output, verify_output);
242
243		let AndCheckOutput {
244			z_challenge,
245			rerand,
246		} = verify_output;
247		let [a_eval, b_eval, c_eval] = rerand.bitand_evals;
248		let eval_point = rerand.eval_point;
249
250		let verifier_univariate_domain = verifier_message_domain.reduce_dim(SKIPPED_VARS);
251
252		let one_bit_mlvs = [first_mlv, second_mlv, third_mlv];
253
254		let verifier_lagrange_evals = verifier_univariate_domain.lagrange_evals(&z_challenge);
255		let folder = BitAxisFolder::new(&verifier_lagrange_evals);
256		for (i, eval) in [a_eval, b_eval, c_eval].iter().enumerate() {
257			let folded: FieldBuffer<B128> = folder.fold(&GlobalAllocator, &one_bit_mlvs[i]);
258			assert_eq!(evaluate(&folded, &eval_point), *eval);
259		}
260	}
261}