use crate::kem::ntru_prime::constants::{P, Q};
pub const N: usize = P; pub const Q_VAL: i32 = Q as i32;
pub const NTT_SIZE: usize = 1024;
pub const ROOT_OF_UNITY: i32 = 8;
const ROOT_POWERS: [i32; 872] = [
1, 8, 64, 512, 502, 474, 401, 618, 297, 225, 317, 224, 448, 727, 462, 128, 153, 265, 533, 283,
573, 507, 646, 157, 765, 724, 493, 358, 282, 680, 1537, 1217, 429, 483, 267, 615, 698, 367,
409, 780, 823, 192, 1537, 1293, 1089, 177, 1416, 1332, 1514, 719, 318, 326, 115, 1241, 1429,
1081, 951, 499, 401, 495, 376, 416, 838, 1179, 1469, 1203, 155, 1269, 1570, 1436, 1529, 803,
102, 816, 934, 1549, 275, 272, 581, 783, 1054, 886, 1200, 1253, 1406, 929, 134, 1071, 1391,
891, 1643, 438, 443, 701, 911, 1370, 1449, 855, 713, 1209, 1205, 1589, 795, 1170, 1339, 441,
594, 1319, 598, 1251, 1220, 1395, 1035, 265, 925, 1053, 1479, 566, 1373, 907, 1083, 649, 1401,
975, 1173, 1311, 1231, 1069, 821, 1249, 1587, 1483, 1319, 1109, 837, 1089, 1367, 715, 721,
1121, 1365, 1135, 1243, 1221, 869, 671, 795, 849, 1319, 1629, 1091, 955, 1201, 1255, 1197,
1559, 1049, 959, 1401, 1631, 1455, 1447, 889, 735, 1039, 915, 1289, 1577, 1127, 1157, 1599,
1013, 895, 1313, 1287, 1285, 1251, 1375, 1551, 1677, 1293, 1419, 1463, 1327, 1247, 1185, 1133,
1193, 1185, 1235, 1309, 1419, 1455, 1505, 1413, 1361, 1395, 1491, 1389, 1283, 1309, 1383, 1371,
1357, 1319, 1361, 1373, 1323, 1297, 1345, 1375, 1367, 1347, 1351, 1381, 1415, 1407, 1391, 1415,
1421, 1409, 1409, 1423, 1425, 1419, 1425, 1427, 1421, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425,
1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425,
1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427, 1425, 1425, 1427,
];
#[inline(always)]
pub fn freeze_ntt(a: i32) -> i32 {
const Q: i32 = 4_591;
let b = a - Q * ((228 * a) >> 20);
let c = b - Q * ((58_470 * b + 134_217_728) >> 28);
c
}
#[inline(always)]
pub fn mod_mul_ntt(a: i32, b: i32) -> i32 {
freeze_ntt(a * b)
}
pub fn mod_pow(mut base: i32, mut exp: i32, modulus: i32) -> i32 {
let mut result = 1;
base = ((base % modulus) + modulus) % modulus;
while exp > 0 {
if exp & 1 == 1 {
result = mod_mul_ntt(result, base) % modulus;
}
exp >>= 1;
base = mod_mul_ntt(base, base) % modulus;
}
result
}
#[inline(always)]
pub fn mod_inv(a: i32) -> i32 {
const Q: i32 = 4_591;
mod_pow(a, Q - 2, Q)
}
pub fn karatsuba_mul(a: &[i16], b: &[i16]) -> Vec<i16> {
const THRESHOLD: usize = 64;
let n = a.len();
if n <= THRESHOLD {
let mut result = vec![0i16; n];
for i in 0..n {
let mut acc = 0i32;
for j in 0..=i {
acc += (a[j] as i32) * (b[i - j] as i32);
}
result[i] = freeze_ntt(acc) as i16;
}
return result;
}
let half = n / 2;
let a_low = &a[..half];
let a_high = &a[half..];
let b_low = &b[..half];
let b_high = &b[half..];
let mut a_low_plus_high = vec![0i16; half.max(a_high.len())];
let mut b_low_plus_high = vec![0i16; half.max(b_high.len())];
for i in 0..a_low.len() {
a_low_plus_high[i] = mod_add_ntt(a_low[i], if i < a_high.len() { a_high[i] } else { 0 });
}
for i in a_low.len()..a_high.len() {
a_low_plus_high[i] = a_high[i];
}
for i in 0..b_low.len() {
b_low_plus_high[i] = mod_add_ntt(b_low[i], if i < b_high.len() { b_high[i] } else { 0 });
}
for i in b_low.len()..b_high.len() {
b_low_plus_high[i] = b_high[i];
}
let z0 = karatsuba_mul(a_low, b_low);
let z2 = karatsuba_mul(a_high, b_high);
let z1_temp = karatsuba_mul(&a_low_plus_high, &b_low_plus_high);
let mut z1 = vec![0i16; z0.len().max(z2.len())];
for i in 0..z1.len() {
let mut val = 0i32;
if i < z1_temp.len() {
val += z1_temp[i] as i32;
}
if i < z0.len() {
val -= z0[i] as i32;
}
if i < z2.len() {
val -= z2[i] as i32;
}
z1[i] = freeze_ntt(val) as i16;
}
let mut result = vec![0i16; n];
for i in 0..z0.len().min(n) {
result[i] = z0[i];
}
for i in 0..z1.len() {
if half + i < n {
result[half + i] = mod_add_ntt(result[half + i], z1[i]);
}
}
for i in 0..z2.len() {
if 2 * half + i < n {
result[2 * half + i] = mod_add_ntt(result[2 * half + i], z2[i]);
}
}
result
}
#[inline(always)]
fn mod_add_ntt(a: i16, b: i16) -> i16 {
freeze_ntt((a as i32 + b as i32)) as i16
}
#[inline(always)]
fn mod_sub_ntt(a: i16, b: i16) -> i16 {
freeze_ntt((a as i32 - b as i32)) as i16
}
pub fn ntt_forward_naive(a: &[i32]) -> Vec<i32> {
let n = a.len();
let mut result = vec![0i32; n];
for k in 0..n {
let mut sum = 0i32;
for j in 0..n {
let twiddle = mod_pow(ROOT_OF_UNITY, (j * k) as i32, Q_VAL);
sum = freeze_ntt(sum + mod_mul_ntt(a[j], twiddle));
}
result[k] = sum;
}
result
}
pub fn ntt_inverse_naive(a: &[i32]) -> Vec<i32> {
let n = a.len();
let n_inv = mod_inv(n as i32);
let mut result = vec![0i32; n];
for k in 0..n {
let mut sum = 0i32;
for j in 0..n {
let twiddle = mod_pow(
ROOT_OF_UNITY,
Q_VAL - 1 - ((j * k) as i32 % (Q_VAL - 1)),
Q_VAL,
);
sum = freeze_ntt(sum + mod_mul_ntt(a[j], twiddle));
}
result[k] = mod_mul_ntt(sum, n_inv);
}
result
}
pub fn bit_reverse(a: &mut [i32]) {
let n = a.len();
let log_n = n.ilog2();
for i in 0..n {
let mut rev = 0u32;
let mut temp = i as u32;
for _ in 0..log_n {
rev = (rev << 1) | (temp & 1);
temp >>= 1;
}
let rev = rev as usize;
if i < rev {
a.swap(i, rev);
}
}
}
#[inline(always)]
fn butterfly(a: i32, b: i32, twiddle: i32) -> (i32, i32) {
let b_twiddle = mod_mul_ntt(b, twiddle);
(freeze_ntt(a + b_twiddle), freeze_ntt(a - b_twiddle))
}
pub fn ntt_forward_inplace(a: &mut [i32]) {
let n = a.len();
bit_reverse(a);
let mut len = 2;
while len <= n {
let half_len = len / 2;
let step = n / len;
for i in (0..n).step_by(len) {
for j in 0..half_len {
let twiddle = mod_pow(ROOT_OF_UNITY, (step * j) as i32, Q_VAL);
let (u, v) = butterfly(a[i + j], a[i + j + half_len], twiddle);
a[i + j] = u;
a[i + j + half_len] = v;
}
}
len *= 2;
}
}
pub fn ntt_inverse_inplace(a: &mut [i32]) {
let n = a.len();
let n_inv = mod_inv(n as i32);
bit_reverse(a);
let mut len = 2;
while len <= n {
let half_len = len / 2;
let step = n / len;
for i in (0..n).step_by(len) {
for j in 0..half_len {
let twiddle = mod_inv(mod_pow(ROOT_OF_UNITY, (step * j) as i32, Q_VAL));
let (u, v) = butterfly(a[i + j], a[i + j + half_len], twiddle);
a[i + j] = u;
a[i + j + half_len] = v;
}
}
len *= 2;
}
for coeff in a.iter_mut() {
*coeff = mod_mul_ntt(*coeff, n_inv);
}
}
pub fn ntt_forward(a: &[i32]) -> Vec<i32> {
let mut result = a.to_vec();
ntt_forward_inplace(&mut result);
result
}
pub fn ntt_inverse(a: &[i32]) -> Vec<i32> {
let mut result = a.to_vec();
ntt_inverse_inplace(&mut result);
result
}
pub fn pointwise_mul_ntt(a: &[i32], b: &[i32]) -> Vec<i32> {
a.iter()
.zip(b.iter())
.map(|(&x, &y)| mod_mul_ntt(x, y))
.collect()
}
pub fn poly_mul_ntt(f: &[i32], g: &[i32]) -> Vec<i32> {
let n = f.len();
let m = g.len();
let result_len = n + m - 1;
let padded_len = result_len.next_power_of_two();
let mut f_pad = vec![0i32; padded_len];
let mut g_pad = vec![0i32; padded_len];
f_pad[..n].copy_from_slice(f);
g_pad[..m].copy_from_slice(g);
ntt_forward_inplace(&mut f_pad);
ntt_forward_inplace(&mut g_pad);
let mut result_ntt = pointwise_mul_ntt(&f_pad, &g_pad);
ntt_inverse_inplace(&mut result_ntt);
result_ntt.truncate(result_len);
result_ntt
}
pub fn ntru_poly_mul_opt(f: &[i16], g: &[i16]) -> Vec<i16> {
let mut fg = vec![0i32; 2 * P - 1];
if P >= 64 {
let result = karatsuba_mul(f, g);
for (i, &val) in result.iter().enumerate().take(2 * P - 1) {
fg[i] = val as i32;
}
} else {
for i in 0..P {
let mut acc = 0i32;
for j in 0..=i {
acc += (f[j] as i32) * (g[i - j] as i32);
}
fg[i] = acc;
}
for i in P..(2 * P - 1) {
let mut acc = 0i32;
for j in (i - P + 1)..P {
acc += (f[j] as i32) * (g[i - j] as i32);
}
fg[i] = acc;
}
}
let mut result = vec![0i16; P];
for i in (P..(2 * P - 1)).rev() {
let coeff = freeze_ntt(fg[i]);
let target1 = i - P;
let target2 = i - P + 1;
fg[target1] += coeff as i32;
if target2 < P {
fg[target2] += coeff as i32;
}
}
for i in 0..P {
result[i] = freeze_ntt(fg[i]) as i16;
}
result
}
pub fn ntru_poly_mul_karatsuba(f: &[i16], g: &[i16]) -> Vec<i16> {
ntru_poly_mul_opt(f, g)
}
pub fn ntru_poly_mul_ntt(f: &[i16], g: &[i16]) -> Vec<i16> {
ntru_poly_mul_opt(f, g)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_freeze_ntt() {
assert_eq!(freeze_ntt(0), 0);
assert_eq!(freeze_ntt(4591), 0);
assert_eq!(freeze_ntt(4591 * 2), 0);
assert_eq!(freeze_ntt(2295), 2295);
}
#[test]
fn test_mod_pow() {
assert_eq!(mod_pow(2, 0, 4591), 1);
assert_eq!(mod_pow(2, 1, 4591), 2);
assert_eq!(mod_pow(2, 10, 4591), 1024);
}
#[test]
fn test_bit_reverse() {
let mut a = vec![0, 1, 2, 3, 4, 5, 6, 7];
let original = a.clone();
bit_reverse(&mut a);
assert_eq!(a[0], original[0]);
assert_eq!(a[1], original[4]);
assert_eq!(a[2], original[2]);
assert_eq!(a[3], original[6]);
}
#[test]
fn test_ntt_roundtrip() {
let a = vec![1i32, 2, 3, 4, 5, 6, 7, 8];
let forward = ntt_forward(&a);
let back = ntt_inverse(&forward);
for (orig, recovered) in a.iter().zip(back.iter()) {
let diff = (*orig - *recovered).rem_euclid(4591);
assert_eq!(diff, 0);
}
}
#[test]
fn test_poly_mul_ntt_simple() {
let f = vec![1i32, 2, 3];
let g = vec![2i32, 3];
let result = poly_mul_ntt(&f, &g);
assert_eq!(result[0], 2);
assert_eq!(result[1], 7);
assert_eq!(result[2], 12);
assert_eq!(result[3], 9);
}
#[test]
fn test_ntru_poly_mul_ntt() {
let f: Vec<i16> = (0..761).map(|i| (i % 3 - 1) as i16).collect();
let g: Vec<i16> = (0..761).map(|i| (i % 5 - 2) as i16).collect();
let result = ntru_poly_mul_ntt(&f, &g);
assert_eq!(result.len(), 761);
for &val in &result {
assert!(val >= -2295 && val < 2296);
}
}
#[test]
fn test_karatsuba_mul() {
let f: Vec<i16> = (0..761).map(|i| (i % 3 - 1) as i16).collect();
let g: Vec<i16> = (0..761).map(|i| (i % 5 - 2) as i16).collect();
let result = karatsuba_mul(&f, &g);
assert_eq!(result.len(), 761);
for &val in &result {
assert!(val >= -2295 && val < 2296);
}
}
#[test]
fn test_ntru_poly_mul_opt() {
let f: Vec<i16> = (0..761).map(|i| (i % 3 - 1) as i16).collect();
let g: Vec<i16> = (0..761).map(|i| (i % 5 - 2) as i16).collect();
let result = ntru_poly_mul_opt(&f, &g);
assert_eq!(result.len(), 761);
for &val in &result {
assert!(val >= -2295 && val < 2296);
}
}
#[test]
fn test_karatsuba_small() {
let f = vec![1i16, 2, 3, 4, 5, 6, 7, 8];
let g = vec![2i16, 3, 4, 5, 6, 7, 8, 9];
let result = karatsuba_mul(&f, &g);
assert_eq!(result[0], 2);
assert_eq!(result[1], 7);
}
}