Skip to main content

binius_spartan_frontend/
circuits.rs

1// Copyright 2025 Irreducible Inc.
2
3use std::{iter, iter::successors};
4
5use binius_field::{Field, arithmetic_traits::InvertOrZero};
6
7use crate::circuit_builder::CircuitBuilder;
8
9pub fn extrapolate_line<Builder: CircuitBuilder>(
10	builder: &mut Builder,
11	y0: Builder::Wire,
12	y1: Builder::Wire,
13	z: Builder::Wire,
14) -> Builder::Wire {
15	// y(z) = y0 + (y1 - y0) * z
16	// In binary fields, subtraction is addition (XOR)
17	let diff = builder.add(y1, y0);
18	let scaled = builder.mul(diff, z);
19	builder.add(y0, scaled)
20}
21
22pub fn evaluate_univariate<Builder: CircuitBuilder>(
23	builder: &mut Builder,
24	coeffs: &[Builder::Wire],
25	z: Builder::Wire,
26) -> Builder::Wire {
27	// Use Horner's method: p(z) = a0 + z(a1 + z(a2 + z(...)))
28	// Start from highest degree coefficient and work backwards
29	if coeffs.is_empty() {
30		return builder.constant(Builder::Field::ZERO);
31	}
32
33	coeffs[..coeffs.len() - 1]
34		.iter()
35		.rev()
36		.fold(coeffs[coeffs.len() - 1], |acc, &coeff| {
37			let temp = builder.mul(acc, z);
38			builder.add(temp, coeff)
39		})
40}
41
42pub fn evaluate_multilinear<Builder: CircuitBuilder>(
43	builder: &mut Builder,
44	coeffs: &[Builder::Wire],
45	coords: &[Builder::Wire],
46) -> Vec<Builder::Wire> {
47	// coords has length n, coeffs has length 2^n
48	// Evaluation algorithm: fold over each coordinate in reverse order
49	// For each coordinate, interpolate between pairs: lo + coord * (hi - lo)
50	coords
51		.iter()
52		.rev()
53		.fold(coeffs.to_vec(), |current, &coord| {
54			let (evals_0, evals_1) = current.split_at(current.len() / 2);
55			iter::zip(evals_0, evals_1)
56				.map(|(&eval_0, &eval_1)| extrapolate_line(builder, eval_0, eval_1, coord))
57				.collect()
58		})
59}
60
61pub fn powers<Builder: CircuitBuilder>(
62	builder: &mut Builder,
63	x: Builder::Wire,
64	n: usize,
65) -> Vec<Builder::Wire> {
66	// return a vector of n wires containing the values of x^i for i in [1, n]
67	successors(Some(x), |&prev| Some(builder.mul(prev, x)))
68		.take(n)
69		.collect()
70}
71
72pub fn square<Builder: CircuitBuilder>(builder: &mut Builder, x: Builder::Wire) -> Builder::Wire {
73	builder.mul(x, x)
74}
75
76pub fn invert<Builder: CircuitBuilder>(builder: &mut Builder, x: Builder::Wire) -> Builder::Wire {
77	let [inv] = builder.hint([x], |vals| {
78		let [x_val] = vals;
79		[x_val.invert_or_zero()]
80	});
81
82	let one = builder.constant(Builder::Field::ONE);
83	let prod = builder.mul(x, inv);
84	builder.assert_eq(prod, one);
85
86	inv
87}
88
89pub fn invert_or_zero<Builder: CircuitBuilder>(
90	builder: &mut Builder,
91	x: Builder::Wire,
92) -> Builder::Wire {
93	let [inv] = builder.hint([x], |vals| {
94		let [x_val] = vals;
95		[x_val.invert_or_zero()]
96	});
97
98	let one = builder.constant(Builder::Field::ONE);
99	let prod = builder.mul(x, inv);
100	let prod_sub_one = builder.sub(prod, one);
101
102	let prod_sub_one_or_x = builder.mul(prod_sub_one, x);
103	builder.assert_zero(prod_sub_one_or_x);
104
105	let prod_sub_one_or_inv = builder.mul(prod_sub_one, inv);
106	builder.assert_zero(prod_sub_one_or_inv);
107
108	inv
109}
110
111pub fn assert_is_bit<Builder: CircuitBuilder>(builder: &mut Builder, val: Builder::Wire) {
112	let val_sq = square(builder, val);
113	builder.assert_eq(val_sq, val);
114}
115
116#[cfg(test)]
117mod tests {
118	use std::{array, iter};
119
120	use binius_field::{Field, Ghash128b as B128, Random, arithmetic_traits::Square};
121	use binius_math::{
122		line::extrapolate_line as extrapolate_line_math,
123		multilinear::evaluate::evaluate,
124		test_utils::{random_field_buffer, random_scalars},
125		univariate,
126	};
127	use rand::{SeedableRng, rngs::StdRng};
128
129	use super::*;
130	use crate::{
131		circuit_builder::{ConstraintBuilder, WitnessError, WitnessGenerator},
132		compiler::compile,
133	};
134
135	trait TestCircuit<const N_INOUT: usize> {
136		fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; N_INOUT]);
137	}
138
139	fn test_helper<C: TestCircuit<N_INOUT>, const N_INOUT: usize>(
140		inout_vals: [B128; N_INOUT],
141	) -> Result<(), WitnessError> {
142		let mut constraint_builder = ConstraintBuilder::new();
143		let inout_wires = array::from_fn(|_| constraint_builder.alloc_inout());
144		C::build(&mut constraint_builder, inout_wires);
145		let (cs, layout) = compile(constraint_builder);
146
147		let mut witness_gen = WitnessGenerator::new(&layout);
148		let inout_assigned =
149			array::from_fn(|i| witness_gen.write_inout(inout_wires[i], inout_vals[i]));
150		C::build(&mut witness_gen, inout_assigned);
151		let witness = witness_gen.build()?;
152
153		cs.validate(&witness);
154		Ok(())
155	}
156
157	#[test]
158	fn test_square() {
159		struct SquareCircuit;
160
161		impl TestCircuit<2> for SquareCircuit {
162			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 2]) {
163				let [x, expected] = inout;
164				let result = square(builder, x);
165				builder.assert_eq(result, expected);
166			}
167		}
168
169		let mut rng = StdRng::seed_from_u64(0);
170		let x_val = B128::random(&mut rng);
171		let expected = Square::square(x_val);
172
173		test_helper::<SquareCircuit, 2>([x_val, expected]).unwrap();
174	}
175
176	#[test]
177	fn test_assert_is_bit_zero() {
178		struct BitCheckCircuit;
179
180		impl TestCircuit<1> for BitCheckCircuit {
181			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 1]) {
182				assert_is_bit(builder, inout[0]);
183			}
184		}
185
186		// Test that 0 is a valid bit
187		test_helper::<BitCheckCircuit, 1>([B128::ZERO]).unwrap();
188	}
189
190	#[test]
191	fn test_assert_is_bit_one() {
192		struct BitCheckCircuit;
193
194		impl TestCircuit<1> for BitCheckCircuit {
195			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 1]) {
196				assert_is_bit(builder, inout[0]);
197			}
198		}
199
200		// Test that 1 is a valid bit
201		test_helper::<BitCheckCircuit, 1>([B128::ONE]).unwrap();
202	}
203
204	#[test]
205	fn test_extrapolate_line() {
206		struct ExtrapolateLineCircuit;
207
208		impl TestCircuit<4> for ExtrapolateLineCircuit {
209			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 4]) {
210				let [y0, y1, z, expected] = inout;
211				let result = extrapolate_line(builder, y0, y1, z);
212				builder.assert_eq(result, expected);
213			}
214		}
215
216		let mut rng = StdRng::seed_from_u64(0);
217		let y0_val = B128::random(&mut rng);
218		let y1_val = B128::random(&mut rng);
219		let z_val = B128::random(&mut rng);
220		let expected = extrapolate_line_math(y0_val, y1_val, z_val);
221
222		test_helper::<ExtrapolateLineCircuit, 4>([y0_val, y1_val, z_val, expected]).unwrap();
223	}
224
225	#[test]
226	fn test_evaluate_univariate() {
227		struct UnivariateCircuit;
228
229		impl TestCircuit<5> for UnivariateCircuit {
230			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 5]) {
231				let (coeffs, rest) = inout.split_at(3);
232				let [z, expected] = rest else { unreachable!() };
233				let result = evaluate_univariate(builder, coeffs, *z);
234				builder.assert_eq(result, *expected);
235			}
236		}
237
238		let mut rng = StdRng::seed_from_u64(0);
239		let coeffs_vals = [
240			B128::random(&mut rng),
241			B128::random(&mut rng),
242			B128::random(&mut rng),
243		];
244		let z_val = B128::random(&mut rng);
245		let expected = univariate::evaluate_univariate(&coeffs_vals, &z_val);
246
247		test_helper::<UnivariateCircuit, 5>([
248			coeffs_vals[0],
249			coeffs_vals[1],
250			coeffs_vals[2],
251			z_val,
252			expected,
253		])
254		.unwrap();
255	}
256
257	#[test]
258	fn test_evaluate_multilinear() {
259		struct MultilinearCircuit;
260
261		impl TestCircuit<7> for MultilinearCircuit {
262			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 7]) {
263				let (coeffs, rest) = inout.split_at(4);
264				let (coords, expected_slice) = rest.split_at(2);
265				let expected = expected_slice[0];
266				let result = evaluate_multilinear(builder, coeffs, coords);
267				assert_eq!(result.len(), 1);
268				builder.assert_eq(result[0], expected);
269			}
270		}
271
272		let mut rng = StdRng::seed_from_u64(0);
273		let coeffs_vals = random_field_buffer::<B128>(&mut rng, 2);
274		let coords_vals = random_scalars(&mut rng, 2);
275		let expected = evaluate(&coeffs_vals, &coords_vals);
276
277		test_helper::<MultilinearCircuit, 7>([
278			coeffs_vals.get(0),
279			coeffs_vals.get(1),
280			coeffs_vals.get(2),
281			coeffs_vals.get(3),
282			coords_vals[0],
283			coords_vals[1],
284			expected,
285		])
286		.unwrap();
287	}
288
289	#[test]
290	fn test_powers() {
291		struct PowersCircuit;
292
293		impl<const N_INOUT: usize> TestCircuit<N_INOUT> for PowersCircuit {
294			fn build<Builder: CircuitBuilder>(
295				builder: &mut Builder,
296				inout: [Builder::Wire; N_INOUT],
297			) {
298				let x = inout[0];
299				let expected = &inout[1..];
300				let result = powers(builder, x, N_INOUT - 1);
301				assert_eq!(result.len(), expected.len());
302				for (r, e) in iter::zip(&result, expected) {
303					builder.assert_eq(*r, *e);
304				}
305			}
306		}
307
308		let mut rng = StdRng::seed_from_u64(0);
309		let x_val = B128::random(&mut rng);
310		let expected_vals = binius_field::util::powers(x_val)
311			.skip(1)
312			.take(4)
313			.collect::<Vec<_>>();
314
315		test_helper::<PowersCircuit, 5>([
316			x_val,
317			expected_vals[0],
318			expected_vals[1],
319			expected_vals[2],
320			expected_vals[3],
321		])
322		.unwrap();
323	}
324
325	#[test]
326	fn test_invert() {
327		struct InvertCircuit;
328
329		impl TestCircuit<2> for InvertCircuit {
330			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 2]) {
331				let [x, expected] = inout;
332				let result = invert(builder, x);
333				builder.assert_eq(result, expected);
334			}
335		}
336
337		let mut rng = StdRng::seed_from_u64(0);
338		let x_val = B128::random(&mut rng);
339		let expected = x_val.invert_or_zero();
340
341		test_helper::<InvertCircuit, 2>([x_val, expected]).unwrap();
342		assert!(test_helper::<InvertCircuit, 2>([B128::ZERO, B128::ZERO]).is_err());
343	}
344
345	#[test]
346	fn test_invert_or_zero_nonzero() {
347		struct InvertOrZeroCircuit;
348
349		impl TestCircuit<2> for InvertOrZeroCircuit {
350			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 2]) {
351				let [x, expected] = inout;
352				let result = invert_or_zero(builder, x);
353				builder.assert_eq(result, expected);
354			}
355		}
356
357		let mut rng = StdRng::seed_from_u64(0);
358		let x_val = B128::random(&mut rng);
359		let expected = x_val.invert_or_zero();
360
361		test_helper::<InvertOrZeroCircuit, 2>([x_val, expected]).unwrap();
362	}
363
364	#[test]
365	fn test_invert_or_zero_zero() {
366		struct InvertOrZeroCircuit;
367
368		impl TestCircuit<2> for InvertOrZeroCircuit {
369			fn build<Builder: CircuitBuilder>(builder: &mut Builder, inout: [Builder::Wire; 2]) {
370				let [x, expected] = inout;
371				let result = invert_or_zero(builder, x);
372				builder.assert_eq(result, expected);
373			}
374		}
375
376		test_helper::<InvertOrZeroCircuit, 2>([B128::ZERO, B128::ZERO]).unwrap();
377	}
378}