Skip to main content

ic_ec/
field.rs

1//! Arithmetic in GF(2^255 - 19).
2//!
3//! Field elements are five 51-bit limbs in a `u64`, the representation used by
4//! every high-quality Curve25519 implementation: products fit a `u128` without
5//! overflow, and carries propagate in a fixed pattern with no data-dependent
6//! branches. Nothing here inspects a limb value to decide control flow, so the
7//! whole module is constant-time with respect to secrets.
8
9//! Indexed loops over fixed-size limb and word arrays are used throughout; they
10//! mirror the index algebra in the specifications these routines implement, so
11//! `needless_range_loop` is allowed rather than obscuring the correspondence.
12#![allow(clippy::needless_range_loop)]
13
14use ic_core::ct::Choice;
15
16/// A field element modulo 2^255 - 19.
17#[derive(Clone, Copy, Debug, PartialEq, Eq)]
18pub struct Fe(pub [u64; 5]);
19
20const MASK: u64 = (1 << 51) - 1;
21
22/// `2 * p`, used so subtraction never goes negative.
23// Limb 0 is 2*(2^51 - 19) = 2^52 - 38; the rest are 2*(2^51 - 1) = 2^52 - 2.
24// The digits are grouped to show all thirteen nibbles of each 52-bit limb
25// rather than in fours, which would obscure the limb boundary.
26#[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    /// The additive identity.
37    pub const ZERO: Fe = Fe([0, 0, 0, 0, 0]);
38    /// The multiplicative identity.
39    pub const ONE: Fe = Fe([1, 0, 0, 0, 0]);
40
41    /// A small integer as a field element.
42    pub const fn from_u64(v: u64) -> Fe {
43        Fe([v & MASK, v >> 51, 0, 0, 0])
44    }
45
46    /// Field addition.
47    #[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    /// Field subtraction, via `self + 2p - other` so limbs stay non-negative.
57    #[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    /// Field negation.
67    #[inline]
68    pub fn neg(&self) -> Fe {
69        Fe::ZERO.sub(self)
70    }
71
72    /// One pass of carry propagation.
73    ///
74    /// Leaves every limb just above 2^51 rather than strictly below it -- each
75    /// keeps its own carry-in, and limb 0 keeps nineteen times the one that
76    /// came off the top. That is well inside the 2^52 a multiplication
77    /// accepts, and not paying a second pass to tidy it is the point.
78    #[inline]
79    fn weak_reduce(self) -> Fe {
80        // The five carries taken at once, for the same reason `carry_reduce`
81        // takes its at once: walking the limbs in order puts all five on one
82        // dependency chain, and nothing about the reduction requires that.
83        //
84        // This one runs on every `sub`, and `sub` is three of the operations
85        // in a point doubling, so it is on the hot path of every scalar
86        // multiplication on this curve.
87        //
88        // The bound: `sub` forms `self + 2p - other` with limbs below 2^52, so
89        // each input limb is below 2^53 and each carry below 4. Nineteen times
90        // the top carry is under 80, which is why the scaled carry cannot
91        // overflow here even without a second pass.
92        let r = self.0;
93        let c = [r[0] >> 51, r[1] >> 51, r[2] >> 51, r[3] >> 51, r[4] >> 51];
94        Fe([
95            (r[0] & MASK) + c[4].wrapping_mul(19),
96            (r[1] & MASK) + c[0],
97            (r[2] & MASK) + c[1],
98            (r[3] & MASK) + c[2],
99            (r[4] & MASK) + c[3],
100        ])
101    }
102
103    /// Field multiplication.
104    // What has already been tried on the field multiply, so it is not tried
105    // again. Curve25519 is doublings almost entirely and a doubling is four
106    // multiplications and four squarings, so this function is most of Ed25519
107    // and X25519.
108    //
109    // - **`#[inline(always)]` on `mul`, `square`, `add`, `sub` and
110    //   `carry_reduce`.** The per-crate assembly shows `Point::double` calling
111    //   them rather than inlining, with a 480-byte frame, which looks like the
112    //   problem. Measured: Ed25519 verify went from 1.59x behind dalek to
113    //   1.96x, reproducibly. The bodies are large enough that forcing them
114    //   inline costs more in spills than the calls cost. Under the benchmark's
115    //   profile LLVM already inlines them; the out-of-line copies are for other
116    //   callers.
117    // - **Removing `weak_reduce` from `sub`.** Also looks like waste: a
118    //   doubling does five of them, each a serial carry chain. Measured with
119    //   the diagnostic in `ed25519.rs`: `sub` is 2.14ns against `add` at
120    //   0.92ns, so all five are about 6ns of a 96ns doubling. Not where the
121    //   time is.
122    //
123    // - **A four-wide AVX2 field multiply.** This one was built and measured
124    //   rather than argued about, because the primitive really is faster:
125    //   24.35ns for four multiplications against 43.83ns for four scalar ones,
126    //   1.80x, verified lane for lane against `mul`. AVX2 has no 64x64
127    //   multiply, so it needs ten limbs alternating 26 and 25 bits instead of
128    //   five of 51 -- ten of those come to exactly 255, which keeps the
129    //   reduction constant at 19 -- and the products were generated from that
130    //   layout and checked against integer arithmetic mod p before any of it
131    //   was written.
132    //
133    //   It was not kept, because the primitive is not the point layer. 77% of
134    //   a doubling is its four multiplications and four squarings; the rest is
135    //   additions and subtractions across coordinates, and in a four-lane
136    //   layout those do not disappear, they become lane shuffles. Two `mul4`
137    //   calls are 48.7ns against the 73.8ns of field work they replace, so
138    //   even a shuffle cost of 20ns -- optimistic -- leaves a doubling at
139    //   68.7ns against 96ns, and verification at about 24us against dalek's
140    //   19.8. It narrows the gap and does not close it, and it would cost this
141    //   crate its `forbid(unsafe_code)`, which no other crate here has.
142    //
143    //   What dalek gets from its own AVX2 backend is 19.8 -> 16.6us, so the
144    //   target moves too. Closing this properly means their whole point layer,
145    //   not a faster multiply.
146    //
147    // Where the time is: 160 `mul` instructions in a doubling against 775
148    // `mov`. Five `u128` accumulators and two five-limb operands do not fit in
149    // sixteen registers, so the schoolbook spills, and that is a property of
150    // the shape rather than of the instruction selection. `-C
151    // target-cpu=native` turns the `mul`s into `mulx` and drops the moves to
152    // 475 -- 18% fewer instructions -- and buys 2% end to end. dalek gains 16%
153    // from the same flag, because it has an AVX2 field backend that engages
154    // there and this does not. That backend, not scheduling, is the remaining
155    // gap on Ed25519.
156    #[inline]
157    pub fn mul(&self, other: &Fe) -> Fe {
158        let a = &self.0;
159        let b = &other.0;
160
161        // The 19s are computed in 64 bits, and every product below is
162        // `u64 * u64 -> u128`.
163        //
164        // They used to be `(b[i] as u128) * 19`, which makes the scaled value a
165        // 128-bit quantity with no bound the compiler can see under 2^64 -- so
166        // each product became a full 128x128 multiply, three instructions and
167        // some adds where one `mul` would do. A limb is below 2^52 and 19 times
168        // it is below 2^57, so the scaling belongs in 64 bits and the compiler
169        // can then see that both operands of every product fit.
170        let b1_19 = b[1] * 19;
171        let b2_19 = b[2] * 19;
172        let b3_19 = b[3] * 19;
173        let b4_19 = b[4] * 19;
174
175        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);
176        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);
177        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);
178        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);
179        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]);
180
181        carry_reduce(r0, r1, r2, r3, r4)
182    }
183
184    /// Field squaring.
185    #[inline]
186    pub fn square(&self) -> Fe {
187        // Twenty-five limb products become fifteen.
188        //
189        // In `a * b` every pair `(i, j)` is distinct, so all twenty-five
190        // appear. Squaring pairs `i` with `j` and `j` with `i` to the same
191        // product, so each off-diagonal term is computed once and doubled --
192        // which is a shift, not a multiplication. The 19s are the same
193        // reduction the general multiply uses, folding `2^255 = 19` back into
194        // the low limbs, pre-multiplied into the doubled coefficients where
195        // both apply.
196        //
197        // This is on the hot path everywhere: four squarings per Edwards
198        // doubling, four per Montgomery ladder step, and a few hundred in a
199        // field inversion.
200        // Scalings in 64 bits, products as `u64 * u64 -> u128`; see `mul`.
201        // A limb is below 2^52, so 38 times it is below 2^58.
202        let a = &self.0;
203        let a0_2 = a[0] * 2;
204        let a1_2 = a[1] * 2;
205        let a1_38 = a[1] * 38;
206        let a2_38 = a[2] * 38;
207        let a3_38 = a[3] * 38;
208        let a3_19 = a[3] * 19;
209        let a4_19 = a[4] * 19;
210
211        // r_k = sum_{i+j=k} a_i a_j + 19 * sum_{i+j=k+5} a_i a_j
212        let r0 = m(a[0], a[0]) + m(a1_38, a[4]) + m(a2_38, a[3]);
213        let r1 = m(a0_2, a[1]) + m(a2_38, a[4]) + m(a3_19, a[3]);
214        let r2 = m(a0_2, a[2]) + m(a[1], a[1]) + m(a3_38, a[4]);
215        let r3 = m(a0_2, a[3]) + m(a1_2, a[2]) + m(a4_19, a[4]);
216        let r4 = m(a0_2, a[4]) + m(a1_2, a[3]) + m(a[2], a[2]);
217
218        carry_reduce(r0, r1, r2, r3, r4)
219    }
220
221    /// Repeated squaring, `self^(2^n)`.
222    #[inline]
223    pub fn square_n(&self, n: usize) -> Fe {
224        let mut r = *self;
225        for _ in 0..n {
226            r = r.square();
227        }
228        r
229    }
230
231    /// Multiplication by the Montgomery ladder constant `a24 = 121666`.
232    #[inline]
233    pub fn mul121666(&self) -> Fe {
234        let m = |x: u64| (x as u128) * 121_666;
235        let a = &self.0;
236        carry_reduce(m(a[0]), m(a[1]), m(a[2]), m(a[3]), m(a[4]))
237    }
238
239    /// Multiplicative inverse, `self^(p-2)`, with `inverse(0) == 0`.
240    ///
241    /// Uses the standard addition chain: 254 squarings and 11 multiplications.
242    pub fn invert(&self) -> Fe {
243        let z2 = self.square();
244        let z9 = z2.square_n(2).mul(self);
245        let z11 = z9.mul(&z2);
246        let z2_5_0 = z11.square().mul(&z9);
247        let z2_10_0 = z2_5_0.square_n(5).mul(&z2_5_0);
248        let z2_20_0 = z2_10_0.square_n(10).mul(&z2_10_0);
249        let z2_40_0 = z2_20_0.square_n(20).mul(&z2_20_0);
250        let z2_50_0 = z2_40_0.square_n(10).mul(&z2_10_0);
251        let z2_100_0 = z2_50_0.square_n(50).mul(&z2_50_0);
252        let z2_200_0 = z2_100_0.square_n(100).mul(&z2_100_0);
253        let z2_250_0 = z2_200_0.square_n(50).mul(&z2_50_0);
254        z2_250_0.square_n(5).mul(&z11)
255    }
256
257    /// `self^((p-5)/8)`, the exponent used to take square roots.
258    pub fn pow22523(&self) -> Fe {
259        let z2 = self.square();
260        let z9 = z2.square_n(2).mul(self);
261        let z11 = z9.mul(&z2);
262        let z2_5_0 = z11.square().mul(&z9);
263        let z2_10_0 = z2_5_0.square_n(5).mul(&z2_5_0);
264        let z2_20_0 = z2_10_0.square_n(10).mul(&z2_10_0);
265        let z2_40_0 = z2_20_0.square_n(20).mul(&z2_20_0);
266        let z2_50_0 = z2_40_0.square_n(10).mul(&z2_10_0);
267        let z2_100_0 = z2_50_0.square_n(50).mul(&z2_50_0);
268        let z2_200_0 = z2_100_0.square_n(100).mul(&z2_100_0);
269        let z2_250_0 = z2_200_0.square_n(50).mul(&z2_50_0);
270        z2_250_0.square_n(2).mul(self)
271    }
272
273    /// Decode 32 little-endian bytes, ignoring the top bit as RFC 7748 requires.
274    pub fn from_bytes(bytes: &[u8; 32]) -> Fe {
275        let load = |i: usize| -> u64 {
276            let mut v = [0u8; 8];
277            v.copy_from_slice(&bytes[i..i + 8]);
278            u64::from_le_bytes(v)
279        };
280        let l0 = load(0) & MASK;
281        let l1 = (load(6) >> 3) & MASK;
282        let l2 = (load(12) >> 6) & MASK;
283        let l3 = (load(19) >> 1) & MASK;
284        let l4 = (load(24) >> 12) & MASK;
285        Fe([l0, l1, l2, l3, l4])
286    }
287
288    /// Encode as 32 little-endian bytes, fully reduced modulo p.
289    pub fn to_bytes(&self) -> [u8; 32] {
290        // Three passes leave every limb strictly below 2^51: each pass can
291        // push at most a 19 back into limb 0, so the residue shrinks each time.
292        let mut t = self.weak_reduce().weak_reduce().weak_reduce().0;
293
294        // Conditionally subtract p, in constant time.
295        // q is 1 exactly when t >= p.
296        let mut q = (t[0] + 19) >> 51;
297        for i in 1..5 {
298            q = (t[i] + q) >> 51;
299        }
300        t[0] += 19 * q;
301        let mut carry = t[0] >> 51;
302        t[0] &= MASK;
303        for i in 1..5 {
304            t[i] += carry;
305            carry = t[i] >> 51;
306            t[i] &= MASK;
307        }
308        // Drop the bit that overflowed past 2^255.
309        t[4] &= (1 << 51) - 1;
310
311        let mut out = [0u8; 32];
312        let mut acc: u128 = 0;
313        let mut acc_bits = 0usize;
314        let mut idx = 0usize;
315        for limb in t.iter() {
316            acc |= (*limb as u128) << acc_bits;
317            acc_bits += 51;
318            while acc_bits >= 8 && idx < 32 {
319                out[idx] = acc as u8;
320                acc >>= 8;
321                acc_bits -= 8;
322                idx += 1;
323            }
324        }
325        while idx < 32 {
326            out[idx] = acc as u8;
327            acc >>= 8;
328            idx += 1;
329        }
330        out
331    }
332
333    /// Constant-time conditional swap.
334    #[inline]
335    pub fn cswap(a: &mut Fe, b: &mut Fe, choice: Choice) {
336        let mask = (choice.unwrap_u8() as u64).wrapping_neg();
337        for i in 0..5 {
338            let t = mask & (a.0[i] ^ b.0[i]);
339            a.0[i] ^= t;
340            b.0[i] ^= t;
341        }
342    }
343
344    /// Constant-time conditional move: `a = b` when `choice` is true.
345    #[inline]
346    pub fn cmov(a: &mut Fe, b: &Fe, choice: Choice) {
347        let mask = (choice.unwrap_u8() as u64).wrapping_neg();
348        for i in 0..5 {
349            a.0[i] ^= mask & (a.0[i] ^ b.0[i]);
350        }
351    }
352
353    /// Constant-time test for zero.
354    pub fn is_zero(&self) -> Choice {
355        ic_core::ct::is_zero(&self.to_bytes())
356    }
357
358    /// Constant-time equality.
359    pub fn ct_eq(&self, other: &Fe) -> Choice {
360        ic_core::ct::eq(&self.to_bytes(), &other.to_bytes())
361    }
362
363    /// The least significant bit of the canonical encoding — the "sign" used by
364    /// Edwards point compression.
365    pub fn is_negative(&self) -> Choice {
366        Choice::from_u8(self.to_bytes()[0] & 1)
367    }
368}
369
370/// One 64x64 multiplication, widened.
371///
372/// Written out so both operands are visibly `u64`: that is what lets the
373/// compiler emit a single widening multiply instead of a 128-bit one.
374#[inline(always)]
375fn m(x: u64, y: u64) -> u128 {
376    (x as u128) * (y as u128)
377}
378
379/// Fold five 128-bit products back into 51-bit limbs.
380#[inline]
381#[allow(clippy::too_many_arguments)]
382fn carry_reduce(r0: u128, r1: u128, r2: u128, r3: u128, r4: u128) -> Fe {
383    // Taken as five values rather than an array: an array is passed by value,
384    // which is eighty bytes through the stack on every multiplication and
385    // squaring. Measured at 22.3 -> 21.1 microseconds on the double-scalar
386    // multiplication, which is the only reason the signature is this shape.
387    let r = [r0, r1, r2, r3, r4];
388    // The five carries are extracted independently, not chained.
389    //
390    // This used to walk the limbs in order, each iteration adding the previous
391    // carry before computing its own -- five 128-bit shift-and-mask steps on a
392    // single dependency chain, on the hot path of every multiplication and
393    // squaring. Nothing about the reduction requires that order: each `r[i]`
394    // already holds its full product sum, so every carry can be taken at once
395    // and delivered to its neighbour afterwards, which is five independent
396    // shifts the scheduler can overlap instead of five it cannot.
397    //
398    // The bound that makes it safe: both operands of a product have limbs
399    // below 2^52, so `r[i] < 5 * 19 * 2^104 < 2^110.3` and `c[i] < 2^59.3`.
400    // The largest quantity below is `c[4] * 19 < 2^63.6`, which is why the
401    // scaled carry still fits in a u64. `limbs_at_their_maximum_do_not_carry_
402    // out_of_a_u64` drives that worst case, and an arithmetic overflow there
403    // is a panic in a debug build rather than a wrong answer in a release one.
404    let c: [u64; 5] = [
405        (r[0] >> 51) as u64,
406        (r[1] >> 51) as u64,
407        (r[2] >> 51) as u64,
408        (r[3] >> 51) as u64,
409        (r[4] >> 51) as u64,
410    ];
411    let mut out: [u64; 5] = [
412        (r[0] as u64 & MASK) + c[4] * 19,
413        (r[1] as u64 & MASK) + c[0],
414        (r[2] as u64 & MASK) + c[1],
415        (r[3] as u64 & MASK) + c[2],
416        (r[4] as u64 & MASK) + c[3],
417    ];
418
419    // One short 64-bit pass to settle what those additions carried. Serial,
420    // but over small numbers and only once.
421    let mut carry = out[0] >> 51;
422    out[0] &= MASK;
423    for slot in out.iter_mut().skip(1) {
424        *slot += carry;
425        carry = *slot >> 51;
426        *slot &= MASK;
427    }
428    out[0] += carry * 19;
429    Fe(out)
430}
431
432#[cfg(test)]
433mod tests {
434    use super::*;
435
436    /// Squaring must agree with multiplying a value by itself.
437    ///
438    /// The dedicated formula reaches the same answer by a different route --
439    /// fifteen products where the general one has twenty-five, with the
440    /// off-diagonal terms doubled rather than recomputed -- so agreement is
441    /// the whole correctness argument. `mul` is what the RFC 7748 and RFC 8032
442    /// vectors validate, which makes it the oracle.
443    ///
444    /// The values include zero, one, the largest limbs the representation
445    /// holds unreduced, and values that carry out of every limb, because the
446    /// doubling is where this formula can overflow if the bounds are wrong.
447    #[test]
448    fn squaring_agrees_with_multiplication() {
449        let mut cases = std::vec![
450            Fe::ZERO,
451            Fe::ONE,
452            Fe([1, 1, 1, 1, 1]),
453            Fe([(1u64 << 51) - 1; 5]),
454            Fe([(1u64 << 51) - 1, 0, (1u64 << 51) - 1, 0, (1u64 << 51) - 1]),
455            Fe([0, (1u64 << 51) - 1, 0, (1u64 << 51) - 1, 0]),
456        ];
457        // And a spread of pseudo-random field elements.
458        let mut x = Fe([0x51a2, 0x9e37, 0x79b9, 0x7f4a, 0x7c15]);
459        for _ in 0..16 {
460            x = x.mul(&Fe([3, 5, 7, 11, 13])).add(&Fe::ONE);
461            cases.push(x);
462        }
463
464        let mut checked = 0;
465        for f in &cases {
466            assert_eq!(
467                f.square().to_bytes(),
468                f.mul(f).to_bytes(),
469                "square and mul-by-self differ"
470            );
471            checked += 1;
472        }
473        assert_eq!(checked, 22, "the comparison did not run");
474    }
475
476    fn fe(v: u64) -> Fe {
477        Fe::from_u64(v)
478    }
479
480    #[test]
481    fn encode_decode_roundtrip() {
482        for v in [0u64, 1, 2, 19, 1 << 51, u64::MAX] {
483            let a = fe(v);
484            assert_eq!(Fe::from_bytes(&a.to_bytes()).to_bytes(), a.to_bytes());
485        }
486    }
487
488    #[test]
489    fn small_arithmetic() {
490        assert_eq!(fe(2).add(&fe(3)).to_bytes(), fe(5).to_bytes());
491        assert_eq!(fe(5).sub(&fe(3)).to_bytes(), fe(2).to_bytes());
492        assert_eq!(fe(6).mul(&fe(7)).to_bytes(), fe(42).to_bytes());
493        assert_eq!(fe(9).square().to_bytes(), fe(81).to_bytes());
494    }
495
496    #[test]
497    fn subtraction_wraps_into_the_field() {
498        // 0 - 1 == p - 1, whose encoding is ec ff .. 7f.
499        let r = Fe::ZERO.sub(&Fe::ONE).to_bytes();
500        assert_eq!(r[0], 0xec);
501        assert_eq!(r[31], 0x7f);
502        for b in &r[1..31] {
503            assert_eq!(*b, 0xff);
504        }
505    }
506
507    #[test]
508    fn p_encodes_as_zero() {
509        // p itself must reduce to 0.
510        let mut p_bytes = [0xffu8; 32];
511        p_bytes[0] = 0xed;
512        p_bytes[31] = 0x7f;
513        assert_eq!(Fe::from_bytes(&p_bytes).to_bytes(), [0u8; 32]);
514    }
515
516    #[test]
517    fn inversion_is_correct() {
518        for v in [1u64, 2, 3, 19, 12345, u32::MAX as u64] {
519            let a = fe(v);
520            assert_eq!(a.mul(&a.invert()).to_bytes(), Fe::ONE.to_bytes(), "1/{v}");
521        }
522        assert_eq!(Fe::ZERO.invert().to_bytes(), [0u8; 32]);
523    }
524
525    #[test]
526    fn multiplication_is_associative_and_distributive() {
527        let a = Fe::from_bytes(&[0x11; 32]);
528        let b = Fe::from_bytes(&[0x7a; 32]);
529        let c = Fe::from_bytes(&[0xc3; 32]);
530        assert_eq!(a.mul(&b).mul(&c).to_bytes(), a.mul(&b.mul(&c)).to_bytes());
531        assert_eq!(
532            a.mul(&b.add(&c)).to_bytes(),
533            a.mul(&b).add(&a.mul(&c)).to_bytes()
534        );
535    }
536
537    #[test]
538    fn pow22523_gives_a_square_root() {
539        // For a square x, x^((p-5)/8) * x is a square root up to a factor of i.
540        let x = fe(4);
541        let r = x.pow22523().mul(&x);
542        let sq = r.square();
543        // r^2 is either x or -x.
544        assert!(
545            bool::from(sq.ct_eq(&x)) || bool::from(sq.ct_eq(&x.neg())),
546            "square root property"
547        );
548    }
549
550    #[test]
551    fn cswap_and_cmov_are_conditional() {
552        let mut a = fe(1);
553        let mut b = fe(2);
554        Fe::cswap(&mut a, &mut b, Choice::FALSE);
555        assert_eq!(a.to_bytes(), fe(1).to_bytes());
556        Fe::cswap(&mut a, &mut b, Choice::TRUE);
557        assert_eq!(a.to_bytes(), fe(2).to_bytes());
558
559        let mut c = fe(5);
560        Fe::cmov(&mut c, &fe(9), Choice::FALSE);
561        assert_eq!(c.to_bytes(), fe(5).to_bytes());
562        Fe::cmov(&mut c, &fe(9), Choice::TRUE);
563        assert_eq!(c.to_bytes(), fe(9).to_bytes());
564    }
565
566    #[test]
567    fn high_bit_of_input_is_ignored() {
568        let mut a = [0x42u8; 32];
569        let mut b = a;
570        a[31] &= 0x7f;
571        b[31] |= 0x80;
572        assert_eq!(Fe::from_bytes(&a).to_bytes(), Fe::from_bytes(&b).to_bytes());
573    }
574
575    /// The serial carry chain `carry_reduce` replaced, kept as an oracle.
576    ///
577    /// It is the implementation the RFC 7748 and 8032 vectors were passing
578    /// against before the parallel form went in, so agreeing with it on
579    /// arbitrary limb patterns is the evidence that the rewrite changed the
580    /// schedule and not the arithmetic.
581    fn carry_reduce_serial(r: [u128; 5]) -> Fe {
582        let mut out = [0u64; 5];
583        let mut carry: u128 = 0;
584        for (slot, limb) in out.iter_mut().zip(r) {
585            let v = limb + carry;
586            carry = v >> 51;
587            *slot = (v & MASK as u128) as u64;
588        }
589        out[0] += (carry as u64) * 19;
590        let mut c = out[0] >> 51;
591        out[0] &= MASK;
592        for slot in out.iter_mut().skip(1) {
593            *slot += c;
594            c = *slot >> 51;
595            *slot &= MASK;
596        }
597        out[0] += c * 19;
598        Fe(out)
599    }
600
601    /// Build the limb products the way `mul` does, without reducing.
602    fn raw_products(a: &[u64; 5], b: &[u64; 5]) -> [u128; 5] {
603        let m = |x: u64, y: u64| (x as u128) * (y as u128);
604        let (b1, b2, b3, b4) = (b[1] * 19, b[2] * 19, b[3] * 19, b[4] * 19);
605        [
606            m(a[0], b[0]) + m(a[1], b4) + m(a[2], b3) + m(a[3], b2) + m(a[4], b1),
607            m(a[0], b[1]) + m(a[1], b[0]) + m(a[2], b4) + m(a[3], b3) + m(a[4], b2),
608            m(a[0], b[2]) + m(a[1], b[1]) + m(a[2], b[0]) + m(a[3], b4) + m(a[4], b3),
609            m(a[0], b[3]) + m(a[1], b[2]) + m(a[2], b[1]) + m(a[3], b[0]) + m(a[4], b4),
610            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]),
611        ]
612    }
613
614    /// The worst case the safety argument rests on.
615    ///
616    /// `add` does not reduce, so a limb reaching `mul` can be as large as
617    /// `2^52 - 2`. Every limb is put there at once, which maximises every
618    /// product sum simultaneously -- a state the curve itself may never reach,
619    /// which is the point of testing it rather than arguing about it. In a
620    /// debug build the scaled carry overflowing a u64 panics here.
621    #[test]
622    fn limbs_at_their_maximum_do_not_carry_out_of_a_u64() {
623        let max = [(1u64 << 52) - 2; 5];
624        let r = raw_products(&max, &max);
625        assert_eq!(
626            carry_reduce(r[0], r[1], r[2], r[3], r[4]).to_bytes(),
627            carry_reduce_serial(r).to_bytes(),
628            "parallel and serial carry disagree at the limb maximum"
629        );
630    }
631
632    /// The two carry forms agree on arbitrary limb patterns.
633    #[test]
634    fn parallel_carry_agrees_with_the_serial_one() {
635        let mut state = 0x243f_6a88_85a3_08d3u64;
636        let mut next = || {
637            // xorshift64*, enough to walk the limb space without a dependency.
638            state ^= state >> 12;
639            state ^= state << 25;
640            state ^= state >> 27;
641            state.wrapping_mul(0x2545_f491_4f6c_dd1d)
642        };
643        for _ in 0..20_000 {
644            let mut a = [0u64; 5];
645            let mut b = [0u64; 5];
646            for i in 0..5 {
647                // The full range `mul` promises to accept, endpoints included.
648                a[i] = next() % (1 << 52);
649                b[i] = next() % (1 << 52);
650            }
651            let r = raw_products(&a, &b);
652            assert_eq!(
653                carry_reduce(r[0], r[1], r[2], r[3], r[4]).to_bytes(),
654                carry_reduce_serial(r).to_bytes(),
655                "disagreement on a={a:?} b={b:?}"
656            );
657        }
658    }
659}