use std::cmp::Ordering;
use num::{BigInt, Zero};
pub(crate) fn is_negative(a: &[u64]) -> bool {
matches!(a.last(), Some(top) if top >> 63 == 1)
}
pub(crate) fn neg_into(out: &mut [u64], a: &[u64]) {
debug_assert_eq!(out.len(), a.len());
let mut carry = 1u128;
for i in 0..out.len() {
let v = (!a[i]) as u128 + carry;
out[i] = v as u64;
carry = v >> 64;
}
}
pub(crate) fn neg_in_place(a: &mut [u64]) {
let mut carry = 1u128;
for w in a.iter_mut() {
let v = (!*w) as u128 + carry;
*w = v as u64;
carry = v >> 64;
}
}
pub(crate) fn add_into(out: &mut [u64], a: &[u64], b: &[u64]) {
debug_assert!(out.len() == a.len() && a.len() == b.len());
let mut carry = 0u128;
for i in 0..out.len() {
let s = a[i] as u128 + b[i] as u128 + carry;
out[i] = s as u64;
carry = s >> 64;
}
}
pub(crate) fn add_into_self(dst: &mut [u64], src: &[u64]) {
debug_assert_eq!(dst.len(), src.len());
let mut carry = 0u128;
for i in 0..dst.len() {
let s = dst[i] as u128 + src[i] as u128 + carry;
dst[i] = s as u64;
carry = s >> 64;
}
}
pub(crate) fn sub_into_self(dst: &mut [u64], src: &[u64]) {
debug_assert_eq!(dst.len(), src.len());
let mut borrow = 0i128;
for i in 0..dst.len() {
let v = dst[i] as i128 - src[i] as i128 - borrow;
dst[i] = v as u64;
borrow = if v < 0 { 1 } else { 0 };
}
}
pub(crate) fn sub_into(out: &mut [u64], a: &[u64], b: &[u64]) {
debug_assert!(out.len() == a.len() && a.len() == b.len());
let mut borrow = 0i128;
for i in 0..out.len() {
let v = a[i] as i128 - b[i] as i128 - borrow;
out[i] = v as u64;
borrow = if v < 0 { 1 } else { 0 };
}
}
pub(crate) fn mul_u32_into(out: &mut [u64], a: &[u64], v: u32) {
debug_assert_eq!(out.len(), a.len());
let mut carry = 0u128;
for i in 0..out.len() {
let prod = a[i] as u128 * v as u128 + carry;
out[i] = prod as u64;
carry = prod >> 64;
}
}
pub(crate) fn umul_into(out: &mut [u64], a: &[u64], b: &[u64]) {
for w in out.iter_mut() {
*w = 0;
}
let lo = out.len();
for (i, a_i) in a.iter().enumerate() {
if i >= lo {
break;
}
let ai = *a_i as u128;
if ai == 0 {
continue;
}
let mut carry = 0u128;
let mut idx = i;
for &bj in b.iter() {
if idx >= lo {
carry = 0;
break;
}
let prod = ai * bj as u128 + out[idx] as u128 + carry;
out[idx] = prod as u64;
carry = prod >> 64;
idx += 1;
}
while carry != 0 && idx < lo {
let v = out[idx] as u128 + carry;
out[idx] = v as u64;
carry = v >> 64;
idx += 1;
}
}
}
#[cfg(test)]
pub(crate) fn mul_into(out: &mut [u64], a: &[u64], b: &[u64]) {
let na = is_negative(a);
let nb = is_negative(b);
let mut ma = a.to_vec();
let mut mb = b.to_vec();
if na {
neg_in_place(&mut ma);
}
if nb {
neg_in_place(&mut mb);
}
umul_into(out, &ma, &mb);
if na ^ nb {
neg_in_place(out);
}
}
pub(crate) fn shl_into(out: &mut [u64], a: &[u64], bits: u32) {
debug_assert_eq!(out.len(), a.len());
let len = out.len();
let limb = (bits / 64) as usize;
let bit = bits % 64;
for (i, out_i) in out.iter_mut().enumerate().take(len) {
let src = i as isize - limb as isize;
if src < 0 {
*out_i = 0;
continue;
}
let src = src as usize;
let mut word = a[src] as u128;
if bit != 0 {
let lower = if src >= 1 { a[src - 1] as u128 } else { 0 };
word = (word << bit) | (lower >> (64 - bit));
}
*out_i = word as u64;
}
}
pub(crate) fn shr_unsigned_into(out: &mut [u64], a: &[u64], bits: u32) {
debug_assert_eq!(out.len(), a.len());
let len = out.len();
let limb = (bits / 64) as usize;
let bit = bits % 64;
for (i, out_i) in out.iter_mut().enumerate().take(len) {
let src = i + limb;
let mut word = if src < len { a[src] as u128 } else { 0 };
if bit != 0 {
let hi = if src + 1 < len { a[src + 1] as u128 } else { 0 };
word = (word >> bit) | (hi << (64 - bit));
word &= u64::MAX as u128;
}
*out_i = word as u64;
}
}
pub(crate) fn shr_to_i128(a: &[u64], shift: u32) -> i128 {
let len = a.len();
let sign = if is_negative(a) { u64::MAX } else { 0 };
let limb = (shift / 64) as usize;
let bit = shift % 64;
let get = |idx: usize| -> u64 {
if idx < len {
a[idx]
} else {
sign
}
};
let mut lo = 0u128;
for out_idx in 0..2 {
let src = limb + out_idx;
let mut word = get(src) as u128;
if bit != 0 {
let hi = get(src + 1) as u128;
word = (word >> bit) | (hi << (64 - bit));
word &= u64::MAX as u128;
}
lo |= word << (64 * out_idx);
}
lo as i128
}
pub(crate) fn ucmp(a: &[u64], b: &[u64]) -> Ordering {
debug_assert_eq!(a.len(), b.len());
for i in (0..a.len()).rev() {
match a[i].cmp(&b[i]) {
Ordering::Equal => continue,
ord => return ord,
}
}
Ordering::Equal
}
pub(crate) fn bit_length(a: &[u64]) -> u64 {
if is_negative(a) {
let mut mag = vec![0u64; a.len()];
neg_into(&mut mag, a);
bit_length_unsigned(&mag)
} else {
bit_length_unsigned(a)
}
}
fn bit_length_unsigned(m: &[u64]) -> u64 {
for i in (0..m.len()).rev() {
if m[i] != 0 {
return (i as u64) * 64 + (64 - m[i].leading_zeros() as u64);
}
}
0
}
pub(crate) fn from_bigint_into(out: &mut [u64], x: &BigInt) {
let neg = x.sign() == num::bigint::Sign::Minus;
let mag = if neg { -x } else { x.clone() };
let (_, words) = mag.to_u64_digits();
debug_assert!(
words.len() <= out.len(),
"MultiwordInt too narrow: value needs {} limbs ({} bits), have {}",
words.len(),
x.bits(),
out.len()
);
for w in out.iter_mut() {
*w = 0;
}
for (i, w) in words.iter().take(out.len()).enumerate() {
out[i] = *w;
}
if neg {
neg_in_place(out);
}
}
pub(crate) fn to_bigint(a: &[u64]) -> BigInt {
let neg = is_negative(a);
let mut mag = a.to_vec();
if neg {
neg_in_place(&mut mag);
}
let mut acc = BigInt::zero();
for i in (0..a.len()).rev() {
acc <<= 64;
acc += BigInt::from(mag[i]);
}
if neg {
-acc
} else {
acc
}
}
#[derive(Clone, PartialEq, Eq, Debug)]
pub(crate) struct MultiwordInt {
limbs: Vec<u64>,
}
impl MultiwordInt {
pub(crate) fn zero(len: usize) -> Self {
Self {
limbs: vec![0; len],
}
}
pub(crate) fn len(&self) -> usize {
self.limbs.len()
}
pub(crate) fn limbs(&self) -> &[u64] {
&self.limbs
}
pub(crate) fn from_i128(v: i128, len: usize) -> Self {
debug_assert!(len >= 2, "MultiwordInt needs ≥ 2 limbs for a 128-bit input");
let lo = v as u128;
let ext = if v < 0 { u64::MAX } else { 0 };
let mut limbs = vec![ext; len];
limbs[0] = lo as u64;
limbs[1] = (lo >> 64) as u64;
Self { limbs }
}
pub(crate) fn sub(&self, other: &Self) -> Self {
let mut out = vec![0u64; self.len()];
sub_into(&mut out, &self.limbs, &other.limbs);
Self { limbs: out }
}
fn add(&self, other: &Self) -> Self {
let mut out = vec![0u64; self.len()];
add_into(&mut out, &self.limbs, &other.limbs);
Self { limbs: out }
}
fn mul_u32(&self, v: u32) -> Self {
let mut out = vec![0u64; self.len()];
mul_u32_into(&mut out, &self.limbs, v);
Self { limbs: out }
}
#[cfg(test)]
pub(crate) fn mul(&self, other: &Self, out_len: usize) -> Self {
let mut out = vec![0u64; out_len];
mul_into(&mut out, &self.limbs, &other.limbs);
Self { limbs: out }
}
fn shr_unsigned(&self, bits: u32) -> Self {
let mut out = vec![0u64; self.len()];
shr_unsigned_into(&mut out, &self.limbs, bits);
Self { limbs: out }
}
pub(crate) fn from_garner_digits(digits: &[u32], primes: &[u32], len: usize) -> Self {
debug_assert_eq!(digits.len(), primes.len());
let mut acc = Self::zero(len);
let mut modulus = {
let mut l = vec![0u64; len];
l[0] = 1;
Self { limbs: l }
};
for i in 0..digits.len() {
acc = acc.add(&modulus.mul_u32(digits[i]));
modulus = modulus.mul_u32(primes[i]);
}
let half = modulus.shr_unsigned(1);
if ucmp(&acc.limbs, &half.limbs) == Ordering::Greater {
acc.sub(&modulus)
} else {
acc
}
}
}
impl std::ops::SubAssign<&MultiwordInt> for MultiwordInt {
fn sub_assign(&mut self, rhs: &MultiwordInt) {
sub_into_self(&mut self.limbs, &rhs.limbs);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::rns::{NttPrimeList, NttPrimes24Bit8, Rns};
use num::{BigInt, Zero};
use rand::{rngs::StdRng, RngExt, SeedableRng};
impl MultiwordInt {
pub(crate) fn bit_length(&self) -> u64 {
bit_length(&self.limbs)
}
pub(crate) fn shr_to_i128(&self, shift: u32) -> i128 {
shr_to_i128(&self.limbs, shift)
}
pub(crate) fn shl(&self, bits: u32) -> Self {
debug_assert!(
self.bit_length() + bits as u64 <= 64 * self.len() as u64,
"MultiwordInt shl overflow: {}-bit value << {bits} in {} limbs",
self.bit_length(),
self.len()
);
let mut out = vec![0u64; self.len()];
shl_into(&mut out, &self.limbs, bits);
Self { limbs: out }
}
pub(crate) fn from_rns<const K: usize, P: NttPrimeList<K>>(
r: &Rns<K, P>,
len: usize,
) -> Self {
Self::from_garner_digits(&r.to_garner(), &P::PRIMES, len)
}
pub(crate) fn from_bigint(x: &BigInt, len: usize) -> Self {
let mut limbs = vec![0u64; len];
from_bigint_into(&mut limbs, x);
Self { limbs }
}
pub(crate) fn to_bigint(&self) -> BigInt {
to_bigint(&self.limbs)
}
}
const LEN: usize = 4;
fn rand_value(rng: &mut StdRng) -> (MultiwordInt, BigInt) {
let hi = rng.random::<i64>() as i128;
let lo = rng.random::<i128>() & ((1i128 << 96) - 1);
let p = MultiwordInt::from_i128(hi, LEN)
.shl(96)
.add(&MultiwordInt::from_i128(lo, LEN));
let b = p.to_bigint();
(p, b)
}
#[test]
fn from_to_i128_roundtrip() {
let mut rng = StdRng::seed_from_u64(1);
for _ in 0..1000 {
let v = rng.random::<i128>();
assert_eq!(MultiwordInt::from_i128(v, LEN).to_bigint(), BigInt::from(v));
}
}
#[test]
fn bigint_roundtrip() {
let mut rng = StdRng::seed_from_u64(2);
for _ in 0..1000 {
let (p, b) = rand_value(&mut rng);
assert_eq!(MultiwordInt::from_bigint(&b, LEN), p);
assert_eq!(p.to_bigint(), b);
}
}
#[test]
fn bit_length_matches_bigint() {
let mut rng = StdRng::seed_from_u64(3);
for _ in 0..1000 {
let (p, b) = rand_value(&mut rng);
assert_eq!(p.bit_length(), b.bits(), "value {b}");
}
assert_eq!(MultiwordInt::zero(LEN).bit_length(), BigInt::zero().bits());
}
#[test]
fn sub_matches_bigint() {
let mut rng = StdRng::seed_from_u64(4);
for _ in 0..1000 {
let (pa, ba) = rand_value(&mut rng);
let (pb, bb) = rand_value(&mut rng);
assert_eq!(pa.sub(&pb).to_bigint(), &ba - &bb);
let mut acc = pa.clone();
acc -= &pb;
assert_eq!(acc.to_bigint(), &ba - &bb);
}
}
#[test]
fn shl_matches_bigint() {
let mut rng = StdRng::seed_from_u64(5);
for _ in 0..500 {
let lo = rng.random::<i64>() as i128;
let p = MultiwordInt::from_i128(lo, LEN);
let b = BigInt::from(lo);
for &s in &[0u32, 1, 7, 31, 63, 64, 65, 100, 130] {
assert_eq!(p.shl(s).to_bigint(), &b << s, "v={b} s={s}");
}
}
}
#[test]
fn shr_to_i128_matches_bigint() {
let mut rng = StdRng::seed_from_u64(6);
for _ in 0..1000 {
let (p, b) = rand_value(&mut rng);
for &s in &[0u32, 1, 33, 64, 90, 128] {
let want = &b >> s;
if want.bits() <= 126 {
assert_eq!(BigInt::from(p.shr_to_i128(s)), want, "v={b} s={s}");
}
}
}
}
#[test]
fn mul_matches_bigint() {
let mut rng = StdRng::seed_from_u64(8);
const OUT: usize = 8;
for _ in 0..2000 {
let (pa, ba) = rand_value(&mut rng);
let (pb, bb) = rand_value(&mut rng);
assert_eq!(pa.mul(&pb, OUT).to_bigint(), &ba * &bb);
}
}
#[test]
fn from_rns_reconstructs_signed() {
let mut rng = StdRng::seed_from_u64(7);
for _ in 0..2000 {
let v = BigInt::from(rng.random::<i128>()) >> 1; let r = Rns::<8, NttPrimes24Bit8>::from_bigint(&v);
assert_eq!(MultiwordInt::from_rns(&r, LEN).to_bigint(), v);
}
}
}