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