use core::ops::{Range, RangeInclusive};
use crate::RngSource;
mod sealed {
pub trait Sealed {}
}
pub trait RandInt: sealed::Sealed + Copy + PartialOrd {
#[doc(hidden)]
fn lemire_bounded(rng: &mut RngSource, range_size: Self) -> Self;
#[doc(hidden)]
fn lemire_inclusive(rng: &mut RngSource, low: Self, high: Self) -> Self;
#[doc(hidden)]
fn add(self, n: Self) -> Self;
#[doc(hidden)]
fn sub(self, n: Self) -> Self;
#[doc(hidden)]
fn one() -> Self;
}
macro_rules! impl_rand_int_via_u32 {
($($t:ty),* $(,)?) => {
$(
impl sealed::Sealed for $t {}
impl RandInt for $t {
#[inline]
fn lemire_bounded(rng: &mut RngSource, range_size: Self) -> Self {
lemire_u32(rng, range_size as u32) as Self
}
#[inline]
fn lemire_inclusive(rng: &mut RngSource, low: Self, high: Self) -> Self {
let span = (high as u64) - (low as u64) + 1;
let n = if span == (1u64 << 32) { 0u32 } else { span as u32 };
low.wrapping_add(lemire_u32(rng, n) as Self)
}
#[inline]
fn add(self, n: Self) -> Self { self + n }
#[inline]
fn sub(self, n: Self) -> Self { self - n }
#[inline]
fn one() -> Self { 1 }
}
)*
};
}
impl_rand_int_via_u32!(u8, u16, u32);
impl sealed::Sealed for u64 {}
impl RandInt for u64 {
#[inline]
fn lemire_bounded(rng: &mut RngSource, range_size: Self) -> Self {
lemire_u64(rng, range_size)
}
#[inline]
fn lemire_inclusive(rng: &mut RngSource, low: Self, high: Self) -> Self {
let span = (high as u128) - (low as u128) + 1;
let n = if span == (1u128 << 64) {
0u64
} else {
span as u64
};
low.wrapping_add(lemire_u64(rng, n))
}
#[inline]
fn add(self, n: Self) -> Self {
self + n
}
#[inline]
fn sub(self, n: Self) -> Self {
self - n
}
#[inline]
fn one() -> Self {
1
}
}
impl sealed::Sealed for usize {}
impl RandInt for usize {
#[inline]
fn lemire_bounded(rng: &mut RngSource, range_size: Self) -> Self {
#[cfg(target_pointer_width = "32")]
{
lemire_u32(rng, range_size as u32) as usize
}
#[cfg(target_pointer_width = "64")]
{
lemire_u64(rng, range_size as u64) as usize
}
}
#[inline]
fn lemire_inclusive(rng: &mut RngSource, low: Self, high: Self) -> Self {
#[cfg(target_pointer_width = "32")]
{
let span = (high as u64) - (low as u64) + 1;
let n = if span == (1u64 << 32) {
0u32
} else {
span as u32
};
low.wrapping_add(lemire_u32(rng, n) as usize)
}
#[cfg(target_pointer_width = "64")]
{
let span = (high as u128) - (low as u128) + 1;
let n = if span == (1u128 << 64) {
0u64
} else {
span as u64
};
low.wrapping_add(lemire_u64(rng, n) as usize)
}
}
#[inline]
fn add(self, n: Self) -> Self {
self + n
}
#[inline]
fn sub(self, n: Self) -> Self {
self - n
}
#[inline]
fn one() -> Self {
1
}
}
pub fn gen_range<T: RandInt>(rng: &mut RngSource, range: Range<T>) -> T {
debug_assert!(range.start < range.end, "gen_range: empty range");
if range.start >= range.end {
return range.start;
}
let size = range.end.sub(range.start);
range.start.add(T::lemire_bounded(rng, size))
}
pub fn gen_range_inclusive<T: RandInt>(rng: &mut RngSource, range: RangeInclusive<T>) -> T {
let (start, end) = range.into_inner();
debug_assert!(start <= end, "gen_range_inclusive: inverted range");
T::lemire_inclusive(rng, start, end)
}
#[inline]
fn lemire_u32(rng: &mut RngSource, n: u32) -> u32 {
let mut buf = [0u8; 4];
if n == 0 {
rng.fill_bytes(&mut buf);
return u32::from_le_bytes(buf);
}
loop {
rng.fill_bytes(&mut buf);
let x = u32::from_le_bytes(buf);
let m = u64::from(x) * u64::from(n);
let l = m as u32;
if l < n {
let t = n.wrapping_neg() % n;
if l < t {
continue;
}
}
return (m >> 32) as u32;
}
}
#[inline]
fn lemire_u64(rng: &mut RngSource, n: u64) -> u64 {
let mut buf = [0u8; 8];
if n == 0 {
rng.fill_bytes(&mut buf);
return u64::from_le_bytes(buf);
}
loop {
rng.fill_bytes(&mut buf);
let x = u64::from_le_bytes(buf);
let m = u128::from(x) * u128::from(n);
let l = m as u64;
if l < n {
let t = n.wrapping_neg() % n;
if l < t {
continue;
}
}
return (m >> 64) as u64;
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
fn rng() -> RngSource {
RngSource::from_seed(&[0x5Au8; 32])
}
#[test]
fn full_range_u8_does_not_panic_and_covers_extremes() {
let mut r = rng();
let mut saw_min = false;
let mut saw_max = false;
for _ in 0..100_000 {
let v = gen_range_inclusive(&mut r, 0u8..=255u8);
if v == 0 {
saw_min = true;
}
if v == 255 {
saw_max = true;
}
}
assert!(saw_min, "full u8 range must be able to yield 0");
assert!(saw_max, "full u8 range must be able to yield 255");
}
#[test]
fn full_range_u16_does_not_panic_and_covers_extremes() {
let mut r = rng();
let mut saw_min = false;
let mut saw_max = false;
for _ in 0..500_000 {
let v = gen_range_inclusive(&mut r, 0u16..=65535u16);
if v == 0 {
saw_min = true;
}
if v == 65535 {
saw_max = true;
}
}
assert!(saw_min, "full u16 range must be able to yield 0");
assert!(saw_max, "full u16 range must be able to yield 65535");
}
#[test]
fn full_range_u32_does_not_panic() {
let mut r = rng();
for _ in 0..10_000 {
let _ = gen_range_inclusive(&mut r, 0u32..=u32::MAX);
}
}
#[test]
fn full_range_u64_does_not_panic() {
let mut r = rng();
let mut saw_high = false;
for _ in 0..10_000 {
let v = gen_range_inclusive(&mut r, 0u64..=u64::MAX);
if v > (u64::MAX >> 1) {
saw_high = true;
}
}
assert!(saw_high, "full u64 range must reach the upper half");
}
#[test]
fn narrow_inclusive_range_stays_in_bounds() {
let mut r = rng();
for _ in 0..10_000 {
let v = gen_range_inclusive(&mut r, 1u8..=6u8);
assert!((1..=6).contains(&v));
}
}
#[cfg(not(debug_assertions))]
#[test]
fn empty_exclusive_range_short_circuits_without_consuming_rng() {
let mut used = rng();
let mut untouched = rng();
let v = gen_range(&mut used, 7u32..7u32);
assert_eq!(v, 7u32, "empty range must return start");
let mut a = [0u8; 32];
let mut b = [0u8; 32];
used.fill_bytes(&mut a);
untouched.fill_bytes(&mut b);
assert_eq!(a, b, "empty range must not consume RNG bytes");
let mut r = rng();
assert_eq!(gen_range(&mut r, 9u64..9u64), 9u64);
}
#[cfg(debug_assertions)]
#[test]
#[should_panic(expected = "empty range")]
fn empty_exclusive_range_panics_in_debug() {
let mut r = rng();
let _ = gen_range(&mut r, 7u32..7u32);
}
}