1use serde::{Deserialize, Serialize};
22
23use crate::error::CoreError;
24
25const SPLITMIX64_GAMMA: u64 = 0x9e37_79b9_7f4a_7c15;
27
28fn splitmix64(state: &mut u64) -> u64 {
32 *state = state.wrapping_add(SPLITMIX64_GAMMA);
33 mix64(*state)
34}
35
36fn mix64(mut z: u64) -> u64 {
39 z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
40 z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
41 z ^ (z >> 31)
42}
43
44#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
51#[serde(try_from = "SeededRngData", into = "SeededRngData")]
52pub struct SeededRng {
53 state: [u64; 4],
54 spare_normal: Option<f64>,
55}
56
57#[derive(Serialize, Deserialize)]
58#[serde(deny_unknown_fields)]
59struct SeededRngData {
60 state: [String; 4],
61 #[serde(default)]
62 spare_normal: Option<f64>,
63}
64
65fn parse_state_word(text: &str) -> Result<u64, CoreError> {
67 text.strip_prefix("0x")
68 .filter(|digits| {
69 !digits.is_empty()
70 && digits.len() <= 16
71 && digits.bytes().all(|b| b.is_ascii_hexdigit())
72 })
73 .and_then(|digits| u64::from_str_radix(digits, 16).ok())
74 .ok_or(CoreError::InvalidRandomState)
75}
76
77impl TryFrom<SeededRngData> for SeededRng {
78 type Error = CoreError;
79
80 fn try_from(data: SeededRngData) -> Result<Self, CoreError> {
81 let mut state = [0; 4];
82 for (word, text) in state.iter_mut().zip(&data.state) {
83 *word = parse_state_word(text)?;
84 }
85 if state == [0; 4] {
86 return Err(CoreError::InvalidRandomState);
87 }
88 if let Some(spare) = data.spare_normal
89 && !spare.is_finite()
90 {
91 return Err(CoreError::Domain {
92 what: "spare normal deviate",
93 value: spare,
94 });
95 }
96 Ok(SeededRng {
97 state,
98 spare_normal: data.spare_normal,
99 })
100 }
101}
102
103impl From<SeededRng> for SeededRngData {
104 fn from(rng: SeededRng) -> Self {
105 SeededRngData {
106 state: rng.state.map(|word| format!("{word:#018x}")),
107 spare_normal: rng.spare_normal,
108 }
109 }
110}
111
112impl SeededRng {
113 pub fn seed_from_u64(seed: u64) -> Self {
116 let mut sm = seed;
117 let state = [
118 splitmix64(&mut sm),
119 splitmix64(&mut sm),
120 splitmix64(&mut sm),
121 splitmix64(&mut sm),
122 ];
123 SeededRng {
124 state,
125 spare_normal: None,
126 }
127 }
128
129 pub fn for_stream(seed: u64, keys: &[u64]) -> Self {
142 let hash = keys.iter().fold(seed, |h, &key| {
143 mix64(h ^ mix64(key.wrapping_add(SPLITMIX64_GAMMA)))
144 });
145 Self::seed_from_u64(hash)
146 }
147
148 pub fn next_u64(&mut self) -> u64 {
150 let s = &mut self.state;
151 let result = s[0].wrapping_add(s[3]).rotate_left(23).wrapping_add(s[0]);
152 let t = s[1] << 17;
153 s[2] ^= s[0];
154 s[3] ^= s[1];
155 s[1] ^= s[2];
156 s[0] ^= s[3];
157 s[2] ^= t;
158 s[3] = s[3].rotate_left(45);
159 result
160 }
161
162 pub fn uniform(&mut self) -> f64 {
165 const SCALE: f64 = 1.0 / 9_007_199_254_740_992.0;
167 let top = (self.next_u64() >> 11) as f64;
169 top * SCALE
170 }
171
172 pub fn standard_normal(&mut self) -> f64 {
177 if let Some(spare) = self.spare_normal.take() {
178 return spare;
179 }
180 loop {
181 let u = 2.0 * self.uniform() - 1.0;
182 let v = 2.0 * self.uniform() - 1.0;
183 let s = u * u + v * v;
184 if s > 0.0 && s < 1.0 {
185 let factor = (-2.0 * s.ln() / s).sqrt();
186 self.spare_normal = Some(v * factor);
187 return u * factor;
188 }
189 }
190 }
191}
192
193#[cfg(test)]
194mod tests {
195 use rand_core::{Rng, SeedableRng};
196 use rand_xoshiro::{SplitMix64, Xoshiro256PlusPlus};
197
198 use super::*;
199
200 #[test]
203 fn stream_matches_the_rand_xoshiro_implementation() {
204 for seed in [0, 1, 42, 0x0123_4567_89ab_cdef, u64::MAX] {
205 let mut ours = SeededRng::seed_from_u64(seed);
206 let mut theirs = Xoshiro256PlusPlus::seed_from_u64(seed);
207 for _ in 0..10_000 {
208 assert_eq!(ours.next_u64(), theirs.next_u64(), "seed {seed}");
209 }
210 }
211 }
212
213 #[test]
214 fn seeding_is_four_splitmix64_outputs() {
215 let mut theirs = SplitMix64::seed_from_u64(7);
216 let mut sm = 7;
217 for _ in 0..100 {
218 assert_eq!(splitmix64(&mut sm), theirs.next_u64());
219 }
220 }
221
222 #[test]
223 fn a_stream_depends_only_on_its_seed_and_keys() {
224 assert_eq!(SeededRng::for_stream(42, &[]), SeededRng::seed_from_u64(42));
226 let mut sm = 7_u64;
228 let inner = splitmix64(&mut sm);
229 let mut outer = (42 ^ inner).wrapping_sub(SPLITMIX64_GAMMA);
230 assert_eq!(
231 SeededRng::for_stream(42, &[7]),
232 SeededRng::seed_from_u64(splitmix64(&mut outer))
233 );
234 let mut first = SeededRng::for_stream(1, &[3, 5]);
236 let mut again = SeededRng::for_stream(1, &[3, 5]);
237 let a: Vec<u64> = (0..4).map(|_| first.next_u64()).collect();
238 let b: Vec<u64> = (0..4).map(|_| again.next_u64()).collect();
239 assert_eq!(a, b);
240 let others = [
241 SeededRng::for_stream(2, &[3, 5]),
242 SeededRng::for_stream(1, &[3, 6]),
243 SeededRng::for_stream(1, &[5, 3]),
244 SeededRng::for_stream(1, &[3]),
245 SeededRng::for_stream(1, &[3, 5, 0]),
246 ];
247 for mut other in others {
248 assert_ne!(other.next_u64(), a[0]);
249 }
250 }
251
252 #[test]
253 fn neighbouring_streams_are_uncorrelated() {
254 let n = 20_000_u64;
257 let u: Vec<f64> = (0..n)
258 .map(|k| SeededRng::for_stream(2026, &[k]).uniform())
259 .collect();
260 let count = n as f64;
261 let mean = u.iter().sum::<f64>() / count;
262 assert!(
263 (mean - 0.5).abs() < 4.0 * (1.0 / 12.0 / count).sqrt(),
264 "{mean}"
265 );
266 let lag: f64 = u
267 .windows(2)
268 .map(|w| (w[0] - 0.5) * (w[1] - 0.5))
269 .sum::<f64>()
270 / (count - 1.0)
271 * 12.0;
272 assert!(lag.abs() < 4.0 / (count - 1.0).sqrt(), "{lag}");
273 }
274
275 #[test]
276 fn uniform_is_in_the_half_open_unit_interval() {
277 let mut rng = SeededRng::seed_from_u64(3);
278 let n = 1_000_000;
279 let mut sum = 0.0;
280 for _ in 0..n {
281 let x = rng.uniform();
282 assert!((0.0..1.0).contains(&x));
283 sum += x;
284 }
285 let mean = sum / f64::from(n);
287 assert!((mean - 0.5).abs() < 5.0 * (1.0 / 12.0 / f64::from(n)).sqrt());
288 }
289
290 #[test]
294 fn standard_normal_has_the_normal_moments_and_cdf() {
295 let mut rng = SeededRng::seed_from_u64(2026);
296 let n = 1_000_000;
297 let points = [-2.0, -1.0, 0.0, 1.0, 2.0];
298 let phi = [
299 0.022_750_131_948_179,
300 0.158_655_253_931_457,
301 0.5,
302 0.841_344_746_068_543,
303 0.977_249_868_051_821,
304 ];
305 let mut below = [0_u32; 5];
306 let (mut sum, mut sum_sq, mut sum_4) = (0.0, 0.0, 0.0);
307 for _ in 0..n {
308 let x = rng.standard_normal();
309 sum += x;
310 sum_sq += x * x;
311 sum_4 += x * x * x * x;
312 for (count, &point) in below.iter_mut().zip(&points) {
313 if x <= point {
314 *count += 1;
315 }
316 }
317 }
318 let nf = f64::from(n);
319 let mean = sum / nf;
320 let variance = sum_sq / nf - mean * mean;
321 assert!(mean.abs() < 5.0 / nf.sqrt(), "mean {mean}");
322 assert!(
324 (variance - 1.0).abs() < 5.0 * (2.0 / nf).sqrt(),
325 "variance {variance}"
326 );
327 assert!((sum_4 / nf - 3.0).abs() < 5.0 * (96.0 / nf).sqrt());
329 for ((&count, &p), &point) in below.iter().zip(&phi).zip(&points) {
330 let fraction = f64::from(count) / nf;
331 let standard_error = (p * (1.0 - p) / nf).sqrt();
332 assert!(
333 (fraction - p).abs() < 5.0 * standard_error,
334 "CDF at {point}: {fraction} vs {p}"
335 );
336 }
337 }
338
339 #[test]
340 fn same_seed_same_normals_and_serde_resumes_the_stream() {
341 let mut a = SeededRng::seed_from_u64(99);
342 let mut b = SeededRng::seed_from_u64(99);
343 for _ in 0..1001 {
344 assert_eq!(a.standard_normal().to_bits(), b.standard_normal().to_bits());
345 }
346 assert!(a.spare_normal.is_some());
348 let json = serde_json::to_string(&a).unwrap();
349 let mut resumed: SeededRng = serde_json::from_str(&json).unwrap();
350 for _ in 0..1000 {
351 assert_eq!(
352 a.standard_normal().to_bits(),
353 resumed.standard_normal().to_bits()
354 );
355 }
356 }
357
358 #[test]
359 fn state_serializes_as_hex_words_and_bad_states_are_rejected() {
360 let rng = SeededRng::seed_from_u64(1);
361 let value = serde_json::to_value(&rng).unwrap();
362 let words = value["state"].as_array().unwrap();
363 assert_eq!(words.len(), 4);
364 for (word, &expected) in words.iter().zip(&rng.state) {
365 let text = word.as_str().unwrap();
366 assert_eq!(text.len(), 18);
367 assert_eq!(parse_state_word(text).unwrap(), expected);
368 }
369 let zero = r#"{"state":["0x0","0x0","0x0","0x00"],"spare_normal":null}"#;
370 assert_eq!(
371 serde_json::from_str::<SeededRng>(zero)
372 .unwrap_err()
373 .to_string(),
374 CoreError::InvalidRandomState.to_string()
375 );
376 for bad in [
377 r#"{"state":[1,2,3,4]}"#,
378 r#"{"state":["1","0x2","0x3","0x4"]}"#,
379 r#"{"state":["0x","0x2","0x3","0x4"]}"#,
380 r#"{"state":["0x10000000000000000","0x2","0x3","0x4"]}"#,
381 r#"{"state":["0xg","0x2","0x3","0x4"]}"#,
382 r#"{"state":["0x+1","0x2","0x3","0x4"]}"#,
383 ] {
384 assert!(serde_json::from_str::<SeededRng>(bad).is_err(), "{bad}");
385 }
386 let short = r#"{"state":["0x1","0x2","0x3","0x4"]}"#;
387 assert_eq!(
388 serde_json::from_str::<SeededRng>(short).unwrap().state,
389 [1, 2, 3, 4]
390 );
391 }
392}