Skip to main content

primitives/algebra/field/mersenne/
m107.rs

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