Skip to main content

primitives/algebra/field/mersenne/
m107.rs

1use std::{
2    iter::{Product, Sum},
3    ops::{Add, AddAssign, Mul, MulAssign, Neg, Sub, SubAssign},
4};
5
6use crypto_bigint::rand_core::RngCore;
7use ff::{Field, PrimeField};
8use hybrid_array::Array;
9use rand::Rng;
10use serde::{Deserialize, Serialize};
11use subtle::{Choice, ConditionallySelectable, ConstantTimeEq, CtOption};
12use typenum::{U1, U14, U16};
13use wincode::{SchemaRead, SchemaWrite};
14
15use crate::{
16    algebra::{
17        field::FieldExtension,
18        ops::{AccReduce, ReduceWide},
19        uniform_bytes::FromUniformBytes,
20    },
21    random::{CryptoRngCore, Random},
22    types::{HeapArray, Positive},
23};
24
25mod ff_impl {
26    use ff::PrimeField;
27    use serde::{Deserialize, Serialize};
28
29    #[derive(PrimeField, Serialize, Deserialize)]
30    #[PrimeFieldModulus = "162259276829213363391578010288127"]
31    #[PrimeFieldGenerator = "3"]
32    #[PrimeFieldReprEndianness = "little"]
33    pub struct Mersenne107FF([u64; 2]);
34}
35
36#[derive(
37    Copy, Clone, Default, Debug, SchemaRead, SchemaWrite, PartialEq, Eq, Hash, Ord, PartialOrd,
38)]
39#[repr(C)]
40pub struct Mersenne107(pub(super) u128);
41
42impl Serialize for Mersenne107 {
43    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
44    where
45        S: serde::Serializer,
46    {
47        self.as_le_array().serialize(serializer)
48    }
49}
50
51impl<'de> Deserialize<'de> for Mersenne107 {
52    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
53    where
54        D: serde::Deserializer<'de>,
55    {
56        let arr = <[u8; 14]>::deserialize(deserializer)?;
57        Self::from_canonical_bytes(&arr).ok_or_else(|| {
58            serde::de::Error::custom("Invalid Mersenne107 canonical byte representation")
59        })
60    }
61}
62
63impl Mersenne107 {
64    pub const NUM_BITS: usize = 107;
65    pub const MODULUS: u128 = (1u128 << Self::NUM_BITS) - 1;
66    pub const MAX: u128 = Self::MODULUS - 1;
67
68    fn as_le_array(&self) -> [u8; 14] {
69        let mut arr = [0u8; 14];
70        arr[..14].copy_from_slice(&self.0.to_le_bytes()[..14]);
71        arr
72    }
73
74    fn from_canonical_bytes(arr: &[u8; 14]) -> Option<Self> {
75        let mut tmp = [0u8; 16];
76        tmp[..14].copy_from_slice(arr);
77        let val = u128::from_le_bytes(tmp);
78        (val < Self::MODULUS).then_some(Self(val))
79    }
80}
81
82///////////////////////////////////////////////////////////////////////////////////////////////////
83// Arithmetic ops - multiplication
84///////////////////////////////////////////////////////////////////////////////////////////////////
85
86#[macros::op_variants(owned)]
87impl<'a> MulAssign<&'a Mersenne107> for Mersenne107 {
88    #[inline]
89    fn mul_assign(&mut self, rhs: &'a Mersenne107) {
90        self.0 = super::m107_ops::mul(self.0, rhs.0);
91    }
92}
93
94#[macros::op_variants(owned)]
95impl<'a> Mul<&'a Mersenne107> for Mersenne107 {
96    type Output = Self;
97    #[inline]
98    fn mul(self, rhs: &'a Mersenne107) -> Self::Output {
99        let mut res = self;
100        res.mul_assign(rhs);
101        res
102    }
103}
104
105///////////////////////////////////////////////////////////////////////////////////////////////////
106// Arithmetic ops - addition
107///////////////////////////////////////////////////////////////////////////////////////////////////
108
109#[macros::op_variants(owned)]
110impl<'a> AddAssign<&'a Mersenne107> for Mersenne107 {
111    #[inline]
112    fn add_assign(&mut self, rhs: &'a Mersenne107) {
113        self.0 += rhs.0;
114        super::m107_ops::reduce_mod_1bit_inplace(&mut self.0);
115    }
116}
117#[macros::op_variants(owned)]
118impl<'a> Add<&'a Mersenne107> for Mersenne107 {
119    type Output = Self;
120
121    #[inline]
122    fn add(self, rhs: &'a Mersenne107) -> Self::Output {
123        let mut res = self;
124        res.add_assign(rhs);
125        res
126    }
127}
128
129///////////////////////////////////////////////////////////////////////////////////////////////////
130// Arithmetic ops - substraction
131///////////////////////////////////////////////////////////////////////////////////////////////////
132
133#[macros::op_variants(owned)]
134impl<'a> SubAssign<&'a Mersenne107> for Mersenne107 {
135    #[inline]
136    fn sub_assign(&mut self, rhs: &'a Mersenne107) {
137        self.0 += Self::MODULUS - rhs.0;
138        super::m107_ops::reduce_mod_1bit_inplace(&mut self.0);
139    }
140}
141
142#[macros::op_variants(owned)]
143impl<'a> Sub<&'a Mersenne107> for Mersenne107 {
144    type Output = Self;
145
146    #[inline]
147    fn sub(mut self, rhs: &'a Mersenne107) -> Self::Output {
148        self.sub_assign(rhs);
149        self
150    }
151}
152
153///////////////////////////////////////////////////////////////////////////////////////////////////
154// Arithmetic ops - negation
155///////////////////////////////////////////////////////////////////////////////////////////////////
156
157#[macros::op_variants(borrowed)]
158impl Neg for Mersenne107 {
159    type Output = Mersenne107;
160
161    fn neg(self) -> Self::Output {
162        Self(super::m107_ops::reduce_mod_1bit(Self::MODULUS - self.0))
163    }
164}
165
166///////////////////////////////////////////////////////////////////////////////////////////////////
167// Constant time
168///////////////////////////////////////////////////////////////////////////////////////////////////
169
170impl ConditionallySelectable for Mersenne107 {
171    fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
172        Self(u128::conditional_select(&a.0, &b.0, choice))
173    }
174}
175
176impl ConstantTimeEq for Mersenne107 {
177    fn ct_eq(&self, other: &Self) -> Choice {
178        self.0.ct_eq(&other.0)
179    }
180}
181
182///////////////////////////////////////////////////////////////////////////////////////////////////
183// Iterator operations
184///////////////////////////////////////////////////////////////////////////////////////////////////
185
186impl Sum for Mersenne107 {
187    fn sum<I: Iterator<Item = Self>>(iter: I) -> Self {
188        let t = iter.fold(<Self as AccReduce>::zero_wide(), |mut acc, x| {
189            Self::acc(&mut acc, &x);
190            acc
191        });
192        Self::reduce_mod_order(t)
193    }
194}
195
196impl<'a> Sum<&'a Self> for Mersenne107 {
197    fn sum<I: Iterator<Item = &'a Self>>(iter: I) -> Self {
198        let t = iter.fold(<Self as AccReduce>::zero_wide(), |mut acc, x| {
199            Self::acc(&mut acc, x);
200            acc
201        });
202        Self::reduce_mod_order(t)
203    }
204}
205
206impl<'a> Product<&'a Self> for Mersenne107 {
207    fn product<I: Iterator<Item = &'a Self>>(iter: I) -> Self {
208        iter.fold(Self::ONE, |acc, x| acc * x)
209    }
210}
211
212impl Product for Mersenne107 {
213    fn product<I: Iterator<Item = Self>>(iter: I) -> Self {
214        iter.fold(Self::ONE, |acc, x| acc * x)
215    }
216}
217
218// Fixed-exponent helpers used by the `sqrt`/`sqrt_ratio` closed forms (p ≡ 3 mod 4).
219impl Mersenne107 {
220    /// Returns `self^(2^n)`, i.e. `n` repeated squarings.
221    #[inline]
222    fn pow2_pow(self, n: u32) -> Self {
223        (0..n).fold(self, |acc, _| acc.square())
224    }
225
226    /// Returns `self^(2^k - 1)` via an Itoh–Tsujii-style addition chain: ≈`k` squarings and
227    /// `O(log k)` multiplications, versus the ~`2k` operations of plain square-and-multiply on the
228    /// all-ones exponent.
229    fn pow2_minus_1(self, k: u32) -> Self {
230        debug_assert!(k >= 1);
231        // Invariant: `acc == self^(2^cur - 1)`; grow `cur` from 1 to `k` following the bits of `k`.
232        let mut acc = self;
233        let mut cur = 1u32;
234        for i in (0..(u32::BITS - 1 - k.leading_zeros())).rev() {
235            // self^(2^{2·cur} - 1) = (self^(2^cur - 1))^(2^cur) · self^(2^cur - 1)
236            acc = acc.pow2_pow(cur) * acc;
237            cur *= 2;
238            if (k >> i) & 1 == 1 {
239                // self^(2^{cur+1} - 1) = (self^(2^cur - 1))^2 · self
240                acc = acc.square() * self;
241                cur += 1;
242            }
243        }
244        debug_assert_eq!(cur, k);
245        acc
246    }
247}
248
249///////////////////////////////////////////////////////////////////////////////////////////////////
250// Field trait
251///////////////////////////////////////////////////////////////////////////////////////////////////
252
253impl Field for Mersenne107 {
254    const ZERO: Self = Mersenne107(0);
255    const ONE: Self = Mersenne107(1);
256
257    fn random(mut rng: impl RngCore) -> Self {
258        let tmp = rng.gen::<u128>();
259        Self(super::m107_ops::reduce_mod(tmp)) // the probability is skewed because M107 modulus
260                                               // does not
261                                               // divide 2^128 exactly
262    }
263
264    fn square(&self) -> Self {
265        *self * self // TODO: optimize ?
266    }
267
268    fn double(&self) -> Self {
269        Self(super::m107_ops::reduce_mod_1bit(self.0 << 1))
270    }
271
272    fn invert(&self) -> CtOption<Self> {
273        // Fallback to ff implementation
274        // TODO: see if we can optimize this without ff
275        let val: ff_impl::Mersenne107FF = self.into();
276        let inv = val.invert();
277        inv.map(|v| v.into())
278    }
279
280    fn sqrt_ratio(num: &Self, div: &Self) -> (Choice, Self) {
281        // p = 2^107 - 1 ≡ 3 (mod 4), so sqrt(num/div) has the closed form
282        //     y = num·div·(num·div³)^((p-3)/4)
283        // computed with a single fixed exponentiation and no modular inversion (unlike the generic
284        // Tonelli–Shanks `sqrt_ratio`). `y` is a true square root of num/div iff y²·div == num. The
285        // exponent (p-3)/4 = 2^105 - 1.
286        let uv = *num * div;
287        let uv3 = div.square() * uv;
288        let y = uv * uv3.pow2_minus_1(105);
289        let is_qr = (y.square() * div).ct_eq(num);
290        (is_qr, y)
291    }
292
293    fn sqrt(&self) -> CtOption<Self> {
294        // p ≡ 3 (mod 4): sqrt(a) = a^((p+1)/4) = a^(2^105), i.e. 105 repeated squarings. The result
295        // is a valid square root iff it squares back to `a` (rejects non-residues; zero maps to
296        // zero).
297        let root = self.pow2_pow(105);
298        CtOption::new(root, root.square().ct_eq(self))
299    }
300}
301
302///////////////////////////////////////////////////////////////////////////////////////////////////
303// Field extension trait
304///////////////////////////////////////////////////////////////////////////////////////////////////
305
306impl FieldExtension for Mersenne107 {
307    type Subfield = Self;
308    type Degree = U1;
309    type FieldBitSize = typenum::U<{ Mersenne107::NUM_BITS }>;
310    type FieldBytesSize = U14;
311
312    fn to_subfield_elements(&self) -> Array<Self::Subfield, Self::Degree> {
313        Array([*self])
314    }
315
316    fn from_subfield_elements(elems: Array<Self::Subfield, Self::Degree>) -> Self {
317        elems[0]
318    }
319
320    fn to_le_bytes(&self) -> Array<u8, Self::FieldBytesSize> {
321        self.as_le_array().into()
322    }
323
324    fn from_le_bytes(bytes: &[u8]) -> Option<Self> {
325        if bytes.len() == 14 {
326            let arr: &[u8; 14] = bytes.try_into().expect("This should never fail");
327            Self::from_canonical_bytes(arr)
328        } else {
329            None
330        }
331    }
332
333    fn mul_by_subfield(&self, other: &Self::Subfield) -> Self {
334        *self * other
335    }
336
337    fn generator() -> Self {
338        Self(3u128) // 3^(2^107-2) == 1 mod 2^107-1
339    }
340}
341
342impl Random for Mersenne107 {
343    fn random(mut rng: impl CryptoRngCore) -> Self {
344        let tmp = rng.gen::<u128>();
345        Self(super::m107_ops::reduce_mod(tmp))
346    }
347
348    fn random_array<M: Positive>(mut rng: impl CryptoRngCore) -> HeapArray<Self, M> {
349        let mut buf = HeapArray::<Self, M>::default().into_box_bytes();
350        rng.fill_bytes(&mut buf);
351        let mut tmp = HeapArray::from_box_bytes(buf);
352        tmp.iter_mut()
353            .for_each(|v: &mut Self| super::m107_ops::reduce_mod_inplace(&mut v.0));
354        tmp
355    }
356}
357
358unsafe impl bytemuck::Zeroable for Mersenne107 {}
359unsafe impl bytemuck::Pod for Mersenne107 {}
360
361impl FromUniformBytes for Mersenne107 {
362    type UniformBytes = U16;
363    fn from_uniform_bytes(bytes: &hybrid_array::Array<u8, Self::UniformBytes>) -> Self {
364        let mut val = u128::from_le_bytes(bytes.0);
365        super::m107_ops::reduce_mod_inplace(&mut val);
366        Self(val)
367    }
368}
369
370impl From<u64> for Mersenne107 {
371    fn from(val: u64) -> Self {
372        Self(val as u128)
373    }
374}
375
376impl From<u128> for Mersenne107 {
377    fn from(val: u128) -> Self {
378        Self(super::m107_ops::reduce_mod(val))
379    }
380}
381
382impl From<ff_impl::Mersenne107FF> for Mersenne107 {
383    fn from(val: ff_impl::Mersenne107FF) -> Self {
384        Self::from_le_bytes(&val.to_repr().as_ref()[..14]).unwrap()
385    }
386}
387
388impl<'a> From<&'a Mersenne107> for ff_impl::Mersenne107FF {
389    fn from(val: &'a Mersenne107) -> Self {
390        Self::from_repr(ff_impl::Mersenne107FFRepr(val.0.to_le_bytes())).unwrap()
391    }
392}
393
394#[cfg(test)]
395mod test {
396    use ff::Field;
397    use num_bigint::BigInt;
398    use typenum::Unsigned;
399
400    use crate::{
401        algebra::field::{
402            mersenne::{m107::Mersenne107, test::bigint_to_m107},
403            FieldExtension,
404        },
405        random::test_rng,
406    };
407
408    type M = typenum::U1000;
409
410    #[test]
411    fn test_neg() {
412        fn test_internal(a: Mersenne107) {
413            let exp = bigint_to_m107(-BigInt::from(a.0));
414            let act = -a;
415            assert_eq!(exp, act, "a = {a:?}");
416        }
417
418        let mut rng = test_rng();
419        for _ in 0..M::to_usize() {
420            let a = Mersenne107::random(&mut rng);
421            test_internal(a);
422        }
423
424        // Corner cases
425        test_internal(Mersenne107::ZERO);
426        test_internal(Mersenne107::ONE);
427        test_internal(Mersenne107(Mersenne107::MAX));
428    }
429
430    #[test]
431    fn test_invert() {
432        fn test_internal(a: Mersenne107) {
433            let a_inv = a.invert().unwrap();
434            let act = a * a_inv;
435            assert_eq!(Mersenne107::ONE, act, "a = {a:?}");
436        }
437
438        let mut rng = test_rng();
439        for _ in 0..M::to_usize() {
440            let a = Mersenne107::random(&mut rng);
441            if a == Mersenne107::ZERO {
442                continue;
443            }
444            test_internal(a);
445        }
446
447        // Corner cases
448        test_internal(Mersenne107::ONE);
449        test_internal(Mersenne107(Mersenne107::MAX));
450    }
451
452    #[test]
453    fn test_sqrt() {
454        fn test_internal(a: Mersenne107) {
455            let a_sqrt = a.sqrt();
456            if a_sqrt.into_option().is_none() {
457                return;
458            }
459
460            let a_sqrt = a_sqrt.unwrap();
461            let act = a_sqrt * a_sqrt;
462            assert_eq!(a, act, "a = {a:?}");
463        }
464
465        let mut rng = test_rng();
466        for _ in 0..M::to_usize() {
467            let a = Mersenne107::random(&mut rng);
468            test_internal(a);
469        }
470
471        // Corner cases
472        test_internal(Mersenne107::ZERO);
473        test_internal(Mersenne107::ONE);
474        test_internal(Mersenne107(Mersenne107::MAX));
475    }
476
477    #[test]
478    fn test_sqrt_ratio() {
479        fn test_internal(num: Mersenne107, div: Mersenne107) {
480            let (is_qr, y) = Mersenne107::sqrt_ratio(&num, &div);
481
482            // QR-ness must agree with the generic ff reference implementation.
483            let num_ff: super::ff_impl::Mersenne107FF = (&num).into();
484            let div_ff: super::ff_impl::Mersenne107FF = (&div).into();
485            let (is_qr_ref, _) = super::ff_impl::Mersenne107FF::sqrt_ratio(&num_ff, &div_ff);
486            assert_eq!(
487                bool::from(is_qr),
488                bool::from(is_qr_ref),
489                "QR mismatch for num={num:?} div={div:?}"
490            );
491
492            // When num/div is a square, y must satisfy y² · div == num.
493            if bool::from(is_qr) {
494                assert_eq!(
495                    y.square() * div,
496                    num,
497                    "wrong root for num={num:?} div={div:?}"
498                );
499            }
500        }
501
502        let mut rng = test_rng();
503        for _ in 0..M::to_usize() {
504            let num = Mersenne107::random(&mut rng);
505            let div = Mersenne107::random(&mut rng);
506            if div == Mersenne107::ZERO {
507                continue;
508            }
509            test_internal(num, div);
510            // The daBits use case: inverse square root sqrt_ratio(1, y) == y^{-1/2}.
511            test_internal(Mersenne107::ONE, div);
512        }
513
514        // Corner cases
515        test_internal(Mersenne107::ONE, Mersenne107::ONE);
516        test_internal(Mersenne107::ZERO, Mersenne107::ONE);
517        test_internal(Mersenne107(Mersenne107::MAX), Mersenne107::ONE);
518    }
519
520    #[test]
521    fn test_canonical_bytes_decoding() {
522        let value = Mersenne107::from(123456789u64);
523        let bytes = value.to_le_bytes();
524        assert_eq!(Mersenne107::from_le_bytes(&bytes), Some(value));
525
526        // Encoding of the modulus p = 2^107 - 1 is non-canonical and must be rejected.
527        let modulus_bytes = [
528            0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x07,
529        ];
530        assert_eq!(Mersenne107::from_le_bytes(&modulus_bytes), None);
531    }
532
533    macro_rules! test_op {
534        ($op:tt) => {
535                fn test_internal(a: Mersenne107, b: Mersenne107) {
536                    let exp = bigint_to_m107(BigInt::from(a.0) $op BigInt::from(b.0));
537                    let act = a $op b;
538                    assert_eq!(exp, act, "a = {a:?}, b = {b:?}");
539                }
540
541                let mut rng = test_rng();
542                for _ in 0..M::to_usize() {
543                    let a = Mersenne107::random(&mut rng);
544                    let b = Mersenne107::random(&mut rng);
545                    test_internal(a, b);
546                }
547
548                // Corner cases
549                test_internal(Mersenne107::ZERO, Mersenne107::ZERO);
550                test_internal(Mersenne107::ZERO, Mersenne107::ONE);
551                test_internal(Mersenne107::ONE, Mersenne107::ZERO);
552                test_internal(Mersenne107::ONE, Mersenne107::ONE);
553                test_internal(Mersenne107::ZERO, Mersenne107(Mersenne107::MAX));
554                test_internal(Mersenne107::ONE, Mersenne107(Mersenne107::MAX));
555                test_internal(Mersenne107(Mersenne107::MAX), Mersenne107::ZERO);
556                test_internal(Mersenne107(Mersenne107::MAX), Mersenne107::ONE);
557                test_internal(Mersenne107(Mersenne107::MAX), Mersenne107(Mersenne107::MAX));
558            }
559    }
560
561    #[test]
562    fn test_mul() {
563        test_op!(*);
564    }
565    #[test]
566    fn test_add() {
567        test_op!(+);
568    }
569    #[test]
570    fn test_sub() {
571        test_op!(-);
572    }
573}