Skip to main content

binius_spartan_prover/
lib.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! Spartan-based proof generation for Binius64 constraint systems.
5//!
6//! This crate provides the [`Prover`] struct for generating zero-knowledge proofs
7//! using the Spartan protocol adapted for Binius64's constraint system. It is the
8//! prover-side counterpart to `binius_spartan_verifier`.
9//!
10//! # When to use this crate
11//!
12//! Use this crate when you have a constraint system built with `binius_spartan_frontend`
13//! and need to generate a Spartan-based proof. This is an alternative to the main
14//! `binius_prover` crate.
15//!
16//! # Key types
17//!
18//! - [`Prover`] - Main proving interface; call [`Prover::setup`] with a verifier, then
19//!   [`Prover::prove`] with witness data
20//! - [`IOPProver`] - Core IOP proving logic, independent of the compilation strategy
21//!
22//! # Related crates
23//!
24//! - `binius_spartan_verifier` - Verification counterpart
25//! - `binius_spartan_frontend` - Constraint system builder for Spartan
26//! - `binius_prover` - Alternative proving backend
27
28#![warn(rustdoc::missing_crate_level_docs)]
29
30mod error;
31mod wiring;
32pub mod wrapper;
33
34use std::{
35	iter::{repeat_n, repeat_with},
36	marker::PhantomData,
37	ops::Deref,
38};
39
40use binius_compute::{Allocator, BufferPool, VecLike};
41use binius_field::{BinaryField, Field, PackedField};
42use binius_hash_prover::ParallelHashSuite;
43use binius_iop_prover::{basefold::compiler::BaseFoldProverCompiler, channel::IOPProverChannel};
44use binius_ip_prover::{
45	channel::IPProverChannel,
46	sumcheck::{quadratic_mlecheck_prover, zk_mlecheck},
47};
48use binius_math::{
49	FieldBuffer, FieldSlice, FieldVec,
50	inner_product::inner_product_buffers,
51	multilinear::eq::eq_ind_partial_eval,
52	ntt::{NeighborsLastMultiThread, domain_context::GaoMateerPreExpanded},
53	univariate::evaluate_univariate,
54};
55use binius_spartan_frontend::constraint_system::{
56	MulConstraint, Witness, WitnessIndex, WitnessSegment,
57};
58use binius_spartan_verifier::{
59	Verifier,
60	constraint_system::{BlindingInfo, ConstraintSystemPadded},
61	wiring::evaluate_wiring_mle_public,
62};
63use binius_transcript::{ProverTranscript, fiat_shamir::Challenger};
64use binius_utils::{SerializeBytes, checked_arithmetics::checked_log_2, rayon::prelude::*};
65use digest::Output;
66pub use error::*;
67use itertools::chain;
68use rand::CryptoRng;
69
70use crate::wiring::{WiringTranspose, fold_constraints};
71
72type ProverNTT<F> = NeighborsLastMultiThread<GaoMateerPreExpanded<F>>;
73
74/// IOP prover for a particular constraint system.
75///
76/// This struct encapsulates the constraint system and pre-computed wiring transpose,
77/// providing the core proving logic independent of the specific IOP compilation strategy.
78/// Most users should use [`Prover`] instead, which wraps this with a BaseFold compiler.
79#[derive(Debug)]
80pub struct IOPProver<F: Field> {
81	constraint_system: ConstraintSystemPadded<F>,
82	precommit_wiring_transpose: WiringTranspose,
83	private_wiring_transpose: WiringTranspose,
84}
85
86/// Struct for proving instances of a particular constraint system.
87///
88/// The [`Self::setup`] constructor pre-processes reusable structures for proving instances of the
89/// given constraint system. Then [`Self::prove`] is called one or more times with individual
90/// instances.
91pub struct Prover<P, H>
92where
93	P: PackedField<Scalar: BinaryField>,
94	H: ParallelHashSuite,
95{
96	iop_prover: IOPProver<P::Scalar>,
97	basefold_compiler: BaseFoldProverCompiler<P, ProverNTT<P::Scalar>>,
98	/// The pool that recycles this prover's working buffers. It lives for the prover's lifetime,
99	/// so blocks freed by one `prove` call are reused by the next.
100	pool: BufferPool,
101	/// The prover creates its Merkle transcript channels with the hash suite `H`.
102	_hash_marker: PhantomData<H>,
103}
104
105impl<F: Field> IOPProver<F> {
106	/// Constructs an IOP prover for a constraint system.
107	pub fn new(constraint_system: ConstraintSystemPadded<F>) -> Self {
108		let precommit_wiring_transpose = WiringTranspose::transpose(
109			WitnessSegment::Precommit,
110			constraint_system.precommit_size(),
111			constraint_system.mul_constraints(),
112		);
113		let private_wiring_transpose = WiringTranspose::transpose(
114			WitnessSegment::Private,
115			constraint_system.private_size(),
116			constraint_system.mul_constraints(),
117		);
118		Self {
119			constraint_system,
120			precommit_wiring_transpose,
121			private_wiring_transpose,
122		}
123	}
124
125	pub const fn constraint_system(&self) -> &ConstraintSystemPadded<F> {
126		&self.constraint_system
127	}
128
129	/// Packs and commits the precommit segment of a witness on the channel.
130	///
131	/// This must be called before [`Self::prove`], and the returned oracle handle and packed
132	/// buffer must be passed into `prove`. Callers that wrap the IOP (e.g. the ZK wrapper) can
133	/// invoke this separately so the precommit oracle handle is available before the rest of
134	/// the protocol runs.
135	pub fn commit_precommit<P, Channel, A>(
136		&self,
137		witness: &Witness<F>,
138		rng: &mut impl CryptoRng,
139		channel: &mut Channel,
140		alloc: &A,
141	) -> (Channel::Oracle, FieldVec<P, A>)
142	where
143		F: BinaryField,
144		P: PackedField<Scalar = F>,
145		Channel: IOPProverChannel<P, A>,
146		A: Allocator,
147	{
148		let cs = &self.constraint_system;
149		let precommit_blinding = *cs.blinding_info();
150		let precommit_packed = pack_and_blind_witness::<_, _, P>(
151			alloc,
152			cs.log_precommit() as usize,
153			witness.precommit(),
154			cs.n_precommit() as usize,
155			&precommit_blinding,
156			rng,
157		);
158		let precommit_oracle = channel.send_oracle(precommit_packed.as_view());
159		(precommit_oracle, precommit_packed)
160	}
161
162	/// Proves using an IOP channel interface.
163	///
164	/// This is the core proving logic, independent of the specific IOP compilation strategy.
165	/// For most users, [`Prover::prove`] is the simpler interface.
166	///
167	/// # Arguments
168	///
169	/// * `witness` - The witness values for the constraint system
170	/// * `precommit_oracle` - Oracle handle obtained from [`Self::commit_precommit`]
171	/// * `precommit_packed` - Packed precommit buffer obtained from [`Self::commit_precommit`]
172	/// * `rng` - Random number generator for blinding
173	/// * `channel` - The IOP prover channel (public input must be observed on transcript before
174	///   creating the channel; the precommit oracle must already have been committed on it via
175	///   [`Self::commit_precommit`])
176	pub fn prove<P, Channel, A>(
177		&self,
178		witness: &Witness<F>,
179		precommit_oracle: Channel::Oracle,
180		precommit_packed: FieldVec<P, A>,
181		mut rng: impl CryptoRng,
182		channel: &mut Channel,
183		alloc: &A,
184	) -> Result<(), Error>
185	where
186		F: BinaryField,
187		P: PackedField<Scalar = F>,
188		Channel: IOPProverChannel<P, A>,
189		A: Allocator,
190	{
191		let _prove_guard =
192			tracing::info_span!("Prove", operation = "prove", perfetto_category = "operation")
193				.entered();
194
195		let cs = &self.constraint_system;
196
197		// Check that the witness segments have the expected sizes
198		let expected_public_size = 1 << cs.log_public() as usize;
199		let expected_precommit_size = cs.precommit_size();
200		let expected_private_size = cs.private_size();
201		if witness.public().len() != expected_public_size {
202			return Err(Error::ArgumentError {
203				arg: "witness".to_string(),
204				msg: format!(
205					"public segment has {} elements, expected {}",
206					witness.public().len(),
207					expected_public_size
208				),
209			});
210		}
211		if witness.precommit().len() != expected_precommit_size {
212			return Err(Error::ArgumentError {
213				arg: "witness".to_string(),
214				msg: format!(
215					"precommit segment has {} elements, expected {}",
216					witness.precommit().len(),
217					expected_precommit_size
218				),
219			});
220		}
221		if witness.private().len() != expected_private_size {
222			return Err(Error::ArgumentError {
223				arg: "witness".to_string(),
224				msg: format!(
225					"private segment has {} elements, expected {}",
226					witness.private().len(),
227					expected_private_size
228				),
229			});
230		}
231
232		let log_mul_constraints = checked_log_2(cs.mul_constraints().len());
233
234		// Create mask buffer for the ZK mulcheck mask polynomial.
235		let (m_n, m_d) = cs.mask_dims();
236		let mask_degree = 2; // quadratic composition
237		let log_masks_buffer_size = m_n + m_d;
238
239		let masks_buffer = {
240			// Growing a pooled buffer past the block it was handed would reallocate and free that
241			// block at the element's alignment rather than the pool's, so the fill below takes
242			// exactly the allocated count.
243			let packed_len = 1 << log_masks_buffer_size.saturating_sub(P::LOG_WIDTH);
244			let mut values = alloc.alloc::<P>(packed_len);
245			values.extend(repeat_with(|| P::random(&mut rng)).take(packed_len));
246			FieldBuffer::new(log_masks_buffer_size, values)
247		};
248
249		let mulcheck_mask =
250			zk_mlecheck::Mask::new(log_mul_constraints, mask_degree, masks_buffer.as_view());
251
252		// Pack private witness into field elements and add blinding
253		let blinding_info = cs.blinding_info();
254		let private_packed = pack_and_blind_witness::<_, _, P>(
255			alloc,
256			cs.log_private() as usize,
257			witness.private(),
258			cs.n_private() as usize,
259			blinding_info,
260			&mut rng,
261		);
262
263		// Send the private and mask oracles to the channel. The precommit oracle was committed
264		// by the caller via `commit_precommit` and passed in as `precommit_oracle`.
265		let private_oracle = channel.send_oracle(private_packed.as_view());
266		let mask_oracle = channel.send_oracle(masks_buffer.as_view());
267
268		// Prove the multiplication constraints
269		let (mulcheck_evals, mask_eval, r_x) = prove_mulcheck::<F, P, _, _>(
270			cs.mul_constraints(),
271			witness.public(),
272			precommit_packed.as_view(),
273			private_packed.as_view(),
274			mulcheck_mask,
275			&mut *channel,
276			alloc,
277		);
278
279		// λ is the batching challenge for the constraint operands
280		let lambda = channel.sample();
281
282		// Batch together the constraint operand evaluation claims.
283		let batched_sum = evaluate_univariate(&mulcheck_evals, &lambda);
284
285		// Compute eq indicator tensor for r_x (shared across all segment evaluations)
286		let r_x_tensor = eq_ind_partial_eval::<F>(&r_x);
287
288		// Compute rₓ^⊤ (M_A + λ M_B + λ² M_C) x
289		let public_eval = evaluate_wiring_mle_public(
290			cs.mul_constraints(),
291			witness.public(),
292			&lambda,
293			r_x_tensor.as_ref(),
294		);
295
296		// Compute the precommit segment's contribution to the wiring check.
297		// The prover sends this as a scalar; the oracle relation then verifies it.
298		let precommit_wiring_poly =
299			fold_constraints(alloc, &self.precommit_wiring_transpose, lambda, r_x_tensor.as_ref());
300		let precommit_claim = inner_product_buffers(&precommit_packed, &precommit_wiring_poly);
301		channel.send_one(precommit_claim);
302
303		let private_claim = batched_sum - public_eval - precommit_claim;
304
305		// Fold private wiring constraints
306		let private_wiring_poly =
307			fold_constraints(alloc, &self.private_wiring_transpose, lambda, r_x_tensor.as_ref());
308
309		// Compute the mask folding polynomial (libra_eval tensor)
310		let n_vars = r_x.len();
311		let libra_eval_tensor =
312			zk_mlecheck::expand_libra_eval::<A, P>(alloc, &r_x, n_vars, mask_degree, m_n, m_d);
313
314		// Prove all oracle relations, handing the channel each committed buffer for the combined
315		// opening.
316		channel.prove_oracle_relation(
317			precommit_oracle.clone(),
318			precommit_wiring_poly.into(),
319			precommit_claim,
320		);
321		channel.finalize_oracle(precommit_oracle, precommit_packed);
322		channel.prove_oracle_relation(
323			private_oracle.clone(),
324			private_wiring_poly.into(),
325			private_claim,
326		);
327		channel.finalize_oracle(private_oracle, private_packed);
328		channel.prove_oracle_relation(mask_oracle.clone(), libra_eval_tensor.into(), mask_eval);
329		channel.finalize_oracle(mask_oracle, masks_buffer);
330
331		Ok(())
332	}
333}
334
335impl<F, P, H> Prover<P, H>
336where
337	F: BinaryField,
338	P: PackedField<Scalar = F>,
339	H: ParallelHashSuite,
340	Output<H::LeafHash>: SerializeBytes,
341{
342	/// Constructs a prover corresponding to a constraint system verifier.
343	///
344	/// See [`Prover`] struct documentation for details.
345	pub fn setup(verifier: &Verifier<F, H>) -> Result<Self, Error> {
346		let log_num_shares = binius_utils::rayon::current_num_threads().ilog2() as usize;
347
348		// Rebuild the verifier's evaluation domain, which its compiler fixed as the Gao-Mateer
349		// basis of that dimension.
350		let domain_context =
351			GaoMateerPreExpanded::generate(verifier.iop_compiler().max_log_domain_size());
352		let ntt = NeighborsLastMultiThread::new(domain_context, log_num_shares);
353
354		// Create the BaseFold ZK compiler from verifier compiler (reuses oracle_specs and
355		// fri_params)
356		let basefold_compiler =
357			BaseFoldProverCompiler::from_verifier_compiler(verifier.iop_compiler(), ntt);
358
359		let iop_prover = IOPProver::new(verifier.constraint_system().clone());
360
361		Ok(Prover {
362			iop_prover,
363			basefold_compiler,
364			pool: BufferPool::new(),
365			_hash_marker: PhantomData,
366		})
367	}
368
369	/// Returns a reference to the IOP prover.
370	pub const fn iop_prover(&self) -> &IOPProver<P::Scalar> {
371		&self.iop_prover
372	}
373
374	/// Returns a reference to the BaseFold ZK prover compiler.
375	pub const fn iop_compiler(&self) -> &BaseFoldProverCompiler<P, ProverNTT<F>> {
376		&self.basefold_compiler
377	}
378
379	/// Generates a proof for a witness against the constraint system.
380	///
381	/// # Arguments
382	///
383	/// * `witness` - The witness values for the constraint system
384	/// * `rng` - Random number generator for blinding
385	/// * `transcript` - The prover transcript for Fiat-Shamir
386	///
387	/// # Preconditions
388	///
389	/// * The witness length must match the constraint system size
390	pub fn prove<Challenger_: Challenger>(
391		&self,
392		witness: &Witness<F>,
393		mut rng: impl CryptoRng,
394		transcript: &mut ProverTranscript<Challenger_>,
395	) -> Result<(), Error> {
396		// Prover observes the public input (includes it in Fiat-Shamir).
397		let public = witness.public();
398		transcript.observe().write_slice(public);
399
400		// Working buffers for this proof are drawn from the prover's pool, recycling blocks freed
401		// by earlier proofs. The channel gets the same pool, so the Merkle trees it commits draw
402		// their nodes from it too.
403		let alloc = &self.pool;
404		// Create ZK channel (owns the RNG for mask generation), commit the precommit oracle,
405		// and delegate to the IOP prover.
406		let mut channel = self
407			.basefold_compiler
408			.create_channel_from_transcript::<H, Challenger_, _, _>(transcript, &mut rng, alloc);
409		let (precommit_oracle, precommit_packed) =
410			self.iop_prover
411				.commit_precommit::<P, _, _>(witness, &mut rng, &mut channel, &alloc);
412		// The IOP prover only queues the oracle relations; `finish` runs the single combined
413		// opening.
414		self.iop_prover.prove::<P, _, _>(
415			witness,
416			precommit_oracle,
417			precommit_packed,
418			rng,
419			&mut channel,
420			&alloc,
421		)?;
422		channel.finish();
423		Ok(())
424	}
425}
426
427fn prove_mulcheck<F, P, Channel, A>(
428	mul_constraints: &[MulConstraint<WitnessIndex>],
429	public: &[F],
430	precommit_packed: FieldSlice<'_, P>,
431	private_packed: FieldSlice<'_, P>,
432	mask: zk_mlecheck::Mask<P, impl Deref<Target = [P]>>,
433	channel: &mut Channel,
434	alloc: &A,
435) -> ([F; 3], F, Vec<F>)
436where
437	F: BinaryField,
438	P: PackedField<Scalar = F>,
439	Channel: IPProverChannel<F>,
440	A: Allocator,
441{
442	let mulcheck_witness = wiring::build_mulcheck_witness(
443		alloc,
444		mul_constraints,
445		public,
446		precommit_packed,
447		private_packed,
448	);
449
450	// Sample random evaluation point for mulcheck
451	let r_mulcheck = channel.sample_many(mask.n_vars());
452
453	// Prove the mul-gate zerocheck a * b - c = 0 over the shared store.
454	let mlecheck_prover = quadratic_mlecheck_prover(
455		alloc,
456		[mulcheck_witness.a, mulcheck_witness.b, mulcheck_witness.c],
457		|[a, b, c]| a * b - c, // composition
458		|[a, b, _c]| a * b,    // infinity_composition (quadratic term only)
459		r_mulcheck,
460		F::ZERO, // eval_claim: zerocheck
461	);
462
463	// Run the ZK MLE-check protocol
464	let mlecheck_output = zk_mlecheck::prove(mlecheck_prover, mask, channel);
465
466	// Extract the reduced evaluation point and multilinear evaluations
467	let mut r_x = mlecheck_output.challenges;
468	r_x.reverse(); // Match verifier's order
469
470	let [a_eval, b_eval, c_eval]: [F; 3] = mlecheck_output
471		.multilinear_evals
472		.try_into()
473		.expect("mlecheck returns 3 evaluations");
474
475	// Write the multilinear evaluations to channel
476	channel.send_many(&[a_eval, b_eval, c_eval]);
477
478	let mulcheck_evals = [a_eval, b_eval, c_eval];
479	let mask_eval = mlecheck_output.mask_eval;
480
481	(mulcheck_evals, mask_eval, r_x)
482}
483
484/// Packs witness values into a [`FieldBuffer`] and adds blinding values for dummy wires.
485fn pack_and_blind_witness<A: Allocator, F: Field, P: PackedField<Scalar = F>>(
486	alloc: &A,
487	log_private: usize,
488	private: &[F],
489	n_private: usize,
490	blinding_info: &BlindingInfo,
491	mut rng: impl CryptoRng,
492) -> FieldVec<P, A> {
493	// Growing a pooled buffer past the block it was handed would reallocate and free that block at
494	// the element's alignment rather than the pool's, so the fill below must fit exactly.
495	let packed_len = 1 << log_private.saturating_sub(P::LOG_WIDTH);
496	let mut packed = alloc.alloc::<P>(packed_len);
497	if log_private < P::LOG_WIDTH {
498		// The whole segment lives in one packed element's low lanes.
499		debug_assert_eq!(packed_len, 1);
500		let elems_iter = private.iter().copied();
501		let zeros_iter = repeat_n(F::ZERO, (1 << log_private) - private.len());
502
503		packed.push(P::from_scalars(chain!(elems_iter, zeros_iter)));
504	} else {
505		// Zero the block once, then overwrite the prefix holding real scalars. The zip stops at the
506		// last real chunk, so the zero tail is the padding the buffer wants anyway. Collecting into
507		// a `Vec` first and copying that in would cost a second allocation and a full memcpy — the
508		// very thing pooling is here to remove.
509		debug_assert!(private.len() <= 1 << log_private);
510		packed.resize(packed_len, P::zero());
511		private
512			.par_chunks(P::WIDTH)
513			.zip(packed.par_iter_mut())
514			.for_each(|(chunk, out)| *out = P::from_scalars(chunk.iter().copied()));
515	}
516
517	let mut buffer = FieldBuffer::new(log_private, packed);
518
519	// Add blinding values after the actual private wires
520	// Set random values for non-constraint dummy wires
521	for i in 0..blinding_info.n_dummy_wires {
522		buffer.set(n_private + i, F::random(&mut rng));
523	}
524
525	// Set random values for dummy constraint wires (A * B = C)
526	let constraint_wire_base = n_private + blinding_info.n_dummy_wires;
527	for i in 0..blinding_info.n_dummy_constraints {
528		let a = F::random(&mut rng);
529		let b = F::random(&mut rng);
530		let c = a * b;
531
532		buffer.set(constraint_wire_base + 3 * i, a);
533		buffer.set(constraint_wire_base + 3 * i + 1, b);
534		buffer.set(constraint_wire_base + 3 * i + 2, c);
535	}
536
537	buffer
538}