Skip to main content

snarkrs_gpu_layout/
lib.rs

1//! The packed layouts that cross the Rust/GPU boundary, defined once for every backend.
2//!
3//! # One definition, four languages
4//!
5//! Rust here, MSL in `crates/metal/src/shaders/bn254_fr.metal`, CUDA C in
6//! `crates/gpu-kernels/src/kernels/bn254_fr.cuh`, and the WGSL that `crates/wgpu/src/gen`
7//! emits. Every `#[repr(C)]` type below has a byte-for-byte twin on each of those sides.
8//! Metal and CUDA retype the constants as kernel text, so each kernel side carries a
9//! `static_assert` on `sizeof` and each of those two backends keeps a test that greps its
10//! own source for the exact constant lines; this side carries a `const` assertion on
11//! `size_of`. `snarkrs-wgpu` reads the constants out of this crate at run time and generates
12//! its WGSL from them, so it has no copy that can drift, but it does depend on the struct
13//! strides and on the infinity sentinel. If you change a struct here you are changing a
14//! wire format three backends read; change the MSL and the CUDA header in the same commit.
15//!
16//! This crate exists so that the three backends cannot disagree about what a point is.
17//! They run on different hardware with different compilers, and a benchmark that compares
18//! them is only meaningful if they are fed bit-identical inputs.
19//!
20//! # Why repacking is not optional
21//!
22//! Measured on this workspace's arkworks version: `size_of::<ark_ec::G1Affine>()` is 72,
23//! not 64, because the affine type carries an `infinity: bool` plus padding, and
24//! `size_of::<ark_ec::G2Affine>()` is 136, not 128. Neither `ark_ff::Fp` nor the affine
25//! types are `repr(C)`, so their field order is not even guaranteed. Byte-casting a
26//! `Vec<G1Affine>` into a device buffer hands the GPU a 72-byte stride while the kernel
27//! reads 64, and every point after the first is garbage. There is no way around an
28//! explicit repack, so this module makes the repack the only path.
29//!
30//! # Montgomery form, and the asymmetry that silently breaks proofs
31//!
32//! arkworks stores an `Fp` internally in Montgomery form with `R = 2^256`, which is
33//! exactly the radix both GPU CIOS routines use. So:
34//!
35//! * [`PackedFr`] and [`PackedFq`] hold **Montgomery** limbs, copied straight out of
36//!   `Fp::0.0` with no conversion on either side. This is what field kernels want: the
37//!   NTT, the coset shift and `H = A*B - C` are all multiplications and additions, and
38//!   Montgomery form is closed under both.
39//! * [`PackedScalar`] holds **standard** limbs, via `into_bigint()`. This is what an MSM
40//!   wants, because Pippenger slices a scalar into window digits and a window digit of a
41//!   Montgomery representative is a digit of `a*R mod n`, which is a different number.
42//!
43//! Getting this backwards produces a proof that is wrong by a factor of `R` and fails
44//! verification with nothing else to go on, so the two types are deliberately distinct
45//! and neither converts into the other.
46use core::mem::{align_of, size_of};
47
48use ark_ff::{BigInt, PrimeField};
49use snarkrs_field::{Fq, Fq2, Fr, G1Affine, G2Affine};
50
51/// The GLV twiddle table behind `ptau prepare`'s group FFT: the 36-byte decomposed-scalar
52/// layout and the host-side lattice work that fills it, shared by the Metal and CUDA FFT
53/// drivers so the two cannot disagree about what a twiddle is.
54pub mod glv;
55
56pub use glv::PackedGlv;
57
58/// Limbs per field element. 32-bit limbs, not 64.
59///
60/// Justification, in order of weight:
61///
62/// 1. Measured on this M2 Max during the scouting phase: an 8 x u32 CIOS Montgomery
63///    multiply that widens both factors to `ulong` only at the multiply-accumulate site
64///    ran at 4.51 G mul/s, against 3.49 G mul/s for the same algorithm written with
65///    `mulhi` and 32-bit carry chains. The `ulong` form is 29% faster because the Metal
66///    compiler lowers `(ulong)x * (ulong)y` with both factors 32-bit-widened onto the
67///    native 32x32->64 path, whereas a genuine 64x64 product is emulated.
68/// 2. Both Apple-targeted references agree on 8 x u32: zkmopro/gpu-acceleration
69///    (`shader/misc/types.metal`, `NUM_LIMBS 8`, `LOG_LIMB_SIZE 32`) and
70///    zkonduit/metal-msm (`UnsignedInteger<8>`). The one project that uses 4 x u64
71///    (lambdaworks) hand-emulates a `u128` class underneath, paying the emulated 64-bit
72///    multiply on every limb product.
73/// 3. `R = 2^(32 * 8) = 2^256` then coincides exactly with arkworks' own Montgomery
74///    radix, which is what makes the zero-conversion packing above possible.
75pub const LIMBS: usize = 8;
76
77/// BN254 scalar field modulus `r`, little-endian 32-bit limbs.
78///
79/// `r = 21888242871839275222246405745257275088548364400416034343698204186575808495617`
80pub const FR_MODULUS: [u32; LIMBS] = [
81    0xf000_0001,
82    0x43e1_f593,
83    0x79b9_7091,
84    0x2833_e848,
85    0x8181_585d,
86    0xb850_45b6,
87    0xe131_a029,
88    0x3064_4e72,
89];
90
91/// `-r^{-1} mod 2^32`, the CIOS per-limb reduction multiplier for `Fr`.
92///
93/// Checked against `ark-ff` in [`tests::montgomery_constants_match_ark`] rather than
94/// trusted: a wrong `N0` produces a multiply that is wrong for almost every input, which
95/// this crate's GPU-vs-host test would catch, but a wrong `N0` that is *right* for the
96/// handful of small vectors a human tries by hand is exactly how this bug ships.
97pub const FR_N0: u32 = 0xefff_ffff;
98
99/// BN254 base field modulus `q`, little-endian 32-bit limbs.
100pub const FQ_MODULUS: [u32; LIMBS] = [
101    0xd87c_fd47,
102    0x3c20_8c16,
103    0x6871_ca8d,
104    0x9781_6a91,
105    0x8181_585d,
106    0xb850_45b6,
107    0xe131_a029,
108    0x3064_4e72,
109];
110
111/// `-q^{-1} mod 2^32`. Equals 3834012553, the same value zkmopro and zkonduit ship for
112/// BN254 `Fq`, which is a cheap independent cross-check on the derivation.
113pub const FQ_N0: u32 = 0xe486_6389;
114
115/// BN254 scalar field element, 32 bytes, **Montgomery form**, little-endian 32-bit limbs.
116///
117/// MSL twin: `struct Fr { uint v[8]; }` in `shaders/bn254_fr.metal`.
118#[repr(C)]
119#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
120pub struct PackedFr {
121    pub v: [u32; LIMBS],
122}
123
124/// BN254 scalar as an integer in `[0, r)`, 32 bytes, **standard form**, little-endian
125/// 32-bit limbs. For MSM window decomposition only. See the module docs.
126///
127/// MSL twin: `struct Scalar { uint v[8]; }`.
128#[repr(C)]
129#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
130pub struct PackedScalar {
131    pub v: [u32; LIMBS],
132}
133
134/// BN254 base field element, 32 bytes, **Montgomery form**, little-endian 32-bit limbs.
135///
136/// MSL twin: `struct Fq { uint v[8]; }`.
137#[repr(C)]
138#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
139pub struct PackedFq {
140    pub v: [u32; LIMBS],
141}
142
143/// `Fq2 = Fq[u]/(u^2 + 1)`, 64 bytes, `c0` then `c1`, each Montgomery.
144///
145/// MSL twin: `struct Fq2 { Fq c0; Fq c1; }`.
146#[repr(C)]
147#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
148pub struct PackedFq2 {
149    pub c0: PackedFq,
150    pub c1: PackedFq,
151}
152
153/// G1 affine point, exactly 64 bytes: `x` then `y`, no flag word.
154///
155/// # How infinity is signalled
156///
157/// The all-zero encoding, `x == 0 && y == 0`, means the point at infinity. This is
158/// unambiguous rather than a convention: BN254 G1 is `y^2 = x^3 + 3`, so `(0, 0)` fails
159/// the curve equation (`0 != 3`) and can never be a real point. Three things follow, and
160/// all three are why the sentinel beats a 65th byte:
161///
162/// * The stride stays a power of two, 64 bytes, so a base vector is naturally aligned and
163///   a thread's load of one point is two 32-byte segments rather than a straddle.
164/// * A freshly allocated device buffer is zero-filled on all three backends, so an
165///   accumulator array starts at infinity with no memset kernel and no host staging
166///   buffer.
167/// * The test is `x | y == 0` over 16 words, which is branch-free.
168///
169/// This matters in practice and is not a theoretical case: snarkjs zkeys really do
170/// contain points at infinity in the A, B and C query vectors, and the reference Metal
171/// MSM implementations that ignore the arkworks `infinity` flag lift `(0, 0)` into a live
172/// non-identity projective point and poison the bucket it lands in.
173#[repr(C)]
174#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
175pub struct PackedG1Affine {
176    pub x: PackedFq,
177    pub y: PackedFq,
178}
179
180/// G2 affine point, exactly 128 bytes: `x` then `y`, each an [`PackedFq2`]. Infinity is
181/// the all-zero encoding, for the same reason as [`PackedG1Affine`] (BN254 G2 is
182/// `y^2 = x^3 + 3/(9 + u)`, whose constant term is nonzero, so `(0, 0)` is off-curve).
183#[repr(C)]
184#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
185pub struct PackedG2Affine {
186    pub x: PackedFq2,
187    pub y: PackedFq2,
188}
189
190// The wire format. A future arkworks bump, a stray `#[repr(align)]`, or someone adding a
191// field cannot get past these.
192const _: () = {
193    assert!(size_of::<PackedFr>() == 32);
194    assert!(size_of::<PackedScalar>() == 32);
195    assert!(size_of::<PackedFq>() == 32);
196    assert!(size_of::<PackedFq2>() == 64);
197    assert!(size_of::<PackedG1Affine>() == 64);
198    assert!(size_of::<PackedG2Affine>() == 128);
199    // MSL's `uint` is 4-byte aligned and Metal requires buffer offsets to be a multiple
200    // of 4. Anything larger here would mean Rust inserted padding the shader does not
201    // know about.
202    assert!(align_of::<PackedFr>() == 4);
203    assert!(align_of::<PackedG1Affine>() == 4);
204    assert!(align_of::<PackedG2Affine>() == 4);
205    // arkworks' backing integer is 4 x u64 = 32 bytes for both BN254 fields. If this ever
206    // stops holding, `from_*`/`to_*` below are reinterpreting the wrong number of words.
207    assert!(size_of::<BigInt<4>>() == 32);
208};
209
210#[inline]
211fn split(limbs: [u64; 4]) -> [u32; LIMBS] {
212    let mut out = [0u32; LIMBS];
213    for (i, w) in limbs.iter().enumerate() {
214        out[2 * i] = *w as u32;
215        out[2 * i + 1] = (*w >> 32) as u32;
216    }
217    out
218}
219
220#[inline]
221fn join(v: [u32; LIMBS]) -> [u64; 4] {
222    let mut out = [0u64; 4];
223    for (i, w) in out.iter_mut().enumerate() {
224        *w = u64::from(v[2 * i]) | (u64::from(v[2 * i + 1]) << 32);
225    }
226    out
227}
228
229impl PackedFr {
230    pub const ZERO: Self = Self { v: [0; LIMBS] };
231
232    /// Montgomery limbs, taken verbatim from arkworks' internal representation. No
233    /// conversion happens on either side of the boundary.
234    #[inline]
235    pub fn from_fr(x: &Fr) -> Self {
236        Self { v: split(x.0 .0) }
237    }
238
239    /// Inverse of [`Self::from_fr`]. Uses `new_unchecked`, which stores the limbs as the
240    /// Montgomery representative rather than converting into it. Going through
241    /// `Fr::from_bigint` here instead would multiply by `R` a second time and every value
242    /// would come back wrong by a factor of `R`.
243    #[inline]
244    pub fn to_fr(&self) -> Fr {
245        Fr::new_unchecked(BigInt::new(join(self.v)))
246    }
247
248    pub fn pack_slice(xs: &[Fr]) -> Vec<Self> {
249        xs.iter().map(Self::from_fr).collect()
250    }
251
252    /// Packs directly into caller memory. On Metal that pointer is `MTLBuffer::contents()`
253    /// and unified memory makes this the whole upload; on CUDA and WebGPU a transfer still
254    /// follows.
255    ///
256    /// Panics if the destination is not exactly as long as the source, because a short
257    /// destination would leave the tail of a device buffer holding whatever was there
258    /// before, and a proof built on it would fail verification with no other symptom.
259    pub fn pack_into(xs: &[Fr], out: &mut [Self]) {
260        assert_eq!(xs.len(), out.len(), "packed destination length mismatch");
261        for (dst, src) in out.iter_mut().zip(xs) {
262            *dst = Self::from_fr(src);
263        }
264    }
265
266    pub fn unpack_slice(xs: &[Self]) -> Vec<Fr> {
267        xs.iter().map(Self::to_fr).collect()
268    }
269}
270
271impl PackedScalar {
272    pub const ZERO: Self = Self { v: [0; LIMBS] };
273
274    /// Standard form, `x` as an integer in `[0, r)`.
275    #[inline]
276    pub fn from_fr(x: &Fr) -> Self {
277        Self {
278            v: split(x.into_bigint().0),
279        }
280    }
281
282    /// `None` if the limbs are not a canonical residue, which for a scalar is a real
283    /// possibility (unlike a Montgomery representative, which is canonical by
284    /// construction) and must not be papered over with a silent reduction.
285    #[inline]
286    pub fn to_fr(&self) -> Option<Fr> {
287        Fr::from_bigint(BigInt::new(join(self.v)))
288    }
289
290    pub fn pack_slice(xs: &[Fr]) -> Vec<Self> {
291        xs.iter().map(Self::from_fr).collect()
292    }
293
294    /// Panics if the destination is not exactly as long as the source, because a short
295    /// destination would leave the tail of a device buffer holding whatever was there
296    /// before, and a proof built on it would fail verification with no other symptom.
297    pub fn pack_into(xs: &[Fr], out: &mut [Self]) {
298        assert_eq!(xs.len(), out.len(), "packed destination length mismatch");
299        for (dst, src) in out.iter_mut().zip(xs) {
300            *dst = Self::from_fr(src);
301        }
302    }
303}
304
305impl PackedFq {
306    pub const ZERO: Self = Self { v: [0; LIMBS] };
307
308    #[inline]
309    pub fn from_fq(x: &Fq) -> Self {
310        Self { v: split(x.0 .0) }
311    }
312
313    #[inline]
314    pub fn to_fq(&self) -> Fq {
315        Fq::new_unchecked(BigInt::new(join(self.v)))
316    }
317
318    #[inline]
319    fn is_zero(&self) -> bool {
320        self.v.iter().fold(0u32, |a, b| a | b) == 0
321    }
322}
323
324impl PackedFq2 {
325    pub const ZERO: Self = Self {
326        c0: PackedFq::ZERO,
327        c1: PackedFq::ZERO,
328    };
329
330    #[inline]
331    pub fn from_fq2(x: &Fq2) -> Self {
332        Self {
333            c0: PackedFq::from_fq(&x.c0),
334            c1: PackedFq::from_fq(&x.c1),
335        }
336    }
337
338    #[inline]
339    pub fn to_fq2(&self) -> Fq2 {
340        Fq2::new(self.c0.to_fq(), self.c1.to_fq())
341    }
342
343    #[inline]
344    fn is_zero(&self) -> bool {
345        self.c0.is_zero() && self.c1.is_zero()
346    }
347}
348
349impl PackedG1Affine {
350    /// The point at infinity, and also what a zeroed device buffer already contains.
351    pub const INFINITY: Self = Self {
352        x: PackedFq::ZERO,
353        y: PackedFq::ZERO,
354    };
355
356    #[inline]
357    pub fn from_affine(p: &G1Affine) -> Self {
358        // Read the flag, do not read x and y and hope. An arkworks identity has
359        // x = 0, y = 0, infinity = true, so for the identity specifically the two paths
360        // agree; for anything else the flag is the only correct source of truth and
361        // reading it costs one branch at prepare time, once per base, forever.
362        if p.infinity {
363            Self::INFINITY
364        } else {
365            Self {
366                x: PackedFq::from_fq(&p.x),
367                y: PackedFq::from_fq(&p.y),
368            }
369        }
370    }
371
372    #[inline]
373    pub fn is_infinity(&self) -> bool {
374        self.x.is_zero() && self.y.is_zero()
375    }
376
377    /// Does not check the curve equation: these come back from a kernel that was handed
378    /// on-curve inputs, and re-checking every point would cost more than the kernel.
379    #[inline]
380    pub fn to_affine(&self) -> G1Affine {
381        if self.is_infinity() {
382            G1Affine::identity()
383        } else {
384            G1Affine::new_unchecked(self.x.to_fq(), self.y.to_fq())
385        }
386    }
387
388    pub fn pack_slice(ps: &[G1Affine]) -> Vec<Self> {
389        ps.iter().map(Self::from_affine).collect()
390    }
391
392    /// Panics if the destination is not exactly as long as the source, because a short
393    /// destination would leave the tail of a device buffer holding whatever was there
394    /// before, and a proof built on it would fail verification with no other symptom.
395    pub fn pack_into(ps: &[G1Affine], out: &mut [Self]) {
396        assert_eq!(ps.len(), out.len(), "packed destination length mismatch");
397        for (dst, src) in out.iter_mut().zip(ps) {
398            *dst = Self::from_affine(src);
399        }
400    }
401}
402
403impl PackedG2Affine {
404    pub const INFINITY: Self = Self {
405        x: PackedFq2::ZERO,
406        y: PackedFq2::ZERO,
407    };
408
409    #[inline]
410    pub fn from_affine(p: &G2Affine) -> Self {
411        if p.infinity {
412            Self::INFINITY
413        } else {
414            Self {
415                x: PackedFq2::from_fq2(&p.x),
416                y: PackedFq2::from_fq2(&p.y),
417            }
418        }
419    }
420
421    #[inline]
422    pub fn is_infinity(&self) -> bool {
423        self.x.is_zero() && self.y.is_zero()
424    }
425
426    #[inline]
427    pub fn to_affine(&self) -> G2Affine {
428        if self.is_infinity() {
429            G2Affine::identity()
430        } else {
431            G2Affine::new_unchecked(self.x.to_fq2(), self.y.to_fq2())
432        }
433    }
434
435    pub fn pack_slice(ps: &[G2Affine]) -> Vec<Self> {
436        ps.iter().map(Self::from_affine).collect()
437    }
438
439    /// Panics if the destination is not exactly as long as the source, because a short
440    /// destination would leave the tail of a device buffer holding whatever was there
441    /// before, and a proof built on it would fail verification with no other symptom.
442    pub fn pack_into(ps: &[G2Affine], out: &mut [Self]) {
443        assert_eq!(ps.len(), out.len(), "packed destination length mismatch");
444        for (dst, src) in out.iter_mut().zip(ps) {
445            *dst = Self::from_affine(src);
446        }
447    }
448}
449
450/// Marker for a type whose in-memory bytes are exactly what the GPU should see.
451///
452/// # Safety
453///
454/// Implementors must be `#[repr(C)]`, contain nothing but `u32` (directly or through
455/// other `Packed` types), have no padding, and treat every bit pattern as valid. All of
456/// that is enforced for the types below by the `const` block above plus inspection: each
457/// one is a `[u32; 8]` or a tuple of them, so there is nowhere for padding to hide.
458pub unsafe trait Packed: Copy {}
459
460unsafe impl Packed for PackedFr {}
461unsafe impl Packed for PackedScalar {}
462unsafe impl Packed for PackedFq {}
463unsafe impl Packed for PackedFq2 {}
464unsafe impl Packed for PackedG1Affine {}
465unsafe impl Packed for PackedG2Affine {}
466// A bare `u32` meets the contract trivially, and each backend's MSM has a word buffer
467// that goes through the same reader its points do (the host-tail combine's spill row
468// indices).
469unsafe impl Packed for u32 {}
470
471/// Byte view of a packed slice, for a device upload.
472pub fn as_bytes<T: Packed>(items: &[T]) -> &[u8] {
473    // SAFETY: `Packed` promises no padding and no invalid bit patterns, so every byte of
474    // the slice is initialised and readable. The result borrows `items`, so the lifetime
475    // is the source's.
476    unsafe {
477        core::slice::from_raw_parts(items.as_ptr().cast::<u8>(), size_of::<T>() * items.len())
478    }
479}
480
481#[cfg(test)]
482pub(crate) mod tests {
483    use super::*;
484    use ark_ff::{One, Zero};
485    use snarkrs_field::{CurveGroup, G1Projective, G2Projective, PrimeGroup};
486
487    use crate::testrng::SplitMix64;
488
489    /// `R mod r`, computed independently of arkworks. This is what the Montgomery
490    /// representative of 1 must be, and asserting it pins down that `LIMBS = 8` with
491    /// `R = 2^256` really is arkworks' radix and not a coincidence that happens to
492    /// round-trip.
493    const FR_R_MOD_N: [u32; LIMBS] = [
494        0x4fff_fffb,
495        0xac96_341c,
496        0x9f60_cd29,
497        0x36fc_7695,
498        0x7879_462e,
499        0x666e_a36f,
500        0x9a07_df2f,
501        0x0e0a_77c1,
502    ];
503
504    /// Every magic number in this file, re-derived from or checked against `ark-ff`
505    /// rather than trusted.
506    #[test]
507    fn montgomery_constants_match_ark() {
508        for (name, modulus, n0, ark_mod) in [
509            ("Fr", FR_MODULUS, FR_N0, Fr::MODULUS.0),
510            ("Fq", FQ_MODULUS, FQ_N0, Fq::MODULUS.0),
511        ] {
512            assert_eq!(
513                join(modulus),
514                ark_mod,
515                "{name}: modulus limbs disagree with ark-ff"
516            );
517            // n0 == -m^{-1} mod 2^32. Everything above limb 0 is irrelevant mod 2^32, so
518            // this is the whole condition, not a weakening of it.
519            let low = u64::from(modulus[0]).wrapping_mul(u64::from(n0)) as u32;
520            assert_eq!(low, u32::MAX, "{name}: N0 is not -m^-1 mod 2^32");
521        }
522        // The value both Apple-targeted reference implementations ship for BN254 Fq,
523        // reached here by an independent derivation.
524        assert_eq!(FQ_N0, 3_834_012_553);
525    }
526
527    #[test]
528    fn fr_round_trips_through_montgomery_limbs() {
529        let mut rng = SplitMix64(0xC0FF_EE00);
530        for x in [Fr::zero(), Fr::one(), -Fr::one(), Fr::from(2u64)]
531            .into_iter()
532            .chain((0..4096).map(|_| rng.next_fr()))
533        {
534            assert_eq!(PackedFr::from_fr(&x).to_fr(), x);
535            assert_eq!(PackedScalar::from_fr(&x).to_fr(), Some(x));
536        }
537    }
538
539    /// The two Fr encodings must genuinely differ, otherwise the module docs are a lie
540    /// and someone will use them interchangeably.
541    #[test]
542    fn montgomery_and_standard_encodings_are_not_the_same() {
543        let one = Fr::one();
544        assert_eq!(PackedScalar::from_fr(&one).v, [1, 0, 0, 0, 0, 0, 0, 0]);
545        assert_eq!(PackedFr::from_fr(&one).v, FR_R_MOD_N);
546        assert_ne!(PackedFr::from_fr(&one).v, PackedScalar::from_fr(&one).v);
547    }
548
549    /// A scalar buffer that is not a canonical residue must be rejected, not silently
550    /// reduced. `r` itself is the obvious case a fuzzer or a corrupt buffer produces.
551    #[test]
552    fn a_non_canonical_scalar_is_rejected() {
553        assert_eq!(PackedScalar { v: FR_MODULUS }.to_fr(), None);
554        assert_eq!(
555            PackedScalar {
556                v: [u32::MAX; LIMBS]
557            }
558            .to_fr(),
559            None
560        );
561    }
562
563    fn rand_g1(rng: &mut SplitMix64) -> G1Affine {
564        (G1Projective::generator() * rng.next_fr()).into_affine()
565    }
566
567    fn rand_g2(rng: &mut SplitMix64) -> G2Affine {
568        (G2Projective::generator() * rng.next_fr()).into_affine()
569    }
570
571    #[test]
572    fn g1_packs_to_64_bytes_and_round_trips_including_infinity() {
573        let mut rng = SplitMix64(7);
574        assert_eq!(as_bytes(&[PackedG1Affine::INFINITY]).len(), 64);
575
576        let inf = G1Affine::identity();
577        assert!(PackedG1Affine::from_affine(&inf).is_infinity());
578        assert_eq!(PackedG1Affine::from_affine(&inf).to_affine(), inf);
579
580        for _ in 0..256 {
581            let p = rand_g1(&mut rng);
582            let packed = PackedG1Affine::from_affine(&p);
583            assert!(!packed.is_infinity(), "a random point is not infinity");
584            assert_eq!(packed.to_affine(), p);
585        }
586        // A real point can never collide with the infinity sentinel, because (0, 0) is
587        // off-curve for y^2 = x^3 + 3.
588        assert!(!G1Affine::new_unchecked(Fq::zero(), Fq::zero()).is_on_curve());
589    }
590
591    #[test]
592    fn g2_packs_to_128_bytes_and_round_trips_including_infinity() {
593        let mut rng = SplitMix64(9);
594        assert_eq!(as_bytes(&[PackedG2Affine::INFINITY]).len(), 128);
595
596        let inf = G2Affine::identity();
597        assert!(PackedG2Affine::from_affine(&inf).is_infinity());
598        assert_eq!(PackedG2Affine::from_affine(&inf).to_affine(), inf);
599
600        for _ in 0..64 {
601            let p = rand_g2(&mut rng);
602            let packed = PackedG2Affine::from_affine(&p);
603            assert!(!packed.is_infinity());
604            assert_eq!(packed.to_affine(), p);
605        }
606        assert!(!G2Affine::new_unchecked(Fq2::zero(), Fq2::zero()).is_on_curve());
607    }
608
609    /// The arkworks strides this module exists to avoid. If these ever become 64 and 128
610    /// the repack is still correct, just no longer load-bearing; if they change to some
611    /// third value the comment at the top of this file is stale.
612    #[test]
613    fn ark_affine_strides_are_still_the_ones_documented() {
614        assert_eq!(size_of::<G1Affine>(), 72);
615        assert_eq!(size_of::<G2Affine>(), 136);
616        assert_ne!(size_of::<G1Affine>(), size_of::<PackedG1Affine>());
617    }
618
619    #[test]
620    fn packed_slices_have_the_declared_stride() {
621        assert_eq!(as_bytes(&vec![PackedG1Affine::INFINITY; 5]).len(), 5 * 64);
622        assert_eq!(as_bytes(&vec![PackedG2Affine::INFINITY; 5]).len(), 5 * 128);
623        assert_eq!(as_bytes(&vec![PackedFr::ZERO; 5]).len(), 5 * 32);
624        // and the bytes really are the limbs, low word first
625        let x = PackedFr::from_fr(&Fr::one());
626        assert_eq!(&as_bytes(&[x])[..4], &FR_R_MOD_N[0].to_le_bytes());
627    }
628
629    /// Packing must be a pure function of the input, and the bulk paths must agree with
630    /// the single-element one. A `pack_into` that wrote a stale tail would be invisible
631    /// until a proof failed to verify.
632    #[test]
633    fn packing_is_a_pure_function() {
634        let mut rng = SplitMix64(11);
635        let xs: Vec<Fr> = (0..64).map(|_| rng.next_fr()).collect();
636        assert_eq!(PackedFr::pack_slice(&xs), PackedFr::pack_slice(&xs));
637        let mut out = vec![PackedFr::ZERO; xs.len()];
638        PackedFr::pack_into(&xs, &mut out);
639        assert_eq!(out, PackedFr::pack_slice(&xs));
640        assert_eq!(PackedFr::unpack_slice(&out), xs);
641    }
642}
643
644/// A deterministic PRNG for tests in the GPU backend crates.
645///
646/// SplitMix64, Steele et al. Dependency-free on purpose: the backends that consume this
647/// crate exist to have very few dependencies, and a test RNG is not worth one.
648pub mod testrng {
649    use snarkrs_field::Fr;
650
651    pub struct SplitMix64(pub u64);
652
653    impl SplitMix64 {
654        pub fn next_u64(&mut self) -> u64 {
655            self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
656            let mut z = self.0;
657            z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
658            z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
659            z ^ (z >> 31)
660        }
661
662        /// Uniform on `[0, r)` up to the usual negligible bias of reducing 256 random
663        /// bits mod a 254-bit modulus.
664        pub fn next_fr(&mut self) -> Fr {
665            use ark_ff::PrimeField;
666            let mut b = [0u8; 32];
667            for c in b.chunks_mut(8) {
668                c.copy_from_slice(&self.next_u64().to_le_bytes());
669            }
670            Fr::from_le_bytes_mod_order(&b)
671        }
672    }
673}