dcrypt_algorithms/poly/fft/
mod.rs1#![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
23const 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
32fn 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#[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#[inline]
60fn project_to_2power(x: Scalar) -> Scalar {
61 pow_vartime_u64x4(x, &FR_ODD_PART)
62}
63
64fn 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 debug_assert!(false, "two_adicity: element not in μ_{{2^S}}");
75 TWO_ADICITY_FR
76}
77
78fn 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
108fn 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
196fn 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
213fn 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
235pub 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
248pub 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
266pub 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
283pub 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;