Skip to main content

kyn_vdf/
math.rs

1//! Imaginary Quadratic Class Group arithmetic with Shanks' NUCOMP and NUDUPL algorithms.
2//!
3//! This module implements binary quadratic forms $(a, b, c)$ over a negative fundamental
4//! discriminant $D = b^2 - 4ac < 0$. The group law is given by Shanks' sub-quadratic
5//! composition algorithms NUCOMP and NUDUPL, which use a partial Euclidean reduction
6//! step controlled by a threshold $L = \lfloor |D|^{1/4} \rfloor$ to bound intermediate
7//! coefficient sizes.
8//!
9//! ## References
10//! - Daniel Shanks (1989): *On Gauss and Composition I, II*
11//! - Henri Cohen (1993): *A Course in Computational Algebraic Number Theory*, GTM 138, §5.4
12//! - Lipa Long: chiavdf C++ reference implementation
13
14use num_bigint::BigInt;
15use num_traits::{Zero, One, Signed};
16
17/// Computes $\lfloor \sqrt{\sqrt{n}} \rfloor = \lfloor |D|^{1/4} \rfloor$.
18///
19/// Used as the Shanks NUCOMP/NUDUPL threshold $L$ to bound intermediate form
20/// coefficients during partial Euclidean reduction steps.
21///
22/// # Panics
23/// Does not panic. Returns `BigInt::zero()` for `n = 0`.
24#[inline]
25pub fn isqrt_fourth(n: &BigInt) -> BigInt {
26    let s1 = n.sqrt();
27    s1.sqrt()
28}
29
30/// Partial Extended Euclidean Algorithm (partial XGCD).
31///
32/// Runs the standard Extended GCD algorithm on `(r2, r1)` but **stops early**
33/// once the remainder `r1` drops at or below the threshold `l`.
34///
35/// This is the core primitive used by both NUCOMP and NUDUPL to achieve
36/// sub-quadratic composition via partial Euclidean reduction.
37///
38/// # Parameters
39/// - `r2`: The larger initial value (typically the form's `a` coefficient).
40/// - `r1`: The smaller initial value (typically the intermediate `k`).
41/// - `l`: The stopping threshold $L = \lfloor |D|^{1/4} \rfloor$.
42///
43/// # Returns
44/// A tuple `(co2, co1, r2, r1)` where:
45/// - `co2`, `co1` are the Bézout coefficients at the stopping point:
46///   `co2·(original r2) + co1·(original r1) = r2` (at stop).
47/// - `r2` is the last remainder *before* the threshold was crossed.
48/// - `r1` is the last remainder *at or below* the threshold (may be zero if
49///   the GCD terminated early).
50pub fn xgcd_partial(r2: &BigInt, r1: &BigInt, l: &BigInt) -> (BigInt, BigInt, BigInt, BigInt) {
51    let mut r2_cur = r2.clone();
52    let mut r1_cur = r1.clone();
53    let mut co2 = BigInt::zero();
54    let mut co1 = BigInt::from(-1);
55
56    while !r1_cur.is_zero() && &r1_cur > l {
57        let q = &r2_cur / &r1_cur;
58        let t1 = &r2_cur - &q * &r1_cur;
59        let t2 = &co2 - &q * &co1;
60        r2_cur = r1_cur;
61        r1_cur = t1;
62        co2 = co1;
63        co1 = t2;
64    }
65    (co2, co1, r2_cur, r1_cur)
66}
67
68/// A binary quadratic form $f = ax^2 + bxy + cy^2$ over an imaginary quadratic field.
69///
70/// Forms are elements of the class group $\text{Cl}(D)$ of a negative fundamental
71/// discriminant $D = b^2 - 4ac < 0$. The group law is composition, and the identity
72/// element is the principal form $(1, 1, (1-D)/4)$.
73///
74/// All arithmetic operations (composition, squaring, exponentiation) produce
75/// **reduced** forms, meaning:
76/// - $-a < b \le a$
77/// - $a \le c$
78/// - if $a = c$ then $b \ge 0$
79#[derive(Clone, Debug, PartialEq, Eq, Hash)]
80pub struct Form {
81    /// The leading coefficient $a > 0$.
82    pub a: BigInt,
83    /// The cross-term coefficient $b$ (satisfies $-a < b \le a$ when reduced).
84    pub b: BigInt,
85    /// The constant coefficient $c > 0$ (satisfies $c \ge a$ when reduced).
86    pub c: BigInt,
87}
88
89impl Form {
90    /// Creates a new form $(a, b, c)$ without any validation or reduction.
91    ///
92    /// Prefer [`Form::from_abd`] or [`Form::generator`] when constructing from
93    /// external data, as they validate the discriminant identity.
94    pub fn new(a: BigInt, b: BigInt, c: BigInt) -> Self {
95        Self { a, b, c }
96    }
97
98    /// Returns the principal identity element $(1, 1, (1 - D) / 4)$ for discriminant $D$.
99    ///
100    /// The identity satisfies $e \circ f = f$ for any form $f$ in the same class group.
101    /// Requires $D \equiv 3 \pmod 4$ (standard for imaginary quadratic fields with
102    /// fundamental discriminant).
103    pub fn identity(d: &BigInt) -> Self {
104        let one = BigInt::one();
105        let four = BigInt::from(4);
106        let c = (&one - d) / four;
107        Self {
108            a: one.clone(),
109            b: one,
110            c,
111        }
112    }
113
114    /// Constructs a form $(a, b, c)$ from $a$, $b$, and discriminant $D$,
115    /// computing $c = (b^2 - D) / (4a)$ and verifying exact divisibility.
116    ///
117    /// Returns `None` if:
118    /// - `a` is zero (degenerate form)
119    /// - $(b^2 - D)$ is not exactly divisible by $4a$
120    pub fn from_abd(a: &BigInt, b: &BigInt, d: &BigInt) -> Option<Self> {
121        if a.is_zero() {
122            return None;
123        }
124        let num = b * b - d;
125        let den = a * BigInt::from(4);
126        if &num % &den != BigInt::zero() {
127            return None;
128        }
129        let c = num / den;
130        Some(Self {
131            a: a.clone(),
132            b: b.clone(),
133            c,
134        })
135    }
136
137    /// Returns the canonical generator form $(2, 1, (1 - D) / 8)$.
138    ///
139    /// Valid for prime discriminants $D = -p$ where $p \equiv 7 \pmod 8$.
140    /// In this case $2$ splits in $\mathbb{Q}(\sqrt{D})$ and the form $(2, 1, \cdot)$
141    /// generates a subgroup of the class group of order $\text{ord}(2)$ in $\text{Cl}(D)$.
142    ///
143    /// Returns `None` if $(1 - D)$ is not divisible by 8, i.e. $D$ is not a valid
144    /// prime discriminant of the required form.
145    pub fn generator(d: &BigInt) -> Option<Self> {
146        let num = BigInt::one() - d;
147        let eight = BigInt::from(8);
148        if &num % &eight != BigInt::zero() {
149            return None;
150        }
151        let c = num / eight;
152        Some(Self {
153            a: BigInt::from(2),
154            b: BigInt::one(),
155            c,
156        })
157    }
158
159    /// Returns `true` if the form is in reduced normal form.
160    ///
161    /// A form $(a, b, c)$ is **reduced** if and only if:
162    /// - $a > 0$ and $c > 0$
163    /// - $-a < b \le a$ (normalization condition)
164    /// - $a \le c$ (minimality condition)
165    /// - if $a = c$, then $b \ge 0$ (uniqueness condition)
166    pub fn is_reduced(&self) -> bool {
167        if self.a <= BigInt::zero() || self.c <= BigInt::zero() {
168            return false;
169        }
170        if self.b <= -(&self.a) || self.b > self.a {
171            return false;
172        }
173        if self.a > self.c {
174            return false;
175        }
176        if self.a == self.c && self.b < BigInt::zero() {
177            return false;
178        }
179        true
180    }
181
182    /// Reduces this form in place using the Euclidean-style Gauss reduction algorithm.
183    ///
184    /// After reduction the form satisfies the standard reduced form conditions
185    /// (see [`Form::is_reduced`]). Every form class has a unique reduced representative.
186    ///
187    /// # Algorithm
188    /// Iterates two steps until the form is reduced:
189    /// 1. **Normalize** $b$: compute $s = \lfloor (a - b) / 2a \rfloor$, then
190    ///    $b' = b + 2as$, $c' = (b'^2 - D) / 4a$.
191    /// 2. **Minimize** $a$: if $a > c'$, swap and negate — set $(a, b, c) = (c', -b', a)$ — and loop.
192    ///
193    /// Terminates because the $a$-coefficient strictly decreases each swap.
194    pub fn reduce(&mut self, d: &BigInt) {
195        use num_integer::Integer;
196        let two = BigInt::from(2);
197        let four = BigInt::from(4);
198
199        loop {
200            // Step 1: normalize b into (-a, a]
201            let a2 = &self.a * &two;
202            let s = (&self.a - &self.b).div_floor(&a2);
203            let b_new = &self.b + &a2 * &s;
204            let c_new = (&b_new * &b_new - d) / (&self.a * &four);
205            let a_new = self.a.clone();
206
207            // Step 2: if a > c after normalization, swap (a, c) and negate b, then loop
208            if a_new > c_new {
209                self.a = c_new;
210                self.c = a_new;
211                self.b = -b_new;
212                continue;
213            }
214
215            // Step 3: uniqueness — if a == c ensure b ≥ 0
216            if a_new == c_new && b_new.is_negative() {
217                self.b = -b_new;
218            } else {
219                self.b = b_new;
220            }
221            self.a = a_new;
222            self.c = c_new;
223            break;
224        }
225    }
226
227    /// Shanks' **NUDUPL** algorithm — fast squaring of a binary quadratic form.
228    ///
229    /// Computes $f^2 = f \circ f$ in the class group without full reduction of
230    /// intermediate results, using a partial Euclidean step controlled by threshold `l`.
231    ///
232    /// # Algorithm (Shanks, 1989 — Cohen §5.4.2)
233    /// Given form $(a_1, b_1, c_1)$ and threshold $L$:
234    /// 1. Compute $s = \gcd(b_1, a_1)$ via extended GCD and co-factor $k = -xc_1$.
235    /// 2. If $s > 1$: divide $a_1$ by $s$, multiply $c_1$ by $s$.
236    /// 3. Reduce $k \pmod{a_1}$.
237    /// 4. If $a_1 < L$: use direct multiplication (no partial reduction).
238    /// 5. Otherwise: run partial XGCD on $(a_1, k)$ stopping at $L$, then compute the
239    ///    new $(a, b, c)$ from the Bézout coefficients `(co2, co1)` and remainders `(r2, r1)`.
240    ///
241    /// # Parameters
242    /// - `d`: The negative fundamental discriminant.
243    /// - `l`: The Shanks threshold $L = \lfloor |D|^{1/4} \rfloor$.
244    ///
245    /// # Returns
246    /// A (possibly unreduced) form. Callers must call [`Form::reduce`] to obtain the
247    /// canonical representative. Use [`Form::square`] for the combined squaring+reduction.
248    pub fn nudupl(&self, d: &BigInt, l: &BigInt) -> Form {
249        use num_integer::Integer;
250        let two = BigInt::from(2);
251        let four = BigInt::from(4);
252
253        let mut a1 = self.a.clone();
254        let mut c1 = self.c.clone();
255
256        // Extended GCD of |b| and a to find the co-factor s = gcd(b, a)
257        // and Bézout coefficient x such that x·|b| + y·a = s
258        let ext = if self.b.is_negative() {
259            let b_abs = -(&self.b);
260            let e = b_abs.extended_gcd(&a1);
261            (-e.x, e.gcd)
262        } else {
263            let e = self.b.extended_gcd(&a1);
264            (e.x, e.gcd)
265        };
266
267        // k = -x·c1 (initial value before partial reduction)
268        let mut k = -(&ext.0 * &c1);
269        let s = ext.1; // s = gcd(b, a)
270
271        // If s > 1, divide out the common factor
272        if s != BigInt::one() {
273            a1 /= &s;
274            c1 *= &s;
275        }
276
277        // Normalize k into [0, a1)
278        k = k.mod_floor(&a1);
279
280        if a1 < *l {
281            // Small a1: direct composition (no partial XGCD needed)
282            let t = &a1 * &k;
283            let res_a = &a1 * &a1;
284            let cb = &two * &t + &self.b;
285            let res_c = ((&self.b + &t) * &k + &c1) / &a1;
286            Form::new(res_a, cb, res_c)
287        } else {
288            // Large a1: partial XGCD stops when remainder ≤ L
289            // Returns Bézout coefficients (co2, co1) and remainders (r2, r1) at stop
290            let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, l);
291
292            // Compute auxiliary value m2 = (b·r1 - c1·co1) / a1
293            let m2 = (&self.b * &r1 - &c1 * &co1) / &a1;
294
295            // New a coefficient: r1² - co1·m2 (sign adjusted to ensure positivity)
296            let mut res_a = &r1 * &r1 - &co1 * &m2;
297            if !co1.is_negative() {
298                res_a = -res_a;
299            }
300
301            // Recover b from the partial XGCD Bézout relation
302            let cb_num = &two * (&a1 * &r1 - &res_a * &co2);
303            let cb = (cb_num / &co1 - &self.b).mod_floor(&(&res_a * &two));
304
305            // Compute c from the discriminant identity: c = (b² - D) / 4a
306            let mut res_c = (&cb * &cb - d) / (&res_a * &four);
307
308            // Ensure a > 0 (normalize sign)
309            if res_a.is_negative() {
310                res_a = -res_a;
311                res_c = -res_c;
312            }
313
314            Form::new(res_a, cb, res_c)
315        }
316    }
317
318    /// Shanks' **NUCOMP** algorithm — fast composition of two binary quadratic forms.
319    ///
320    /// Computes $f_1 \circ f_2$ in the class group using a partial Euclidean reduction
321    /// step, keeping intermediate coefficients bounded by threshold `l`.
322    ///
323    /// # Algorithm (Shanks, 1989 — Cohen §5.4.1)
324    /// Given forms $f_1 = (a_1, b_1, c_1)$ and $f_2 = (a_2, b_2, c_2)$ with $a_1 \le a_2$:
325    /// 1. Compute $ss = (b_1 + b_2)/2$, $m = (b_1 - b_2)/2$.
326    /// 2. Compute $sp = \gcd(a_2 \bmod a_1, a_1)$ via extended GCD.
327    /// 3. If $sp = 1$: set $k = m \cdot v_1 \bmod a_1$.
328    /// 4. If $sp > 1$: reduce through a second extended GCD of $ss$ and $sp$,
329    ///    divide $a_1, a_2$ by $s = \gcd(ss, sp)$ and scale $c_2$.
330    /// 5. Apply partial XGCD or direct multiplication depending on $a_1$ vs $L$.
331    ///
332    /// # Parameters
333    /// - `other`: The second form $f_2$ to compose with.
334    /// - `d`: The negative fundamental discriminant.
335    /// - `l`: The Shanks threshold $L = \lfloor |D|^{1/4} \rfloor$.
336    ///
337    /// # Returns
338    /// A (possibly unreduced) form. Use [`Form::compose`] for the combined compose+reduction.
339    pub fn nucomp(&self, other: &Form, d: &BigInt, l: &BigInt) -> Form {
340        use num_integer::Integer;
341        // Enforce a1 ≤ a2 for the algorithm's precondition
342        if self.a > other.a {
343            return other.nucomp(self, d, l);
344        }
345
346        let two = BigInt::from(2);
347        let four = BigInt::from(4);
348
349        let mut a1 = self.a.clone();
350        let mut a2 = other.a.clone();
351        let mut c2 = other.c.clone();
352
353        // ss = (b1 + b2) / 2,  m = (b1 - b2) / 2
354        let ss = (&self.b + &other.b) / &two;
355        let m = (&self.b - &other.b) / &two;
356
357        // Compute sp = gcd(a2 mod a1, a1) and Bézout coefficient v1
358        let t = a2.mod_floor(&a1);
359        let (v1, sp) = if t.is_zero() {
360            (BigInt::zero(), a1.clone())
361        } else {
362            let e = t.extended_gcd(&a1);
363            (e.x, e.gcd)
364        };
365
366        // Initial k = m·v1 mod a1
367        let mut k = (&m * &v1).mod_floor(&a1);
368
369        if sp != BigInt::one() {
370            // sp > 1: second GCD step to remove common factor from (ss, sp)
371            let e2 = ss.extended_gcd(&sp);
372            let v2 = e2.x; // Bézout: v2·ss + u2·sp = s
373            let u2 = e2.y;
374            let s = e2.gcd;
375            // Merge: k = k·u2 - v2·c2
376            k = &k * &u2 - &v2 * &c2;
377            if s != BigInt::one() {
378                // Divide out common factor s from both a's and scale c2
379                a1 /= &s;
380                a2 /= &s;
381                c2 *= &s;
382            }
383            k = k.mod_floor(&a1);
384        }
385
386        if a1 < *l {
387            // Small a1: direct multiplication (no partial XGCD needed)
388            let t_val = &a2 * &k;
389            let ca = &a2 * &a1;
390            let cb = &two * &t_val + &other.b;
391            let cc = ((&other.b + &t_val) * &k + &c2) / &a1;
392            Form::new(ca, cb, cc)
393        } else {
394            // Large a1: partial XGCD on (a1, k) stopping at L
395            let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, l);
396
397            // Auxiliary values from Bézout coefficients
398            let m1 = (&m * &co1 + &a2 * &r1) / &a1;
399            let m2 = (&ss * &r1 - &c2 * &co1) / &a1;
400
401            // New a coefficient: r1·m1 - co1·m2 (sign adjusted for positivity)
402            let mut ca = &r1 * &m1 - &co1 * &m2;
403            if !co1.is_negative() {
404                ca = -ca;
405            }
406
407            // Recover b from the Bézout relation
408            let t_val = &a2 * &r1;
409            let cb_num = &two * (&t_val - &ca * &co2);
410            let cb = (cb_num / &co1 - &other.b).mod_floor(&(&ca * &two));
411
412            // Compute c from the discriminant identity: c = (b² - D) / 4a
413            let mut cc = (&cb * &cb - d) / (&ca * &four);
414
415            // Ensure a > 0
416            if ca.is_negative() {
417                ca = -ca;
418                cc = -cc;
419            }
420
421            Form::new(ca, cb, cc)
422        }
423    }
424
425    /// Fast binary exponentiation using NUDUPL and NUCOMP with threshold-based reduction.
426    ///
427    /// Computes $f^n$ using the standard left-to-right binary method:
428    /// - Start with $\text{result} = f$ (from the leading bit of $n$).
429    /// - For each remaining bit $i$ from high to low:
430    ///   - Square: $\text{result} \leftarrow \text{NUDUPL}(\text{result})$
431    ///   - If bit $i$ is set: compose $\text{result} \leftarrow \text{NUCOMP}(\text{result}, f)$
432    ///   - Opportunistically reduce when `result.a.bits() > |D|.bits() / 2`
433    ///     to prevent coefficient blowup between full reductions.
434    /// - Final full Gauss reduction before return.
435    ///
436    /// Returns the identity form for `exp = 0`.
437    ///
438    /// # Parameters
439    /// - `exp`: Non-negative exponent as `BigUint`.
440    /// - `d`: The negative fundamental discriminant.
441    /// - `l`: The Shanks threshold $L = \lfloor |D|^{1/4} \rfloor$.
442    pub fn fast_pow(&self, exp: &num_bigint::BigUint, d: &BigInt, l: &BigInt) -> Form {
443        use num_traits::Zero;
444        if exp.is_zero() {
445            return Self::identity(d);
446        }
447
448        let mut res = self.clone();
449        // Opportunistic reduction threshold: reduce when a exceeds half the discriminant bit-width
450        let max_bits = d.abs().bits() / 2;
451        let num_bits = exp.bits();
452
453        if num_bits > 1 {
454            for i in (0..num_bits - 1).rev() {
455                res = res.nudupl(d, l);
456                // Bound growth: reduce if a coefficient is getting too large
457                if res.a.bits() > max_bits {
458                    res.reduce(d);
459                }
460
461                if exp.bit(i) {
462                    res = res.nucomp(self, d, l);
463                }
464            }
465        }
466
467        // Final canonical reduction
468        res.reduce(d);
469        res
470    }
471
472    /// Raises this form to the power of `exp` using [`Form::fast_pow`].
473    ///
474    /// Computes the Shanks threshold $L = \lfloor |D|^{1/4} \rfloor$ internally.
475    /// Returns the identity form when `exp = 0`.
476    pub fn pow(&self, exp: &num_bigint::BigUint, d: &BigInt) -> Self {
477        let l = isqrt_fourth(&d.abs());
478        self.fast_pow(exp, d, &l)
479    }
480
481    /// Composes this form with `other` and returns the reduced result.
482    ///
483    /// This is the primary group operation $f_1 \circ f_2$ of the class group $\text{Cl}(D)$.
484    /// Internally calls [`Form::nucomp`] followed by [`Form::reduce`].
485    pub fn compose(&self, other: &Form, d_disc: &BigInt) -> Form {
486        let l = isqrt_fourth(&d_disc.abs());
487        let mut form = self.nucomp(other, d_disc, &l);
488        form.reduce(d_disc);
489        form
490    }
491
492    /// Squares this form and returns the reduced result.
493    ///
494    /// Equivalent to `self.compose(self, d_disc)` but uses the faster
495    /// [`Form::nudupl`] specialization for self-composition.
496    pub fn square(&self, d_disc: &BigInt) -> Form {
497        let l = isqrt_fourth(&d_disc.abs());
498        let mut form = self.nudupl(d_disc, &l);
499        form.reduce(d_disc);
500        form
501    }
502}
503
504#[cfg(test)]
505mod tests {
506    use super::*;
507
508    #[test]
509    fn test_reduction_valid_negative_discriminant() {
510        let d = BigInt::from(-71);
511        let mut form = Form::new(BigInt::from(2), BigInt::from(1), BigInt::from(9));
512        form.reduce(&d);
513        assert_eq!(form.a, BigInt::from(2));
514        assert_eq!(form.b, BigInt::from(1));
515        assert_eq!(form.c, BigInt::from(9)); // Already in reduced form
516    }
517
518    #[test]
519    fn test_compose_and_square() {
520        let d = BigInt::from(-71);
521        let id = Form::identity(&d); // (1, 1, 18)
522        let f2 = Form::new(BigInt::from(2), BigInt::from(1), BigInt::from(9));
523
524        // Compose with identity must equal self
525        let comp1 = id.compose(&f2, &d);
526        assert_eq!(comp1, f2);
527
528        // Square
529        let sq = f2.square(&d);
530        assert_eq!(sq, Form::new(BigInt::from(4), BigInt::from(-3), BigInt::from(5)));
531
532        // Compose f2 with itself must equal square
533        let comp2 = f2.compose(&f2, &d);
534        assert_eq!(comp2, sq);
535
536        // f2 * f2^-1 must equal identity
537        let f2_inv = Form::new(f2.a.clone(), -f2.b.clone(), f2.c.clone());
538        let comp3 = f2.compose(&f2_inv, &d);
539        assert_eq!(comp3, id);
540    }
541}