#![allow(clippy::needless_range_loop)]
use ic_core::ct::Choice;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Fe(pub [u64; 5]);
const MASK: u64 = (1 << 51) - 1;
#[allow(clippy::unusual_byte_groupings)]
const TWO_P: [u64; 5] = [
0xFFFFFFFFFFFDA,
0xFFFFFFFFFFFFE,
0xFFFFFFFFFFFFE,
0xFFFFFFFFFFFFE,
0xFFFFFFFFFFFFE,
];
impl Fe {
pub const ZERO: Fe = Fe([0, 0, 0, 0, 0]);
pub const ONE: Fe = Fe([1, 0, 0, 0, 0]);
pub const fn from_u64(v: u64) -> Fe {
Fe([v & MASK, v >> 51, 0, 0, 0])
}
#[inline]
pub fn add(&self, other: &Fe) -> Fe {
let mut r = [0u64; 5];
for i in 0..5 {
r[i] = self.0[i] + other.0[i];
}
Fe(r)
}
#[inline]
pub fn sub(&self, other: &Fe) -> Fe {
let mut r = [0u64; 5];
for i in 0..5 {
r[i] = self.0[i] + TWO_P[i] - other.0[i];
}
Fe(r).weak_reduce()
}
#[inline]
pub fn neg(&self) -> Fe {
Fe::ZERO.sub(self)
}
#[inline]
fn weak_reduce(self) -> Fe {
let mut r = self.0;
let mut carry = r[0] >> 51;
r[0] &= MASK;
for i in 1..5 {
r[i] += carry;
carry = r[i] >> 51;
r[i] &= MASK;
}
r[0] += carry.wrapping_mul(19);
Fe(r)
}
#[inline]
pub fn mul(&self, other: &Fe) -> Fe {
let a = &self.0;
let b = &other.0;
let b1_19 = b[1] * 19;
let b2_19 = b[2] * 19;
let b3_19 = b[3] * 19;
let b4_19 = b[4] * 19;
let r0 = m(a[0], b[0]) + m(a[1], b4_19) + m(a[2], b3_19) + m(a[3], b2_19) + m(a[4], b1_19);
let r1 = m(a[0], b[1]) + m(a[1], b[0]) + m(a[2], b4_19) + m(a[3], b3_19) + m(a[4], b2_19);
let r2 = m(a[0], b[2]) + m(a[1], b[1]) + m(a[2], b[0]) + m(a[3], b4_19) + m(a[4], b3_19);
let r3 = m(a[0], b[3]) + m(a[1], b[2]) + m(a[2], b[1]) + m(a[3], b[0]) + m(a[4], b4_19);
let r4 = m(a[0], b[4]) + m(a[1], b[3]) + m(a[2], b[2]) + m(a[3], b[1]) + m(a[4], b[0]);
carry_reduce([r0, r1, r2, r3, r4])
}
#[inline]
pub fn square(&self) -> Fe {
let a = &self.0;
let a0_2 = a[0] * 2;
let a1_2 = a[1] * 2;
let a1_38 = a[1] * 38;
let a2_38 = a[2] * 38;
let a3_38 = a[3] * 38;
let a3_19 = a[3] * 19;
let a4_19 = a[4] * 19;
let r0 = m(a[0], a[0]) + m(a1_38, a[4]) + m(a2_38, a[3]);
let r1 = m(a0_2, a[1]) + m(a2_38, a[4]) + m(a3_19, a[3]);
let r2 = m(a0_2, a[2]) + m(a[1], a[1]) + m(a3_38, a[4]);
let r3 = m(a0_2, a[3]) + m(a1_2, a[2]) + m(a4_19, a[4]);
let r4 = m(a0_2, a[4]) + m(a1_2, a[3]) + m(a[2], a[2]);
carry_reduce([r0, r1, r2, r3, r4])
}
#[inline]
pub fn square_n(&self, n: usize) -> Fe {
let mut r = *self;
for _ in 0..n {
r = r.square();
}
r
}
#[inline]
pub fn mul121666(&self) -> Fe {
let mut r = [0u128; 5];
for i in 0..5 {
r[i] = (self.0[i] as u128) * 121_666;
}
carry_reduce(r)
}
pub fn invert(&self) -> Fe {
let z2 = self.square();
let z9 = z2.square_n(2).mul(self);
let z11 = z9.mul(&z2);
let z2_5_0 = z11.square().mul(&z9);
let z2_10_0 = z2_5_0.square_n(5).mul(&z2_5_0);
let z2_20_0 = z2_10_0.square_n(10).mul(&z2_10_0);
let z2_40_0 = z2_20_0.square_n(20).mul(&z2_20_0);
let z2_50_0 = z2_40_0.square_n(10).mul(&z2_10_0);
let z2_100_0 = z2_50_0.square_n(50).mul(&z2_50_0);
let z2_200_0 = z2_100_0.square_n(100).mul(&z2_100_0);
let z2_250_0 = z2_200_0.square_n(50).mul(&z2_50_0);
z2_250_0.square_n(5).mul(&z11)
}
pub fn pow22523(&self) -> Fe {
let z2 = self.square();
let z9 = z2.square_n(2).mul(self);
let z11 = z9.mul(&z2);
let z2_5_0 = z11.square().mul(&z9);
let z2_10_0 = z2_5_0.square_n(5).mul(&z2_5_0);
let z2_20_0 = z2_10_0.square_n(10).mul(&z2_10_0);
let z2_40_0 = z2_20_0.square_n(20).mul(&z2_20_0);
let z2_50_0 = z2_40_0.square_n(10).mul(&z2_10_0);
let z2_100_0 = z2_50_0.square_n(50).mul(&z2_50_0);
let z2_200_0 = z2_100_0.square_n(100).mul(&z2_100_0);
let z2_250_0 = z2_200_0.square_n(50).mul(&z2_50_0);
z2_250_0.square_n(2).mul(self)
}
pub fn from_bytes(bytes: &[u8; 32]) -> Fe {
let load = |i: usize| -> u64 {
let mut v = [0u8; 8];
v.copy_from_slice(&bytes[i..i + 8]);
u64::from_le_bytes(v)
};
let l0 = load(0) & MASK;
let l1 = (load(6) >> 3) & MASK;
let l2 = (load(12) >> 6) & MASK;
let l3 = (load(19) >> 1) & MASK;
let l4 = (load(24) >> 12) & MASK;
Fe([l0, l1, l2, l3, l4])
}
pub fn to_bytes(&self) -> [u8; 32] {
let mut t = self.weak_reduce().weak_reduce().weak_reduce().0;
let mut q = (t[0] + 19) >> 51;
for i in 1..5 {
q = (t[i] + q) >> 51;
}
t[0] += 19 * q;
let mut carry = t[0] >> 51;
t[0] &= MASK;
for i in 1..5 {
t[i] += carry;
carry = t[i] >> 51;
t[i] &= MASK;
}
t[4] &= (1 << 51) - 1;
let mut out = [0u8; 32];
let mut acc: u128 = 0;
let mut acc_bits = 0usize;
let mut idx = 0usize;
for limb in t.iter() {
acc |= (*limb as u128) << acc_bits;
acc_bits += 51;
while acc_bits >= 8 && idx < 32 {
out[idx] = acc as u8;
acc >>= 8;
acc_bits -= 8;
idx += 1;
}
}
while idx < 32 {
out[idx] = acc as u8;
acc >>= 8;
idx += 1;
}
out
}
#[inline]
pub fn cswap(a: &mut Fe, b: &mut Fe, choice: Choice) {
let mask = (choice.unwrap_u8() as u64).wrapping_neg();
for i in 0..5 {
let t = mask & (a.0[i] ^ b.0[i]);
a.0[i] ^= t;
b.0[i] ^= t;
}
}
#[inline]
pub fn cmov(a: &mut Fe, b: &Fe, choice: Choice) {
let mask = (choice.unwrap_u8() as u64).wrapping_neg();
for i in 0..5 {
a.0[i] ^= mask & (a.0[i] ^ b.0[i]);
}
}
pub fn is_zero(&self) -> Choice {
ic_core::ct::is_zero(&self.to_bytes())
}
pub fn ct_eq(&self, other: &Fe) -> Choice {
ic_core::ct::eq(&self.to_bytes(), &other.to_bytes())
}
pub fn is_negative(&self) -> Choice {
Choice::from_u8(self.to_bytes()[0] & 1)
}
}
#[inline(always)]
fn m(x: u64, y: u64) -> u128 {
(x as u128) * (y as u128)
}
#[inline]
fn carry_reduce(r: [u128; 5]) -> Fe {
let c: [u64; 5] = [
(r[0] >> 51) as u64,
(r[1] >> 51) as u64,
(r[2] >> 51) as u64,
(r[3] >> 51) as u64,
(r[4] >> 51) as u64,
];
let mut out: [u64; 5] = [
(r[0] as u64 & MASK) + c[4] * 19,
(r[1] as u64 & MASK) + c[0],
(r[2] as u64 & MASK) + c[1],
(r[3] as u64 & MASK) + c[2],
(r[4] as u64 & MASK) + c[3],
];
let mut carry = out[0] >> 51;
out[0] &= MASK;
for slot in out.iter_mut().skip(1) {
*slot += carry;
carry = *slot >> 51;
*slot &= MASK;
}
out[0] += carry * 19;
Fe(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn squaring_agrees_with_multiplication() {
let mut cases = std::vec![
Fe::ZERO,
Fe::ONE,
Fe([1, 1, 1, 1, 1]),
Fe([(1u64 << 51) - 1; 5]),
Fe([(1u64 << 51) - 1, 0, (1u64 << 51) - 1, 0, (1u64 << 51) - 1]),
Fe([0, (1u64 << 51) - 1, 0, (1u64 << 51) - 1, 0]),
];
let mut x = Fe([0x51a2, 0x9e37, 0x79b9, 0x7f4a, 0x7c15]);
for _ in 0..16 {
x = x.mul(&Fe([3, 5, 7, 11, 13])).add(&Fe::ONE);
cases.push(x);
}
let mut checked = 0;
for f in &cases {
assert_eq!(
f.square().to_bytes(),
f.mul(f).to_bytes(),
"square and mul-by-self differ"
);
checked += 1;
}
assert_eq!(checked, 22, "the comparison did not run");
}
fn fe(v: u64) -> Fe {
Fe::from_u64(v)
}
#[test]
fn encode_decode_roundtrip() {
for v in [0u64, 1, 2, 19, 1 << 51, u64::MAX] {
let a = fe(v);
assert_eq!(Fe::from_bytes(&a.to_bytes()).to_bytes(), a.to_bytes());
}
}
#[test]
fn small_arithmetic() {
assert_eq!(fe(2).add(&fe(3)).to_bytes(), fe(5).to_bytes());
assert_eq!(fe(5).sub(&fe(3)).to_bytes(), fe(2).to_bytes());
assert_eq!(fe(6).mul(&fe(7)).to_bytes(), fe(42).to_bytes());
assert_eq!(fe(9).square().to_bytes(), fe(81).to_bytes());
}
#[test]
fn subtraction_wraps_into_the_field() {
let r = Fe::ZERO.sub(&Fe::ONE).to_bytes();
assert_eq!(r[0], 0xec);
assert_eq!(r[31], 0x7f);
for b in &r[1..31] {
assert_eq!(*b, 0xff);
}
}
#[test]
fn p_encodes_as_zero() {
let mut p_bytes = [0xffu8; 32];
p_bytes[0] = 0xed;
p_bytes[31] = 0x7f;
assert_eq!(Fe::from_bytes(&p_bytes).to_bytes(), [0u8; 32]);
}
#[test]
fn inversion_is_correct() {
for v in [1u64, 2, 3, 19, 12345, u32::MAX as u64] {
let a = fe(v);
assert_eq!(a.mul(&a.invert()).to_bytes(), Fe::ONE.to_bytes(), "1/{v}");
}
assert_eq!(Fe::ZERO.invert().to_bytes(), [0u8; 32]);
}
#[test]
fn multiplication_is_associative_and_distributive() {
let a = Fe::from_bytes(&[0x11; 32]);
let b = Fe::from_bytes(&[0x7a; 32]);
let c = Fe::from_bytes(&[0xc3; 32]);
assert_eq!(a.mul(&b).mul(&c).to_bytes(), a.mul(&b.mul(&c)).to_bytes());
assert_eq!(
a.mul(&b.add(&c)).to_bytes(),
a.mul(&b).add(&a.mul(&c)).to_bytes()
);
}
#[test]
fn pow22523_gives_a_square_root() {
let x = fe(4);
let r = x.pow22523().mul(&x);
let sq = r.square();
assert!(
bool::from(sq.ct_eq(&x)) || bool::from(sq.ct_eq(&x.neg())),
"square root property"
);
}
#[test]
fn cswap_and_cmov_are_conditional() {
let mut a = fe(1);
let mut b = fe(2);
Fe::cswap(&mut a, &mut b, Choice::FALSE);
assert_eq!(a.to_bytes(), fe(1).to_bytes());
Fe::cswap(&mut a, &mut b, Choice::TRUE);
assert_eq!(a.to_bytes(), fe(2).to_bytes());
let mut c = fe(5);
Fe::cmov(&mut c, &fe(9), Choice::FALSE);
assert_eq!(c.to_bytes(), fe(5).to_bytes());
Fe::cmov(&mut c, &fe(9), Choice::TRUE);
assert_eq!(c.to_bytes(), fe(9).to_bytes());
}
#[test]
fn high_bit_of_input_is_ignored() {
let mut a = [0x42u8; 32];
let mut b = a;
a[31] &= 0x7f;
b[31] |= 0x80;
assert_eq!(Fe::from_bytes(&a).to_bytes(), Fe::from_bytes(&b).to_bytes());
}
fn carry_reduce_serial(r: [u128; 5]) -> Fe {
let mut out = [0u64; 5];
let mut carry: u128 = 0;
for (slot, limb) in out.iter_mut().zip(r) {
let v = limb + carry;
carry = v >> 51;
*slot = (v & MASK as u128) as u64;
}
out[0] += (carry as u64) * 19;
let mut c = out[0] >> 51;
out[0] &= MASK;
for slot in out.iter_mut().skip(1) {
*slot += c;
c = *slot >> 51;
*slot &= MASK;
}
out[0] += c * 19;
Fe(out)
}
fn raw_products(a: &[u64; 5], b: &[u64; 5]) -> [u128; 5] {
let m = |x: u64, y: u64| (x as u128) * (y as u128);
let (b1, b2, b3, b4) = (b[1] * 19, b[2] * 19, b[3] * 19, b[4] * 19);
[
m(a[0], b[0]) + m(a[1], b4) + m(a[2], b3) + m(a[3], b2) + m(a[4], b1),
m(a[0], b[1]) + m(a[1], b[0]) + m(a[2], b4) + m(a[3], b3) + m(a[4], b2),
m(a[0], b[2]) + m(a[1], b[1]) + m(a[2], b[0]) + m(a[3], b4) + m(a[4], b3),
m(a[0], b[3]) + m(a[1], b[2]) + m(a[2], b[1]) + m(a[3], b[0]) + m(a[4], b4),
m(a[0], b[4]) + m(a[1], b[3]) + m(a[2], b[2]) + m(a[3], b[1]) + m(a[4], b[0]),
]
}
#[test]
fn limbs_at_their_maximum_do_not_carry_out_of_a_u64() {
let max = [(1u64 << 52) - 2; 5];
let r = raw_products(&max, &max);
assert_eq!(
carry_reduce(r).to_bytes(),
carry_reduce_serial(r).to_bytes(),
"parallel and serial carry disagree at the limb maximum"
);
}
#[test]
fn parallel_carry_agrees_with_the_serial_one() {
let mut state = 0x243f_6a88_85a3_08d3u64;
let mut next = || {
state ^= state >> 12;
state ^= state << 25;
state ^= state >> 27;
state.wrapping_mul(0x2545_f491_4f6c_dd1d)
};
for _ in 0..20_000 {
let mut a = [0u64; 5];
let mut b = [0u64; 5];
for i in 0..5 {
a[i] = next() % (1 << 52);
b[i] = next() % (1 << 52);
}
let r = raw_products(&a, &b);
assert_eq!(
carry_reduce(r).to_bytes(),
carry_reduce_serial(r).to_bytes(),
"disagreement on a={a:?} b={b:?}"
);
}
}
}