Skip to main content

dcrypt_algorithms/poly/fft/
mod.rs

1// Path: dcrypt/crates/algorithms/src/poly/fft/mod.rs
2//! Fast Fourier Transform (FFT) over the BLS12-381 Scalar Field
3//!
4//! This module implements a Number Theoretic Transform (NTT), which is an FFT
5//! adapted for finite fields. It operates on vectors of `Scalar` elements from the
6//! BLS12-381 curve, enabling O(n log n) polynomial multiplication and interpolation.
7//!
8//! This is the high-performance engine required for schemes like Verkle trees that
9//! rely on polynomial commitments over large prime fields.
10
11#![allow(clippy::needless_range_loop)]
12
13#[cfg(feature = "alloc")]
14extern crate alloc;
15#[cfg(feature = "alloc")]
16use crate::alloc_prelude::*;
17
18use crate::ec::bls12_381::Bls12_381Scalar as Scalar;
19use crate::error::{Error, Result};
20
21const FFT_SIZE: usize = 256;
22
23// --- field-specific 2-adicity and odd cofactor for BLS12-381 Fr ---
24const TWO_ADICITY_FR: u32 = 32;
25const FR_ODD_PART: [u64; 4] = [
26    0xfffe_5bfe_ffff_ffff,
27    0x09a1_d805_53bd_a402,
28    0x299d_7d48_3339_d808,
29    0x0000_0000_73ed_a753,
30];
31
32// The original hardcoded constant (kept as a seed candidate).
33fn get_root_of_unity() -> Scalar {
34    Scalar::from_raw([
35        0x4253_d252_a210_b619,
36        0x81c3_5f15_01a0_2431,
37        0xb734_6a32_008b_0320,
38        0x0a16_14a8_64b3_09e1,
39    ])
40}
41
42// --- NEW: small helpers ---
43
44#[inline]
45fn pow_vartime_u64x4(base: Scalar, by: &[u64; 4]) -> Scalar {
46    let mut res = Scalar::one();
47    for e in by.iter().rev() {
48        for i in (0..64).rev() {
49            res = res.square();
50            if ((*e >> i) & 1) == 1 {
51                res *= base;
52            }
53        }
54    }
55    res
56}
57
58/// Project an arbitrary element into μ_{2^S}: x ↦ x^T
59#[inline]
60fn project_to_2power(x: Scalar) -> Scalar {
61    pow_vartime_u64x4(x, &FR_ODD_PART)
62}
63
64/// Compute the 2-adic order k of an element r ∈ μ_{2^S}:
65/// the smallest k ≥ 1 such that r^(2^k) = 1.
66fn two_adicity(mut r: Scalar) -> u32 {
67    for k in 1..=TWO_ADICITY_FR {
68        r = r.square();
69        if r == Scalar::one() {
70            return k;
71        }
72    }
73    // FIX: Escape the curly braces in the format string.
74    debug_assert!(false, "two_adicity: element not in μ_{{2^S}}");
75    TWO_ADICITY_FR
76}
77
78/// Deterministically pick a seed in μ_{2^S} whose 2-adic order k ≥ min_k.
79fn select_2power_seed(min_k: u32) -> (Scalar, u32) {
80    let bases: [Scalar; 12] = [
81        get_root_of_unity(),
82        Scalar::from(5u64),
83        Scalar::from(7u64),
84        Scalar::from(2u64),
85        Scalar::from(3u64),
86        Scalar::from(11u64),
87        Scalar::from(13u64),
88        Scalar::from(17u64),
89        Scalar::from(19u64),
90        Scalar::from(29u64),
91        Scalar::from(31u64),
92        Scalar::from(37u64),
93    ];
94
95    for base in bases.iter() {
96        let seed = project_to_2power(*base);
97        if !bool::from(seed.is_zero()) {
98            let k = two_adicity(seed);
99            if k >= min_k {
100                return (seed, k);
101            }
102        }
103    }
104
105    panic!("Could not find a suitable 2-power root of unity seed");
106}
107
108// --- Derived roots built from a consistent seed ---
109
110fn get_fft_n_root() -> Scalar {
111    let need = FFT_SIZE.trailing_zeros();
112    let (seed, k) = select_2power_seed(need);
113
114    let mut w_n = seed;
115    for _ in 0..(k - need) {
116        w_n = w_n.square();
117    }
118
119    #[cfg(debug_assertions)]
120    {
121        let mut t = w_n;
122        for _ in 0..need {
123            t = t.square();
124        }
125        debug_assert_eq!(t, Scalar::one(), "w_N^N must be 1");
126
127        let mut half = w_n;
128        for _ in 0..(need - 1) {
129            half = half.square();
130        }
131        debug_assert_eq!(half, -Scalar::one(), "w_N^(N/2) must be -1");
132    }
133    w_n
134}
135
136fn get_roots_of_unity() -> Vec<Scalar> {
137    let w_n = get_fft_n_root();
138    let mut roots = vec![Scalar::one(); FFT_SIZE];
139    for i in 1..FFT_SIZE {
140        roots[i] = roots[i - 1] * w_n;
141    }
142    roots
143}
144
145fn get_inverse_roots_of_unity() -> Vec<Scalar> {
146    let inv_w_n = get_fft_n_root().invert().unwrap();
147    let mut roots = vec![Scalar::one(); FFT_SIZE];
148    for i in 1..FFT_SIZE {
149        roots[i] = roots[i - 1] * inv_w_n;
150    }
151    roots
152}
153
154fn get_n_inv() -> Scalar {
155    Scalar::from(FFT_SIZE as u64).invert().unwrap()
156}
157
158fn get_primitive_2n_root() -> Scalar {
159    let need = FFT_SIZE.trailing_zeros();
160    let (seed, k) = select_2power_seed(need + 1);
161
162    let mut g = seed;
163    for _ in 0..(k - (need + 1)) {
164        g = g.square();
165    }
166
167    debug_assert_eq!(g.square(), get_fft_n_root(), "g^2 must equal w_N");
168
169    let mut gn = g;
170    for _ in 0..need {
171        gn = gn.square();
172    }
173    debug_assert_eq!(gn, -Scalar::one(), "g^N must be -1");
174
175    g
176}
177
178fn get_twist_factors() -> Vec<Scalar> {
179    let g = get_primitive_2n_root();
180    let mut factors = vec![Scalar::one(); FFT_SIZE];
181    for i in 1..FFT_SIZE {
182        factors[i] = factors[i - 1] * g;
183    }
184    factors
185}
186
187fn get_inverse_twist_factors() -> Vec<Scalar> {
188    let inv_g = get_primitive_2n_root().invert().unwrap();
189    let mut factors = vec![Scalar::one(); FFT_SIZE];
190    for i in 1..FFT_SIZE {
191        factors[i] = factors[i - 1] * inv_g;
192    }
193    factors
194}
195
196/// Performs a bit-reversal permutation on the input slice in-place.
197fn bit_reverse_permutation<T>(data: &mut [T]) {
198    let n = data.len();
199    let mut j = 0;
200    for i in 1..n {
201        let mut bit = n >> 1;
202        while (j & bit) != 0 {
203            j ^= bit;
204            bit >>= 1;
205        }
206        j ^= bit;
207        if i < j {
208            data.swap(i, j);
209        }
210    }
211}
212
213/// Core Cooley-Tukey FFT/NTT algorithm.
214fn fft_cooley_tukey(coeffs: &mut [Scalar], roots: &[Scalar]) {
215    let n = coeffs.len();
216    let mut len = 2;
217    while len <= n {
218        let half_len = len >> 1;
219        let step = roots.len() / len;
220        let root = roots[step];
221        for i in (0..n).step_by(len) {
222            let mut w = Scalar::one();
223            for j in 0..half_len {
224                let u = coeffs[i + j];
225                let v = coeffs[i + j + half_len] * w;
226                coeffs[i + j] = u + v;
227                coeffs[i + j + half_len] = u - v;
228                w *= root;
229            }
230        }
231        len <<= 1;
232    }
233}
234
235/// Computes the forward Fast Fourier Transform (NTT) of a polynomial for **cyclic** convolution.
236pub fn fft(coeffs: &mut [Scalar]) -> Result<()> {
237    if coeffs.len() != FFT_SIZE || !coeffs.len().is_power_of_two() {
238        return Err(Error::Parameter {
239            name: "coeffs".into(),
240            reason: "FFT length must be a power of two (256)".into(),
241        });
242    }
243    bit_reverse_permutation(coeffs);
244    fft_cooley_tukey(coeffs, &get_roots_of_unity());
245    Ok(())
246}
247
248/// Computes the inverse Fast Fourier Transform (iNTT) for **cyclic** convolution.
249pub fn ifft(evals: &mut [Scalar]) -> Result<()> {
250    if evals.len() != FFT_SIZE || !evals.len().is_power_of_two() {
251        return Err(Error::Parameter {
252            name: "evals".into(),
253            reason: "FFT length must be a power of two (256)".into(),
254        });
255    }
256    bit_reverse_permutation(evals);
257    fft_cooley_tukey(evals, &get_inverse_roots_of_unity());
258
259    let n_inv = get_n_inv();
260    for c in evals.iter_mut() {
261        *c *= n_inv;
262    }
263    Ok(())
264}
265
266/// Computes the forward **negacyclic** NTT.
267pub fn fft_negacyclic(coeffs: &mut [Scalar]) -> Result<()> {
268    if coeffs.len() != FFT_SIZE {
269        return Err(Error::Parameter {
270            name: "coeffs".into(),
271            reason: "Negacyclic FFT requires length 256".into(),
272        });
273    }
274
275    let twists = get_twist_factors();
276    for i in 0..FFT_SIZE {
277        coeffs[i] *= twists[i];
278    }
279
280    fft(coeffs)
281}
282
283/// Computes the inverse **negacyclic** NTT.
284pub fn ifft_negacyclic(evals: &mut [Scalar]) -> Result<()> {
285    if evals.len() != FFT_SIZE {
286        return Err(Error::Parameter {
287            name: "evals".into(),
288            reason: "Negacyclic IFFT requires length 256".into(),
289        });
290    }
291
292    ifft(evals)?;
293
294    let inv_twists = get_inverse_twist_factors();
295    for i in 0..FFT_SIZE {
296        evals[i] *= inv_twists[i];
297    }
298
299    Ok(())
300}
301
302#[cfg(test)]
303mod tests;