1use 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 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 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
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 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_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_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}