1use std::{
24 cell::RefCell,
25 fmt,
26 iter::{Product, Sum},
27 ops::{Add, AddAssign, Mul, MulAssign, Neg, Sub, SubAssign},
28 rc::{Rc, Weak},
29};
30
31use binius_field::{
32 ExtensionField, Field,
33 arithmetic_traits::{InvertOrZero, Square},
34 field::FieldOps,
35};
36use binius_spartan_frontend::circuit_builder::CircuitBuilder;
37
38use super::gadgets;
39
40pub enum CircuitElem<F: Field, B: CircuitBuilder<Field = F>> {
46 Constant(F),
47 Wire {
48 builder: Weak<RefCell<B>>,
49 wire: B::Wire,
50 },
51}
52
53impl<F: Field, B: CircuitBuilder<Field = F>> fmt::Debug for CircuitElem<F, B>
56where
57 B::Wire: fmt::Debug,
58{
59 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
60 match self {
61 Self::Constant(c) => f.debug_tuple("Constant").field(c).finish(),
62 Self::Wire { wire, .. } => f.debug_struct("Wire").field("wire", wire).finish(),
63 }
64 }
65}
66
67impl<F: Field, B: CircuitBuilder<Field = F>> Clone for CircuitElem<F, B> {
70 fn clone(&self) -> Self {
71 match self {
72 Self::Constant(c) => Self::Constant(*c),
73 Self::Wire { builder, wire } => Self::Wire {
74 builder: builder.clone(),
75 wire: *wire,
76 },
77 }
78 }
79}
80
81impl<F, B> CircuitElem<F, B>
82where
83 F: Field,
84 B: CircuitBuilder<Field = F>,
85{
86 pub fn wire(builder: &Rc<RefCell<B>>, wire: B::Wire) -> Self {
88 Self::Wire {
89 builder: Rc::downgrade(builder),
90 wire,
91 }
92 }
93
94 pub fn to_wire(&self, builder: &mut B) -> B::Wire {
99 match self {
100 Self::Constant(val) => builder.constant(*val),
101 Self::Wire { wire, .. } => *wire,
102 }
103 }
104
105 #[allow(clippy::option_if_let_else)]
111 pub fn combine<const IN: usize, const OUT: usize>(
112 elems: [&Self; IN],
113 f_op: impl Fn([F; IN]) -> [F; OUT],
114 builder_op: impl Fn(&mut B, [B::Wire; IN]) -> [B::Wire; OUT],
115 ) -> [Self; OUT] {
116 let builder = elems.iter().find_map(|elem| match elem {
117 Self::Wire { builder, .. } => Some(builder),
118 _ => None,
119 });
120
121 if let Some(builder_ptr) = builder {
122 let Some(builder) = builder_ptr.upgrade() else {
123 panic!("combine cannot be called on a CircuitElem after the channel is closed");
124 };
125 let mut builder = builder.borrow_mut();
126 let inner_wires = elems.map(|elem| match elem {
127 Self::Constant(val) => builder.constant(*val),
128 Self::Wire {
129 builder: other_builder_ptr,
130 wire,
131 } => {
132 assert!(
133 Weak::ptr_eq(builder_ptr, other_builder_ptr),
134 "all combined CircuitElems must come from the same channel"
135 );
136 *wire
137 }
138 });
139 builder_op(&mut builder, inner_wires).map(|wire| Self::Wire {
140 builder: builder_ptr.clone(),
141 wire,
142 })
143 } else {
144 let inner_constants = elems.map(|elem| {
145 let Self::Constant(val) = elem else {
146 unreachable!(
147 "the enum has only two variants; none of them are Wire; thus all must be Constant"
148 );
149 };
150 *val
151 });
152 f_op(inner_constants).map(Self::Constant)
153 }
154 }
155
156 #[allow(clippy::option_if_let_else)]
161 pub fn combine_varlen(
162 elems: &[&Self],
163 n_out: usize,
164 f_op: impl FnOnce(&[F]) -> Vec<F>,
165 builder_op: impl FnOnce(&mut B, &[B::Wire]) -> Vec<B::Wire>,
166 ) -> Vec<Self> {
167 let builder = elems.iter().find_map(|elem| match elem {
168 Self::Wire { builder, .. } => Some(builder),
169 _ => None,
170 });
171
172 if let Some(builder_ptr) = builder {
173 let Some(builder) = builder_ptr.upgrade() else {
174 panic!(
175 "combine_varlen cannot be called on a CircuitElem after the channel is closed"
176 );
177 };
178 let mut builder = builder.borrow_mut();
179 let inner_wires = elems
180 .iter()
181 .map(|elem| match elem {
182 Self::Constant(val) => builder.constant(*val),
183 Self::Wire {
184 builder: other_builder_ptr,
185 wire,
186 } => {
187 assert!(
188 Weak::ptr_eq(builder_ptr, other_builder_ptr),
189 "all combined CircuitElems must come from the same channel"
190 );
191 *wire
192 }
193 })
194 .collect::<Vec<_>>();
195 let result = builder_op(&mut builder, &inner_wires);
196 debug_assert_eq!(result.len(), n_out);
197 result
198 .into_iter()
199 .map(|wire| Self::Wire {
200 builder: builder_ptr.clone(),
201 wire,
202 })
203 .collect()
204 } else {
205 let inner_constants = elems
206 .iter()
207 .map(|elem| {
208 let Self::Constant(val) = elem else {
209 unreachable!(
210 "no Wire variant exists in elems; all entries must be Constant"
211 );
212 };
213 *val
214 })
215 .collect::<Vec<_>>();
216 let result = f_op(&inner_constants);
217 debug_assert_eq!(result.len(), n_out);
218 result.into_iter().map(Self::Constant).collect()
219 }
220 }
221}
222
223impl<F: Field, B: CircuitBuilder<Field = F>> Neg for CircuitElem<F, B> {
226 type Output = Self;
227
228 fn neg(self) -> Self {
229 self
230 }
231}
232
233impl<F: Field, B: CircuitBuilder<Field = F>> Add for CircuitElem<F, B> {
234 type Output = Self;
235
236 fn add(self, rhs: Self) -> Self {
237 self + &rhs
238 }
239}
240
241impl<F: Field, B: CircuitBuilder<Field = F>> Sub for CircuitElem<F, B> {
242 type Output = Self;
243
244 fn sub(self, rhs: Self) -> Self {
245 self - &rhs
246 }
247}
248
249impl<F: Field, B: CircuitBuilder<Field = F>> Mul for CircuitElem<F, B> {
250 type Output = Self;
251
252 fn mul(self, rhs: Self) -> Self {
253 self * &rhs
254 }
255}
256
257impl<F, B> Add<&Self> for CircuitElem<F, B>
260where
261 F: Field,
262 B: CircuitBuilder<Field = F>,
263{
264 type Output = Self;
265
266 fn add(self, rhs: &Self) -> Self {
267 &self + rhs
268 }
269}
270
271impl<F, B> Sub<&Self> for CircuitElem<F, B>
272where
273 F: Field,
274 B: CircuitBuilder<Field = F>,
275{
276 type Output = Self;
277
278 fn sub(self, rhs: &Self) -> Self {
279 &self - rhs
280 }
281}
282
283impl<F, B> Mul<&Self> for CircuitElem<F, B>
284where
285 F: Field,
286 B: CircuitBuilder<Field = F>,
287{
288 type Output = Self;
289
290 fn mul(self, rhs: &Self) -> Self {
291 &self * rhs
292 }
293}
294
295impl<F, B> Add for &CircuitElem<F, B>
296where
297 F: Field,
298 B: CircuitBuilder<Field = F>,
299{
300 type Output = CircuitElem<F, B>;
301
302 fn add(self, rhs: Self) -> Self::Output {
303 let [ret] = CircuitElem::combine(
304 [self, rhs],
305 |[lhs, rhs]| [lhs + rhs],
306 |builder, [lhs, rhs]| [builder.add(lhs, rhs)],
307 );
308 ret
309 }
310}
311
312impl<F, B> Sub for &CircuitElem<F, B>
313where
314 F: Field,
315 B: CircuitBuilder<Field = F>,
316{
317 type Output = CircuitElem<F, B>;
318
319 fn sub(self, rhs: Self) -> Self::Output {
320 let [ret] = CircuitElem::combine(
321 [self, rhs],
322 |[lhs, rhs]| [lhs - rhs],
323 |builder, [lhs, rhs]| [builder.sub(lhs, rhs)],
324 );
325 ret
326 }
327}
328
329impl<F, B> Mul for &CircuitElem<F, B>
330where
331 F: Field,
332 B: CircuitBuilder<Field = F>,
333{
334 type Output = CircuitElem<F, B>;
335
336 fn mul(self, rhs: Self) -> Self::Output {
337 if matches!(self, CircuitElem::Constant(c) if *c == F::ZERO)
340 || matches!(rhs, CircuitElem::Constant(c) if *c == F::ZERO)
341 {
342 return CircuitElem::Constant(F::ZERO);
343 }
344 let [ret] = CircuitElem::combine(
345 [self, rhs],
346 |[lhs, rhs]| [lhs * rhs],
347 |builder, [lhs, rhs]| [builder.mul(lhs, rhs)],
348 );
349 ret
350 }
351}
352
353impl<F: Field, B: CircuitBuilder<Field = F>> AddAssign for CircuitElem<F, B> {
356 fn add_assign(&mut self, rhs: Self) {
357 *self = &*self + &rhs;
358 }
359}
360
361impl<F: Field, B: CircuitBuilder<Field = F>> SubAssign for CircuitElem<F, B> {
362 fn sub_assign(&mut self, rhs: Self) {
363 *self = &*self - &rhs;
364 }
365}
366
367impl<F: Field, B: CircuitBuilder<Field = F>> MulAssign for CircuitElem<F, B> {
368 fn mul_assign(&mut self, rhs: Self) {
369 *self = &*self * &rhs;
370 }
371}
372
373impl<F: Field, B: CircuitBuilder<Field = F>> AddAssign<&Self> for CircuitElem<F, B> {
374 fn add_assign(&mut self, rhs: &Self) {
375 *self = &*self + rhs;
376 }
377}
378
379impl<F: Field, B: CircuitBuilder<Field = F>> SubAssign<&Self> for CircuitElem<F, B> {
380 fn sub_assign(&mut self, rhs: &Self) {
381 *self = &*self - rhs;
382 }
383}
384
385impl<F: Field, B: CircuitBuilder<Field = F>> MulAssign<&Self> for CircuitElem<F, B> {
386 fn mul_assign(&mut self, rhs: &Self) {
387 *self = &*self * rhs;
388 }
389}
390
391impl<F: Field, B: CircuitBuilder<Field = F>> Sum for CircuitElem<F, B> {
394 fn sum<I: Iterator<Item = Self>>(iter: I) -> Self {
395 iter.fold(CircuitElem::Constant(F::ZERO), |acc, x| acc + x)
396 }
397}
398
399impl<'a, F: Field, B: CircuitBuilder<Field = F>> Sum<&'a Self> for CircuitElem<F, B> {
400 fn sum<I: Iterator<Item = &'a Self>>(iter: I) -> Self {
401 iter.fold(Self::Constant(F::ZERO), |acc, x| acc + x)
402 }
403}
404
405impl<F: Field, B: CircuitBuilder<Field = F>> Product for CircuitElem<F, B> {
406 fn product<I: Iterator<Item = Self>>(iter: I) -> Self {
407 iter.fold(Self::Constant(F::ONE), |acc, x| acc * x)
408 }
409}
410
411impl<'a, F: Field, B: CircuitBuilder<Field = F>> Product<&'a Self> for CircuitElem<F, B> {
412 fn product<I: Iterator<Item = &'a Self>>(iter: I) -> Self {
413 iter.fold(Self::Constant(F::ONE), |acc, x| acc * x)
414 }
415}
416
417impl<F: Field, B: CircuitBuilder<Field = F>> Square for CircuitElem<F, B> {
418 fn square(self) -> Self {
419 let [ret] = Self::combine([&self], |[x]| [x.square()], |builder, [x]| [builder.mul(x, x)]);
420 ret
421 }
422}
423
424impl<F: Field, B: CircuitBuilder<Field = F>> InvertOrZero for CircuitElem<F, B> {
425 fn invert_or_zero(self) -> Self {
434 unimplemented!(
435 "the wrapper inverts only values argued non-zero; use `invert` (see its safety \
436 contract), or implement this if a zero-admitting inverse is ever needed"
437 )
438 }
439
440 unsafe fn invert(self) -> Self {
447 let [ret] = Self::combine(
448 [&self],
449 |[x]| [unsafe { x.invert() }],
451 |builder, [x]| {
452 let [inv] = builder.hint([x], |[v]| [v.invert_or_zero()]);
453 let one = builder.constant(F::ONE);
454 let product = builder.mul(x, inv);
455 builder.assert_eq(product, one);
456 [inv]
457 },
458 );
459 ret
460 }
461}
462
463impl<F: Field, B: CircuitBuilder<Field = F>> FieldOps for CircuitElem<F, B> {
464 type Scalar = F;
465
466 fn zero() -> Self {
467 Self::Constant(F::ZERO)
468 }
469
470 fn one() -> Self {
471 Self::Constant(F::ONE)
472 }
473
474 fn square_transpose<FSub: Field>(elems: &mut [Self])
475 where
476 Self::Scalar: ExtensionField<FSub>,
477 {
478 let degree = F::DEGREE;
479 assert_eq!(elems.len(), degree);
480
481 if degree == 1 {
482 return;
483 }
484
485 let inputs = elems.iter().collect::<Vec<_>>();
486 let outputs = Self::combine_varlen(
487 &inputs,
488 degree,
489 |vals| {
490 let mut out = vals.to_vec();
491 <F as ExtensionField<FSub>>::square_transpose(&mut out);
492 out
493 },
494 |builder, wires| gadgets::square_transpose::<_, FSub>(builder, wires),
495 );
496 for (e, out) in elems.iter_mut().zip(outputs) {
497 *e = out;
498 }
499 }
500}
501
502impl<F: Field, B: CircuitBuilder<Field = F>> From<F> for CircuitElem<F, B> {
503 fn from(val: F) -> Self {
504 Self::Constant(val)
505 }
506}