#![allow(clippy::needless_range_loop)]
use ic_core::ct::Choice;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Fe([u32; 10]);
const fn width(i: usize) -> u32 {
26 - (i as u32 & 1)
}
const TWO_P: [u32; 10] = [
0x7FF_FFDA, 0x3FF_FFFE, 0x7FF_FFFE, 0x3FF_FFFE, 0x7FF_FFFE, 0x3FF_FFFE, 0x7FF_FFFE, 0x3FF_FFFE,
0x7FF_FFFE, 0x3FF_FFFE,
];
impl Fe {
pub const ZERO: Fe = Fe([0; 10]);
pub const ONE: Fe = Fe([1, 0, 0, 0, 0, 0, 0, 0, 0, 0]);
pub const fn from_limbs51(l: [u64; 5]) -> Fe {
let mut bytes = [0u8; 32];
let mut i = 0;
while i < 32 {
let bit = 8 * i;
let limb = bit / 51;
let off = bit % 51;
let mut v = l[limb] >> off;
if off > 43 && limb < 4 {
v |= l[limb + 1] << (51 - off);
}
bytes[i] = v as u8;
i += 1;
}
Fe::from_bytes_const(&bytes)
}
#[cfg(test)]
pub const fn from_u64(v: u64) -> Fe {
let mut bytes = [0u8; 32];
let le = v.to_le_bytes();
let mut i = 0;
while i < 8 {
bytes[i] = le[i];
i += 1;
}
Fe::from_bytes_const(&bytes)
}
#[inline]
pub fn add(&self, other: &Fe) -> Fe {
let mut r = [0u32; 10];
for i in 0..10 {
r[i] = self.0[i] + other.0[i];
}
Fe(r).weak_reduce()
}
#[inline]
pub fn sub(&self, other: &Fe) -> Fe {
let mut r = [0u32; 10];
for i in 0..10 {
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 = 0u32;
for i in 0..10 {
r[i] += carry;
carry = r[i] >> width(i);
r[i] &= (1 << width(i)) - 1;
}
r[0] += 19 * carry;
let c = r[0] >> 26;
r[0] &= (1 << 26) - 1;
r[1] += c;
Fe(r)
}
#[inline]
pub fn mul(&self, other: &Fe) -> Fe {
let a = &self.0;
let b = &other.0;
let mut b19 = [0u32; 10];
for j in 0..10 {
b19[j] = 19 * b[j];
}
let mut a2 = [0u32; 10];
for i in 0..10 {
a2[i] = a[i] << (i & 1);
}
let z0 = m(a[0], b[0])
+ m(a2[1], b19[9])
+ m(a[2], b19[8])
+ m(a2[3], b19[7])
+ m(a[4], b19[6])
+ m(a2[5], b19[5])
+ m(a[6], b19[4])
+ m(a2[7], b19[3])
+ m(a[8], b19[2])
+ m(a2[9], b19[1]);
let z1 = m(a[0], b[1])
+ m(a[1], b[0])
+ m(a[2], b19[9])
+ m(a[3], b19[8])
+ m(a[4], b19[7])
+ m(a[5], b19[6])
+ m(a[6], b19[5])
+ m(a[7], b19[4])
+ m(a[8], b19[3])
+ m(a[9], b19[2]);
let z2 = m(a[0], b[2])
+ m(a2[1], b[1])
+ m(a[2], b[0])
+ m(a2[3], b19[9])
+ m(a[4], b19[8])
+ m(a2[5], b19[7])
+ m(a[6], b19[6])
+ m(a2[7], b19[5])
+ m(a[8], b19[4])
+ m(a2[9], b19[3]);
let z3 = m(a[0], b[3])
+ m(a[1], b[2])
+ m(a[2], b[1])
+ m(a[3], b[0])
+ m(a[4], b19[9])
+ m(a[5], b19[8])
+ m(a[6], b19[7])
+ m(a[7], b19[6])
+ m(a[8], b19[5])
+ m(a[9], b19[4]);
let z4 = m(a[0], b[4])
+ m(a2[1], b[3])
+ m(a[2], b[2])
+ m(a2[3], b[1])
+ m(a[4], b[0])
+ m(a2[5], b19[9])
+ m(a[6], b19[8])
+ m(a2[7], b19[7])
+ m(a[8], b19[6])
+ m(a2[9], b19[5]);
let z5 = m(a[0], b[5])
+ m(a[1], b[4])
+ m(a[2], b[3])
+ m(a[3], b[2])
+ m(a[4], b[1])
+ m(a[5], b[0])
+ m(a[6], b19[9])
+ m(a[7], b19[8])
+ m(a[8], b19[7])
+ m(a[9], b19[6]);
let z6 = m(a[0], b[6])
+ m(a2[1], b[5])
+ m(a[2], b[4])
+ m(a2[3], b[3])
+ m(a[4], b[2])
+ m(a2[5], b[1])
+ m(a[6], b[0])
+ m(a2[7], b19[9])
+ m(a[8], b19[8])
+ m(a2[9], b19[7]);
let z7 = m(a[0], b[7])
+ m(a[1], b[6])
+ m(a[2], b[5])
+ m(a[3], b[4])
+ m(a[4], b[3])
+ m(a[5], b[2])
+ m(a[6], b[1])
+ m(a[7], b[0])
+ m(a[8], b19[9])
+ m(a[9], b19[8]);
let z8 = m(a[0], b[8])
+ m(a2[1], b[7])
+ m(a[2], b[6])
+ m(a2[3], b[5])
+ m(a[4], b[4])
+ m(a2[5], b[3])
+ m(a[6], b[2])
+ m(a2[7], b[1])
+ m(a[8], b[0])
+ m(a2[9], b19[9]);
let z9 = m(a[0], b[9])
+ m(a[1], b[8])
+ m(a[2], b[7])
+ m(a[3], b[6])
+ m(a[4], b[5])
+ m(a[5], b[4])
+ m(a[6], b[3])
+ m(a[7], b[2])
+ m(a[8], b[1])
+ m(a[9], b[0]);
carry_reduce([z0, z1, z2, z3, z4, z5, z6, z7, z8, z9])
}
#[inline]
pub fn square(&self) -> Fe {
let a = &self.0;
let mut a19 = [0u32; 10];
for j in 0..10 {
a19[j] = 19 * a[j];
}
let mut d = [0u32; 10];
let mut q = [0u32; 10];
for i in 0..10 {
d[i] = a[i] << 1;
q[i] = a[i] << 2;
}
let z0 = m(a[0], a[0])
+ m(q[1], a19[9])
+ m(d[2], a19[8])
+ m(q[3], a19[7])
+ m(d[4], a19[6])
+ m(d[5], a19[5]);
let z1 =
m(d[0], a[1]) + m(d[2], a19[9]) + m(d[3], a19[8]) + m(d[4], a19[7]) + m(d[5], a19[6]);
let z2 = m(d[0], a[2])
+ m(d[1], a[1])
+ m(q[3], a19[9])
+ m(d[4], a19[8])
+ m(q[5], a19[7])
+ m(a[6], a19[6]);
let z3 =
m(d[0], a[3]) + m(d[1], a[2]) + m(d[4], a19[9]) + m(d[5], a19[8]) + m(d[6], a19[7]);
let z4 = m(d[0], a[4])
+ m(q[1], a[3])
+ m(a[2], a[2])
+ m(q[5], a19[9])
+ m(d[6], a19[8])
+ m(d[7], a19[7]);
let z5 = m(d[0], a[5]) + m(d[1], a[4]) + m(d[2], a[3]) + m(d[6], a19[9]) + m(d[7], a19[8]);
let z6 = m(d[0], a[6])
+ m(q[1], a[5])
+ m(d[2], a[4])
+ m(d[3], a[3])
+ m(q[7], a19[9])
+ m(a[8], a19[8]);
let z7 = m(d[0], a[7]) + m(d[1], a[6]) + m(d[2], a[5]) + m(d[3], a[4]) + m(d[8], a19[9]);
let z8 = m(d[0], a[8])
+ m(q[1], a[7])
+ m(d[2], a[6])
+ m(q[3], a[5])
+ m(a[4], a[4])
+ m(d[9], a19[9]);
let z9 = m(d[0], a[9]) + m(d[1], a[8]) + m(d[2], a[7]) + m(d[3], a[6]) + m(d[4], a[5]);
carry_reduce([z0, z1, z2, z3, z4, z5, z6, z7, z8, z9])
}
#[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 z = [0u64; 10];
for i in 0..10 {
z[i] = m(self.0[i], 121_666);
}
carry_reduce(z)
}
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 {
Fe::from_bytes_const(bytes)
}
const fn from_bytes_const(bytes: &[u8; 32]) -> Fe {
let mut r = [0u32; 10];
let mut i = 0;
let mut bit = 0usize;
while i < 10 {
let mut v = 0u64;
let mut k = 0;
while k < 5 {
let at = bit / 8 + k;
if at < 32 {
v |= (bytes[at] as u64) << (8 * k);
}
k += 1;
}
r[i] = ((v >> (bit % 8)) as u32) & ((1 << width(i)) - 1);
bit += width(i) as usize;
i += 1;
}
Fe(r)
}
pub fn to_bytes(self) -> [u8; 32] {
let mut t = self.weak_reduce().0;
let mut q = (t[0] + 19) >> 26;
for i in 1..10 {
q = (t[i] + q) >> width(i);
}
t[0] += 19 * q;
let mut carry = 0u32;
for i in 0..10 {
t[i] += carry;
carry = t[i] >> width(i);
t[i] &= (1 << width(i)) - 1;
}
let words: [u32; 8] = [
t[0] | t[1] << 26,
t[1] >> 6 | t[2] << 19,
t[2] >> 13 | t[3] << 13,
t[3] >> 19 | t[4] << 6,
t[5] | t[6] << 25,
t[6] >> 7 | t[7] << 19,
t[7] >> 13 | t[8] << 12,
t[8] >> 20 | t[9] << 6,
];
let mut out = [0u8; 32];
for (chunk, w) in out.chunks_exact_mut(4).zip(words.iter()) {
chunk.copy_from_slice(&w.to_le_bytes());
}
out
}
#[inline]
pub fn cswap(a: &mut Fe, b: &mut Fe, choice: Choice) {
let mask = (choice.unwrap_u8() as u32).wrapping_neg();
for i in 0..10 {
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 u32).wrapping_neg();
for i in 0..10 {
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: u32, y: u32) -> u64 {
(x as u64) * (y as u64)
}
#[inline]
fn carry_reduce(mut z: [u64; 10]) -> Fe {
let mut carry = 0u64;
for i in 0..10 {
z[i] += carry;
carry = z[i] >> width(i);
z[i] &= (1 << width(i)) - 1;
}
z[0] += 19 * carry;
let c = z[0] >> 26;
z[0] &= (1 << 26) - 1;
z[1] += c;
let mut r = [0u32; 10];
for i in 0..10 {
r[i] = z[i] as u32;
}
Fe(r)
}
#[cfg(test)]
mod tests {
use super::*;
const P: [u32; 10] = [
0x3FF_FFED, 0x1FF_FFFF, 0x3FF_FFFF, 0x1FF_FFFF, 0x3FF_FFFF, 0x1FF_FFFF, 0x3FF_FFFF,
0x1FF_FFFF, 0x3FF_FFFF, 0x1FF_FFFF,
];
fn bytes(seed: u32) -> [u8; 32] {
use ic_core::traits::Digest;
ic_hash::Sha256::digest(&seed.to_le_bytes())
}
fn operands() -> std::vec::Vec<Fe> {
let mut v: std::vec::Vec<Fe> = (0..24).map(|s| Fe::from_bytes(&bytes(s))).collect();
v.push(Fe::ZERO);
v.push(Fe::ONE);
v.push(Fe::ZERO.sub(&Fe::ONE));
let mut max = [0u32; 10];
for i in 0..10 {
max[i] = (1 << width(i)) - 1;
}
v.push(Fe(max));
max[1] += (1 << 18) - 1;
v.push(Fe(max));
v
}
#[cfg(not(any(target_arch = "riscv32", ic_limb32)))]
#[test]
fn agrees_with_the_five_limb_field() {
use crate::field::Fe as Fe51;
let conv = |x: &Fe| Fe51::from_bytes(&x.to_bytes());
let ops = operands();
let mut checked = 0;
for x in &ops {
let (x51, xb) = (conv(x), x.to_bytes());
assert_eq!(x51.to_bytes(), xb, "encoding");
assert_eq!(x.square().to_bytes(), x51.square().to_bytes(), "square");
assert_eq!(x.neg().to_bytes(), x51.neg().to_bytes(), "neg");
assert_eq!(x.mul121666().to_bytes(), x51.mul121666().to_bytes(), "a24");
assert_eq!(x.is_negative().unwrap_u8(), x51.is_negative().unwrap_u8());
assert_eq!(x.is_zero().unwrap_u8(), x51.is_zero().unwrap_u8());
for y in &ops {
let y51 = conv(y);
assert_eq!(x.mul(y).to_bytes(), x51.mul(&y51).to_bytes(), "mul");
assert_eq!(x.add(y).to_bytes(), x51.add(&y51).to_bytes(), "add");
assert_eq!(x.sub(y).to_bytes(), x51.sub(&y51).to_bytes(), "sub");
checked += 1;
}
}
for x in ops.iter().take(6) {
assert_eq!(x.invert().to_bytes(), conv(x).invert().to_bytes());
assert_eq!(x.pow22523().to_bytes(), conv(x).pow22523().to_bytes());
}
assert_eq!(checked, 29 * 29);
}
#[test]
fn p_encodes_as_zero_and_p_minus_one_does_not() {
assert_eq!(Fe(P).to_bytes(), [0u8; 32]);
let mut pm1 = P;
pm1[0] -= 1;
let mut want = [0xFFu8; 32];
want[0] = 0xEC;
want[31] = 0x7F;
assert_eq!(Fe(pm1).to_bytes(), want);
assert_eq!(Fe::ZERO.sub(&Fe::ONE).to_bytes(), want);
}
#[test]
fn encoding_round_trips_and_ignores_the_top_bit() {
for s in 0..64 {
let mut b = bytes(s);
b[31] &= 0x7F;
if b[31] == 0x7F && b[1..31].iter().all(|&x| x == 0xFF) && b[0] >= 0xED {
continue;
}
assert_eq!(Fe::from_bytes(&b).to_bytes(), b);
let mut hi = b;
hi[31] |= 0x80;
assert_eq!(Fe::from_bytes(&hi), Fe::from_bytes(&b));
}
}
#[test]
fn field_axioms_hold() {
let ops = operands();
for x in &ops {
let inv = x.invert();
let want = if bool::from(x.is_zero()) {
Fe::ZERO
} else {
Fe::ONE
};
assert!(bool::from(x.mul(&inv).ct_eq(&want)), "inverse");
assert!(bool::from(x.square().ct_eq(&x.mul(x))), "square");
assert!(bool::from(x.add(&x.neg()).is_zero()), "x + -x");
assert!(
bool::from(x.mul121666().ct_eq(&x.mul(&Fe::from_u64(121_666)))),
"a24"
);
for y in &ops {
assert!(bool::from(x.mul(y).ct_eq(&y.mul(x))), "commutative");
assert!(bool::from(x.sub(y).add(y).ct_eq(x)), "x - y + y");
let z = x.add(y);
assert!(
bool::from(z.mul(x).ct_eq(&x.square().add(&y.mul(x)))),
"distributive"
);
}
}
}
#[test]
fn the_radix_is_what_the_layout_assumes() {
let two_to = |k: u32| {
let mut b = [0u8; 32];
b[(k / 8) as usize] = 1 << (k % 8);
Fe::from_bytes(&b)
};
assert!(bool::from(
two_to(128).mul(&two_to(127)).ct_eq(&Fe::from_u64(19))
));
assert!(bool::from(two_to(26).mul(&two_to(77)).ct_eq(&two_to(103))));
assert!(bool::from(
two_to(230)
.square()
.ct_eq(&two_to(205).mul(&Fe::from_u64(19)))
));
}
#[test]
fn cswap_and_cmov_are_conditional() {
let (x, y) = (Fe::from_bytes(&bytes(1)), Fe::from_bytes(&bytes(2)));
let (mut a, mut b) = (x, y);
Fe::cswap(&mut a, &mut b, Choice::from_u8(0));
assert_eq!((a, b), (x, y));
Fe::cswap(&mut a, &mut b, Choice::from_u8(1));
assert_eq!((a, b), (y, x));
let mut c = x;
Fe::cmov(&mut c, &y, Choice::from_u8(0));
assert_eq!(c, x);
Fe::cmov(&mut c, &y, Choice::from_u8(1));
assert_eq!(c, y);
}
#[test]
fn from_limbs51_matches_decoding() {
let x = Fe::from_limbs51([(1 << 51) - 1, 1, 0, 0, 0]);
assert!(bool::from(x.ct_eq(&Fe::from_u64((1 << 52) - 1))));
for s in 0..16 {
let b = Fe::from_bytes(&bytes(s)).to_bytes();
let mut l = [0u64; 5];
for (k, limb) in l.iter_mut().enumerate() {
for bit in 0..51 {
let at = 51 * k + bit;
if at < 256 && (b[at / 8] >> (at % 8)) & 1 == 1 {
*limb |= 1 << bit;
}
}
}
assert_eq!(Fe::from_limbs51(l).to_bytes(), b);
}
}
}