use super::{
field_reduction_constants::FieldReductionConstants,
limbs::{add, gte, sub, sub_5_4},
};
pub trait MontgomeryLimbs: FieldReductionConstants {
fn from_limbs(limbs: [u64; 4]) -> Self;
fn to_limbs(&self) -> &[u64; 4];
}
#[inline]
pub(crate) fn montgomery_reduce_9<F: FieldReductionConstants>(c: &[u64; 9]) -> [u64; 4] {
let mut low8 = [c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]];
let h = c[8];
let mut fold_carry = 0u64;
if h != 0 {
let mut carry = 0u128;
for (limb, &r512_limb) in low8.iter_mut().zip(F::R512_MOD.iter()) {
let prod = (h as u128) * (r512_limb as u128) + (*limb as u128) + carry;
*limb = prod as u64;
carry = prod >> 64;
}
for limb in &mut low8[4..] {
let sum = (*limb as u128) + carry;
*limb = sum as u64;
carry = sum >> 64;
}
fold_carry = carry as u64;
}
debug_assert!(
fold_carry <= 1,
"fold carry must be 0 or 1, got {}",
fold_carry
);
let mut out = montgomery_reduce_8::<F>(&low8);
if fold_carry == 1 {
let (sum, carry) = add::<4>(&out, &F::R_MOD);
out = sum;
if carry == 1 || gte::<4>(&out, &F::MODULUS) {
out = sub::<4>(&out, &F::MODULUS);
}
}
out
}
#[inline]
fn montgomery_reduce_8<F: FieldReductionConstants>(t: &[u64; 8]) -> [u64; 4] {
let mut r = [t[0], t[1], t[2], t[3], t[4], t[5], t[6], t[7], 0u64];
for i in 0..4 {
let q = r[i].wrapping_mul(F::MONT_INV);
let mut carry = 0u128;
for j in 0..4 {
let prod = (q as u128) * (F::MODULUS[j] as u128) + (r[i + j] as u128) + carry;
r[i + j] = prod as u64;
carry = prod >> 64;
}
let sum = (r[i + 4] as u128) + carry;
r[i + 4] = sum as u64;
carry = sum >> 64;
let mut k = i + 5;
while carry != 0 && k < 9 {
let sum = (r[k] as u128) + carry;
r[k] = sum as u64;
carry = sum >> 64;
k += 1;
}
}
let mut x5 = [r[4], r[5], r[6], r[7], r[8]];
if x5[4] == 1 {
x5 = sub_5_4(&x5, &F::MODULUS);
debug_assert!(x5[4] == 0, "after 5th-limb subtract, x5[4] should be 0");
}
let mut out = [x5[0], x5[1], x5[2], x5[3]];
for _ in 0..F::MAX_REDC_SUB_CORRECTIONS {
if gte::<4>(&out, &F::MODULUS) {
out = sub::<4>(&out, &F::MODULUS);
}
}
debug_assert!(
!gte::<4>(&out, &F::MODULUS),
"REDC final reduction failed after {} subtractions",
F::MAX_REDC_SUB_CORRECTIONS
);
out
}
#[cfg(test)]
use super::limbs::mul_4_by_4;
#[cfg(test)]
use ff::PrimeField;
#[cfg(test)]
pub(crate) fn test_r512_mod_impl<F: FieldReductionConstants + MontgomeryLimbs + PrimeField>() {
let mut wide_512 = [0u64; 9];
wide_512[8] = 1;
let reduced = montgomery_reduce_9::<F>(&wide_512);
let reduced_field = F::from_limbs(reduced);
assert!(reduced_field != F::ZERO);
}
#[cfg(test)]
pub(crate) fn test_r512_folding_identity_impl<
F: FieldReductionConstants + MontgomeryLimbs + PrimeField,
>() {
let base_input: [u64; 9] = [0, 0, 0, 0, 0, 0, 0, 0, 1];
let base_reduced = F::from_limbs(montgomery_reduce_9::<F>(&base_input));
for h in [1u64, 2, 0xFF, 0xFFFF_FFFF, 0xFFFF_FFFF_FFFF_FFFF] {
let mut wide = [0u64; 9];
wide[8] = h;
let reduced = F::from_limbs(montgomery_reduce_9::<F>(&wide));
let expected = F::from(h) * base_reduced;
assert_eq!(reduced, expected);
}
}
#[cfg(test)]
pub(crate) fn test_montgomery_round_trip_impl<
F: FieldReductionConstants + MontgomeryLimbs + PrimeField + Copy,
>() {
use rand::{SeedableRng, rngs::StdRng};
let mut rng = StdRng::seed_from_u64(12345);
for _ in 0..100 {
let a = F::random(&mut rng);
let b = F::random(&mut rng);
let expected = a * b;
let product_wide = mul_4_by_4(a.to_limbs(), b.to_limbs());
let mut wide_9 = [0u64; 9];
wide_9[..8].copy_from_slice(&product_wide);
let reduced_limbs = montgomery_reduce_9::<F>(&wide_9);
let result = F::from_limbs(reduced_limbs);
assert_eq!(result, expected);
}
}
#[cfg(test)]
#[macro_export]
macro_rules! test_montgomery {
($mod_name:ident, $field:ty) => {
mod $mod_name {
#[test]
fn r512_mod() {
$crate::big_num::montgomery::test_r512_mod_impl::<$field>();
}
#[test]
fn r512_folding_identity() {
$crate::big_num::montgomery::test_r512_folding_identity_impl::<$field>();
}
#[test]
fn montgomery_round_trip() {
$crate::big_num::montgomery::test_montgomery_round_trip_impl::<$field>();
}
}
};
}