use crate::Integer;
const MR_WITNESSES: [u64; 12] = [2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37];
fn mr_witness(n: &Integer, d: &Integer, r: u64, a: &Integer) -> bool {
let mut x = a.modpow(d, n);
let one = Integer::from(1);
let n_minus_one = n - &one;
if x == one || x == n_minus_one {
return true;
}
for _ in 1..r {
x = (&x * &x).mod_floor(n);
if x == n_minus_one {
return true;
}
}
false
}
pub fn is_prime(n: &Integer) -> bool {
if n < &Integer::from(2) {
return false;
}
if *n == Integer::from(2) || *n == Integer::from(3) {
return true;
}
if n.is_even() {
return false;
}
for &p in &[3u64, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37] {
let pb = Integer::from(p as i64);
if *n == pb {
return true;
}
if n.mod_floor(&pb).is_zero() {
return false;
}
}
let one = Integer::from(1);
let n_minus_one = n - &one;
let mut d = n_minus_one.clone();
let mut r = 0u64;
while d.is_even() {
d >>= 1;
r += 1;
}
for &a in &MR_WITNESSES {
let ab = Integer::from(a as i64);
if ab >= *n {
continue;
}
if !mr_witness(n, &d, r, &ab) {
return false;
}
}
true
}
pub fn next_prime(n: &Integer) -> Integer {
let two = Integer::from(2);
let three = Integer::from(3);
if n < &two {
return two;
}
let mut candidate = n + &Integer::from(1);
if candidate == three {
return three;
}
if candidate.is_even() {
candidate += &Integer::from(1);
}
while !is_prime(&candidate) {
candidate += &two;
}
candidate
}
pub fn primes_from(n: &Integer) -> PrimesFrom {
PrimesFrom {
current: if n < &Integer::from(2) {
Integer::from(2)
} else {
n.clone()
},
}
}
pub struct PrimesFrom {
current: Integer,
}
impl Iterator for PrimesFrom {
type Item = Integer;
fn next(&mut self) -> Option<Integer> {
self.current = next_prime(&self.current);
Some(self.current.clone())
}
}
pub fn mod_inv(a: &Integer, m: &Integer) -> Option<Integer> {
if m <= &Integer::from(1) {
return None;
}
let (g, x, _) = extended_gcd(a, m);
if !g.is_one() {
return None;
}
let mut r = x.mod_floor(m);
if r.is_negative() {
r += m;
}
Some(r)
}
pub fn extended_gcd(a: &Integer, b: &Integer) -> (Integer, Integer, Integer) {
let mut old_r = a.clone();
let mut r = b.clone();
let mut old_s = Integer::from(1);
let mut s = Integer::from(0);
let mut old_t = Integer::from(0);
let mut t = Integer::from(1);
while !r.is_zero() {
let (q, rem) = old_r.div_rem(&r);
old_r = r;
r = rem;
let qs = &q * &s;
let new_s = &old_s - &qs;
old_s = s;
s = new_s;
let qt = &q * &t;
let new_t = &old_t - &qt;
old_t = t;
t = new_t;
}
if old_r.is_negative() {
old_r = -old_r;
old_s = -old_s;
old_t = -old_t;
}
(old_r, old_s, old_t)
}
pub fn symmetric_mod(a: &Integer, m: &Integer) -> Integer {
let half = m / &Integer::from(2);
let mut r = a.mod_floor(m);
if r > half {
r -= m;
}
r
}
pub fn crt(r1: &Integer, m1: &Integer, r2: &Integer, m2: &Integer) -> Option<(Integer, Integer)> {
let (g, p, _q) = extended_gcd(m1, m2);
let diff = r1 - r2;
if !diff.mod_floor(&g).is_zero() {
return None;
}
let lcm = (m1 / &g) * m2;
let step = (r2 - r1) / &g;
let mut r = r1 + &(m1 * &p * &step);
r = r.mod_floor(&lcm);
Some((r, lcm))
}
pub fn legendre(a: &Integer, p: &Integer) -> i8 {
jacobi(a, p)
}
pub fn jacobi(a: &Integer, n: &Integer) -> i8 {
if n.is_zero() || n.is_negative() || n.is_even() {
return 0;
}
let four = Integer::from(4);
let eight = Integer::from(8);
let mut a = a.mod_floor(n);
let mut n = n.clone();
let mut t: i8 = 1;
while !a.is_zero() {
while a.is_even() {
a >>= 1;
let r = n.mod_floor(&eight);
if r == Integer::from(3) || r == Integer::from(5) {
t = -t;
}
}
std::mem::swap(&mut a, &mut n);
if a.mod_floor(&four) == Integer::from(3) && n.mod_floor(&four) == Integer::from(3) {
t = -t;
}
a = a.mod_floor(&n);
}
if n.is_one() { t } else { 0 }
}
pub fn mod_sqrt(a: &Integer, p: &Integer) -> Option<Integer> {
if p <= &Integer::from(2) {
return None;
}
let a = a.mod_floor(p);
if a.is_zero() {
return Some(Integer::from(0));
}
if legendre(&a, p) != 1 {
return None;
}
if p.mod_floor(&Integer::from(4)) == Integer::from(3) {
let exp = (p + &Integer::from(1)) / &Integer::from(4);
let x = a.modpow(&exp, p);
return Some(x);
}
let one = Integer::from(1);
let mut q = p - &one;
let mut s = 0u64;
while q.is_even() {
q >>= 1;
s += 1;
}
let mut z = Integer::from(2);
while legendre(&z, p) != -1 {
z += &one;
}
let mut m = s;
let mut c = z.modpow(&q, p);
let mut t = a.modpow(&q, p);
let mut r = a.modpow(&((&q + &one) / &Integer::from(2)), p);
while !t.is_one() {
let mut i = 0u64;
let mut t2i = t.clone();
while !t2i.is_one() {
t2i = (&t2i * &t2i).mod_floor(p);
i += 1;
if i >= m {
return None;
}
}
let mut b = c.clone();
for _ in 0..(m - i - 1) {
b = (&b * &b).mod_floor(p);
}
m = i;
c = (&b * &b).mod_floor(p);
t = (&t * &c).mod_floor(p);
r = (&r * &b).mod_floor(p);
}
Some(r)
}
#[cfg(test)]
mod tests {
use super::*;
fn b(n: i64) -> Integer {
Integer::from(n)
}
#[test]
fn primality_small() {
let primes = [2, 3, 5, 7, 11, 13, 97, 101, 997, 7919];
for p in primes {
assert!(is_prime(&b(p)), "{p} should be prime");
}
let composites = [0, 1, 4, 6, 8, 9, 15, 21, 25, 100, 1001];
for c in composites {
assert!(!is_prime(&b(c)), "{c} should be composite");
}
}
#[test]
fn primality_carmichael() {
assert!(!is_prime(&b(561)));
assert!(!is_prime(&b(1105)));
assert!(!is_prime(&b(1729)));
}
#[test]
fn primality_large() {
assert!(is_prime(&Integer::from(2_147_483_647_i64)));
let big = Integer::from(2).pow_u32(67) - &Integer::from(1);
assert!(!is_prime(&big));
}
#[test]
fn next_prime_works() {
assert_eq!(next_prime(&b(0)), b(2));
assert_eq!(next_prime(&b(1)), b(2));
assert_eq!(next_prime(&b(2)), b(3));
assert_eq!(next_prime(&b(10)), b(11));
assert_eq!(next_prime(&b(13)), b(17));
assert_eq!(next_prime(&b(100)), b(101));
}
#[test]
fn primes_from_iterator() {
let got: Vec<String> = primes_from(&b(10)).take(5).map(|x| x.to_string()).collect();
assert_eq!(got, vec!["11", "13", "17", "19", "23"]);
}
#[test]
fn modular_inverse() {
assert_eq!(mod_inv(&b(3), &b(11)), Some(b(4))); assert_eq!(mod_inv(&b(7), &b(13)), Some(b(2))); assert_eq!(mod_inv(&b(2), &b(4)), None);
assert_eq!(mod_inv(&b(6), &b(9)), None);
let inv = mod_inv(&b(-3), &b(11)).unwrap();
assert_eq!(inv, b(7));
}
#[test]
fn symmetric_modulo_range() {
for a in -14..=14 {
let r = symmetric_mod(&b(a), &b(7));
assert!(
r > b(-4) && r <= b(3),
"symmetric_mod({a}, 7) = {r} out of range"
);
assert_eq!((&r * &r).mod_floor(&b(7)), (b(a) * b(a)).mod_floor(&b(7)));
}
assert_eq!(symmetric_mod(&b(6), &b(7)), b(-1));
assert_eq!(symmetric_mod(&b(5), &b(7)), b(-2));
assert_eq!(symmetric_mod(&b(3), &b(7)), b(3));
}
#[test]
fn crt_basic() {
let (r, m) = crt(&b(2), &b(3), &b(3), &b(5)).unwrap();
assert_eq!(r, b(8));
assert_eq!(m, b(15));
let (r, m) = crt(&b(2), &b(3), &b(3), &b(5)).unwrap();
let (r, m) = crt(&r, &m, &b(2), &b(7)).unwrap();
assert_eq!(r, b(23));
assert_eq!(m, b(105));
}
#[test]
fn crt_inconsistent() {
assert!(crt(&b(1), &b(4), &b(2), &b(4)).is_none());
}
#[test]
fn crt_non_coprime_compatible() {
let (r, m) = crt(&b(1), &b(4), &b(3), &b(6)).unwrap();
assert_eq!(m, b(12));
assert_eq!(r.mod_floor(&b(4)), b(1));
assert_eq!(r.mod_floor(&b(6)), b(3));
}
#[test]
fn legendre_symbols() {
assert_eq!(legendre(&b(2), &b(7)), 1);
assert_eq!(legendre(&b(3), &b(7)), -1);
assert_eq!(legendre(&b(5), &b(7)), -1);
assert_eq!(legendre(&b(0), &b(7)), 0);
}
#[test]
fn jacobi_reciprocity() {
assert_eq!(jacobi(&b(2), &b(15)), 1);
let primes = [7u64, 11, 13, 17, 23, 41, 101, 1009, 9907];
for &p in &primes {
let pb = Integer::from(p as i64);
assert!(is_prime(&pb), "{p} assumed prime");
let exp = (&pb - &Integer::from(1)) / &Integer::from(2);
for a in 0..p.min(60) {
let ab = b(a as i64);
let expected = match ab.modpow(&exp, &pb).to_string().as_str() {
"0" => 0,
"1" => 1,
_ => -1, };
assert_eq!(jacobi(&ab, &pb), expected, "jacobi({a}/{p}) mismatch");
}
}
}
#[test]
fn mod_sqrt_residue() {
for a in [b(0), b(1), b(2), b(4)] {
let x = mod_sqrt(&a, &b(7)).unwrap();
assert_eq!((&x * &x).mod_floor(&b(7)), a.mod_floor(&b(7)));
}
assert!(mod_sqrt(&b(3), &b(7)).is_none());
assert!(mod_sqrt(&b(5), &b(7)).is_none());
}
#[test]
fn mod_sqrt_fast_path_p3_mod4() {
let x = mod_sqrt(&b(3), &b(11)).unwrap();
assert_eq!((&x * &x).mod_floor(&b(11)), b(3));
}
#[test]
fn mod_sqrt_general_prime() {
let x = mod_sqrt(&b(2), &b(17)).unwrap();
assert_eq!((&x * &x).mod_floor(&b(17)), b(2));
for a in 0..41 {
let ab = b(a);
let r = (&ab * &ab).mod_floor(&b(41));
let s = mod_sqrt(&r, &b(41)).unwrap();
assert_eq!((&s * &s).mod_floor(&b(41)), r);
}
}
#[test]
fn extended_gcd_identity() {
let (g, x, y) = extended_gcd(&b(240), &b(46));
assert_eq!(g, b(2));
assert_eq!(&x * &b(240) + &y * &b(46), b(2));
let (g, _, _) = extended_gcd(&b(17), &b(13));
assert_eq!(g, b(1));
let (g, _, _) = extended_gcd(&b(-240), &b(46));
assert_eq!(g, b(2));
}
}