use crate::internal::subtle::{Choice, ConstantTimeEq, CtOption};
pub(crate) const N: [u64; 4] = [
0xBFD2_5E8C_D036_4141,
0xBAAE_DCE6_AF48_A03B,
0xFFFF_FFFF_FFFF_FFFE,
0xFFFF_FFFF_FFFF_FFFF,
];
const R256: [u64; 4] = [
0x402da1732fc9bebf,
0x4551231950b75fc4,
0x0000000000000001,
0x0000000000000000,
];
const ONE: [u64; 4] = [1, 0, 0, 0];
const EXP_N2: [u64; 4] = [
0xBFD2_5E8C_D036_413F,
0xBAAE_DCE6_AF48_A03B,
0xFFFF_FFFF_FFFF_FFFE,
0xFFFF_FFFF_FFFF_FFFF,
];
#[derive(Clone, Copy)]
pub struct Scalar {
limbs: [u64; 4],
}
impl Scalar {
pub const ZERO: Self = Self {
limbs: [0, 0, 0, 0],
};
pub const ONE: Self = Self { limbs: ONE };
pub const fn from_limbs(limbs: [u64; 4]) -> Self {
Self { limbs }
}
pub fn from_repr(bytes: &[u8; 32]) -> CtOption<Self> {
let mut limbs = [0u64; 4];
for i in 0..4 {
let offset = i * 8;
limbs[3 - i] = u64::from_be_bytes([
bytes[offset],
bytes[offset + 1],
bytes[offset + 2],
bytes[offset + 3],
bytes[offset + 4],
bytes[offset + 5],
bytes[offset + 6],
bytes[offset + 7],
]);
}
let is_valid = is_lt_n(&limbs);
CtOption::new(Self::from_limbs(limbs), is_valid)
}
pub fn from_repr_reduced(bytes: &[u8; 32]) -> Self {
let mut limbs = [0u64; 4];
for i in 0..4 {
let offset = i * 8;
limbs[3 - i] = u64::from_be_bytes([
bytes[offset],
bytes[offset + 1],
bytes[offset + 2],
bytes[offset + 3],
bytes[offset + 4],
bytes[offset + 5],
bytes[offset + 6],
bytes[offset + 7],
]);
}
let val = Self::from_limbs(limbs);
if val.is_gte_n().is_true() {
val.sub(&Self::from_limbs(N))
} else {
val
}
}
pub fn to_bytes(&self) -> [u8; 32] {
let mut bytes = [0u8; 32];
for i in 0..4 {
let limb_bytes = self.limbs[3 - i].to_be_bytes();
let offset = i * 8;
bytes[offset..offset + 8].copy_from_slice(&limb_bytes);
}
bytes
}
pub fn mul(&self, rhs: &Self) -> Self {
let t = product_512(&self.limbs, &rhs.limbs);
Scalar::from_limbs(reduce_512(&t))
}
pub(crate) fn mul_raw(&self, rhs: &Self) -> Self {
self.mul(rhs)
}
pub fn add(&self, rhs: &Self) -> Self {
let mut result = [0u64; 4];
let mut carry = 0u64;
for i in 0..4 {
let (sum1, c1) = self.limbs[i].overflowing_add(rhs.limbs[i]);
let (sum2, c2) = sum1.overflowing_add(carry);
result[i] = sum2;
carry = u64::from(c1) | u64::from(c2);
}
let res = Self::from_limbs(result);
res.reduce_if_carry_or_gte_n(carry)
}
pub fn sub(&self, rhs: &Self) -> Self {
let mut result = [0u64; 4];
let mut borrow = 0u64;
for i in 0..4 {
let (diff1, b1) = self.limbs[i].overflowing_sub(rhs.limbs[i]);
let (diff2, b2) = diff1.overflowing_sub(borrow);
result[i] = diff2;
borrow = u64::from(b1) | u64::from(b2);
}
let res = Self::from_limbs(result);
res.conditional_add_n(borrow)
}
pub fn neg(&self) -> Self {
if self.is_zero().is_true() {
*self
} else {
Self::from_limbs(N).sub(self)
}
}
pub fn is_zero(&self) -> Choice {
self.ct_eq(&Self::ZERO)
}
pub fn invert(&self) -> CtOption<Self> {
let is_zero = self.ct_eq(&Self::ZERO);
let result = self.pow(&EXP_N2);
CtOption::new(result, !is_zero)
}
fn pow(&self, exp: &[u64; 4]) -> Self {
let mut result = Self::ONE;
let mut base = *self;
for word in exp.iter() {
for i in 0..64 {
if (word >> i) & 1 == 1 {
result = (&result).mul(&base);
}
base = (&base).mul(&base);
}
}
result
}
fn reduce_if_carry_or_gte_n(&self, carry: u64) -> Self {
let gte_n = self.is_gte_n() | Choice::from_bool(carry != 0);
let mut borrow = 0u64;
let mut sub_result = [0u64; 4];
for i in 0..4 {
let (n_plus_borrow, n_carry) = N[i].overflowing_add(borrow);
let (diff, b) = self.limbs[i].overflowing_sub(n_plus_borrow);
sub_result[i] = diff;
borrow = u64::from(b) | u64::from(n_carry);
}
let mut result = self.limbs;
for i in 0..4 {
if gte_n.is_true() {
result[i] = sub_result[i];
}
}
Self::from_limbs(result)
}
fn conditional_add_n(&self, borrow: u64) -> Self {
if borrow == 0 {
*self
} else {
let mut result = [0u64; 4];
let mut carry = 0u64;
for i in 0..4 {
let (sum1, c1) = self.limbs[i].overflowing_add(N[i]);
let (sum2, c2) = sum1.overflowing_add(carry);
result[i] = sum2;
carry = u64::from(c1) | u64::from(c2);
}
Self::from_limbs(result)
}
}
fn is_gte_n(&self) -> Choice {
for i in (0..4).rev() {
if self.limbs[i] > N[i] {
return Choice::from_bool(true);
} else if self.limbs[i] < N[i] {
return Choice::from_bool(false);
}
}
Choice::from_bool(true)
}
}
impl ConstantTimeEq for Scalar {
fn ct_eq(&self, other: &Self) -> Choice {
let mut result = 1u8;
for i in 0..4 {
result &= u8::from(self.limbs[i] == other.limbs[i]);
}
Choice(result)
}
}
fn product_512(a: &[u64; 4], b: &[u64; 4]) -> [u64; 8] {
let mut t = [0u64; 8];
for i in 0..4 {
let mut carry = 0u64;
for j in 0..4 {
let product = (a[i] as u128) * (b[j] as u128);
let sum = (t[i + j] as u128) + product + (carry as u128);
t[i + j] = sum as u64;
carry = (sum >> 64) as u64;
}
t[i + 4] = t[i + 4].wrapping_add(carry);
}
t
}
fn reduce_512(t: &[u64; 8]) -> [u64; 4] {
let mut limbs = *t;
loop {
let hi = [limbs[4], limbs[5], limbs[6], limbs[7]];
if hi == [0, 0, 0, 0] {
break;
}
let lo = Scalar::from_limbs([limbs[0], limbs[1], limbs[2], limbs[3]]);
let hi_s = Scalar::from_limbs(hi);
let hi_r = hi_s.mul(&Scalar::from_limbs(R256));
let mut w = [0u64; 5];
let mut c: u64 = 0;
for i in 0..4 {
let (s1, c1) = lo.limbs[i].overflowing_add(hi_r.limbs[i]);
let (s2, c2) = s1.overflowing_add(c);
w[i] = s2;
c = u64::from(c1) | u64::from(c2);
}
w[4] = c;
let red = reduce_257(&w);
limbs = [red[0], red[1], red[2], red[3], 0, 0, 0, 0];
}
reduce_257(&[limbs[0], limbs[1], limbs[2], limbs[3], 0])
}
fn reduce_257(w: &[u64; 5]) -> [u64; 4] {
let ge = if w[4] == 1 {
true
} else {
is_gte4(&[w[0], w[1], w[2], w[3]], &N)
};
if !ge {
return [w[0], w[1], w[2], w[3]];
}
let mut r = [0u64; 5];
let mut borrow = 0u64;
for i in 0..4 {
let (d1, b1) = w[i].overflowing_sub(N[i]);
let (d2, b2) = d1.overflowing_sub(borrow);
r[i] = d2;
borrow = u64::from(b1) | u64::from(b2);
}
let (d4, _) = w[4].overflowing_sub(borrow);
r[4] = d4;
if r[4] == 1 {
let mut r2 = [0u64; 5];
let mut borrow = 0u64;
for i in 0..4 {
let (d1, b1) = r[i].overflowing_sub(N[i]);
let (d2, b2) = d1.overflowing_sub(borrow);
r2[i] = d2;
borrow = u64::from(b1) | u64::from(b2);
}
let (d4b, _) = r[4].overflowing_sub(borrow);
r2[4] = d4b;
[r2[0], r2[1], r2[2], r2[3]]
} else {
[r[0], r[1], r[2], r[3]]
}
}
fn is_gte4(a: &[u64; 4], b: &[u64; 4]) -> bool {
for i in (0..4).rev() {
if a[i] > b[i] {
return true;
} else if a[i] < b[i] {
return false;
}
}
true
}
fn is_lt_n(limbs: &[u64; 4]) -> Choice {
for i in (0..4).rev() {
if limbs[i] > N[i] {
return Choice::from_bool(false);
} else if limbs[i] < N[i] {
return Choice::from_bool(true);
}
}
Choice::from_bool(false)
}
#[cfg(test)]
mod tests {
use super::*;
fn sc(v: u8) -> Scalar {
let mut b = [0u8; 32];
b[31] = v;
Scalar::from_repr(&b).unwrap()
}
#[test]
fn mul_small() {
assert_eq!(sc(2).mul(&sc(3)).to_bytes()[31], 6);
assert_eq!(sc(7).mul(&sc(7)).to_bytes()[31], 49);
}
#[test]
fn add_small() {
assert_eq!(sc(1).add(&sc(1)).to_bytes()[31], 2);
}
#[test]
fn from_repr_identity() {
assert_eq!(sc(1).to_bytes()[31], 1);
assert_eq!(Scalar::ONE.to_bytes()[31], 1);
}
#[test]
fn mul_distributes() {
let a = sc(7);
let b = sc(3);
let c = sc(5);
let lhs = a.add(&b).mul(&c);
let rhs = a.mul(&c).add(&b.mul(&c));
assert_eq!(lhs.to_bytes(), rhs.to_bytes());
}
#[test]
fn invert_roundtrip() {
let a = sc(123);
let inv = a.invert().unwrap();
assert_eq!(a.mul(&inv).to_bytes()[31], 1);
}
#[test]
fn mul_reference() {
let a = Scalar::from_limbs([
10499958131665514998,
14799178230035213023,
1164115433906158532,
2175216119781798972,
]);
let b = Scalar::from_limbs([
14037279428536751484,
8711387064946514083,
7002664860023442459,
3872982626502034966,
]);
let p = a.mul(&b);
assert_eq!(
p.limbs,
[
11082476127163626936,
7929653814064238854,
16743205666577460258,
10419521188029476923,
]
);
let inv = a.invert().unwrap();
assert_eq!(
inv.limbs,
[
7386074044270653371,
14125527112431059711,
3914782851617471626,
18026134145265468931,
]
);
assert_eq!(a.mul(&inv).to_bytes()[31], 1);
}
}