pub(crate) const MAX_LIMBS: usize = 64;
pub(crate) fn n0_inv(n0: u64) -> u64 {
let mut inv = 1u64;
for _ in 0..6 {
inv = inv.wrapping_mul(2u64.wrapping_sub(n0.wrapping_mul(inv)));
}
inv.wrapping_neg()
}
pub(crate) fn geq(a: &[u64], b: &[u64]) -> bool {
debug_assert_eq!(a.len(), b.len());
for j in (0..a.len()).rev() {
if a[j] > b[j] {
return true;
}
if a[j] < b[j] {
return false;
}
}
true
}
pub(crate) fn select(mask: u64, a: &[u64], b: &[u64], out: &mut [u64]) {
debug_assert_eq!(a.len(), b.len());
debug_assert_eq!(a.len(), out.len());
for j in 0..out.len() {
out[j] = b[j] ^ ((a[j] ^ b[j]) & mask);
}
}
pub(crate) fn os2ip_be(bytes: &[u8], out: &mut [u64]) {
for w in out.iter_mut() {
*w = 0;
}
for (i, chunk) in bytes.rchunks(8).enumerate() {
let mut w = [0u8; 8];
w[8 - chunk.len()..].copy_from_slice(chunk);
out[i] = u64::from_be_bytes(w);
}
}
pub(crate) fn i2osp_be(limbs: &[u64], out: &mut [u8]) {
for w in out.iter_mut() {
*w = 0;
}
for (i, chunk) in out.rchunks_mut(8).enumerate() {
let be = limbs[i].to_be_bytes();
chunk.copy_from_slice(&be[8 - chunk.len()..]);
}
}
pub(crate) fn mont_mul(a: &[u64], b: &[u64], n: &[u64], n0: u64, out: &mut [u64]) {
let l = n.len();
debug_assert!(l <= MAX_LIMBS);
debug_assert_eq!(a.len(), l);
debug_assert_eq!(b.len(), l);
let mut prod = [0u64; 2 * MAX_LIMBS + 1];
for i in 0..l {
let ai = a[i] as u128;
let mut carry = 0u128;
for j in 0..l {
let s = (prod[i + j] as u128) + ai * (b[j] as u128) + carry;
prod[i + j] = s as u64;
carry = s >> 64;
}
let mut k = i + l;
while carry > 0 {
let s = (prod[k] as u128) + carry;
prod[k] = s as u64;
carry = s >> 64;
k += 1;
}
}
for i in 0..l {
let m = prod[i].wrapping_mul(n0);
let mut carry = 0u128;
for j in 0..l {
let s = (prod[i + j] as u128) + (m as u128) * (n[j] as u128) + carry;
prod[i + j] = s as u64;
carry = s >> 64;
}
let mut k = i + l;
while carry > 0 {
let s = (prod[k] as u128) + carry;
prod[k] = s as u64;
carry = s >> 64;
k += 1;
}
}
let hi = (prod[2 * l] != 0) as u64;
let mut r = [0u64; MAX_LIMBS];
r[..l].copy_from_slice(&prod[l..2 * l]);
let mut borrow = 0u64;
let mut tmp = [0u64; MAX_LIMBS];
for j in 0..l {
let (v, b1) = r[j].overflowing_sub(n[j]);
let (v, b2) = v.overflowing_sub(borrow);
tmp[j] = v;
borrow = (b1 as u64) | (b2 as u64);
}
let cond = if hi >= 1 {
u64::MAX
} else {
((borrow == 0) as u64).wrapping_neg()
};
for j in 0..l {
out[j] = r[j] ^ ((r[j] ^ tmp[j]) & cond);
}
}
pub(crate) fn to_mont(x: &[u64], r2: &[u64], n: &[u64], n0: u64, out: &mut [u64]) {
mont_mul(x, r2, n, n0, out);
}
pub(crate) fn from_mont(x: &mut [u64], n: &[u64], n0: u64) {
let l = n.len();
let mut one = [0u64; MAX_LIMBS];
one[0] = 1;
let mut tmp = [0u64; MAX_LIMBS];
tmp[..l].copy_from_slice(x);
mont_mul(&tmp[..l], &one[..l], n, n0, x);
}
pub(crate) fn reduce_limbs(a: &[u64], n: &[u64], out: &mut [u64]) {
let l = n.len();
debug_assert!(l <= MAX_LIMBS);
let mut acc = [0u64; MAX_LIMBS + 1];
let mut n_ext = [0u64; MAX_LIMBS + 1];
n_ext[..l].copy_from_slice(n);
let mut tmp = [0u64; MAX_LIMBS + 1];
for i in (0..a.len() * 64).rev() {
let bit = (a[i / 64] >> (i % 64)) & 1;
let mut carry = bit;
for w in acc.iter_mut().take(l + 1) {
let c = *w >> 63;
*w = (*w << 1) | carry;
carry = c;
}
debug_assert_eq!(carry, 0, "2·acc + bit 必须 < 2^(64(l+1))");
let ge = if acc[l] != 0 { true } else { geq(&acc[..l], n) };
if ge {
let mut borrow = 0u64;
for j in 0..=l {
let (v, b) = acc[j].overflowing_sub(n_ext[j]);
let (v, b2) = v.overflowing_sub(borrow);
tmp[j] = v;
borrow = (b as u64) | (b2 as u64);
}
debug_assert_eq!(borrow, 0);
acc = tmp;
}
}
out[..l].copy_from_slice(&acc[..l]);
}
pub(crate) fn compute_r2(n: &[u64]) -> Vec<u64> {
let l = n.len();
let mut bits = vec![0u64; 2 * l + 1];
bits[2 * l] = 1;
let mut r2 = vec![0u64; l];
reduce_limbs(&bits, n, &mut r2);
r2
}
pub(crate) fn mul_full(a: &[u64], b: &[u64]) -> Vec<u64> {
let l = a.len();
let mut out = vec![0u64; 2 * l];
for i in 0..l {
let ai = a[i] as u128;
let mut carry = 0u128;
for j in 0..l {
let s = (out[i + j] as u128) + ai * (b[j] as u128) + carry;
out[i + j] = s as u64;
carry = s >> 64;
}
out[i + l] = carry as u64;
}
out
}
pub(crate) fn add_limbs(a: &[u64], b: &[u64], out: &mut [u64]) {
let mut carry = 0u128;
for j in 0..a.len() {
let s = (a[j] as u128) + (b[j] as u128) + carry;
out[j] = s as u64;
carry = s >> 64;
}
debug_assert_eq!(carry, 0, "加法溢出须由调用方排除");
}
pub(crate) fn sub_limbs(a: &[u64], b: &[u64], out: &mut [u64]) -> u64 {
let mut borrow = 0u64;
for j in 0..a.len() {
let (v, b1) = a[j].overflowing_sub(b[j]);
let (v, b2) = v.overflowing_sub(borrow);
out[j] = v;
borrow = (b1 as u64) | (b2 as u64);
}
borrow
}
fn is_zero(x: &[u64]) -> bool {
x.iter().all(|&w| w == 0)
}
fn is_one(x: &[u64]) -> bool {
x[0] == 1 && x[1..].iter().all(|&w| w == 0)
}
fn shr1(x: &mut [u64]) {
let mut carry = 0u64;
for w in x.iter_mut().rev() {
let c = *w & 1;
*w = (*w >> 1) | (carry << 63);
carry = c;
}
}
fn half_mod(x: &mut [u64], n: &[u64]) {
if x[0] & 1 == 0 {
shr1(x);
return;
}
let mut carry = 0u64;
for (w, nw) in x.iter_mut().zip(n.iter()) {
let (s, c1) = w.overflowing_add(*nw);
let (s, c2) = s.overflowing_add(carry);
*w = s;
carry = (c1 as u64) | (c2 as u64);
}
let mut cin = carry;
for w in x.iter_mut().rev() {
let c = *w & 1;
*w = (*w >> 1) | (cin << 63);
cin = c;
}
}
fn sub_mod(x: &mut [u64], y: &[u64], n: &[u64]) {
let mut tmp = vec![0u64; x.len()];
let borrow = sub_limbs(x, y, &mut tmp);
if borrow == 1 {
let mut carry = 0u64;
for (w, nw) in tmp.iter_mut().zip(n.iter()) {
let (s, c1) = w.overflowing_add(*nw);
let (s, c2) = s.overflowing_add(carry);
*w = s;
carry = (c1 as u64) | (c2 as u64);
}
}
x.copy_from_slice(&tmp);
}
pub(crate) fn mod_inverse_odd(a: &[u64], n: &[u64]) -> Option<Vec<u64>> {
let l = n.len();
debug_assert_eq!(a.len(), l);
if is_zero(a) {
return None; }
let mut u = a.to_vec();
let mut v = n.to_vec();
let mut x1 = vec![0u64; l];
let mut x2 = vec![0u64; l];
let mut scratch = vec![0u64; l];
x1[0] = 1;
let limit = 4 * 64 * l + 64;
for _ in 0..limit {
if is_one(&u) {
return Some(x1);
}
if is_one(&v) {
return Some(x2);
}
if is_zero(&u) || is_zero(&v) {
return None; }
while u[0] & 1 == 0 {
shr1(&mut u);
half_mod(&mut x1, n);
}
while v[0] & 1 == 0 {
shr1(&mut v);
half_mod(&mut x2, n);
}
if geq(&u, &v) {
sub_limbs(&u, &v, &mut scratch); std::mem::swap(&mut u, &mut scratch);
sub_mod(&mut x1, &x2, n);
} else {
sub_limbs(&v, &u, &mut scratch);
std::mem::swap(&mut v, &mut scratch);
sub_mod(&mut x2, &x1, n);
}
}
None }
pub(crate) fn mont_exp(
base: &[u64],
exp: &[u64],
exp_bits: usize,
n: &[u64],
n0: u64,
r2: &[u64],
out: &mut [u64],
) {
let l = n.len();
let mut one = [0u64; MAX_LIMBS];
one[0] = 1;
let mut result = [0u64; MAX_LIMBS];
to_mont(&one[..l], r2, n, n0, &mut result[..l]);
let mut tmp = [0u64; MAX_LIMBS];
let mut tmp2 = [0u64; MAX_LIMBS];
for i in (0..exp_bits).rev() {
mont_mul(&result[..l], &result[..l], n, n0, &mut tmp[..l]);
let mask = ((exp[i / 64] >> (i % 64)) & 1).wrapping_neg();
mont_mul(&tmp[..l], base, n, n0, &mut tmp2[..l]);
select(mask, &tmp2[..l], &tmp[..l], &mut result[..l]);
}
out[..l].copy_from_slice(&result[..l]);
}
#[cfg(test)]
mod tests {
use super::*;
fn mont_product_is_one(a: &[u64], inv: &[u64], n: &[u64]) -> bool {
let l = n.len();
let n0 = n0_inv(n[0]);
let r2 = compute_r2(n);
let mut am = vec![0u64; l];
let mut im = vec![0u64; l];
to_mont(a, &r2, n, n0, &mut am);
to_mont(inv, &r2, n, n0, &mut im);
let mut t = vec![0u64; l];
mont_mul(&am, &im, n, n0, &mut t);
from_mont(&mut t, n, n0);
is_one(&t)
}
#[test]
fn inverse_small_values() {
assert_eq!(mod_inverse_odd(&[3], &[11]), Some(vec![4]));
assert_eq!(mod_inverse_odd(&[10], &[11]), Some(vec![10]));
assert_eq!(mod_inverse_odd(&[1], &[11]), Some(vec![1]));
}
#[test]
fn inverse_non_coprime_is_none() {
assert_eq!(mod_inverse_odd(&[0], &[11]), None);
assert_eq!(mod_inverse_odd(&[5], &[15]), None); assert_eq!(mod_inverse_odd(&[6], &[9]), None); }
#[test]
fn inverse_large_selfcheck() {
let mut successes = 0usize;
for l in 2..=17usize {
for seed in 0..4u64 {
let mut state = seed ^ (l as u64) << 32;
let mut next = move || {
state = state.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut z = state;
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
z ^ (z >> 31)
};
let mut n = vec![0u64; l];
for w in n.iter_mut() {
*w = next();
}
n[l - 1] |= 1 << 63;
n[0] |= 1;
let mut a = vec![0u64; l];
for w in a.iter_mut() {
*w = next();
}
a[l - 1] &= (1 << 63) - 1; let mut r = vec![0u64; l];
reduce_limbs(&a, &n, &mut r);
if let Some(inv) = mod_inverse_odd(&r, &n) {
assert!(
mont_product_is_one(&r, &inv, &n),
"inv·a ≢ 1 (mod n) for l={l} seed={seed}"
);
successes += 1;
}
}
}
assert!(successes >= 32, "too few coprime cases: {successes}/64");
}
}