use ark_ec::scalar_mul::glv::GLVConfig;
use ark_ff::{BigInteger, PrimeField};
use rayon::prelude::*;
use snarkrs_field::{FftField, Field, Fr};
use crate::Packed;
#[repr(C)]
#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
pub struct PackedGlv {
pub k: [u32; 8],
pub sign: u32,
}
const _: () = {
assert!(core::mem::size_of::<PackedGlv>() == 36);
assert!(core::mem::align_of::<PackedGlv>() == 4);
};
unsafe impl Packed for PackedGlv {}
pub fn twiddle_table<C: GLVConfig<ScalarField = Fr>>(bits: u32, first: Fr) -> Vec<PackedGlv> {
const CHUNK: usize = 256;
assert!(
(1..=Fr::TWO_ADICITY).contains(&bits),
"twiddle table over 2^{bits}, outside the {}-bit two-adic subgroup",
Fr::TWO_ADICITY
);
let half = 1usize << (bits - 1);
let mut root = Fr::TWO_ADIC_ROOT_OF_UNITY;
for _ in bits..Fr::TWO_ADICITY {
root.square_in_place();
}
let mut out = vec![PackedGlv::default(); half];
out.par_chunks_mut(CHUNK)
.enumerate()
.for_each(|(ci, chunk)| {
let mut w = first * root.pow([(ci * CHUNK) as u64]);
for slot in chunk.iter_mut() {
*slot = decompose::<C>(&w);
w *= root;
}
});
out
}
pub fn decompose<C: GLVConfig<ScalarField = Fr>>(k: &Fr) -> PackedGlv {
let ((pos1, k1), (pos2, k2)) = C::scalar_decomposition(*k);
let mut out = PackedGlv {
k: [0; 8],
sign: u32::from(!pos1) | (u32::from(!pos2) << 1),
};
for (half, mag) in [(0usize, k1), (4, k2)] {
let b = mag.into_bigint().0;
assert!(
b[2] == 0 && b[3] == 0 && b[1] >> 63 == 0,
"GLV magnitude is {} bits, past the 127 the table holds",
mag.into_bigint().num_bits()
);
for (i, w) in b[..2].iter().enumerate() {
out.k[half + 2 * i] = *w as u32;
out.k[half + 2 * i + 1] = (*w >> 32) as u32;
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::testrng::SplitMix64;
use snarkrs_field::{AffineRepr, CurveGroup, Fq, G1Affine, G2Affine};
#[test]
fn beta_is_one_cube_root_and_the_eigenvalues_are_two() {
use ark_ff::AdditiveGroup;
let beta = <snarkrs_field::g1::Config as GLVConfig>::ENDO_COEFFS[0];
assert_ne!(beta, Fq::ONE);
assert_eq!(
beta * beta * beta,
Fq::ONE,
"beta is not a cube root of one"
);
let beta2 = <snarkrs_field::g2::Config as GLVConfig>::ENDO_COEFFS[0];
assert_eq!(beta2.c0, beta, "G2 takes a different beta");
assert_eq!(beta2.c1, Fq::ZERO, "G2's beta is not in the prime field");
let l1 = <snarkrs_field::g1::Config as GLVConfig>::LAMBDA;
let l2 = <snarkrs_field::g2::Config as GLVConfig>::LAMBDA;
assert_eq!(l1 * l1, l2, "G2's eigenvalue is not G1's squared");
assert_eq!(l1 * l1 * l1, Fr::ONE);
let p1 = G1Affine::generator();
let mut e1 = p1;
e1.x *= beta;
assert_eq!(e1, (p1 * l1).into_affine(), "phi is not [lambda] on G1");
let p2 = G2Affine::generator();
let mut e2 = p2;
e2.x *= beta2;
assert_eq!(e2, (p2 * l2).into_affine(), "phi is not [lambda^2] on G2");
}
#[test]
fn the_chunked_twiddle_walk_is_the_serial_one() {
fn check<C: GLVConfig<ScalarField = Fr>>(name: &str, bits: u32, first: Fr) {
let mut root = Fr::TWO_ADIC_ROOT_OF_UNITY;
for _ in bits..Fr::TWO_ADICITY {
root.square_in_place();
}
let got = twiddle_table::<C>(bits, first);
assert_eq!(
got.len(),
1usize << (bits - 1),
"{name}: the 2^{bits} table is not half the block"
);
let mut w = first;
for (i, slot) in got.iter().enumerate() {
assert_eq!(*slot, decompose::<C>(&w), "{name}: 2^{bits} twiddle {i}");
w *= root;
}
}
for bits in [1u32, 2, 5, 11] {
let inv_n = Fr::from(1u64 << bits)
.inverse()
.expect("a power of two is a unit mod r");
check::<snarkrs_field::g1::Config>("G1", bits, Fr::ONE);
check::<snarkrs_field::g1::Config>("G1", bits, inv_n);
check::<snarkrs_field::g2::Config>("G2", bits, Fr::ONE);
check::<snarkrs_field::g2::Config>("G2", bits, inv_n);
}
}
#[test]
fn a_decomposed_twiddle_still_means_the_same_scalar() {
fn check<C: GLVConfig<ScalarField = Fr>>(name: &str, ks: &[Fr]) {
let mut widest = 0u32;
for k in ks {
let g = decompose::<C>(k);
let mag = |half: usize| {
let mut b = [0u64; 4];
for (i, w) in b[..2].iter_mut().enumerate() {
*w =
u64::from(g.k[half + 2 * i]) | (u64::from(g.k[half + 2 * i + 1]) << 32);
}
Fr::from(ark_ff::BigInt(b))
};
let k1 = mag(0);
let k2 = mag(4);
widest = widest
.max(k1.into_bigint().num_bits())
.max(k2.into_bigint().num_bits());
let s1 = if g.sign & 1 == 0 { k1 } else { -k1 };
let s2 = if g.sign & 2 == 0 { k2 } else { -k2 };
assert_eq!(
s1 + C::LAMBDA * s2,
*k,
"{name}: decomposition of {k} is not it"
);
assert!(
g.sign < 4,
"{name}: sign word has bits the kernels do not read"
);
}
assert!(
widest <= 127,
"{name}: {widest} bits is past the table's 127"
);
}
let mut ks = vec![
Fr::from(0u64),
Fr::ONE,
-Fr::ONE,
-Fr::from(2u64),
<snarkrs_field::g1::Config as GLVConfig>::LAMBDA,
<snarkrs_field::g2::Config as GLVConfig>::LAMBDA,
Fr::TWO_ADIC_ROOT_OF_UNITY,
];
for bits in 0..=28u32 {
ks.push(
Fr::from(1u64 << bits)
.inverse()
.expect("a power of two is a unit mod r"),
);
}
let mut w = Fr::ONE;
for _ in 0..4096 {
ks.push(w);
w *= Fr::TWO_ADIC_ROOT_OF_UNITY;
}
let mut rng = SplitMix64(0x91d_c0de);
for _ in 0..4096 {
ks.push(rng.next_fr());
}
check::<snarkrs_field::g1::Config>("G1", &ks);
check::<snarkrs_field::g2::Config>("G2", &ks);
}
}