igneum/proto-vdf/src/classgroup.rs
2026-10-03 16:27:52 +00:00

470 lines
17 KiB
Rust

//! Imaginary quadratic class group Cl(D), D < 0, D = 1 mod 8, |D| prime.
//!
//! This is the production choice: no trusted setup, the group order is unknown to everyone.
//! Chia Network's chiavdf uses the same construction for its timelords
//! (vendor/chiavdf/src/create_discriminant.h derives D from a seed with HashPrime forcing
//! bits 0, 1, 2 and the top bit, so -D = 7 mod 8; vendor/chiavdf/src/vdf_new.h uses the
//! generator (2, 1, (1-D)/8) and the identity (1, 1, (1-D)/4)).
//!
//! Elements are reduced primitive positive-definite binary quadratic forms (a, b, c) with
//! b^2 - 4ac = D, -a < b <= a <= c, and b >= 0 when a == c. Reduced forms are unique per class,
//! so equality of reduced forms is equality in the group.
//!
//! Composition is Cohen, A Course in Computational Algebraic Number Theory, Algorithm 5.4.7.
//! Squaring uses the dedicated duplication formula also used by chiavdf's `square` in
//! vdf_new.h (valid because gcd(a, b) = 1 for every reduced form of a prime discriminant).
//! Neither NUCOMP nor NUDUPL (Shanks, Atkin) is implemented here; chiavdf's qfb_nudupl in
//! vendor/chiavdf/src/nucomp.h is the optimised production form. Expect this prototype to be
//! several times slower per squaring than chiavdf. Reduction is the textbook loop, not the
//! Pulmark fast reducer chiavdf uses.
use crate::group::Group;
use crate::hash::{bytes_to_int, hash_prime, signed_to_bytes};
use rug::ops::{DivRounding, NegAssign, RemRounding};
use rug::Integer;
use std::cmp::Ordering;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Form {
pub a: Integer,
pub b: Integer,
pub c: Integer,
}
pub struct ClassGroup {
/// Negative discriminant.
pub d: Integer,
pub d_bits: u32,
/// L = floor(|D|^(1/4)), the NUDUPL partial-gcd bound (chiavdf: root(-D, 4)).
pub l: Integer,
}
impl ClassGroup {
/// D = -p where p is the `bits`-bit prime derived from `seed`, p = 7 mod 8.
/// Mirrors chiavdf CreateDiscriminant(seed, length) = -HashPrime(seed, length, {0,1,2,length-1}).
pub fn from_seed(seed: &[u8], bits: u32) -> ClassGroup {
let p = hash_prime(seed, bits, &[0, 1, 2, bits - 1]);
Self::from_discriminant(-p)
}
pub fn from_discriminant(d: Integer) -> ClassGroup {
assert!(d.cmp0() == Ordering::Less);
let d_bits = d.significant_bits();
let l = Integer::from(d.abs_ref()).root(4);
ClassGroup { d, d_bits, l }
}
/// (a, b, c) from a, b with c = (b^2 - D) / 4a, reduced. None if not integral.
pub fn from_ab(&self, a: Integer, b: Integer) -> Option<Form> {
if a.cmp0() != Ordering::Greater {
return None;
}
let num = Integer::from(&b * &b) - &self.d;
let den = Integer::from(&a << 2u32);
if !num.is_divisible(&den) {
return None;
}
let c = num / den;
let mut f = Form { a, b, c };
self.reduce(&mut f);
Some(f)
}
/// The chiavdf generator (2, 1, (1-D)/8). Needs D = 1 mod 8.
pub fn generator(&self) -> Form {
self.from_ab(Integer::from(2), Integer::from(1))
.expect("D must be 1 mod 8 for the generator (2, 1, (1-D)/8)")
}
pub fn inverse(&self, f: &Form) -> Form {
let mut g = Form { a: f.a.clone(), b: Integer::from(-&f.b), c: f.c.clone() };
self.reduce(&mut g);
g
}
pub fn discriminant_of(f: &Form) -> Integer {
Integer::from(&f.b * &f.b) - Integer::from(&f.a * &f.c) * 4u32
}
fn normalize(f: &mut Form) {
// Want -a < b <= a.
let neg_a = Integer::from(-&f.a);
if neg_a < f.b && f.b <= f.a {
return;
}
// r = floor((a - b) / 2a); b' = b + 2ra; c' = c + r(ra + b).
let two_a = Integer::from(&f.a << 1u32);
let (r, _) = Integer::from(&f.a - &f.b).div_rem_floor(two_a);
let ra = Integer::from(&r * &f.a);
f.c += Integer::from(&ra + &f.b) * &r;
f.b += Integer::from(&ra << 1u32);
}
pub fn is_reduced(f: &Form) -> bool {
let neg_a = Integer::from(-&f.a);
if !(neg_a < f.b && f.b <= f.a) {
return false;
}
match f.a.cmp(&f.c) {
Ordering::Greater => false,
Ordering::Equal => f.b.cmp0() != Ordering::Less,
Ordering::Less => true,
}
}
pub fn reduce(&self, f: &mut Form) {
Self::normalize(f);
loop {
let swap = match f.a.cmp(&f.c) {
Ordering::Greater => true,
Ordering::Equal => f.b.cmp0() == Ordering::Less,
Ordering::Less => false,
};
if !swap {
break;
}
// (a, b, c) <- (c, -b, a)
std::mem::swap(&mut f.a, &mut f.c);
f.b.neg_assign();
Self::normalize(f);
}
}
/// Cohen Algorithm 5.4.7.
pub fn compose(&self, f1: &Form, f2: &Form) -> Form {
let (f1, f2) = if f1.a > f2.a { (f2, f1) } else { (f1, f2) };
// Step 1.
let s = Integer::from(&f1.b + &f2.b) >> 1u32;
let n = Integer::from(&f2.b - &s);
// Step 2: u a2 + v a1 = d.
let (y1, d) = if f1.a.is_divisible(&f2.a) {
(Integer::new(), f2.a.clone())
} else {
let (d, u, _v) = f2.a.clone().extended_gcd(f1.a.clone(), Integer::new());
(u, d)
};
// Step 3: u s + v d = d1.
let (x2, y2, d1) = if s.is_divisible(&d) {
(Integer::new(), Integer::from(-1), d)
} else {
let (d1, u, v) = s.clone().extended_gcd(d, Integer::new());
(u, -v, d1)
};
// Step 4.
let v1 = Integer::from(&f1.a / &d1);
let v2 = Integer::from(&f2.a / &d1);
let r = Integer::from(&y1 * &y2) * &n - Integer::from(&x2 * &f2.c);
let r = r.rem_euc(&v1);
let v2r = Integer::from(&v2 * &r);
let b3 = Integer::from(&f2.b + Integer::from(&v2r << 1u32));
let a3 = Integer::from(&v1 * &v2);
let c3 = (Integer::from(&f2.c * &d1) + Integer::from(&f2.b + &v2r) * &r) / &v1;
let mut out = Form { a: a3, b: b3, c: c3 };
self.reduce(&mut out);
out
}
/// Duplication, the formula chiavdf uses in vdf_new.h `square`:
/// with s b + t a = g = gcd(a, b) (g = 1 here), u = (c s) mod a,
/// A = a^2, B = b - 2au, C = u^2 - (bu - c)/a. Falls back to compose if g != 1.
pub fn square_form(&self, f: &mut Form) {
let (g, s, _t) = f.b.clone().extended_gcd(f.a.clone(), Integer::new());
if g != 1 {
let r = self.compose(f, &f.clone());
*f = r;
return;
}
let u = Integer::from(&f.c * &s).rem_euc(&f.a);
let au = Integer::from(&f.a * &u);
let bu_c = Integer::from(&f.b * &u) - &f.c;
let q = bu_c / &f.a; // exact: bu = c mod a
let a2 = Integer::from(&f.a * &f.a);
let b2 = Integer::from(&f.b - Integer::from(&au << 1u32));
let c2 = Integer::from(&u * &u) - q;
f.a = a2;
f.b = b2;
f.c = c2;
self.reduce(f);
}
/// NUDUPL (Shanks; Atkin's variant), ported from chiavdf `qfb_nudupl` in
/// vendor/chiavdf/src/nucomp.h (itself from William Hart's FLINT/Antic qfb code).
/// The partial extended gcd stops once the remainder drops below L = |D|^(1/4), so the
/// intermediate coefficients stay near sqrt(|D|) instead of growing to |D|, and the final
/// reduction is one or two steps. Output is reduced.
pub fn nudupl(&self, f: &mut Form) {
let mut a1 = f.a.clone();
let mut c1 = f.c.clone();
// s = gcd(|b|, a) with v2 |b| = s (mod a); fix the sign so v2 b = s (mod a).
let b_abs = Integer::from(f.b.abs_ref());
let (s, mut v2, _) = b_abs.extended_gcd(a1.clone(), Integer::new());
if f.b.cmp0() == Ordering::Less {
v2.neg_assign();
}
// k = -(c inv(b)) mod a
let mut k = Integer::from(&v2 * &c1);
k.neg_assign();
if s != 1 {
a1 /= &s;
c1 *= &s;
}
let k = k.rem_floor(&a1);
if a1 < self.l {
let t = Integer::from(&a1 * &k);
let new_a = Integer::from(&a1 * &a1);
let cb = Integer::from(&t << 1u32) + &f.b;
let new_c = (Integer::from(&f.b + &t) * &k + &c1).div_floor(&a1);
f.a = new_a;
f.b = cb;
f.c = new_c;
} else {
let mut r2 = a1.clone();
let mut r1 = k;
let (co2, co1) = xgcd_partial(&mut r2, &mut r1, &self.l);
// m2 = (b r1 - c1 co1) / a1
let mut m2 = Integer::from(&f.b * &r1);
m2 -= Integer::from(&c1 * &co1);
let m2 = m2.div_exact(&a1);
// new_a = r1^2 - co1 m2, negated when co1 >= 0
let mut new_a = Integer::from(&r1 * &r1);
new_a -= Integer::from(&co1 * &m2);
if co1.cmp0() != Ordering::Less {
new_a.neg_assign();
}
// cb = 2 (a1 r1 - new_a co2) / co1 - b, then mod 2 new_a (floor remainder)
let mut cb = Integer::from(&a1 * &r1);
cb -= Integer::from(&new_a * &co2);
cb <<= 1u32;
let mut cb = cb.div_exact(&co1);
cb -= &f.b;
let temp = Integer::from(&new_a << 1u32);
let cb = cb.rem_floor(&temp);
// new_c = (cb^2 - D) / (4 new_a)
let mut new_c = Integer::from(&cb * &cb);
new_c -= &self.d;
let mut new_c = new_c.div_exact(&new_a);
new_c >>= 2u32;
if new_a.cmp0() == Ordering::Less {
new_a.neg_assign();
new_c.neg_assign();
}
f.a = new_a;
f.b = cb;
f.c = new_c;
}
self.reduce(f);
}
/// NUCOMP (Shanks; Atkin's variant), ported from chiavdf `qfb_nucomp` in
/// vendor/chiavdf/src/nucomp.h (William Hart, FLINT/Antic). Same idea as NUDUPL for two
/// different forms. Output is reduced.
pub fn nucomp(&self, f: &Form, g: &Form) -> Form {
let (f, g) = if f.a > g.a { (g, f) } else { (f, g) };
let mut a1 = f.a.clone();
let mut a2 = g.a.clone();
let mut c2 = g.c.clone();
let ss = Integer::from(&f.b + &g.b) >> 1u32;
let m = Integer::from(&f.b - &g.b) >> 1u32;
let t = Integer::from(&a2).rem_floor(&a1);
let (sp, v1) = if t.cmp0() == Ordering::Equal {
(a1.clone(), Integer::new())
} else {
let (sp, v1, _) = t.extended_gcd(a1.clone(), Integer::new());
(sp, v1)
};
let mut k = Integer::from(&m * &v1).rem_floor(&a1);
if sp != 1 {
// s = v2 ss + u2 sp
let (s, v2, u2) = ss.clone().extended_gcd(sp.clone(), Integer::new());
k *= &u2;
k -= Integer::from(&v2 * &c2);
if s != 1 {
a1 = a1.div_exact(&s);
a2 = a2.div_exact(&s);
c2 *= &s;
}
k = k.rem_floor(&a1);
}
let mut out = if a1 < self.l {
let t = Integer::from(&a2 * &k);
let ca = Integer::from(&a2 * &a1);
let cb = Integer::from(&t << 1u32) + &g.b;
let cc = (Integer::from(&g.b + &t) * &k + &c2).div_exact(&a1);
Form { a: ca, b: cb, c: cc }
} else {
let mut r2 = a1.clone();
let mut r1 = k;
let (co2, co1) = xgcd_partial(&mut r2, &mut r1, &self.l);
let t = Integer::from(&a2 * &r1);
let m1 = (Integer::from(&m * &co1) + &t).div_exact(&a1);
let m2 = (Integer::from(&ss * &r1) - Integer::from(&c2 * &co1)).div_exact(&a1);
let r1m1 = Integer::from(&r1 * &m1);
let co1m2 = Integer::from(&co1 * &m2);
let mut ca = if co1.cmp0() == Ordering::Less { r1m1 - co1m2 } else { co1m2 - r1m1 };
let mut cb = Integer::from(&t - Integer::from(&ca * &co2));
cb <<= 1u32;
let mut cb = cb.div_exact(&co1);
cb -= &g.b;
let temp = Integer::from(&ca << 1u32);
let cb = cb.rem_floor(&temp);
let mut cc = Integer::from(&cb * &cb);
cc -= &self.d;
let mut cc = cc.div_exact(&ca);
cc >>= 2u32;
if ca.cmp0() == Ordering::Less {
ca.neg_assign();
cc.neg_assign();
}
Form { a: ca, b: cb, c: cc }
};
self.reduce(&mut out);
out
}
fn coeff_width(&self) -> usize {
// Reduced forms have a, |b| <= sqrt(|D|/3) < 2^(d_bits/2 + 1). Keep a full d_bits
// width so unreduced but valid forms still round-trip during tests.
((self.d_bits + 7) / 8) as usize
}
}
/// Partial extended Euclid, plain-division form. Reference for the test suite only.
/// On return co2 r1 - co1 r2 = +-(original r2), and r1 <= L.
pub fn xgcd_partial_simple(r2: &mut Integer, r1: &mut Integer, l: &Integer) -> (Integer, Integer) {
let mut co2 = Integer::new();
let mut co1 = Integer::from(-1);
while r1.cmp0() != Ordering::Equal && *r1 > *l {
let (q, r) = r2.clone().div_rem_floor(r1.clone());
*r2 = std::mem::replace(r1, r);
co2 -= Integer::from(&co1 * &q);
std::mem::swap(&mut co2, &mut co1);
}
(co2, co1)
}
/// Top 63 bits of a non-negative integer, from bit `shift` upward.
fn top_word(x: &Integer, shift: u32) -> i64 {
if shift == 0 {
return x.to_u64_wrapping() as i64;
}
Integer::from(x >> shift).to_u64_wrapping() as i64
}
/// Partial extended Euclid with Lehmer acceleration, ported from chiavdf's
/// `mpz_xgcd_partial` (vendor/chiavdf/src/xgcd_partial.c, William Hart, FLINT). Each outer
/// round runs as many Euclid steps as the top 63 bits of r2, r1 allow in machine words
/// (Collins/Jebelean exit tests), then applies the 2x2 cofactor matrix to the big numbers.
/// Produces the same (r2, r1, co2, co1) as `xgcd_partial_simple`.
pub fn xgcd_partial(r2: &mut Integer, r1: &mut Integer, l: &Integer) -> (Integer, Integer) {
let mut co2 = Integer::new();
let mut co1 = Integer::from(-1);
while r1.cmp0() != Ordering::Equal && *r1 > *l {
let bits2 = r2.significant_bits() as i64;
let bits1 = r1.significant_bits() as i64;
let bits = (bits2.max(bits1) - 64 + 1).max(0) as u32;
let mut rr2 = top_word(r2, bits);
let mut rr1 = top_word(r1, bits);
let bb = top_word(l, bits);
let (mut aa2, mut aa1, mut bb2, mut bb1): (i64, i64, i64, i64) = (0, 1, 1, 0);
let mut i: u32 = 0;
while rr1 != 0 && rr1 > bb {
let qq = rr2 / rr1;
let t1 = rr2 - qq * rr1;
let t2 = aa2 - qq * aa1;
let t3 = bb2 - qq * bb1;
if i & 1 == 1 {
if t1 < -t3 || rr1 - t1 < t2 - aa1 {
break;
}
} else if t1 < -t2 || rr1 - t1 < t3 - bb1 {
break;
}
rr2 = rr1;
rr1 = t1;
aa2 = aa1;
aa1 = t2;
bb2 = bb1;
bb1 = t3;
i += 1;
}
if i == 0 {
let (q, r) = r2.clone().div_rem_floor(r1.clone());
*r2 = std::mem::replace(r1, r);
co2 -= Integer::from(&co1 * &q);
std::mem::swap(&mut co2, &mut co1);
} else {
let r = Integer::from(&*r2 * bb2) + Integer::from(&*r1 * aa2);
let new_r1 = Integer::from(&*r1 * aa1) + Integer::from(&*r2 * bb1);
*r2 = r;
*r1 = new_r1;
let r = Integer::from(&co2 * bb2) + Integer::from(&co1 * aa2);
let new_co1 = Integer::from(&co1 * aa1) + Integer::from(&co2 * bb1);
co2 = r;
co1 = new_co1;
if r1.cmp0() == Ordering::Less {
co1.neg_assign();
r1.neg_assign();
}
if r2.cmp0() == Ordering::Less {
co2.neg_assign();
r2.neg_assign();
}
}
}
(co2, co1)
}
impl Group for ClassGroup {
type Elem = Form;
fn name(&self) -> String {
format!("class group, {}-bit prime discriminant", self.d_bits)
}
fn identity(&self) -> Form {
self.from_ab(Integer::from(1), Integer::from(1)).expect("D = 1 mod 4")
}
fn square(&self, x: &mut Form) {
self.nudupl(x);
}
fn mul(&self, a: &Form, b: &Form) -> Form {
self.nucomp(a, b)
}
/// a and b, each as sign byte + fixed-width magnitude. c is recomputed.
fn serialize(&self, x: &Form) -> Vec<u8> {
let w = self.coeff_width();
let mut out = signed_to_bytes(&x.a, w);
out.extend_from_slice(&signed_to_bytes(&x.b, w));
out
}
fn deserialize(&self, bytes: &[u8]) -> Option<Form> {
let w = self.coeff_width();
if bytes.len() != 2 * (w + 1) {
return None;
}
let read = |s: &[u8]| -> Integer {
let mag = bytes_to_int(&s[1..]);
if s[0] == 1 {
-mag
} else {
mag
}
};
let a = read(&bytes[..w + 1]);
let b = read(&bytes[w + 1..]);
let f = self.from_ab(a, b)?;
if self.is_valid(&f) {
Some(f)
} else {
None
}
}
fn is_valid(&self, x: &Form) -> bool {
x.a.cmp0() == Ordering::Greater && Self::is_reduced(x) && Self::discriminant_of(x) == self.d
}
}