Skip to main content

binius_spartan_frontend/
constraint_system.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{
5	cmp::Ordering,
6	collections::HashMap,
7	mem,
8	ops::{Index, IndexMut},
9};
10
11use binius_field::Field;
12use binius_utils::checked_arithmetics::log2_ceil_usize;
13use smallvec::{SmallVec, smallvec};
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
16pub enum WireKind {
17	Constant,
18	InOut,
19	/// A wire whose value is a pure function of public-derivable inputs (constants, inout, and
20	/// other derived wires). Derived wires live in the public segment and emit no constraint; the
21	/// verifier recomputes them itself via the `InstanceGenerator`.
22	Derived,
23	Precommit,
24	Private,
25}
26
27impl WireKind {
28	/// The witness segment a wire of this kind lives in: constants, inout, and derived wires occupy
29	/// the public segment; precommit and private wires occupy their own segments.
30	pub const fn segment(self) -> WitnessSegment {
31		match self {
32			WireKind::Constant | WireKind::InOut | WireKind::Derived => WitnessSegment::Public,
33			WireKind::Precommit => WitnessSegment::Precommit,
34			WireKind::Private => WitnessSegment::Private,
35		}
36	}
37
38	/// Returns whether this wire kind lives in the public segment: its value is determined once the
39	/// public inputs (constants and inout) are known, so the verifier can recompute it without any
40	/// secret data.
41	///
42	/// Shared by all `CircuitBuilder` implementations so they make identical derived-vs-private
43	/// decisions. The output of a binary op (or hint) is [`WireKind::Derived`] iff every input is
44	/// public, and [`WireKind::Private`] otherwise.
45	pub fn is_public(self) -> bool {
46		self.segment() == WitnessSegment::Public
47	}
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
51pub struct ConstraintWire {
52	pub(crate) kind: WireKind,
53	pub(crate) id: u32,
54}
55
56impl ConstraintWire {
57	/// Creates a constraint wire referencing an inout wire by ID.
58	///
59	/// TODO: This is not ideal, and instead we should use some sort of allocator.
60	pub const fn inout(id: u32) -> Self {
61		Self {
62			kind: WireKind::InOut,
63			id,
64		}
65	}
66
67	/// Creates a constraint wire referencing a precommit wire by ID.
68	pub const fn precommit(id: u32) -> Self {
69		Self {
70			kind: WireKind::Precommit,
71			id,
72		}
73	}
74}
75
76#[derive(Debug, Clone)]
77pub struct Operand<W>(SmallVec<[W; 4]>);
78
79impl<W> Default for Operand<W> {
80	fn default() -> Self {
81		Operand(SmallVec::new())
82	}
83}
84
85impl<W: Copy + Ord> Operand<W> {
86	pub fn new(mut term: SmallVec<[W; 4]>) -> Self {
87		term.sort_unstable();
88
89		let has_duplicate_wire = term.windows(2).any(|w| w[0] == w[1]);
90		let term = if has_duplicate_wire {
91			term.chunk_by(|a, b| a == b)
92				.flat_map(|group| {
93					// Group is a slice of wires that are all equal. We want to return an empty
94					// iterator if the group is even length and a singleton iterator otherwise.
95					let last_even_idx = group.len() / 2 * 2;
96					group[last_even_idx..].iter().copied()
97				})
98				.collect()
99		} else {
100			term
101		};
102
103		Self(term)
104	}
105
106	pub fn len(&self) -> usize {
107		self.0.len()
108	}
109
110	pub fn is_empty(&self) -> bool {
111		self.0.is_empty()
112	}
113
114	pub fn wires(&self) -> &[W] {
115		&self.0
116	}
117
118	pub fn merge(&mut self, rhs: &Self) -> (Operand<W>, Operand<W>) {
119		// Classic merge algorithm for sorted vectors, but where duplicate items cancel out.
120		let lhs = mem::take(&mut self.0);
121		let dst = &mut self.0;
122
123		let mut lhs_iter = lhs.into_iter().peekable();
124		let mut rhs_iter = rhs.0.iter().copied().peekable();
125
126		let mut additions = Operand::default();
127		let mut removals = Operand::default();
128
129		loop {
130			match (lhs_iter.peek(), rhs_iter.peek()) {
131				(Some(next_lhs), Some(next_rhs)) => {
132					match next_lhs.cmp(next_rhs) {
133						Ordering::Equal => {
134							// Advance both iterators, but don't push the wires because they cancel.
135							let wire = lhs_iter.next().expect("peek returned Some");
136							let _ = rhs_iter.next().expect("peek returned Some");
137
138							removals.0.push(wire);
139						}
140						Ordering::Less => dst.push(lhs_iter.next().expect("peek returned Some")),
141						Ordering::Greater => {
142							let wire = rhs_iter.next().expect("peek returned Some");
143							additions.0.push(wire);
144							dst.push(wire);
145						}
146					}
147				}
148				(Some(_), None) => dst.push(lhs_iter.next().expect("peek returned Some")),
149				(None, Some(_)) => {
150					let wire = rhs_iter.next().expect("peek returned Some");
151					additions.0.push(wire);
152					dst.push(wire);
153				}
154				(None, None) => break,
155			}
156		}
157
158		(additions, removals)
159	}
160}
161
162impl<W> From<W> for Operand<W> {
163	fn from(value: W) -> Self {
164		Operand(smallvec![value])
165	}
166}
167
168#[derive(Debug, Clone)]
169pub struct MulConstraint<W> {
170	pub a: Operand<W>,
171	pub b: Operand<W>,
172	pub c: Operand<W>,
173}
174
175#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
176pub enum WitnessSegment {
177	/// The public segment contains constant and input/output witness values.
178	Public,
179	/// The precommit segment contains values committed in a separate oracle before the private
180	/// segment. These are zero-knowledge hidden but not prunable or rearrangeable.
181	Precommit,
182	/// The private segment contains the remaining witness values, which are hidden from the
183	/// verifier.
184	Private,
185}
186
187#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
188pub struct WitnessIndex {
189	pub segment: WitnessSegment,
190	pub index: u32,
191}
192
193impl WitnessIndex {
194	pub const fn public(index: u32) -> Self {
195		Self {
196			segment: WitnessSegment::Public,
197			index,
198		}
199	}
200
201	pub const fn precommit(index: u32) -> Self {
202		Self {
203			segment: WitnessSegment::Precommit,
204			index,
205		}
206	}
207
208	pub const fn private(index: u32) -> Self {
209		Self {
210			segment: WitnessSegment::Private,
211			index,
212		}
213	}
214}
215
216pub struct Witness<F> {
217	public: Vec<F>,
218	precommit: Vec<F>,
219	private: Vec<F>,
220}
221
222impl<F> Witness<F> {
223	pub const fn new(public: Vec<F>, precommit: Vec<F>, private: Vec<F>) -> Self {
224		Self {
225			public,
226			precommit,
227			private,
228		}
229	}
230
231	pub fn public(&self) -> &[F] {
232		&self.public
233	}
234
235	pub fn precommit(&self) -> &[F] {
236		&self.precommit
237	}
238
239	pub fn private(&self) -> &[F] {
240		&self.private
241	}
242}
243
244impl<F> Index<WitnessIndex> for Witness<F> {
245	type Output = F;
246
247	fn index(&self, index: WitnessIndex) -> &Self::Output {
248		match index.segment {
249			WitnessSegment::Public => &self.public[index.index as usize],
250			WitnessSegment::Precommit => &self.precommit[index.index as usize],
251			WitnessSegment::Private => &self.private[index.index as usize],
252		}
253	}
254}
255
256impl<F> IndexMut<WitnessIndex> for Witness<F> {
257	fn index_mut(&mut self, index: WitnessIndex) -> &mut Self::Output {
258		match index.segment {
259			WitnessSegment::Public => &mut self.public[index.index as usize],
260			WitnessSegment::Precommit => &mut self.precommit[index.index as usize],
261			WitnessSegment::Private => &mut self.private[index.index as usize],
262		}
263	}
264}
265
266/// A constraint system with multiplication constraints over witness indices.
267///
268/// Contains multiplication constraints of the form `A * B = C` where A, B, C are operands
269/// (XOR combinations of witness values). Constraints directly reference [`WitnessIndex`]
270/// positions in the witness array.
271///
272/// This struct does not guarantee power-of-two constraint counts or witness size.
273#[derive(Debug, Clone)]
274pub struct ConstraintSystem<F: Field> {
275	constants: Vec<F>,
276	n_inout: u32,
277	n_precommit: u32,
278	n_private: u32,
279	log_public: u32,
280	mul_constraints: Vec<MulConstraint<WitnessIndex>>,
281	one_wire_index: u32,
282}
283
284impl<F: Field> ConstraintSystem<F> {
285	/// Create a new constraint system.
286	pub const fn new(
287		constants: Vec<F>,
288		n_inout: u32,
289		n_precommit: u32,
290		n_private: u32,
291		log_public: u32,
292		mul_constraints: Vec<MulConstraint<WitnessIndex>>,
293		one_wire_index: u32,
294	) -> Self {
295		Self {
296			constants,
297			n_inout,
298			n_precommit,
299			n_private,
300			log_public,
301			mul_constraints,
302			one_wire_index,
303		}
304	}
305
306	pub fn constants(&self) -> &[F] {
307		&self.constants
308	}
309
310	pub const fn n_inout(&self) -> u32 {
311		self.n_inout
312	}
313
314	pub const fn n_precommit(&self) -> u32 {
315		self.n_precommit
316	}
317
318	pub const fn n_private(&self) -> u32 {
319		self.n_private
320	}
321
322	pub const fn log_public(&self) -> u32 {
323		self.log_public
324	}
325
326	pub const fn n_public(&self) -> u32 {
327		1 << self.log_public
328	}
329
330	pub fn mul_constraints(&self) -> &[MulConstraint<WitnessIndex>] {
331		&self.mul_constraints
332	}
333
334	pub const fn one_wire(&self) -> WitnessIndex {
335		WitnessIndex {
336			segment: WitnessSegment::Public,
337			index: self.one_wire_index,
338		}
339	}
340
341	/// Validate that a witness satisfies all multiplication constraints.
342	pub fn validate(&self, witness: &Witness<F>) {
343		let operand_val = |operand: &Operand<WitnessIndex>| {
344			operand.wires().iter().map(|&idx| witness[idx]).sum::<F>()
345		};
346
347		for MulConstraint { a, b, c } in &self.mul_constraints {
348			assert_eq!(operand_val(a) * operand_val(b), operand_val(c));
349		}
350	}
351}
352
353/// Random padding appended to a committed segment.
354///
355/// The padding is what keeps the segment's openings independent of its real wires.
356/// Every committed segment carries the same amount, so each is padded on its own.
357#[derive(Debug, Clone, Copy)]
358pub struct BlindingInfo {
359	/// The number of random dummy wires appended after the segment's real wires.
360	pub n_dummy_wires: usize,
361	/// The number of random dummy multiplication constraints that must be added.
362	pub n_dummy_constraints: usize,
363}
364
365/// Dummy multiplication constraints appended to every committed segment.
366///
367/// A wire in no constraint has coefficient zero in the wiring relation.
368/// So no count of plain dummy wires masks a value the relation reveals.
369///
370/// The three wires of a dummy constraint do sit in a constraint, so they reach the relation.
371///
372/// # Why this value
373///
374/// Everything verification publishes touches a committed segment only through that segment's
375/// three operand contributions `(A_S, B_S, C_S)`. The operand evaluations are those plus the
376/// other segments' shares, and the segment's own batched claim is `A_S + lambda * B_S +
377/// lambda^2 * C_S`, already in their span. So three values need masking, not four.
378///
379/// One dummy constraint cannot mask them. Its wires contribute `(alpha * a, alpha * b,
380/// alpha * a * b)`, because the prover pins the third wire to the product of the other two.
381/// Its share of `C_S` is therefore fixed by its shares of `A_S` and `B_S`, and a verifier
382/// holding all three recovers a relation among the real wires.
383///
384/// Two constraints break that: `a` and `b` of each are free, and the distribution they induce
385/// on `(A_S, B_S, C_S)` is statistically close to uniform.
386const N_DUMMY_CONSTRAINTS: usize = 2;
387
388impl BlindingInfo {
389	/// The blinding a committed segment needs when FRI opens it at `n_test_queries` positions.
390	///
391	/// Each query opens one Merkle leaf, revealing one codeword symbol of the segment.
392	/// A symbol is a fixed linear function of the segment.
393	/// So `n_test_queries` random wires make every opened symbol uniform and independent.
394	///
395	/// One further wire covers the leaves that are never opened.
396	/// Their hashes still travel in the authentication paths, and the leaves carry no salt.
397	/// The spare degree of randomness keeps an unopened leaf unguessable, as a salt would.
398	pub const fn for_fri_queries(n_test_queries: usize) -> Self {
399		Self {
400			n_dummy_wires: n_test_queries + 1,
401			n_dummy_constraints: N_DUMMY_CONSTRAINTS,
402		}
403	}
404}
405
406#[derive(Debug, Clone)]
407pub struct WitnessLayout<F: Field> {
408	pub(crate) constants: Vec<F>,
409	n_inout: u32,
410	n_derived: u32,
411	n_precommit: u32,
412	n_private: u32,
413	log_public: u32,
414	log_precommit: u32,
415	log_private: u32,
416	derived_index_map: HashMap<u32, u32>,
417	private_index_map: HashMap<u32, u32>,
418}
419
420impl<F: Field> WitnessLayout<F> {
421	pub fn sparse(
422		constants: Vec<F>,
423		n_inout: u32,
424		n_precommit: u32,
425		derived_alive: &[bool],
426		private_alive: &[bool],
427	) -> Self {
428		let n_constants = constants.len() as u32;
429		let log_precommit = log2_ceil_usize(n_precommit as usize) as u32;
430
431		// Derived wires occupy the public segment after constants and inout. Only derived wires
432		// referenced by a surviving constraint need a slot; intermediates that feed only other
433		// derived wires are computed inline by the generators and get no slot.
434		let derived_index_map = derived_alive
435			.iter()
436			.enumerate()
437			.filter(|(_, alive)| **alive)
438			.enumerate()
439			.map(|(new_idx, (id, _))| (id as u32, new_idx as u32))
440			.collect::<HashMap<_, _>>();
441
442		let n_derived = derived_index_map.len() as u32;
443		let n_public = n_constants + n_inout + n_derived;
444		let log_public = log2_ceil_usize(n_public as usize) as u32;
445
446		let private_index_map = private_alive
447			.iter()
448			.enumerate()
449			.filter(|(_, alive)| **alive)
450			.enumerate()
451			.map(|(new_idx, (id, _))| (id as u32, new_idx as u32))
452			.collect::<HashMap<_, _>>();
453
454		let n_private = private_index_map.len() as u32;
455		let log_private = log2_ceil_usize(n_private as usize) as u32;
456
457		Self {
458			constants,
459			n_inout,
460			n_derived,
461			n_precommit,
462			n_private,
463			log_public,
464			log_precommit,
465			log_private,
466			derived_index_map,
467			private_index_map,
468		}
469	}
470
471	pub fn with_blinding(self, info: BlindingInfo) -> Self {
472		// Both committed segments carry the same blinding:
473		//
474		//     dummy wires                   -> mask the codeword symbols FRI opens
475		//     3 wires per dummy constraint  -> mask the evaluations sent in the clear
476		//
477		// Keep this in sync with the padded constraint system in the verifier crate.
478		let blinding_size = info.n_dummy_wires + 3 * info.n_dummy_constraints;
479
480		let total_precommit = self.n_precommit as usize + blinding_size;
481		let log_precommit = log2_ceil_usize(total_precommit) as u32;
482
483		let total_private = self.n_private as usize + blinding_size;
484		let log_private = log2_ceil_usize(total_private) as u32;
485
486		Self {
487			log_precommit,
488			log_private,
489			..self
490		}
491	}
492
493	pub const fn public_size(&self) -> usize {
494		1 << self.log_public as usize
495	}
496
497	pub const fn precommit_size(&self) -> usize {
498		1 << self.log_precommit as usize
499	}
500
501	pub const fn private_size(&self) -> usize {
502		1 << self.log_private as usize
503	}
504
505	pub const fn n_constants(&self) -> usize {
506		self.constants.len()
507	}
508
509	pub const fn n_inout(&self) -> usize {
510		self.n_inout as usize
511	}
512
513	pub const fn n_derived(&self) -> usize {
514		self.n_derived as usize
515	}
516
517	pub const fn n_precommit(&self) -> usize {
518		self.n_precommit as usize
519	}
520
521	pub const fn n_private(&self) -> usize {
522		self.n_private as usize
523	}
524
525	pub const fn log_public(&self) -> u32 {
526		self.log_public
527	}
528
529	pub const fn log_precommit(&self) -> u32 {
530		self.log_precommit
531	}
532
533	pub const fn log_private(&self) -> u32 {
534		self.log_private
535	}
536
537	pub fn get(&self, wire: &ConstraintWire) -> Option<WitnessIndex> {
538		match wire.kind {
539			WireKind::Constant => {
540				assert!((wire.id as usize) < self.constants.len());
541				Some(WitnessIndex::public(wire.id))
542			}
543			WireKind::InOut => {
544				assert!(wire.id < self.n_inout);
545				Some(WitnessIndex::public(self.constants.len() as u32 + wire.id))
546			}
547			WireKind::Derived => self.derived_index_map.get(&wire.id).map(|&derived_idx| {
548				WitnessIndex::public(self.constants.len() as u32 + self.n_inout + derived_idx)
549			}),
550			WireKind::Precommit => {
551				assert!(wire.id < self.n_precommit);
552				Some(WitnessIndex::precommit(wire.id))
553			}
554			WireKind::Private => self
555				.private_index_map
556				.get(&wire.id)
557				.map(|&id| WitnessIndex::private(id)),
558		}
559	}
560}
561
562#[cfg(test)]
563mod tests {
564	use smallvec::smallvec;
565
566	use super::*;
567
568	#[test]
569	fn test_wires_added_mod2() {
570		// Create 4 wires with different kinds to ensure proper sorting
571		let w = [
572			ConstraintWire {
573				kind: WireKind::Constant,
574				id: 0,
575			},
576			ConstraintWire {
577				kind: WireKind::InOut,
578				id: 0,
579			},
580			ConstraintWire {
581				kind: WireKind::Private,
582				id: 0,
583			},
584			ConstraintWire {
585				kind: WireKind::Private,
586				id: 1,
587			},
588		];
589
590		// Input sequence: w[0], w[2], w[2], w[3], w[3], w[1], w[2], w[1], w[3], w[3]
591		// Counts: w[0]=1, w[1]=2, w[2]=3, w[3]=4
592		// After mod 2: w[0]=1, w[1]=0, w[2]=1, w[3]=0
593		let input = smallvec![w[0], w[2], w[2], w[3], w[3], w[1], w[2], w[1], w[3], w[3]];
594		let operand = Operand::new(input);
595
596		// Expected result: w[0], w[2] (sorted)
597		assert_eq!(operand.wires(), &[w[0], w[2]]);
598	}
599
600	#[test]
601	fn test_sorting_when_no_duplicates() {
602		// Create 4 wires with different kinds to ensure proper sorting
603		let w = [
604			ConstraintWire {
605				kind: WireKind::Constant,
606				id: 0,
607			},
608			ConstraintWire {
609				kind: WireKind::InOut,
610				id: 0,
611			},
612			ConstraintWire {
613				kind: WireKind::Private,
614				id: 0,
615			},
616			ConstraintWire {
617				kind: WireKind::Private,
618				id: 1,
619			},
620		];
621
622		// Input sequence: w[2], w[3], w[0], w[1]
623		let input = smallvec![w[2], w[3], w[0], w[1]];
624		let operand = Operand::new(input);
625
626		// Expected result: w[0], w[1], w[2], w[3] (sorted by WireKind then ID)
627		assert_eq!(operand.wires(), &[w[0], w[1], w[2], w[3]]);
628	}
629
630	#[test]
631	fn test_merge() {
632		// Create 4 wires with different kinds to ensure proper sorting
633		let w = [
634			ConstraintWire {
635				kind: WireKind::Constant,
636				id: 0,
637			},
638			ConstraintWire {
639				kind: WireKind::InOut,
640				id: 0,
641			},
642			ConstraintWire {
643				kind: WireKind::Private,
644				id: 0,
645			},
646			ConstraintWire {
647				kind: WireKind::Private,
648				id: 1,
649			},
650		];
651
652		// Test case 1: merge([w[0]], [])
653		let mut lhs = Operand(smallvec![w[0]]);
654		let rhs = Operand(smallvec![]);
655		let (additions, removals) = lhs.merge(&rhs);
656		assert_eq!(lhs.wires(), &[w[0]]);
657		assert_eq!(additions.wires(), &[]);
658		assert_eq!(removals.wires(), &[]);
659
660		// Test case 2: merge([], [w[0]])
661		let mut lhs = Operand(smallvec![]);
662		let rhs = Operand(smallvec![w[0]]);
663		let (additions, removals) = lhs.merge(&rhs);
664		assert_eq!(lhs.wires(), &[w[0]]);
665		assert_eq!(additions.wires(), &[w[0]]);
666		assert_eq!(removals.wires(), &[]);
667
668		// Test case 3: merge([w[0]], [w[0]])
669		let mut lhs = Operand(smallvec![w[0]]);
670		let rhs = Operand(smallvec![w[0]]);
671		let (additions, removals) = lhs.merge(&rhs);
672		assert_eq!(lhs.wires(), &[]);
673		assert_eq!(additions.wires(), &[]);
674		assert_eq!(removals.wires(), &[w[0]]);
675
676		// Test case 4: merge([w[0]], [w[1]])
677		let mut lhs = Operand(smallvec![w[0]]);
678		let rhs = Operand(smallvec![w[1]]);
679		let (additions, removals) = lhs.merge(&rhs);
680		assert_eq!(lhs.wires(), &[w[0], w[1]]);
681		assert_eq!(additions.wires(), &[w[1]]);
682		assert_eq!(removals.wires(), &[]);
683
684		// Test case 5: merge([w[0]], [w[0], w[1]])
685		let mut lhs = Operand(smallvec![w[0]]);
686		let rhs = Operand(smallvec![w[0], w[1]]);
687		let (additions, removals) = lhs.merge(&rhs);
688		assert_eq!(lhs.wires(), &[w[1]]);
689		assert_eq!(additions.wires(), &[w[1]]);
690		assert_eq!(removals.wires(), &[w[0]]);
691
692		// Test case 6: merge([w[0], w[2]], [w[1], w[3]])
693		let mut lhs = Operand(smallvec![w[0], w[2]]);
694		let rhs = Operand(smallvec![w[1], w[3]]);
695		let (additions, removals) = lhs.merge(&rhs);
696		assert_eq!(lhs.wires(), &[w[0], w[1], w[2], w[3]]);
697		assert_eq!(additions.wires(), &[w[1], w[3]]);
698		assert_eq!(removals.wires(), &[]);
699
700		// Test case 7: merge([w[0], w[2]], [w[0], w[1], w[2], w[3]])
701		let mut lhs = Operand(smallvec![w[0], w[2]]);
702		let rhs = Operand(smallvec![w[0], w[1], w[2], w[3]]);
703		let (additions, removals) = lhs.merge(&rhs);
704		assert_eq!(lhs.wires(), &[w[1], w[3]]);
705		assert_eq!(additions.wires(), &[w[1], w[3]]);
706		assert_eq!(removals.wires(), &[w[0], w[2]]);
707	}
708}