use super::curve163;
use zeroize::Zeroize;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Zeroize)]
pub struct Scalar([u64; 3]);
impl Scalar {
#[must_use]
pub fn from_be_bytes(bytes: &[u8]) -> Self {
Scalar(limbs_from_be_bytes(bytes))
}
#[must_use]
#[allow(clippy::cast_possible_truncation)] pub fn to_be_bytes(self) -> [u8; 21] {
let mut out = [0u8; 21];
for (i, byte) in out.iter_mut().rev().enumerate() {
let limb = i / 8;
let shift = (i % 8) * 8;
*byte = (self.0[limb] >> shift) as u8;
}
out
}
#[must_use]
pub fn is_zero(self) -> bool {
self.0 == [0, 0, 0]
}
fn n() -> [u64; 3] {
limbs_from_be_bytes(&curve163::order())
}
#[cfg(any(feature = "std", feature = "getrandom"))]
#[must_use]
pub(crate) fn from_candidate_bytes(bytes: &[u8; 21]) -> Option<Self> {
let limbs = limbs_from_be_bytes(bytes);
let (_, borrow) = sub3(limbs, Self::n());
let in_range = borrow == 1; let nonzero = limbs != [0, 0, 0];
(in_range && nonzero).then_some(Scalar(limbs))
}
#[must_use]
pub fn multiply(self, other: Self) -> Self {
Scalar(reduce_mod_n(mul3(self.0, other.0)))
}
#[must_use]
pub(crate) fn reduce_wide_bytes(bytes: &[u8]) -> Self {
let n = Self::n();
let mut r = [0u64; 3];
for &byte in bytes {
for bit in (0..8).rev() {
let bit_val = u64::from((byte >> bit) & 1);
r = shl1_or(r, bit_val);
r = cond_sub_if_ge(r, n);
}
}
Scalar(r)
}
}
impl core::ops::Add for Scalar {
type Output = Self;
fn add(self, other: Self) -> Self {
let (sum, _carry) = add3(self.0, other.0);
Scalar(cond_sub_if_ge(sum, Self::n()))
}
}
fn limbs_from_be_bytes(bytes: &[u8]) -> [u64; 3] {
let mut limbs = [0u64; 3];
for (i, &byte) in bytes.iter().rev().enumerate() {
let limb = i / 8;
let shift = (i % 8) * 8;
limbs[limb] |= u64::from(byte) << shift;
}
limbs
}
fn add3(a: [u64; 3], b: [u64; 3]) -> ([u64; 3], u64) {
let mut out = [0u64; 3];
let mut carry = 0u64;
for i in 0..3 {
let (s1, c1) = a[i].overflowing_add(b[i]);
let (s2, c2) = s1.overflowing_add(carry);
out[i] = s2;
carry = u64::from(c1) + u64::from(c2);
}
(out, carry)
}
fn sub3(a: [u64; 3], b: [u64; 3]) -> ([u64; 3], u64) {
let mut out = [0u64; 3];
let mut borrow = 0u64;
for i in 0..3 {
let (d1, b1) = a[i].overflowing_sub(b[i]);
let (d2, b2) = d1.overflowing_sub(borrow);
out[i] = d2;
borrow = u64::from(b1) + u64::from(b2);
}
(out, borrow)
}
fn cond_sub_if_ge(a: [u64; 3], b: [u64; 3]) -> [u64; 3] {
let (diff, borrow) = sub3(a, b);
let mask = borrow.wrapping_sub(1); let mut out = [0u64; 3];
for i in 0..3 {
out[i] = a[i] ^ (mask & (a[i] ^ diff[i]));
}
out
}
#[allow(clippy::cast_possible_truncation)] fn mul3(a: [u64; 3], b: [u64; 3]) -> [u64; 6] {
let mut out = [0u64; 6];
for i in 0..3 {
let mut carry = 0u128;
for j in 0..3 {
let product = u128::from(a[i]) * u128::from(b[j]) + u128::from(out[i + j]) + carry;
out[i + j] = product as u64;
carry = product >> 64;
}
out[i + 3] = carry as u64;
}
out
}
fn reduce_mod_n(product: [u64; 6]) -> [u64; 3] {
let n = Scalar::n();
let mut r = [0u64; 3];
for limb_idx in (0..6).rev() {
for bit in (0..64).rev() {
let bit_val = (product[limb_idx] >> bit) & 1;
r = shl1_or(r, bit_val);
r = cond_sub_if_ge(r, n);
}
}
r
}
fn shl1_or(x: [u64; 3], bit: u64) -> [u64; 3] {
let mut out = [0u64; 3];
let mut carry = bit;
for i in 0..3 {
let next_carry = x[i] >> 63;
out[i] = (x[i] << 1) | carry;
carry = next_carry;
}
out
}
#[cfg(all(test, any(feature = "std", feature = "getrandom")))]
mod from_candidate_bytes_tests {
use super::{curve163, Scalar};
#[test]
fn rejects_zero() {
assert!(Scalar::from_candidate_bytes(&[0u8; 21]).is_none());
}
#[test]
fn rejects_n_itself() {
assert!(Scalar::from_candidate_bytes(&curve163::order()).is_none());
}
#[test]
fn rejects_above_n() {
let mut above_n = curve163::order();
above_n[20] += 1; assert!(Scalar::from_candidate_bytes(&above_n).is_none());
}
#[test]
fn accepts_n_minus_one() {
let mut n_minus_one = curve163::order();
n_minus_one[20] -= 1;
let scalar = Scalar::from_candidate_bytes(&n_minus_one).expect("n - 1 is in [1, n)");
assert_eq!(scalar.to_be_bytes(), n_minus_one);
}
#[test]
fn accepts_one() {
let mut one = [0u8; 21];
one[20] = 1;
let scalar = Scalar::from_candidate_bytes(&one).expect("1 is in [1, n)");
assert_eq!(scalar.to_be_bytes(), one);
}
}