pub const L: [u8; 32] = [
0xed, 0xd3, 0xf5, 0x5c, 0x1a, 0x63, 0x12, 0x58, 0xd6, 0x9c, 0xf7, 0xa2, 0xde, 0xf9, 0xde, 0x14,
0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x10,
];
const L_LIMBS: [u32; 9] = [
0x5cf5_d3ed,
0x5812_631a,
0xa2f7_9cd6,
0x14de_f9de,
0x0000_0000,
0x0000_0000,
0x0000_0000,
0x1000_0000,
0x0000_0000,
];
fn conditional_subtract_l(r: &mut [u32; 9]) {
let mut diff = [0u32; 9];
let mut borrow = 0u64;
for i in 0..9 {
let d = (r[i] as u64)
.wrapping_sub(L_LIMBS[i] as u64)
.wrapping_sub(borrow);
diff[i] = d as u32;
borrow = (d >> 63) & 1;
}
let mask = ((borrow as u32) ^ 1).wrapping_neg();
for i in 0..9 {
r[i] ^= mask & (r[i] ^ diff[i]);
}
}
const NEG_C_DIGITS: [i64; 6] = [666_643, 470_296, 654_183, -997_805, 136_657, -683_901];
const LIMBS: usize = 25;
#[inline]
fn carry(limbs: &mut [i64; LIMBS], i: usize) {
let c = (limbs[i] + (1 << 20)) >> 21;
limbs[i] -= c << 21;
limbs[i + 1] += c;
}
pub fn reduce_wide(input: &[u8; 64]) -> [u8; 32] {
let mut padded = [0u8; 72];
padded[..64].copy_from_slice(input);
let mut limbs = [0i64; LIMBS];
for (i, slot) in limbs.iter_mut().enumerate() {
let bit = i * 21;
let mut b = [0u8; 8];
b.copy_from_slice(&padded[bit / 8..bit / 8 + 8]);
*slot = ((u64::from_le_bytes(b) >> (bit % 8)) & ((1 << 21) - 1)) as i64;
}
for _round in 0..4 {
for k in 0..LIMBS - 1 {
carry(&mut limbs, k);
}
for i in (12..LIMBS).rev() {
let t = limbs[i];
limbs[i] = 0;
for (j, d) in NEG_C_DIGITS.iter().enumerate() {
limbs[i - 12 + j] += d * t;
}
}
}
for _ in 0..2 {
limbs[12] += 1;
for (j, d) in NEG_C_DIGITS.iter().enumerate() {
limbs[j] -= d;
}
}
for k in 0..13 {
carry(&mut limbs, k);
}
let mut u = [0u64; 14];
let mut borrow: i64 = 0;
for (slot, &l) in u.iter_mut().zip(limbs.iter().take(14)) {
let v = l + borrow;
let m = v & ((1 << 21) - 1);
borrow = (v - m) >> 21;
*slot = m as u64;
}
debug_assert!(borrow >= 0, "reduction left a negative value");
let mut wide = [0u8; 40];
for (i, &v) in u.iter().enumerate().take(13) {
let bit = i * 21;
let mut b = [0u8; 8];
b.copy_from_slice(&wide[bit / 8..bit / 8 + 8]);
let merged = u64::from_le_bytes(b) | (v << (bit % 8));
wide[bit / 8..bit / 8 + 8].copy_from_slice(&merged.to_le_bytes());
}
let mut out = [0u8; 32];
out.copy_from_slice(&wide[..32]);
let mut r = [0u32; 9];
for i in 0..8 {
let mut b = [0u8; 4];
b.copy_from_slice(&out[i * 4..i * 4 + 4]);
r[i] = u32::from_le_bytes(b);
}
conditional_subtract_l(&mut r);
conditional_subtract_l(&mut r);
conditional_subtract_l(&mut r);
conditional_subtract_l(&mut r);
for i in 0..8 {
out[i * 4..i * 4 + 4].copy_from_slice(&r[i].to_le_bytes());
}
out
}
#[cfg(test)]
pub fn reduce(input: &[u8; 32]) -> [u8; 32] {
let mut wide = [0u8; 64];
wide[..32].copy_from_slice(input);
reduce_wide(&wide)
}
pub fn mul_add(a: &[u8; 32], b: &[u8; 32], c: &[u8; 32]) -> [u8; 32] {
let al = to_limbs(a);
let bl = to_limbs(b);
let mut prod = [0u64; 16];
for i in 0..8 {
let mut carry = 0u64;
for j in 0..8 {
let t = prod[i + j] + (al[i] as u64) * (bl[j] as u64) + carry;
prod[i + j] = t & 0xFFFF_FFFF;
carry = t >> 32;
}
prod[i + 8] += carry;
}
let mut wide = [0u8; 64];
for i in 0..16 {
wide[i * 4..i * 4 + 4].copy_from_slice(&(prod[i] as u32).to_le_bytes());
}
let mut carry = 0u16;
for i in 0..64 {
let ci = if i < 32 { c[i] as u16 } else { 0 };
let t = wide[i] as u16 + ci + carry;
wide[i] = t as u8;
carry = t >> 8;
}
reduce_wide(&wide)
}
fn to_limbs(bytes: &[u8; 32]) -> [u32; 8] {
let mut l = [0u32; 8];
for i in 0..8 {
l[i] = u32::from_le_bytes([
bytes[i * 4],
bytes[i * 4 + 1],
bytes[i * 4 + 2],
bytes[i * 4 + 3],
]);
}
l
}
#[must_use = "a false return means the scalar encoding was non-canonical"]
pub fn is_canonical(s: &[u8; 32]) -> bool {
let mut borrow = 0u16;
for i in 0..32 {
let d = (s[i] as u16).wrapping_sub(L[i] as u16).wrapping_sub(borrow);
borrow = (d >> 8) & 1;
}
borrow == 1
}
#[cfg(test)]
mod tests {
use super::*;
fn from_u64(v: u64) -> [u8; 32] {
let mut b = [0u8; 32];
b[..8].copy_from_slice(&v.to_le_bytes());
b
}
#[test]
fn small_values_are_unchanged() {
for v in [0u64, 1, 2, 1000, u32::MAX as u64] {
assert_eq!(reduce(&from_u64(v)), from_u64(v), "{v}");
}
}
#[test]
fn l_reduces_to_zero() {
assert_eq!(reduce(&L), [0u8; 32]);
}
#[test]
fn l_plus_one_reduces_to_one() {
let mut l1 = L;
l1[0] += 1;
assert_eq!(reduce(&l1), from_u64(1));
}
#[test]
fn maximum_wide_value_reduces_into_range() {
let r = reduce_wide(&[0xffu8; 64]);
assert!(is_canonical(&r), "reduction must land below L");
}
#[test]
fn mul_add_matches_small_arithmetic() {
let a = from_u64(1_000_003);
let b = from_u64(7_919);
let c = from_u64(65_537);
let expected = from_u64(1_000_003u64 * 7_919 + 65_537);
assert_eq!(mul_add(&a, &b, &c), expected);
}
#[test]
fn mul_add_is_zero_for_multiples_of_l() {
assert_eq!(mul_add(&L, &from_u64(1), &[0u8; 32]), [0u8; 32]);
assert_eq!(mul_add(&[0u8; 32], &from_u64(5), &L), [0u8; 32]);
}
#[test]
fn mul_add_result_is_always_canonical() {
let a = [0xAAu8; 32];
let b = [0x55u8; 32];
let c = [0xF0u8; 32];
assert!(is_canonical(&mul_add(&a, &b, &c)));
}
#[test]
fn canonical_test_matches_the_boundary() {
assert!(!is_canonical(&L), "L itself is not canonical");
let mut below = L;
below[0] -= 1;
assert!(is_canonical(&below));
assert!(is_canonical(&[0u8; 32]));
assert!(!is_canonical(&[0xffu8; 32]));
}
#[test]
fn reduction_commutes_with_multiplication() {
let a = [0x37u8; 32];
let b = [0x91u8; 32];
let direct = mul_add(&a, &b, &[0u8; 32]);
let pre = mul_add(&reduce(&a), &reduce(&b), &[0u8; 32]);
assert_eq!(direct, pre);
}
fn reduce_wide_by_long_division(input: &[u8; 64]) -> [u8; 32] {
let mut r = [0u32; 9];
for bit_index in (0..512).rev() {
let mut carry = 0u32;
for limb in r.iter_mut() {
let next = *limb >> 31;
*limb = (*limb << 1) | carry;
carry = next;
}
r[0] |= ((input[bit_index / 8] >> (bit_index % 8)) & 1) as u32;
conditional_subtract_l(&mut r);
}
let mut out = [0u8; 32];
for i in 0..8 {
out[i * 4..i * 4 + 4].copy_from_slice(&r[i].to_le_bytes());
}
out
}
#[test]
fn the_folding_constants_are_the_digits_of_minus_c() {
let mut c = [0u8; 32];
c.copy_from_slice(&L);
c[31] &= !0x10;
let mut v: i128 = 0;
for (i, b) in c.iter().enumerate().take(16) {
v |= (*b as i128) << (8 * i);
}
let mut v = -v;
let mut got = [0i64; 6];
for slot in got.iter_mut() {
let mut d = (v & ((1 << 21) - 1)) as i64;
if d >= 1 << 20 {
d -= 1 << 21;
}
*slot = d;
v = (v - d as i128) >> 21;
}
assert_eq!(v, 0, "c did not fit in six digits");
assert_eq!(got, NEG_C_DIGITS);
}
#[test]
fn folding_agrees_with_long_division() {
let mut state = 0x2f6d_9c1b_a473_e850u64;
let mut next = || {
state ^= state >> 12;
state ^= state << 25;
state ^= state >> 27;
state.wrapping_mul(0x2545_f491_4f6c_dd1d)
};
for case in 0..3_000 {
let mut input = [0u8; 64];
match case {
0 => {}
1 => input = [0xff; 64],
2 => input[0] = 1,
3 => input[63] = 0x80,
_ => {
for chunk in input.chunks_exact_mut(8) {
chunk.copy_from_slice(&next().to_le_bytes());
}
if case % 5 == 0 {
input[32..].fill(0);
}
}
}
assert_eq!(
reduce_wide(&input),
reduce_wide_by_long_division(&input),
"case {case}, input {input:?}"
);
}
}
}