binius_spartan_verifier/wrapper/gadgets.rs
1// Copyright 2026 The Binius Developers
2
3//! Circuit gadgets — functions generic over [`CircuitBuilder`] that take input wires and return
4//! output wires.
5
6use binius_field::{ExtensionField, Field};
7use binius_spartan_frontend::circuit_builder::CircuitBuilder;
8
9/// Transpose the subfield decomposition of extension field elements.
10///
11/// Given `d` input wires representing extension field elements (where `d` is the degree of
12/// `B::Field` over `FSub`), returns `d` output wires containing the transposed elements.
13///
14/// Each input element decomposes as `input[i] = sum_j coeffs[i][j] * basis(j)` where
15/// `coeffs[i][j]` are in `FSub`. The output satisfies `output[j] = sum_i coeffs[i][j] * basis(i)`,
16/// i.e., `output[j].get_base(i) == input[i].get_base(j)`.
17///
18/// The gadget:
19/// 1. Hints the `d × d` matrix of subfield coefficients
20/// 2. Constrains each coefficient to lie in `FSub` via the Frobenius endomorphism
21/// 3. Constrains that the coefficients reconstruct each input element
22/// 4. Computes the transposed output as basis linear combinations
23///
24/// # Panics
25///
26/// * If `inputs.len() != B::Field::DEGREE`
27///
28/// Instantiating this with a `B::Field` of characteristic other than 2 fails to compile.
29pub fn square_transpose<B: CircuitBuilder, FSub: Field>(
30 builder: &mut B,
31 inputs: &[B::Wire],
32) -> Vec<B::Wire>
33where
34 B::Field: ExtensionField<FSub>,
35{
36 const {
37 assert!(
38 B::Field::CHARACTERISTIC == 2,
39 "square_transpose gadget is only implemented for characteristic 2"
40 );
41 }
42
43 let degree = B::Field::DEGREE;
44 assert_eq!(inputs.len(), degree);
45
46 if degree == 1 {
47 return inputs.to_vec();
48 }
49
50 // An element c of B::Field is in FSub iff c^(2^k) = c, where k = log_2(|FSub|).
51 // Since |B::Field| = 2^ORDER_EXPONENT and [B::Field : FSub] = degree, we have
52 // k = ORDER_EXPONENT / degree.
53 let n_frobenius_squarings = B::Field::ORDER_EXPONENT / degree;
54
55 // Hint the d×d matrix of subfield coefficients.
56 // Each element decomposes as: inputs[i] = sum_j coeffs[i][j] * basis(j)
57 // where coeffs[i][j] is in FSub (embedded in B::Field).
58 let coeffs = (0..degree)
59 .map(|i| {
60 (0..degree)
61 .map(|j| {
62 let [c] = builder
63 .hint([inputs[i]], move |[x]| [B::Field::from(B::Field::get_base(&x, j))]);
64 c
65 })
66 .collect::<Vec<_>>()
67 })
68 .collect::<Vec<_>>();
69
70 // Frobenius subfield membership check: for each coefficient c, verify
71 // c^(2^k) = c by squaring k times and asserting equality.
72 for row in &coeffs {
73 for &c in row {
74 let mut powered = c;
75 for _ in 0..n_frobenius_squarings {
76 powered = builder.mul(powered, powered);
77 }
78 builder.assert_eq(powered, c);
79 }
80 }
81
82 // Reconstruction check: verify that the hinted coefficients actually
83 // decompose each original element.
84 // For each i: sum_j coeffs[i][j] * basis(j) == inputs[i]
85 for i in 0..degree {
86 let reconstructed = basis_linear_combination::<_, FSub>(builder, coeffs[i].iter().copied());
87 builder.assert_eq(reconstructed, inputs[i]);
88 }
89
90 // Compute transposed elements: out[j] = sum_i coeffs[i][j] * basis(i)
91 (0..degree)
92 .map(|j| basis_linear_combination::<_, FSub>(builder, coeffs.iter().map(|row| row[j])))
93 .collect::<Vec<_>>()
94}
95
96/// Compute the linear combination `sum_i scalars[i] * basis(i)` where `basis(i)` are the
97/// extension field basis elements of `B::Field` over `FSub`.
98fn basis_linear_combination<B: CircuitBuilder, FSub: Field>(
99 builder: &mut B,
100 scalars: impl ExactSizeIterator<Item = B::Wire>,
101) -> B::Wire
102where
103 B::Field: ExtensionField<FSub>,
104{
105 assert_eq!(scalars.len(), B::Field::DEGREE);
106
107 // basis(0) is always ONE, so the first term is just scalars[0].
108 let mut scalars = scalars.enumerate();
109 let (_, first) = scalars.next().expect("degree must be at least 1");
110 scalars.fold(first, |sum, (j, scalar)| {
111 let basis = builder.constant(B::Field::basis(j));
112 let term = builder.mul(scalar, basis);
113 builder.add(sum, term)
114 })
115}