1#![allow(clippy::needless_range_loop)]
13
14use ic_core::ct::Choice;
15
16#[derive(Clone, Copy, Debug, PartialEq, Eq)]
18pub struct Fe(pub [u64; 5]);
19
20const MASK: u64 = (1 << 51) - 1;
21
22#[allow(clippy::unusual_byte_groupings)]
27const TWO_P: [u64; 5] = [
28 0xFFFFFFFFFFFDA,
29 0xFFFFFFFFFFFFE,
30 0xFFFFFFFFFFFFE,
31 0xFFFFFFFFFFFFE,
32 0xFFFFFFFFFFFFE,
33];
34
35impl Fe {
36 pub const ZERO: Fe = Fe([0, 0, 0, 0, 0]);
38 pub const ONE: Fe = Fe([1, 0, 0, 0, 0]);
40
41 pub const fn from_u64(v: u64) -> Fe {
43 Fe([v & MASK, v >> 51, 0, 0, 0])
44 }
45
46 #[inline]
48 pub fn add(&self, other: &Fe) -> Fe {
49 let mut r = [0u64; 5];
50 for i in 0..5 {
51 r[i] = self.0[i] + other.0[i];
52 }
53 Fe(r)
54 }
55
56 #[inline]
58 pub fn sub(&self, other: &Fe) -> Fe {
59 let mut r = [0u64; 5];
60 for i in 0..5 {
61 r[i] = self.0[i] + TWO_P[i] - other.0[i];
62 }
63 Fe(r).weak_reduce()
64 }
65
66 #[inline]
68 pub fn neg(&self) -> Fe {
69 Fe::ZERO.sub(self)
70 }
71
72 #[inline]
74 fn weak_reduce(self) -> Fe {
75 let mut r = self.0;
76 let mut carry = r[0] >> 51;
77 r[0] &= MASK;
78 for i in 1..5 {
79 r[i] += carry;
80 carry = r[i] >> 51;
81 r[i] &= MASK;
82 }
83 r[0] += carry.wrapping_mul(19);
84 Fe(r)
85 }
86
87 #[inline]
89 pub fn mul(&self, other: &Fe) -> Fe {
90 let a = &self.0;
91 let b = &other.0;
92
93 let b1_19 = b[1] * 19;
103 let b2_19 = b[2] * 19;
104 let b3_19 = b[3] * 19;
105 let b4_19 = b[4] * 19;
106
107 let r0 = m(a[0], b[0]) + m(a[1], b4_19) + m(a[2], b3_19) + m(a[3], b2_19) + m(a[4], b1_19);
108 let r1 = m(a[0], b[1]) + m(a[1], b[0]) + m(a[2], b4_19) + m(a[3], b3_19) + m(a[4], b2_19);
109 let r2 = m(a[0], b[2]) + m(a[1], b[1]) + m(a[2], b[0]) + m(a[3], b4_19) + m(a[4], b3_19);
110 let r3 = m(a[0], b[3]) + m(a[1], b[2]) + m(a[2], b[1]) + m(a[3], b[0]) + m(a[4], b4_19);
111 let r4 = m(a[0], b[4]) + m(a[1], b[3]) + m(a[2], b[2]) + m(a[3], b[1]) + m(a[4], b[0]);
112
113 carry_reduce([r0, r1, r2, r3, r4])
114 }
115
116 #[inline]
118 pub fn square(&self) -> Fe {
119 let a = &self.0;
135 let a0_2 = a[0] * 2;
136 let a1_2 = a[1] * 2;
137 let a1_38 = a[1] * 38;
138 let a2_38 = a[2] * 38;
139 let a3_38 = a[3] * 38;
140 let a3_19 = a[3] * 19;
141 let a4_19 = a[4] * 19;
142
143 let r0 = m(a[0], a[0]) + m(a1_38, a[4]) + m(a2_38, a[3]);
145 let r1 = m(a0_2, a[1]) + m(a2_38, a[4]) + m(a3_19, a[3]);
146 let r2 = m(a0_2, a[2]) + m(a[1], a[1]) + m(a3_38, a[4]);
147 let r3 = m(a0_2, a[3]) + m(a1_2, a[2]) + m(a4_19, a[4]);
148 let r4 = m(a0_2, a[4]) + m(a1_2, a[3]) + m(a[2], a[2]);
149
150 carry_reduce([r0, r1, r2, r3, r4])
151 }
152
153 #[inline]
155 pub fn square_n(&self, n: usize) -> Fe {
156 let mut r = *self;
157 for _ in 0..n {
158 r = r.square();
159 }
160 r
161 }
162
163 #[inline]
165 pub fn mul121666(&self) -> Fe {
166 let mut r = [0u128; 5];
167 for i in 0..5 {
168 r[i] = (self.0[i] as u128) * 121_666;
169 }
170 carry_reduce(r)
171 }
172
173 pub fn invert(&self) -> Fe {
177 let z2 = self.square();
178 let z9 = z2.square_n(2).mul(self);
179 let z11 = z9.mul(&z2);
180 let z2_5_0 = z11.square().mul(&z9);
181 let z2_10_0 = z2_5_0.square_n(5).mul(&z2_5_0);
182 let z2_20_0 = z2_10_0.square_n(10).mul(&z2_10_0);
183 let z2_40_0 = z2_20_0.square_n(20).mul(&z2_20_0);
184 let z2_50_0 = z2_40_0.square_n(10).mul(&z2_10_0);
185 let z2_100_0 = z2_50_0.square_n(50).mul(&z2_50_0);
186 let z2_200_0 = z2_100_0.square_n(100).mul(&z2_100_0);
187 let z2_250_0 = z2_200_0.square_n(50).mul(&z2_50_0);
188 z2_250_0.square_n(5).mul(&z11)
189 }
190
191 pub fn pow22523(&self) -> Fe {
193 let z2 = self.square();
194 let z9 = z2.square_n(2).mul(self);
195 let z11 = z9.mul(&z2);
196 let z2_5_0 = z11.square().mul(&z9);
197 let z2_10_0 = z2_5_0.square_n(5).mul(&z2_5_0);
198 let z2_20_0 = z2_10_0.square_n(10).mul(&z2_10_0);
199 let z2_40_0 = z2_20_0.square_n(20).mul(&z2_20_0);
200 let z2_50_0 = z2_40_0.square_n(10).mul(&z2_10_0);
201 let z2_100_0 = z2_50_0.square_n(50).mul(&z2_50_0);
202 let z2_200_0 = z2_100_0.square_n(100).mul(&z2_100_0);
203 let z2_250_0 = z2_200_0.square_n(50).mul(&z2_50_0);
204 z2_250_0.square_n(2).mul(self)
205 }
206
207 pub fn from_bytes(bytes: &[u8; 32]) -> Fe {
209 let load = |i: usize| -> u64 {
210 let mut v = [0u8; 8];
211 v.copy_from_slice(&bytes[i..i + 8]);
212 u64::from_le_bytes(v)
213 };
214 let l0 = load(0) & MASK;
215 let l1 = (load(6) >> 3) & MASK;
216 let l2 = (load(12) >> 6) & MASK;
217 let l3 = (load(19) >> 1) & MASK;
218 let l4 = (load(24) >> 12) & MASK;
219 Fe([l0, l1, l2, l3, l4])
220 }
221
222 pub fn to_bytes(&self) -> [u8; 32] {
224 let mut t = self.weak_reduce().weak_reduce().weak_reduce().0;
227
228 let mut q = (t[0] + 19) >> 51;
231 for i in 1..5 {
232 q = (t[i] + q) >> 51;
233 }
234 t[0] += 19 * q;
235 let mut carry = t[0] >> 51;
236 t[0] &= MASK;
237 for i in 1..5 {
238 t[i] += carry;
239 carry = t[i] >> 51;
240 t[i] &= MASK;
241 }
242 t[4] &= (1 << 51) - 1;
244
245 let mut out = [0u8; 32];
246 let mut acc: u128 = 0;
247 let mut acc_bits = 0usize;
248 let mut idx = 0usize;
249 for limb in t.iter() {
250 acc |= (*limb as u128) << acc_bits;
251 acc_bits += 51;
252 while acc_bits >= 8 && idx < 32 {
253 out[idx] = acc as u8;
254 acc >>= 8;
255 acc_bits -= 8;
256 idx += 1;
257 }
258 }
259 while idx < 32 {
260 out[idx] = acc as u8;
261 acc >>= 8;
262 idx += 1;
263 }
264 out
265 }
266
267 #[inline]
269 pub fn cswap(a: &mut Fe, b: &mut Fe, choice: Choice) {
270 let mask = (choice.unwrap_u8() as u64).wrapping_neg();
271 for i in 0..5 {
272 let t = mask & (a.0[i] ^ b.0[i]);
273 a.0[i] ^= t;
274 b.0[i] ^= t;
275 }
276 }
277
278 #[inline]
280 pub fn cmov(a: &mut Fe, b: &Fe, choice: Choice) {
281 let mask = (choice.unwrap_u8() as u64).wrapping_neg();
282 for i in 0..5 {
283 a.0[i] ^= mask & (a.0[i] ^ b.0[i]);
284 }
285 }
286
287 pub fn is_zero(&self) -> Choice {
289 ic_core::ct::is_zero(&self.to_bytes())
290 }
291
292 pub fn ct_eq(&self, other: &Fe) -> Choice {
294 ic_core::ct::eq(&self.to_bytes(), &other.to_bytes())
295 }
296
297 pub fn is_negative(&self) -> Choice {
300 Choice::from_u8(self.to_bytes()[0] & 1)
301 }
302}
303
304#[inline(always)]
309fn m(x: u64, y: u64) -> u128 {
310 (x as u128) * (y as u128)
311}
312
313#[inline]
315fn carry_reduce(r: [u128; 5]) -> Fe {
316 let c: [u64; 5] = [
333 (r[0] >> 51) as u64,
334 (r[1] >> 51) as u64,
335 (r[2] >> 51) as u64,
336 (r[3] >> 51) as u64,
337 (r[4] >> 51) as u64,
338 ];
339 let mut out: [u64; 5] = [
340 (r[0] as u64 & MASK) + c[4] * 19,
341 (r[1] as u64 & MASK) + c[0],
342 (r[2] as u64 & MASK) + c[1],
343 (r[3] as u64 & MASK) + c[2],
344 (r[4] as u64 & MASK) + c[3],
345 ];
346
347 let mut carry = out[0] >> 51;
350 out[0] &= MASK;
351 for slot in out.iter_mut().skip(1) {
352 *slot += carry;
353 carry = *slot >> 51;
354 *slot &= MASK;
355 }
356 out[0] += carry * 19;
357 Fe(out)
358}
359
360#[cfg(test)]
361mod tests {
362 use super::*;
363
364 #[test]
376 fn squaring_agrees_with_multiplication() {
377 let mut cases = std::vec![
378 Fe::ZERO,
379 Fe::ONE,
380 Fe([1, 1, 1, 1, 1]),
381 Fe([(1u64 << 51) - 1; 5]),
382 Fe([(1u64 << 51) - 1, 0, (1u64 << 51) - 1, 0, (1u64 << 51) - 1]),
383 Fe([0, (1u64 << 51) - 1, 0, (1u64 << 51) - 1, 0]),
384 ];
385 let mut x = Fe([0x51a2, 0x9e37, 0x79b9, 0x7f4a, 0x7c15]);
387 for _ in 0..16 {
388 x = x.mul(&Fe([3, 5, 7, 11, 13])).add(&Fe::ONE);
389 cases.push(x);
390 }
391
392 let mut checked = 0;
393 for f in &cases {
394 assert_eq!(
395 f.square().to_bytes(),
396 f.mul(f).to_bytes(),
397 "square and mul-by-self differ"
398 );
399 checked += 1;
400 }
401 assert_eq!(checked, 22, "the comparison did not run");
402 }
403
404 fn fe(v: u64) -> Fe {
405 Fe::from_u64(v)
406 }
407
408 #[test]
409 fn encode_decode_roundtrip() {
410 for v in [0u64, 1, 2, 19, 1 << 51, u64::MAX] {
411 let a = fe(v);
412 assert_eq!(Fe::from_bytes(&a.to_bytes()).to_bytes(), a.to_bytes());
413 }
414 }
415
416 #[test]
417 fn small_arithmetic() {
418 assert_eq!(fe(2).add(&fe(3)).to_bytes(), fe(5).to_bytes());
419 assert_eq!(fe(5).sub(&fe(3)).to_bytes(), fe(2).to_bytes());
420 assert_eq!(fe(6).mul(&fe(7)).to_bytes(), fe(42).to_bytes());
421 assert_eq!(fe(9).square().to_bytes(), fe(81).to_bytes());
422 }
423
424 #[test]
425 fn subtraction_wraps_into_the_field() {
426 let r = Fe::ZERO.sub(&Fe::ONE).to_bytes();
428 assert_eq!(r[0], 0xec);
429 assert_eq!(r[31], 0x7f);
430 for b in &r[1..31] {
431 assert_eq!(*b, 0xff);
432 }
433 }
434
435 #[test]
436 fn p_encodes_as_zero() {
437 let mut p_bytes = [0xffu8; 32];
439 p_bytes[0] = 0xed;
440 p_bytes[31] = 0x7f;
441 assert_eq!(Fe::from_bytes(&p_bytes).to_bytes(), [0u8; 32]);
442 }
443
444 #[test]
445 fn inversion_is_correct() {
446 for v in [1u64, 2, 3, 19, 12345, u32::MAX as u64] {
447 let a = fe(v);
448 assert_eq!(a.mul(&a.invert()).to_bytes(), Fe::ONE.to_bytes(), "1/{v}");
449 }
450 assert_eq!(Fe::ZERO.invert().to_bytes(), [0u8; 32]);
451 }
452
453 #[test]
454 fn multiplication_is_associative_and_distributive() {
455 let a = Fe::from_bytes(&[0x11; 32]);
456 let b = Fe::from_bytes(&[0x7a; 32]);
457 let c = Fe::from_bytes(&[0xc3; 32]);
458 assert_eq!(a.mul(&b).mul(&c).to_bytes(), a.mul(&b.mul(&c)).to_bytes());
459 assert_eq!(
460 a.mul(&b.add(&c)).to_bytes(),
461 a.mul(&b).add(&a.mul(&c)).to_bytes()
462 );
463 }
464
465 #[test]
466 fn pow22523_gives_a_square_root() {
467 let x = fe(4);
469 let r = x.pow22523().mul(&x);
470 let sq = r.square();
471 assert!(
473 bool::from(sq.ct_eq(&x)) || bool::from(sq.ct_eq(&x.neg())),
474 "square root property"
475 );
476 }
477
478 #[test]
479 fn cswap_and_cmov_are_conditional() {
480 let mut a = fe(1);
481 let mut b = fe(2);
482 Fe::cswap(&mut a, &mut b, Choice::FALSE);
483 assert_eq!(a.to_bytes(), fe(1).to_bytes());
484 Fe::cswap(&mut a, &mut b, Choice::TRUE);
485 assert_eq!(a.to_bytes(), fe(2).to_bytes());
486
487 let mut c = fe(5);
488 Fe::cmov(&mut c, &fe(9), Choice::FALSE);
489 assert_eq!(c.to_bytes(), fe(5).to_bytes());
490 Fe::cmov(&mut c, &fe(9), Choice::TRUE);
491 assert_eq!(c.to_bytes(), fe(9).to_bytes());
492 }
493
494 #[test]
495 fn high_bit_of_input_is_ignored() {
496 let mut a = [0x42u8; 32];
497 let mut b = a;
498 a[31] &= 0x7f;
499 b[31] |= 0x80;
500 assert_eq!(Fe::from_bytes(&a).to_bytes(), Fe::from_bytes(&b).to_bytes());
501 }
502
503 fn carry_reduce_serial(r: [u128; 5]) -> Fe {
510 let mut out = [0u64; 5];
511 let mut carry: u128 = 0;
512 for (slot, limb) in out.iter_mut().zip(r) {
513 let v = limb + carry;
514 carry = v >> 51;
515 *slot = (v & MASK as u128) as u64;
516 }
517 out[0] += (carry as u64) * 19;
518 let mut c = out[0] >> 51;
519 out[0] &= MASK;
520 for slot in out.iter_mut().skip(1) {
521 *slot += c;
522 c = *slot >> 51;
523 *slot &= MASK;
524 }
525 out[0] += c * 19;
526 Fe(out)
527 }
528
529 fn raw_products(a: &[u64; 5], b: &[u64; 5]) -> [u128; 5] {
531 let m = |x: u64, y: u64| (x as u128) * (y as u128);
532 let (b1, b2, b3, b4) = (b[1] * 19, b[2] * 19, b[3] * 19, b[4] * 19);
533 [
534 m(a[0], b[0]) + m(a[1], b4) + m(a[2], b3) + m(a[3], b2) + m(a[4], b1),
535 m(a[0], b[1]) + m(a[1], b[0]) + m(a[2], b4) + m(a[3], b3) + m(a[4], b2),
536 m(a[0], b[2]) + m(a[1], b[1]) + m(a[2], b[0]) + m(a[3], b4) + m(a[4], b3),
537 m(a[0], b[3]) + m(a[1], b[2]) + m(a[2], b[1]) + m(a[3], b[0]) + m(a[4], b4),
538 m(a[0], b[4]) + m(a[1], b[3]) + m(a[2], b[2]) + m(a[3], b[1]) + m(a[4], b[0]),
539 ]
540 }
541
542 #[test]
550 fn limbs_at_their_maximum_do_not_carry_out_of_a_u64() {
551 let max = [(1u64 << 52) - 2; 5];
552 let r = raw_products(&max, &max);
553 assert_eq!(
554 carry_reduce(r).to_bytes(),
555 carry_reduce_serial(r).to_bytes(),
556 "parallel and serial carry disagree at the limb maximum"
557 );
558 }
559
560 #[test]
562 fn parallel_carry_agrees_with_the_serial_one() {
563 let mut state = 0x243f_6a88_85a3_08d3u64;
564 let mut next = || {
565 state ^= state >> 12;
567 state ^= state << 25;
568 state ^= state >> 27;
569 state.wrapping_mul(0x2545_f491_4f6c_dd1d)
570 };
571 for _ in 0..20_000 {
572 let mut a = [0u64; 5];
573 let mut b = [0u64; 5];
574 for i in 0..5 {
575 a[i] = next() % (1 << 52);
577 b[i] = next() % (1 << 52);
578 }
579 let r = raw_products(&a, &b);
580 assert_eq!(
581 carry_reduce(r).to_bytes(),
582 carry_reduce_serial(r).to_bytes(),
583 "disagreement on a={a:?} b={b:?}"
584 );
585 }
586 }
587}