use ic_core::ct::Choice;
pub const MAX_LIMBS: usize = 9;
#[cfg(not(any(target_arch = "riscv32", ic_limb32)))]
pub(crate) use wide::{adc, mont_mul, sbb};
#[cfg(any(target_arch = "riscv32", ic_limb32))]
pub(crate) use narrow::{adc, mont_mul, sbb};
#[cfg(any(test, not(any(target_arch = "riscv32", ic_limb32))))]
pub(crate) mod wide {
use super::MAX_LIMBS;
#[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(always)]
pub(crate) const fn mont_mul<const N: usize>(
a: [u64; N],
b: [u64; N],
m: [u64; N],
neg_inv: u64,
) -> ([u64; N], [u64; N], u64) {
let mut t = [0u64; MAX_LIMBS + 2];
let mut i = 0;
while i < N {
let mut carry = 0u128;
let mut j = 0;
while j < N {
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[N] as u128) + carry;
t[N] = sum as u64;
t[N + 1] = (sum >> 64) as u64;
let u = t[0].wrapping_mul(neg_inv);
let sum = (t[0] as u128) + (u as u128) * (m[0] as u128);
let mut carry = sum >> 64;
let mut j = 1;
while j < N {
let sum = (t[j] as u128) + (u as u128) * (m[j] as u128) + carry;
t[j - 1] = sum as u64;
carry = sum >> 64;
j += 1;
}
let sum = (t[N] as u128) + carry;
t[N - 1] = sum as u64;
t[N] = (t[N + 1] as u128 + (sum >> 64)) as u64;
t[N + 1] = 0;
i += 1;
}
let mut lo = [0u64; N];
let mut k = 0;
while k < N {
lo[k] = t[k];
k += 1;
}
let (reduced, borrow) = sbb(lo, m);
(reduced, lo, t[N] | (1 - borrow))
}
}
#[cfg(any(test, target_arch = "riscv32", ic_limb32))]
pub(crate) mod narrow {
use super::MAX_LIMBS;
#[inline(always)]
const fn words<const N: usize>(x: &[u64; N]) -> [u32; 2 * MAX_LIMBS] {
let mut w = [0u32; 2 * MAX_LIMBS];
let mut i = 0;
while i < N {
w[2 * i] = x[i] as u32;
w[2 * i + 1] = (x[i] >> 32) as u32;
i += 1;
}
w
}
#[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 = 0u64;
let mut i = 0;
while i < N {
let lo = (a[i] as u32 as u64) + (b[i] as u32 as u64) + carry;
let hi = (a[i] >> 32) + (b[i] >> 32) + (lo >> 32);
out[i] = (lo as u32 as u64) | (hi << 32);
carry = hi >> 32;
i += 1;
}
(out, carry)
}
#[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 = 0u64;
let mut i = 0;
while i < N {
let lo = (a[i] as u32 as u64)
.wrapping_sub(b[i] as u32 as u64)
.wrapping_sub(borrow);
let hi = (a[i] >> 32).wrapping_sub(b[i] >> 32).wrapping_sub(lo >> 63);
out[i] = (lo as u32 as u64) | (hi << 32);
borrow = hi >> 63;
i += 1;
}
(out, borrow)
}
#[inline]
pub(crate) const fn mont_mul<const N: usize>(
a: [u64; N],
b: [u64; N],
m: [u64; N],
neg_inv: u64,
) -> ([u64; N], [u64; N], u64) {
let m_limbs = m;
let neg_inv = neg_inv as u32;
let n = 2 * N;
let (a, m) = (words(&a), words(&m));
let b = words(&b);
let mut t = [0u32; 2 * MAX_LIMBS + 2];
let mut i = 0;
while i < n {
let bi = b[i] as u64;
let mut carry = 0u64;
let mut j = 0;
while j < n {
let sum = (t[j] as u64) + (a[j] as u64) * bi + carry;
t[j] = sum as u32;
carry = sum >> 32;
j += 1;
}
let sum = (t[n] as u64) + carry;
t[n] = sum as u32;
t[n + 1] = (sum >> 32) as u32;
let u = t[0].wrapping_mul(neg_inv) as u64;
let sum = (t[0] as u64) + u * (m[0] as u64);
let mut carry = sum >> 32;
let mut j = 1;
while j < n {
let sum = (t[j] as u64) + u * (m[j] as u64) + carry;
t[j - 1] = sum as u32;
carry = sum >> 32;
j += 1;
}
let sum = (t[n] as u64) + carry;
t[n - 1] = sum as u32;
t[n] = t[n + 1] + (sum >> 32) as u32;
t[n + 1] = 0;
i += 1;
}
let mut lo = [0u64; N];
let mut k = 0;
while k < N {
lo[k] = (t[2 * k] as u64) | ((t[2 * k + 1] as u64) << 32);
k += 1;
}
let (reduced, borrow) = sbb(lo, m_limbs);
(reduced, lo, t[n] as u64 | (1 - borrow))
}
}
#[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]
pub(crate) fn select_ct<const N: usize>(mask: u64, a: [u64; N], b: [u64; N]) -> [u64; N] {
select(core::hint::black_box(mask), a, b)
}
#[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]);
#[inline(always)]
const fn mont_mul_parts(
a: [u64; $limbs],
b: [u64; $limbs],
) -> ([u64; $limbs], [u64; $limbs], u64) {
$crate::nist::arith::mont_mul(a, b, Self::MODULUS, Self::NEG_INV)
}
#[inline]
fn mont_mul_raw(a: [u64; $limbs], b: [u64; $limbs]) -> [u64; $limbs] {
let (reduced, lo, need) = Self::mont_mul_parts(a, b);
$crate::nist::arith::select_ct(need.wrapping_neg(), reduced, lo)
}
const fn mont_mul_const(a: [u64; $limbs], b: [u64; $limbs]) -> [u64; $limbs] {
let (reduced, lo, need) = Self::mont_mul_parts(a, b);
$crate::nist::arith::select(need.wrapping_neg(), reduced, lo)
}
pub fn to_mont(limbs: [u64; $limbs]) -> Self {
Self(Self::mont_mul_raw(limbs, Self::R2))
}
pub const fn to_mont_const(limbs: [u64; $limbs]) -> Self {
Self(Self::mont_mul_const(limbs, Self::R2))
}
pub 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_const(
{
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_ct(
need.wrapping_neg(),
reduced,
sum,
))
}
#[inline]
fn sub(&self, other: &Self) -> Self {
let (diff, borrow) = $crate::nist::arith::sbb(self.0, other.0);
let mask = core::hint::black_box(borrow.wrapping_neg());
let mut m = Self::MODULUS;
for limb in m.iter_mut() {
*limb &= mask;
}
let (fixed, _) = $crate::nist::arith::adc(diff, m);
Self(fixed)
}
#[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_ct(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_ct(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 narrow_words_agree_with_wide() {
use crate::nist::point::Curve;
fn case<const N: usize>(m: [u64; N], seed: u64) -> usize {
let neg_inv = compute_neg_inv(m[0]);
let mut state = seed;
let mut next = || {
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 ops = std::vec::Vec::new();
let mut one = [0u64; N];
one[0] = 1;
ops.push([0u64; N]);
ops.push(one);
ops.push(wide::sbb(m, one).0);
ops.push([u64::MAX; N]);
for _ in 0..12 {
let mut x = [0u64; N];
for limb in x.iter_mut() {
*limb = next();
}
x[N - 1] %= m[N - 1].max(1);
ops.push(x);
}
let mut checked = 0;
for a in &ops {
for b in &ops {
assert_eq!(wide::adc(*a, *b), narrow::adc(*a, *b), "adc");
assert_eq!(wide::sbb(*a, *b), narrow::sbb(*a, *b), "sbb");
assert_eq!(
wide::mont_mul(*a, *b, m, neg_inv),
narrow::mont_mul(*a, *b, m, neg_inv),
"mont_mul"
);
checked += 1;
}
}
checked
}
let checked = case(<crate::p256::P256 as Curve>::Field::MODULUS, 1)
+ case(<crate::p256::P256 as Curve>::Scalar::MODULUS, 2)
+ case(<crate::p384::P384 as Curve>::Field::MODULUS, 3)
+ case(<crate::p384::P384 as Curve>::Scalar::MODULUS, 4)
+ case(<crate::p521::P521 as Curve>::Field::MODULUS, 5)
+ case(<crate::p521::P521 as Curve>::Scalar::MODULUS, 6);
assert_eq!(checked, 6 * 16 * 16);
}
#[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]);
}
}