use ic_core::ct::Choice;
pub const MAX_LIMBS: usize = 9;
#[inline]
pub(crate) const fn adc<const N: usize>(a: [u64; N], b: [u64; N]) -> ([u64; N], u64) {
let mut out = [0u64; N];
let mut carry = 0u128;
let mut i = 0;
while i < N {
let sum = (a[i] as u128) + (b[i] as u128) + carry;
out[i] = sum as u64;
carry = sum >> 64;
i += 1;
}
(out, carry as u64)
}
#[inline]
pub(crate) const fn sbb<const N: usize>(a: [u64; N], b: [u64; N]) -> ([u64; N], u64) {
let mut out = [0u64; N];
let mut borrow = 0u128;
let mut i = 0;
while i < N {
let diff = (a[i] as u128)
.wrapping_sub(b[i] as u128)
.wrapping_sub(borrow);
out[i] = diff as u64;
borrow = (diff >> 127) & 1;
i += 1;
}
(out, borrow as u64)
}
#[inline]
pub(crate) const fn select<const N: usize>(mask: u64, a: [u64; N], b: [u64; N]) -> [u64; N] {
let mut out = [0u64; N];
let mut i = 0;
while i < N {
out[i] = b[i] ^ (mask & (a[i] ^ b[i]));
i += 1;
}
out
}
#[inline]
const fn double_mod<const N: usize>(x: [u64; N], m: [u64; N]) -> [u64; N] {
let (sum, carry) = adc(x, x);
let (reduced, borrow) = sbb(sum, m);
let need = carry | (1 - borrow);
select(need.wrapping_neg(), reduced, sum)
}
pub(crate) const fn compute_r2<const N: usize>(m: [u64; N]) -> [u64; N] {
let mut x = [0u64; N];
x[0] = 1;
let mut i = 0;
while i < 128 * N {
x = double_mod(x, m);
i += 1;
}
x
}
#[inline]
pub(crate) fn from_be_bytes<const N: usize>(bytes: &[u8]) -> [u64; N] {
let mut limbs = [0u64; N];
for (i, byte) in bytes.iter().rev().enumerate() {
limbs[i / 8] |= (*byte as u64) << (8 * (i % 8));
}
limbs
}
#[inline]
pub(crate) fn to_be_bytes<const N: usize>(limbs: &[u64; N], out: &mut [u8]) {
let n = out.len();
for (i, slot) in out.iter_mut().rev().enumerate() {
*slot = (limbs[i / 8] >> (8 * (i % 8))) as u8;
}
let _ = n;
}
pub(crate) const fn compute_neg_inv(m0: u64) -> u64 {
let mut inv = m0;
let mut i = 0;
while i < 6 {
inv = inv.wrapping_mul(2u64.wrapping_sub(m0.wrapping_mul(inv)));
i += 1;
}
inv.wrapping_neg()
}
pub trait Field: Copy + Clone + core::fmt::Debug + PartialEq + Eq + Sized {
type Bytes: AsRef<[u8]> + AsMut<[u8]> + Copy;
const ZERO: Self;
const ONE: Self;
const BYTE_LEN: usize;
fn add(&self, rhs: &Self) -> Self;
fn sub(&self, rhs: &Self) -> Self;
fn mul(&self, rhs: &Self) -> Self;
fn square(&self) -> Self;
fn double(&self) -> Self;
fn triple(&self) -> Self;
fn neg(&self) -> Self;
fn invert(&self) -> Self;
fn from_bytes(bytes: &Self::Bytes) -> Option<Self>;
fn from_bytes_reduced(bytes: &Self::Bytes) -> Self;
fn to_bytes(&self) -> Self::Bytes;
fn zero_bytes() -> Self::Bytes;
fn is_zero(&self) -> Choice;
fn ct_eq(&self, other: &Self) -> Choice;
fn cmov(a: &mut Self, b: &Self, choice: Choice);
fn is_odd(&self) -> Choice;
}
#[macro_export]
#[doc(hidden)]
macro_rules! mont_field {
($name:ident, $limbs:literal, $bytes:literal, $modulus:expr, $doc:literal) => {
#[doc = $doc]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct $name(pub [u64; $limbs]);
impl $name {
pub const MODULUS: [u64; $limbs] = $modulus;
const R2: [u64; $limbs] = $crate::nist::arith::compute_r2($modulus);
const NEG_INV: u64 = $crate::nist::arith::compute_neg_inv($modulus[0]);
const fn mont_mul_raw(a: [u64; $limbs], b: [u64; $limbs]) -> [u64; $limbs] {
let mut t = [0u64; $crate::nist::arith::MAX_LIMBS + 2];
let mut i = 0;
while i < $limbs {
let mut carry = 0u128;
let mut j = 0;
while j < $limbs {
let sum = (t[j] as u128) + (a[j] as u128) * (b[i] as u128) + carry;
t[j] = sum as u64;
carry = sum >> 64;
j += 1;
}
let sum = (t[$limbs] as u128) + carry;
t[$limbs] = sum as u64;
t[$limbs + 1] = (sum >> 64) as u64;
let u = t[0].wrapping_mul(Self::NEG_INV);
let sum = (t[0] as u128) + (u as u128) * (Self::MODULUS[0] as u128);
let mut carry = sum >> 64;
let mut j = 1;
while j < $limbs {
let sum = (t[j] as u128) + (u as u128) * (Self::MODULUS[j] as u128) + carry;
t[j - 1] = sum as u64;
carry = sum >> 64;
j += 1;
}
let sum = (t[$limbs] as u128) + carry;
t[$limbs - 1] = sum as u64;
t[$limbs] = (t[$limbs + 1] as u128 + (sum >> 64)) as u64;
t[$limbs + 1] = 0;
i += 1;
}
let mut lo = [0u64; $limbs];
let mut k = 0;
while k < $limbs {
lo[k] = t[k];
k += 1;
}
let (reduced, borrow) = $crate::nist::arith::sbb(lo, Self::MODULUS);
let need = t[$limbs] | (1 - borrow);
$crate::nist::arith::select(need.wrapping_neg(), reduced, lo)
}
pub const fn to_mont(limbs: [u64; $limbs]) -> Self {
Self(Self::mont_mul_raw(limbs, Self::R2))
}
pub const fn from_mont(&self) -> [u64; $limbs] {
let mut one = [0u64; $limbs];
one[0] = 1;
Self::mont_mul_raw(self.0, one)
}
pub fn square_n(&self, n: usize) -> Self {
let mut r = *self;
for _ in 0..n {
r = <Self as $crate::nist::arith::Field>::square(&r);
}
r
}
pub fn pow(&self, exponent: &[u64; $limbs]) -> Self {
let mut result = <Self as $crate::nist::arith::Field>::ONE;
for i in (0..$limbs).rev() {
for bit in (0..64).rev() {
result = <Self as $crate::nist::arith::Field>::square(&result);
if (exponent[i] >> bit) & 1 == 1 {
result = <Self as $crate::nist::arith::Field>::mul(&result, self);
}
}
}
result
}
}
impl $crate::nist::arith::Field for $name {
type Bytes = [u8; $bytes];
const ZERO: Self = Self([0u64; $limbs]);
const ONE: Self = Self(Self::mont_mul_raw(
{
let mut one = [0u64; $limbs];
one[0] = 1;
one
},
Self::R2,
));
const BYTE_LEN: usize = $bytes;
#[inline]
fn add(&self, other: &Self) -> Self {
let (sum, carry) = $crate::nist::arith::adc(self.0, other.0);
let (reduced, borrow) = $crate::nist::arith::sbb(sum, Self::MODULUS);
let need = carry | (1 - borrow);
Self($crate::nist::arith::select(
need.wrapping_neg(),
reduced,
sum,
))
}
#[inline]
fn sub(&self, other: &Self) -> Self {
let (diff, borrow) = $crate::nist::arith::sbb(self.0, other.0);
let (fixed, _) = $crate::nist::arith::adc(diff, Self::MODULUS);
Self($crate::nist::arith::select(
borrow.wrapping_neg(),
fixed,
diff,
))
}
#[inline]
fn mul(&self, other: &Self) -> Self {
Self(Self::mont_mul_raw(self.0, other.0))
}
#[inline]
fn square(&self) -> Self {
Self(Self::mont_mul_raw(self.0, self.0))
}
#[inline]
fn double(&self) -> Self {
<Self as $crate::nist::arith::Field>::add(self, self)
}
#[inline]
fn triple(&self) -> Self {
let d = <Self as $crate::nist::arith::Field>::double(self);
<Self as $crate::nist::arith::Field>::add(&d, self)
}
#[inline]
fn neg(&self) -> Self {
<Self as $crate::nist::arith::Field>::sub(
&<Self as $crate::nist::arith::Field>::ZERO,
self,
)
}
fn invert(&self) -> Self {
let mut two = [0u64; $limbs];
two[0] = 2;
let (exp, _) = $crate::nist::arith::sbb(Self::MODULUS, two);
self.pow(&exp)
}
fn from_bytes(bytes: &Self::Bytes) -> Option<Self> {
let limbs: [u64; $limbs] = $crate::nist::arith::from_be_bytes(bytes.as_ref());
let (_, borrow) = $crate::nist::arith::sbb(limbs, Self::MODULUS);
if borrow == 0 {
return None;
}
Some(Self::to_mont(limbs))
}
fn from_bytes_reduced(bytes: &Self::Bytes) -> Self {
let limbs: [u64; $limbs] = $crate::nist::arith::from_be_bytes(bytes.as_ref());
let (reduced, borrow) = $crate::nist::arith::sbb(limbs, Self::MODULUS);
let limbs = $crate::nist::arith::select(borrow.wrapping_neg(), limbs, reduced);
Self::to_mont(limbs)
}
fn to_bytes(&self) -> Self::Bytes {
let limbs = self.from_mont();
let mut out = [0u8; $bytes];
$crate::nist::arith::to_be_bytes(&limbs, &mut out);
out
}
fn zero_bytes() -> Self::Bytes {
[0u8; $bytes]
}
#[inline]
fn is_zero(&self) -> ic_core::ct::Choice {
let mut acc = 0u64;
for limb in self.0.iter() {
acc |= *limb;
}
ic_core::ct::Choice::from_u8(((acc | acc.wrapping_neg()) >> 63) as u8).not()
}
#[inline]
fn ct_eq(&self, other: &Self) -> ic_core::ct::Choice {
let d = <Self as $crate::nist::arith::Field>::sub(self, other);
<Self as $crate::nist::arith::Field>::is_zero(&d)
}
#[inline]
fn cmov(a: &mut Self, b: &Self, choice: ic_core::ct::Choice) {
let mask = (choice.unwrap_u8() as u64).wrapping_neg();
a.0 = $crate::nist::arith::select(mask, b.0, a.0);
}
#[inline]
fn is_odd(&self) -> ic_core::ct::Choice {
ic_core::ct::Choice::from_u8((self.from_mont()[0] & 1) as u8)
}
}
};
}
pub fn sqrt_p3mod4<F: Field, const N: usize>(
x: &F,
modulus: [u64; N],
pow: impl Fn(&F, &[u64; N]) -> F,
) -> F {
let mut one = [0u64; N];
one[0] = 1;
let (sum, _) = adc(modulus, one);
let mut exp = [0u64; N];
let mut i = 0;
while i < N {
let lo = sum[i] >> 2;
let hi = if i + 1 < N { sum[i + 1] << 62 } else { 0 };
exp[i] = lo | hi;
i += 1;
}
pow(x, &exp)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn newton_inverse_satisfies_its_defining_equation() {
for m0 in [
0xffff_ffff_ffff_ffffu64,
0xf3b9_cac2_fc63_2551,
0x0000_0000_ffff_ffff,
0xecec_196a_ccc5_2973,
] {
assert_eq!(m0.wrapping_mul(compute_neg_inv(m0)), u64::MAX, "{m0:#x}");
}
}
#[test]
fn carry_and_borrow_propagate() {
let (sum, carry) = adc([u64::MAX, 0], [1u64, 0]);
assert_eq!(sum, [0, 1]);
assert_eq!(carry, 0);
let (sum, carry) = adc([u64::MAX, u64::MAX], [1u64, 0]);
assert_eq!(sum, [0, 0]);
assert_eq!(carry, 1);
let (diff, borrow) = sbb([0u64, 1], [1u64, 0]);
assert_eq!(diff, [u64::MAX, 0]);
assert_eq!(borrow, 0);
let (diff, borrow) = sbb([0u64, 0], [1u64, 0]);
assert_eq!(diff, [u64::MAX, u64::MAX]);
assert_eq!(borrow, 1);
}
#[test]
fn select_is_branch_free_and_correct() {
assert_eq!(select(u64::MAX, [1u64, 2], [3u64, 4]), [1, 2]);
assert_eq!(select(0, [1u64, 2], [3u64, 4]), [3, 4]);
}
}