1use rand_chacha::ChaCha20Rng;
9use rand_chacha::rand_core::{Rng, SeedableRng};
10
11#[derive(Debug, Clone)]
28pub struct StreamRng {
29 inner: ChaCha20Rng,
30}
31
32impl StreamRng {
33 pub fn new(seed: u64, stream: u64) -> Self {
35 let mut inner = ChaCha20Rng::from_seed(expand_seed(seed));
36 inner.set_stream(stream);
37 Self { inner }
38 }
39
40 pub fn next_u64(&mut self) -> u64 {
42 self.inner.next_u64()
43 }
44
45 pub fn next_open01(&mut self) -> f64 {
49 let k = self.next_u64() >> 11;
50 (k as f64 + 0.5) * (1.0 / (1u64 << 53) as f64)
51 }
52}
53
54fn expand_seed(seed: u64) -> [u8; 32] {
58 let mut state = seed;
59 let mut key = [0u8; 32];
60 for chunk in key.chunks_exact_mut(8) {
61 state = state.wrapping_add(0x9E37_79B9_7F4A_7C15);
62 let mut z = state;
63 z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
64 z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
65 z ^= z >> 31;
66 chunk.copy_from_slice(&z.to_le_bytes());
67 }
68 key
69}
70
71#[cfg(test)]
72mod tests {
73 use super::*;
74 use rayon::prelude::*;
75
76 fn first_draws(seed: u64, stream: u64, n: usize) -> Vec<u64> {
77 let mut rng = StreamRng::new(seed, stream);
78 (0..n).map(|_| rng.next_u64()).collect()
79 }
80
81 #[test]
82 fn same_seed_and_stream_replays() {
83 assert_eq!(first_draws(1, 2, 16), first_draws(1, 2, 16));
84 }
85
86 #[test]
87 fn streams_and_seeds_differ() {
88 assert_ne!(first_draws(1, 0, 4), first_draws(1, 1, 4));
89 assert_ne!(first_draws(1, 0, 4), first_draws(2, 0, 4));
90 }
91
92 #[test]
93 fn output_is_pinned() {
94 assert_eq!(
99 first_draws(42, 0, 3),
100 [
101 693385945204756564,
102 16436763086163553629,
103 3187728548114239752
104 ]
105 );
106 assert_eq!(
107 first_draws(42, 5, 2),
108 [2382546587937276428, 2377724087087505542]
109 );
110 }
111
112 #[test]
113 fn open01_is_strictly_inside() {
114 let mut rng = StreamRng::new(0, 0);
115 for _ in 0..10_000 {
116 let u = rng.next_open01();
117 assert!(u > 0.0 && u < 1.0);
118 }
119 }
120
121 fn simulate(threads: usize) -> Vec<f64> {
123 let pool = rayon::ThreadPoolBuilder::new()
124 .num_threads(threads)
125 .build()
126 .unwrap();
127 pool.install(|| {
128 (0..1_000u64)
129 .into_par_iter()
130 .map(|sim| {
131 let mut rng = StreamRng::new(7, sim);
132 (0..100).map(|_| rng.next_open01()).sum()
133 })
134 .collect()
135 })
136 }
137
138 #[test]
139 fn identical_across_thread_counts() {
140 let one = simulate(1);
141 assert_eq!(one, simulate(4));
142 assert_eq!(one, simulate(16));
143 }
144}