use ruint::Uint;
use super::{CurvePoint, SECP256K1_GENERATOR_X, SECP256K1_GENERATOR_Y};
use crate::math::{
k1_scalar::K1Scalar,
uint::{Limbs, UintDomain},
};
pub const SECP256K1_LAMBDA: Limbs = [
0x1b23bd72, 0xdf02967c, 0x20816678, 0x122e22ea, 0x8812645a, 0xa5261c02, 0xc05c30e0, 0x5363ad4c,
];
pub const SECP256K1_BETA: Limbs = [
0x719501ee, 0xc1396c28, 0x12f58995, 0x9cf04975, 0xac3434e9, 0x6e64479e, 0x657c0710, 0x7ae96a2b,
];
pub fn phi_generator() -> CurvePoint {
CurvePoint::Affine {
x: UintDomain::K1Base.mul(SECP256K1_BETA, SECP256K1_GENERATOR_X),
y: SECP256K1_GENERATOR_Y,
}
}
type Wide = Uint<512, 8>;
fn wide_zero() -> Wide {
Wide::from_limbs([0; 8])
}
fn wide_one() -> Wide {
let mut limbs = [0u64; 8];
limbs[0] = 1;
Wide::from_limbs(limbs)
}
fn limbs_to_wide(limbs: Limbs) -> Wide {
let mut u64_limbs = [0u64; 8];
for i in 0..4 {
u64_limbs[i] = (limbs[2 * i] as u64) | ((limbs[2 * i + 1] as u64) << 32);
}
Wide::from_limbs(u64_limbs)
}
fn wide_to_limbs(v: Wide) -> Limbs {
let u64_limbs = v.as_limbs();
assert!(u64_limbs[4..].iter().all(|&l| l == 0), "GLV magnitude must fit in 256 bits");
core::array::from_fn(|i| {
let word = u64_limbs[i / 2];
if i % 2 == 0 { word as u32 } else { (word >> 32) as u32 }
})
}
#[derive(Clone, Copy)]
struct Signed {
neg: bool,
mag: Wide,
}
impl Signed {
fn new(neg: bool, mag: Wide) -> Self {
if mag == wide_zero() {
Signed { neg: false, mag }
} else {
Signed { neg, mag }
}
}
fn negate(self) -> Self {
Signed::new(!self.neg, self.mag)
}
fn add(self, other: Self) -> Self {
if self.neg == other.neg {
Signed::new(self.neg, self.mag + other.mag)
} else if self.mag >= other.mag {
Signed::new(self.neg, self.mag - other.mag)
} else {
Signed::new(other.neg, other.mag - self.mag)
}
}
fn sub(self, other: Self) -> Self {
self.add(other.negate())
}
fn mul(self, other: Self) -> Self {
Signed::new(self.neg != other.neg, self.mag * other.mag)
}
fn div_round(self, n: Wide) -> Self {
let q = self.mag / n;
let r = self.mag % n;
let q = if r + r >= n { q + wide_one() } else { q };
Signed::new(self.neg, q)
}
}
const GLV_BASIS: [(bool, Limbs); 4] = [
(false, [0x9284eb15, 0xe86c90e4, 0xa7d46bcd, 0x3086d221, 0, 0, 0, 0]),
(true, [0x0abfe4c3, 0x6f547fa9, 0x010e8828, 0xe4437ed6, 0, 0, 0, 0]),
(false, [0x9d44cfd8, 0x57c1108d, 0xa8e2f3f6, 0x14ca50f7, 0x00000001, 0, 0, 0]),
(false, [0x9284eb15, 0xe86c90e4, 0xa7d46bcd, 0x3086d221, 0, 0, 0, 0]),
];
pub fn glv_decompose(k: Limbs) -> [(bool, Limbs); 2] {
let n = limbs_to_wide(K1Scalar::MODULUS);
let [(a1_neg, a1_mag), (b1_neg, b1_mag), (a2_neg, a2_mag), (b2_neg, b2_mag)] = GLV_BASIS;
let a1 = Signed::new(a1_neg, limbs_to_wide(a1_mag));
let b1 = Signed::new(b1_neg, limbs_to_wide(b1_mag));
let a2 = Signed::new(a2_neg, limbs_to_wide(a2_mag));
let b2 = Signed::new(b2_neg, limbs_to_wide(b2_mag));
let k_s = Signed::new(false, limbs_to_wide(k));
let c1 = b2.mul(k_s).div_round(n);
let c2 = b1.negate().mul(k_s).div_round(n);
let k1 = k_s.sub(c1.mul(a1)).sub(c2.mul(a2));
let k2 = c1.negate().mul(b1).sub(c2.mul(b2));
[(k1.neg, wide_to_limbs(k1.mag)), (k2.neg, wide_to_limbs(k2.mag))]
}
pub fn scalar_mul_mod_n(a: Limbs, b: Limbs) -> Limbs {
let n = limbs_to_wide(K1Scalar::MODULUS);
wide_to_limbs((limbs_to_wide(a) * limbs_to_wide(b)) % n)
}
#[cfg(test)]
mod tests {
use super::*;
fn limbs_from_u64(v: u64) -> Limbs {
[v as u32, (v >> 32) as u32, 0, 0, 0, 0, 0, 0]
}
fn recompose(split: [(bool, Limbs); 2]) -> Wide {
let n = limbs_to_wide(K1Scalar::MODULUS);
let lambda = limbs_to_wide(SECP256K1_LAMBDA);
let to_signed = |(neg, mag): (bool, Limbs)| Signed::new(neg, limbs_to_wide(mag));
let a = to_signed(split[0]);
let b = to_signed(split[1]);
let term = Signed::new(false, lambda).mul(b);
let sum = a.add(term);
let mag_mod_n = sum.mag % n;
if sum.neg && mag_mod_n != wide_zero() {
n - mag_mod_n
} else {
mag_mod_n
}
}
#[test]
fn glv_basis_matches_extended_euclid_reduction() {
let n = limbs_to_wide(K1Scalar::MODULUS);
let lambda = limbs_to_wide(SECP256K1_LAMBDA);
let below_sqrt_n = |r: Wide| r * r < n;
let (mut r0, mut r1) = (n, lambda);
let (mut t0, mut t1) = (Signed::new(false, wide_zero()), Signed::new(false, wide_one()));
while !below_sqrt_n(r1) {
let q = r0 / r1;
let r2 = r0 - q * r1;
let t2 = t0.sub(Signed::new(false, q).mul(t1));
(r0, r1, t0, t1) = (r1, r2, t1, t2);
}
let (a1, b1) = (Signed::new(false, r1), t1.negate());
let q = r0 / r1;
let r2 = r0 - q * r1;
let t2 = t0.sub(Signed::new(false, q).mul(t1));
let norm = |r: Wide, t: Wide| r * r + t * t;
let (a2, b2) = if norm(r0, t0.mag) <= norm(r2, t2.mag) {
(Signed::new(false, r0), t0.negate())
} else {
(Signed::new(false, r2), t2.negate())
};
let derived = [a1, b1, a2, b2].map(|s| (s.neg, wide_to_limbs(s.mag)));
assert_eq!(
derived, GLV_BASIS,
"GLV_BASIS is stale relative to the extended-Euclid reduction of (n, lambda)"
);
}
#[test]
fn glv_decompose_recomposes_small_scalars() {
for k in [0u64, 1, 2, 12345, u64::MAX] {
let split = glv_decompose(limbs_from_u64(k));
assert_eq!(recompose(split), limbs_to_wide(limbs_from_u64(k)), "failed for k={k}");
}
}
#[test]
fn glv_decompose_recomposes_full_width_scalar() {
let k: Limbs = [
0x12345678, 0x9abcdef0, 0x0fedcba9, 0x87654321, 0x11223344, 0x55667788, 0x99aabbcc,
0x00112233,
];
let split = glv_decompose(k);
assert_eq!(recompose(split), limbs_to_wide(k));
}
#[test]
fn glv_decompose_halves_are_short() {
let k: Limbs = [
0x12345678, 0x9abcdef0, 0x0fedcba9, 0x87654321, 0x11223344, 0x55667788, 0x99aabbcc,
0x00112233,
];
let mut bound_limbs = [0u32; 8];
bound_limbs[4] = 0x10;
let bound = limbs_to_wide(bound_limbs);
for (_, mag) in glv_decompose(k) {
assert!(limbs_to_wide(mag) < bound, "GLV half is not short: {mag:?}");
}
}
}