use num::{BigInt, Zero};
use crate::rns::{NttPrimeList, Rns};
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub(crate) struct Packed<const L: usize> {
limbs: [u64; L],
}
impl<const L: usize> Packed<L> {
pub(crate) const ZERO: Self = Self { limbs: [0; L] };
#[cfg(test)]
pub(crate) fn from_i128(v: i128) -> Self {
const { assert!(L >= 2, "Packed requires at least 2 limbs (128-bit input)") };
let lo = v as u128;
let ext = if v < 0 { u64::MAX } else { 0 };
let mut limbs = [ext; L];
limbs[0] = lo as u64;
limbs[1] = (lo >> 64) as u64;
Self { limbs }
}
fn is_negative(&self) -> bool {
self.limbs[L - 1] >> 63 == 1
}
fn neg(&self) -> Self {
let mut out = [0u64; L];
let mut carry = 1u128;
for (i, out_i) in out.iter_mut().enumerate().take(L) {
let v = (!self.limbs[i]) as u128 + carry;
*out_i = v as u64;
carry = v >> 64;
}
Self { limbs: out }
}
fn magnitude(&self) -> [u64; L] {
if self.is_negative() {
self.neg().limbs
} else {
self.limbs
}
}
pub(crate) fn bit_length(&self) -> u64 {
let m = self.magnitude();
for i in (0..L).rev() {
if m[i] != 0 {
return (i as u64) * 64 + (64 - m[i].leading_zeros() as u64);
}
}
0
}
pub(crate) fn shr_to_i128(&self, shift: u32) -> i128 {
let sign = if self.is_negative() { u64::MAX } else { 0 };
let limb = (shift / 64) as usize;
let bit = shift % 64;
let get = |idx: usize| -> u64 {
if idx < L {
self.limbs[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 shl(&self, bits: u32) -> Self {
debug_assert!(
self.bit_length() + bits as u64 <= 64 * L as u64,
"Packed<{L}> shl overflow: {}-bit value << {bits}",
self.bit_length()
);
let limb = (bits / 64) as usize;
let bit = bits % 64;
let mut out = [0u64; L];
for (i, out_i) in out.iter_mut().enumerate().take(L) {
let src = i as isize - limb as isize;
if src < 0 {
continue;
}
let src = src as usize;
let mut word = self.limbs[src] as u128;
if bit != 0 {
let lower = if src >= 1 {
self.limbs[src - 1] as u128
} else {
0
};
word = (word << bit) | (lower >> (64 - bit));
}
*out_i = word as u64;
}
Self { limbs: out }
}
pub(crate) fn sub(&self, other: &Self) -> Self {
let mut out = [0u64; L];
let mut borrow = 0i128;
for (i, out_i) in out.iter_mut().enumerate().take(L) {
let v = self.limbs[i] as i128 - other.limbs[i] as i128 - borrow;
*out_i = v as u64;
borrow = if v < 0 { 1 } else { 0 };
}
Self { limbs: out }
}
fn sub_assign(&mut self, other: &Self) {
*self = self.sub(other);
}
fn mul_u32(&self, v: u32) -> Self {
let mut out = [0u64; L];
let mut carry = 0u128;
for (i, out_i) in out.iter_mut().enumerate().take(L) {
let prod = self.limbs[i] as u128 * v as u128 + carry;
*out_i = prod as u64;
carry = prod >> 64;
}
Self { limbs: out }
}
fn ucmp(&self, other: &Self) -> core::cmp::Ordering {
for i in (0..L).rev() {
match self.limbs[i].cmp(&other.limbs[i]) {
core::cmp::Ordering::Equal => continue,
ord => return ord,
}
}
core::cmp::Ordering::Equal
}
pub(crate) fn from_rns<const K: usize, P: NttPrimeList<K>>(r: &Rns<K, P>) -> Self {
let digits = r.to_garner();
let mut acc = Self::ZERO;
let mut modulus = {
let mut l = [0u64; L];
l[0] = 1;
Self { limbs: l }
};
for (i, &digits_i) in digits.iter().enumerate().take(K) {
acc = acc.add_packed(&modulus.mul_u32(digits_i));
modulus = modulus.mul_u32(P::PRIMES[i]);
}
let half = modulus.shr_unsigned(1);
if acc.ucmp(&half) == core::cmp::Ordering::Greater {
acc.sub(&modulus)
} else {
acc
}
}
fn add_packed(&self, other: &Self) -> Self {
let mut out = [0u64; L];
let mut carry = 0u128;
for (i, out_i) in out.iter_mut().enumerate().take(L) {
let s = self.limbs[i] as u128 + other.limbs[i] as u128 + carry;
*out_i = s as u64;
carry = s >> 64;
}
Self { limbs: out }
}
fn shr_unsigned(&self, bits: u32) -> Self {
let limb = (bits / 64) as usize;
let bit = bits % 64;
let mut out = [0u64; L];
for (i, out_i) in out.iter_mut().enumerate().take(L) {
let src = i + limb;
let mut word = if src < L { self.limbs[src] as u128 } else { 0 };
if bit != 0 {
let hi = if src + 1 < L {
self.limbs[src + 1] as u128
} else {
0
};
word = (word >> bit) | (hi << (64 - bit));
word &= u64::MAX as u128;
}
*out_i = word as u64;
}
Self { limbs: out }
}
pub(crate) fn from_bigint(x: &BigInt) -> Self {
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() <= L,
"Packed<{L}> too narrow: value needs {} limbs ({} bits)",
words.len(),
x.bits()
);
let mut limbs = [0u64; L];
for (i, w) in words.iter().take(L).enumerate() {
limbs[i] = *w;
}
let p = Self { limbs };
if neg {
p.neg()
} else {
p
}
}
pub(crate) fn to_bigint(self) -> BigInt {
let neg = self.is_negative();
let m = self.magnitude();
let mut acc = BigInt::zero();
for i in (0..L).rev() {
acc <<= 64;
acc += BigInt::from(m[i]);
}
if neg {
-acc
} else {
acc
}
}
}
impl<const L: usize> std::ops::SubAssign<&Packed<L>> for Packed<L> {
fn sub_assign(&mut self, rhs: &Packed<L>) {
Packed::sub_assign(self, rhs)
}
}
#[cfg(test)]
mod tests {
use crate::rns::{NttPrimes24Bit8, Rns};
use num::{BigInt, Zero};
use rand::{rngs::StdRng, RngExt, SeedableRng};
type Packed = super::Packed<4>;
fn rand_packed(rng: &mut StdRng) -> (Packed, BigInt) {
let hi = rng.random::<i64>() as i128; let lo = rng.random::<i128>();
let p = Packed::from_i128(hi)
.shl(96)
.add_packed(&Packed::from_i128(lo & ((1i128 << 96) - 1)));
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!(Packed::from_i128(v).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_packed(&mut rng);
assert_eq!(Packed::from_bigint(&b), 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_packed(&mut rng);
assert_eq!(p.bit_length(), b.bits(), "value {b}");
}
assert_eq!(Packed::ZERO.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_packed(&mut rng);
let (pb, bb) = rand_packed(&mut rng);
assert_eq!(pa.sub(&pb).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 = Packed::from_i128(lo);
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 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!(Packed::from_rns(&r).to_bigint(), v);
}
}
#[test]
fn shr_to_i128_matches_bigint() {
let mut rng = StdRng::seed_from_u64(6);
for _ in 0..1000 {
let (p, b) = rand_packed(&mut rng);
for &s in &[0u32, 1, 33, 64, 90, 128] {
let want = &b >> s;
let got = BigInt::from(p.shr_to_i128(s));
if want.bits() <= 126 {
assert_eq!(got, want, "v={b} s={s}");
}
}
}
}
}