use std::collections::HashMap;
use super::crt::crt_many;
use super::factor::factor_integer;
use super::mod_inv;
use super::primes::is_prime_bpsw;
use crate::Integer;
fn bsgs_bounded(
base: &Integer,
target: &Integer,
modulus: &Integer,
order_bound: &Integer,
) -> Option<Integer> {
let one = Integer::from(1);
let m = order_bound.sqrt() + &one;
let mut table: HashMap<Integer, Integer> = HashMap::new();
let mut cur = one.mod_floor(modulus);
let mut j = Integer::from(0);
while j < m {
table.entry(cur.clone()).or_insert_with(|| j.clone());
cur = (&cur * base).mod_floor(modulus);
j += &one;
}
let factor = mod_inv(&base.modpow(&m, modulus), modulus)?;
let target_n = target.mod_floor(modulus);
let mut gamma = target_n.clone();
let mut i = Integer::from(0);
while i <= m {
if let Some(j) = table.get(&gamma) {
let x = &(&i * &m) + j;
if base.modpow(&x, modulus) == target_n {
return Some(x);
}
}
gamma = (&gamma * &factor).mod_floor(modulus);
i += &one;
}
None
}
pub fn dlog_bsgs(base: &Integer, target: &Integer, modulus: &Integer) -> Option<Integer> {
let one = Integer::from(1);
if modulus <= &one {
return None;
}
let t = target.mod_floor(modulus);
if t.is_one() {
return Some(Integer::from(0));
}
bsgs_bounded(&base.mod_floor(modulus), &t, modulus, &(modulus - &one))
}
pub fn dlog_pohlig_hellman(base: &Integer, target: &Integer, p: &Integer) -> Option<Integer> {
let one = Integer::from(1);
if !is_prime_bpsw(p) {
return None;
}
let n = p - &one;
let b = base.mod_floor(p);
let t0 = target.mod_floor(p);
if b.is_zero() || t0.is_zero() {
return None;
}
if t0.is_one() {
return Some(Integer::from(0));
}
let mut ord = n.clone();
for (q, _) in factor_integer(&n) {
loop {
let (quot, rem) = ord.div_rem(&q);
if !rem.is_zero() || !b.modpow(", p).is_one() {
break;
}
ord = quot;
}
}
if !t0.modpow(&ord, p).is_one() {
return None;
}
let mut congruences = Vec::new();
for (q, e) in factor_integer(&ord) {
let g = b.modpow(&(&ord / &q), p);
let mut x_i = Integer::from(0);
let mut q_pow = one.clone(); for _ in 0..e {
let inv_bxi = mod_inv(&b.modpow(&x_i, p), p)?;
let h = (&t0 * &inv_bxi).mod_floor(p);
let t = h.modpow(&(&ord / &(&q * &q_pow)), p);
let d = bsgs_bounded(&g, &t, p, &q)?;
x_i = &x_i + &(&d * &q_pow);
q_pow = &q_pow * &q;
}
congruences.push((x_i.mod_floor(&q_pow), q_pow));
}
let (x, _) = crt_many(&congruences)?;
if b.modpow(&x, p) == t0 { Some(x) } else { None }
}
#[cfg(test)]
mod tests {
use super::*;
fn b(n: i64) -> Integer {
Integer::from(n)
}
#[test]
fn bsgs_small_group() {
let p = b(11);
let base = b(2);
let mut target = b(1);
for _ in 0..10 {
let x = dlog_bsgs(&base, &target, &p).unwrap();
assert_eq!(base.modpow(&x, &p), target, "wrong log for {target}");
assert!(x >= b(0) && x < b(10));
target = (&target * &base).mod_floor(&p);
}
assert!(dlog_bsgs(&base, &b(0), &p).is_none());
}
#[test]
fn bsgs_larger_modulus() {
let p = b(1009);
let base = b(11);
let target = base.modpow(&b(377), &p);
let x = dlog_bsgs(&base, &target, &p).unwrap();
assert_eq!(base.modpow(&x, &p), target);
}
#[test]
fn pohlig_hellman_smooth_order() {
let p = b(101);
let base = b(2); for e in [0i64, 1, 7, 42, 83, 99] {
let target = base.modpow(&b(e), &p);
let x = dlog_pohlig_hellman(&base, &target, &p).unwrap();
assert_eq!(x.mod_floor(&b(100)), b(e).mod_floor(&b(100)));
}
}
#[test]
fn pohlig_hellman_safe_prime() {
let p = b(23);
let base = b(5);
let mut target = b(1);
for _ in 0..22 {
let x = dlog_pohlig_hellman(&base, &target, &p).unwrap();
assert_eq!(base.modpow(&x, &p), target);
target = (&target * &base).mod_floor(&p);
}
}
#[test]
fn pohlig_hellman_non_primitive_base() {
let p = b(23);
let base = b(2);
let target = base.modpow(&b(7), &p);
let x = dlog_pohlig_hellman(&base, &target, &p).unwrap();
assert_eq!(base.modpow(&x, &p), target);
assert!(dlog_pohlig_hellman(&base, &b(5), &p).is_none());
}
#[test]
fn pohlig_hellman_rejects_composite_modulus() {
assert!(dlog_pohlig_hellman(&b(2), &b(3), &b(15)).is_none());
}
#[test]
fn pohlig_hellman_larger_prime() {
let p = b(2 * 3 * 5 * 5 * 7 + 1); assert!(is_prime_bpsw(&p));
let mut g = b(2);
loop {
let mut primitive = true;
for q in [b(2), b(3), b(5), b(7)] {
if g.modpow(&(&(&p - &b(1)) / &q), &p).is_one() {
primitive = false;
break;
}
}
if primitive {
break;
}
g += &b(1);
}
for e in [1i64, 100, 777, 1050] {
let target = g.modpow(&b(e), &p);
let x = dlog_pohlig_hellman(&g, &target, &p).unwrap();
assert_eq!(g.modpow(&x, &p), target);
}
}
}