use std::convert::Infallible;
use rand::{rngs::SysRng, Rng, SeedableRng, TryCryptoRng, TryRng};
use vitaminc_protected::Controlled;
use zeroize::{ZeroizeOnDrop, Zeroizing};
pub struct SafeRand(rand::rngs::ChaCha20Rng);
impl ZeroizeOnDrop for SafeRand {}
const _: fn() = assert_zeroize_on_drop::<rand::rngs::ChaCha20Rng>;
fn assert_zeroize_on_drop<T: ZeroizeOnDrop>() {}
impl SafeRand {
pub fn next_below<T>(&mut self, n: T) -> <Self as crate::BoundedRng<T>>::Output
where
Self: crate::BoundedRng<T>,
{
<Self as crate::BoundedRng<T>>::next_below(self, n)
}
#[deprecated(
note = "inclusive `0..=max`; use `next_below(max + 1)`, or `next_below(n)` when you have a length `n`"
)]
pub fn next_bounded_u32(&mut self, max: u32) -> u32 {
crate::bounded::upto_u32(self, max)
}
pub fn from_entropy() -> Result<Self, crate::RandomError> {
Ok(Self::try_from_rng(&mut SysRng)?)
}
pub fn from_controlled_seed<C>(seed: C) -> Self
where
C: Controlled<Inner = [u8; 32]>,
{
let seed = Zeroizing::new(seed.risky_unwrap());
Self(rand::rngs::ChaCha20Rng::from_seed(*seed))
}
}
impl TryCryptoRng for SafeRand {}
impl TryRng for SafeRand {
type Error = Infallible;
#[inline]
fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
Ok(self.0.next_u32())
}
#[inline]
fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
Ok(self.0.next_u64())
}
#[inline]
fn try_fill_bytes(&mut self, dst: &mut [u8]) -> Result<(), Self::Error> {
self.0.fill_bytes(dst);
Ok(())
}
}
impl SeedableRng for SafeRand {
type Seed = [u8; 32];
fn from_seed(seed: Self::Seed) -> Self {
Self(rand::rngs::ChaCha20Rng::from_seed(seed))
}
}
#[cfg(test)]
mod tests {
use super::SafeRand;
use rand::{rngs::ChaCha20Rng, Rng, SeedableRng, TryRng};
const SEED: [u8; 32] = [7u8; 32];
#[test]
fn try_rng_yields_the_chacha20_stream_for_the_seed() {
let mut safe = SafeRand::from_seed(SEED);
let mut reference = ChaCha20Rng::from_seed(SEED);
for _ in 0..4 {
assert_eq!(safe.try_next_u32().unwrap(), reference.next_u32());
}
for _ in 0..4 {
assert_eq!(safe.try_next_u64().unwrap(), reference.next_u64());
}
let mut got = [0u8; 40];
let mut want = [0u8; 40];
safe.try_fill_bytes(&mut got).unwrap();
reference.fill_bytes(&mut want);
assert_eq!(got, want);
assert_ne!(got, [0u8; 40], "fill_bytes must write the buffer");
}
fn reference_below(rng: &mut ChaCha20Rng, n: u64) -> u32 {
((u128::from(rng.next_u64()) * u128::from(n)) >> 64) as u32
}
const BOUNDS: [u32; 14] = [
1,
2,
3,
4,
5,
6,
7,
16,
100,
1000,
(1 << 31) - 1,
1 << 31,
(1 << 31) + 1,
u32::MAX,
];
#[test]
fn next_below_matches_the_lemire_reference() {
for n in BOUNDS {
let mut safe = SafeRand::from_seed(SEED);
let mut reference = ChaCha20Rng::from_seed(SEED);
for _ in 0..256 {
let got = safe.next_below(n);
assert_eq!(
got,
reference_below(&mut reference, u64::from(n)),
"n = {n}"
);
assert!(got < n, "n = {n}");
}
}
}
#[test]
#[allow(deprecated)]
fn next_bounded_u32_is_the_inclusive_form_of_next_below() {
for max in BOUNDS.map(|n| n - 1).into_iter().chain([u32::MAX]) {
let mut safe = SafeRand::from_seed(SEED);
let mut reference = ChaCha20Rng::from_seed(SEED);
for _ in 0..256 {
let got = safe.next_bounded_u32(max);
assert_eq!(
got,
reference_below(&mut reference, u64::from(max) + 1),
"max = {max}"
);
assert!(got <= max, "max = {max}");
}
}
}
#[test]
fn next_below_covers_exactly_the_half_open_range() {
for n in [5usize, 8] {
let mut rng = SafeRand::from_seed(SEED);
let mut seen = vec![false; n + 1];
for _ in 0..512 {
seen[rng.next_below(n as u32) as usize] = true;
}
assert!(seen[..n].iter().all(|&s| s), "n = {n}");
assert!(!seen[n], "n = {n}");
}
}
#[test]
#[should_panic(expected = "range must be non-zero")]
fn next_below_zero_panics() {
SafeRand::from_seed(SEED).next_below(0);
}
#[test]
fn different_seeds_diverge_and_the_same_seed_repeats() {
let mut a = SafeRand::from_seed(SEED);
let mut b = SafeRand::from_seed(SEED);
let mut c = SafeRand::from_seed([8u8; 32]);
let (x, y, z) = (
a.try_next_u64().unwrap(),
b.try_next_u64().unwrap(),
c.try_next_u64().unwrap(),
);
assert_eq!(x, y);
assert_ne!(x, z);
}
#[test]
fn next_below_from_entropy_is_half_open() -> Result<(), crate::RandomError> {
let mut rng = SafeRand::from_entropy()?;
for n in [1, 2, 4, 5, 52, 62, 64, 94, u32::MAX] {
for _ in 0..1000 {
assert!(rng.next_below(n) < n);
}
}
Ok(())
}
#[test]
fn next_below_is_uniform() {
const RANGE: u32 = 5;
const SAMPLES: u32 = 100_000;
let mut rng = SafeRand::from_seed([3u8; 32]);
let mut counts = [0u32; RANGE as usize];
for _ in 0..SAMPLES {
counts[rng.next_below(RANGE) as usize] += 1;
}
let expected = f64::from(SAMPLES) / f64::from(RANGE);
let chi2: f64 = counts
.iter()
.map(|&c| {
let d = f64::from(c) - expected;
d * d / expected
})
.sum();
assert!(chi2 < 18.47, "chi-squared too high: {chi2}");
}
}