urng 1.0.0

Universal Random Number Generator
use wrapn::{wrap, wu32, wu64};

use crate::{_internal::sm64_from_seed32, rng::Rng};

// --- Squares32 ---

/// The Squares random number generator (32-bit output version by Bernard Widynski).
///
/// # Examples
///
/// ```
/// use urng::{Rng, Squares32};
///
/// let mut rng = Squares32::new(1);
/// let _ = rng.nextu();
/// ```
#[repr(C, align(64))]
#[derive(Debug, Clone, Copy)]
pub struct Squares32 {
    pub(crate) c: wu64,
    pub(crate) k: wu64,
}

impl Squares32 {
    /// Creates a new `Squares32` instance seeded with the given value.
    #[inline]
    pub const fn new(seed: u32) -> Self {
        let mut seedgen = sm64_from_seed32!(seed);
        Self {
            c: wrap!(0),
            k: wrap!(seedgen.nextu_const()),
        }
    }

    /// Core computation: 4 rounds of middle-square with counter.
    /// Takes pre-computed y = ctr*key and z = y + key to avoid
    /// redundant multiplication in batch scenarios.
    #[inline(always)]
    pub fn compute_yz(y: wu64, z: wu64) -> wu32 {
        let mut x = y * y + y;
        x = x.rotate_left(32);

        x = x * x + z;
        x = x.rotate_left(32);

        x = x * x + y;
        x = x.rotate_left(32);

        ((x * x + z) >> 32).cast()
    }

    /// Convenience wrapper: compute from counter and key directly.
    #[inline(always)]
    pub fn compute(ctr: wu64, key: wu64) -> wu32 {
        let y = ctr * key;
        let z = y + key;
        Self::compute_yz(y, z)
    }
}

impl Rng for Squares32 {
    type Word = u32;

    #[inline(always)]
    fn nextu(&mut self) -> Self::Word {
        let out = Self::compute(self.c, self.k);
        self.c += 1;
        *out
    }
}

#[cfg(feature = "simd")]
pub use simd::*;

#[cfg(feature = "simd")]
pub mod simd {
    use std::arch::x86_64::*;

    use crate::_internal::{i2f_bits, u2f_01};
    use crate::sm64_from_seed32;

    // C-ABI exports for Squares32

    #[allow(non_upper_case_globals)]
    pub const SQUARES32x8: usize = 8;

    /// A high-throughput Squares random number generator utilizing AVX-512 SIMD instructions.
    /// This implementation processes 8 counters in parallel and is highly optimized with 4-way unrolling.
    ///
    /// # Examples
    ///
    /// ```no_run
    /// use urng::Squares32x8;
    /// unsafe {
    ///     let mut rng = Squares32x8::new(1);
    ///     let _ = rng.nextu();
    /// }
    /// ```
    #[cfg(target_arch = "x86_64")]
    #[repr(C)]
    #[repr(align(64))]
    pub struct Squares32x8 {
        /// 8 counters stored in a 512-bit SIMD register.
        pub c: __m512i,
        /// 8 keys stored in a 512-bit SIMD register.
        pub k: __m512i,
    }

    #[cfg(target_arch = "x86_64")]
    impl Squares32x8 {
        /// Creates a new `Squares32x8` instance from a 32-bit seed.
        /// The seed is used to initialize the counters and keys.
        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f` target feature.
        #[target_feature(enable = "avx512f")]
        pub unsafe fn new(seed: u32) -> Self {
            let mut k = [0u64; SQUARES32x8];
            let mut seedgen = sm64_from_seed32!(seed);
            k.iter_mut().for_each(|v| {
                *v = seedgen.nextu_const();
            });

            unsafe {
                Self {
                    c: _mm512_setr_epi64(0, 1, 2, 3, 4, 5, 6, 7),
                    k: _mm512_loadu_si512(k.as_ptr() as *const _),
                }
            }
        }

        /// Core computation: 4 rounds of middle-square with counter.
        /// Returns 8x u32 random values in the lower 32-bits of each 64-bit lane.
        ///
        /// # Arguments
        /// * `y` - Pre-computed y = ctr * key.
        /// * `z` - Pre-computed z = y + key.
        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn compute_yz(y: __m512i, z: __m512i) -> __m256i {
            let mut x = _mm512_add_epi64(_mm512_mullo_epi64(y, y), y);
            x = _mm512_or_si512(_mm512_slli_epi64(x, 32), _mm512_srli_epi64(x, 32));

            x = _mm512_add_epi64(_mm512_mullo_epi64(x, x), z);
            x = _mm512_or_si512(_mm512_slli_epi64(x, 32), _mm512_srli_epi64(x, 32));

            x = _mm512_add_epi64(_mm512_mullo_epi64(x, x), y);
            x = _mm512_or_si512(_mm512_slli_epi64(x, 32), _mm512_srli_epi64(x, 32));

            _mm512_cvtepi64_epi32(_mm512_srli_epi64(
                _mm512_add_epi64(_mm512_mullo_epi64(x, x), z),
                32,
            ))
        }

        /// Convenience wrapper to compute random values from counter and key directly.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub(crate) unsafe fn compute(c: __m512i, k: __m512i) -> __m256i {
            unsafe {
                let y = _mm512_mullo_epi64(c, k);
                let z = _mm512_add_epi64(y, k);
                Self::compute_yz(y, z)
            }
        }

        /// Generates 8 new `u32` random numbers.
        /// Increments the internal counters by 8.
        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn nextu(&mut self) -> [u32; SQUARES32x8] {
            unsafe {
                let v = Self::compute(self.c, self.k);
                self.c = _mm512_add_epi64(self.c, _mm512_set1_epi64(8));
                std::mem::transmute(v)
            }
        }

        /// Generates 8 random `f32` values in the range [0, 1).
        ///
        /// # Safety
        ///
        /// The caller must ensure the CPU supports the `avx512f,avx512dq` target feature.
        #[target_feature(enable = "avx512f,avx512dq")]
        pub unsafe fn nextf(&mut self) -> [f32; SQUARES32x8] {
            let out = unsafe { self.nextu() };
            let mut dst = [0f32; SQUARES32x8];
            for i in 0..SQUARES32x8 {
                dst[i] = u2f_01!(f32, 32, out[i]);
            }
            dst
        }
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    crate::safe_test! { Squares32 }

    #[cfg(all(
        feature = "simd",
        target_feature = "avx512f",
        target_feature = "avx512dq"
    ))]
    crate::unsafe_test! { Squares32x8 }
}