1#![cfg_attr(not(test), no_std)]
27#![cfg_attr(docsrs, feature(doc_cfg))]
28
29use core::{convert, fmt};
30use rand_core::{Rng, SeedableRng, TryCryptoRng, TryRng};
31
32pub struct ReseedingRng<R, Rsdr> {
91 inner: R,
92 reseeder: Rsdr,
93 threshold: usize,
94 bytes_consumed: usize,
95}
96
97impl<R, Rsdr> ReseedingRng<R, Rsdr>
98where
99 R: SeedableRng,
100 Rsdr: TryRng,
101{
102 #[track_caller]
116 pub fn try_new(threshold: usize, mut reseeder: Rsdr) -> Result<Self, Rsdr::Error> {
117 assert!(threshold > 0, "`threshold` must be greater than zero");
118 R::try_from_rng(&mut reseeder).map(|inner| Self {
124 inner,
125 reseeder,
126 threshold,
127 bytes_consumed: 0,
128 })
129 }
130
131 pub fn try_reseed(&mut self) -> Result<(), Rsdr::Error> {
137 R::try_from_rng(&mut self.reseeder).map(|inner| {
138 self.inner = inner;
139 self.bytes_consumed = 0;
140 })
141 }
142
143 #[cold]
144 fn reset_after_reseed_attempt_at(&mut self, pos: usize) {
145 let _ = self.try_reseed();
149 self.bytes_consumed = pos;
150 }
151}
152
153impl<R, Rsdr> ReseedingRng<R, Rsdr>
154where
155 R: TryRng + SeedableRng,
156 Rsdr: TryRng,
157{
158 #[cold]
159 fn try_fill_bytes_slow(&mut self, mut dst: &mut [u8]) -> Result<(), R::Error> {
160 loop {
161 if self.bytes_consumed < self.threshold {
162 let len = dst.len().min(self.threshold - self.bytes_consumed);
163 self.bytes_consumed += len;
164 self.inner.try_fill_bytes(&mut dst[..len])?;
165 dst = &mut dst[len..];
166 }
167 if dst.is_empty() {
168 break Ok(());
169 } else {
170 let _ = self.try_reseed();
171 self.bytes_consumed = 0;
172 }
173 }
174 }
175}
176
177impl<R, Rsdr> TryRng for ReseedingRng<R, Rsdr>
178where
179 R: TryRng + SeedableRng,
180 Rsdr: TryRng,
181{
182 type Error = R::Error;
183
184 fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
185 let new_bytes_consumed = self.bytes_consumed + 32 / 8;
186 if new_bytes_consumed <= self.threshold {
187 self.bytes_consumed = new_bytes_consumed;
188 } else {
189 self.reset_after_reseed_attempt_at(32 / 8);
190 }
191 self.inner.try_next_u32()
192 }
193
194 fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
195 let new_bytes_consumed = self.bytes_consumed + 64 / 8;
196 if new_bytes_consumed <= self.threshold {
197 self.bytes_consumed = new_bytes_consumed;
198 } else {
199 self.reset_after_reseed_attempt_at(64 / 8);
200 }
201 self.inner.try_next_u64()
202 }
203
204 fn try_fill_bytes(&mut self, dst: &mut [u8]) -> Result<(), Self::Error> {
205 let new_bytes_consumed = self.bytes_consumed + dst.len();
206 if new_bytes_consumed <= self.threshold {
207 self.bytes_consumed = new_bytes_consumed;
208 self.inner.try_fill_bytes(dst)
209 } else {
210 self.try_fill_bytes_slow(dst)
211 }
212 }
213}
214
215impl<R, Rsdr> TryCryptoRng for ReseedingRng<R, Rsdr>
216where
217 R: TryCryptoRng + SeedableRng,
218 Rsdr: TryCryptoRng,
219{
220}
221
222impl<R, Rsdr> Clone for ReseedingRng<R, Rsdr>
224where
225 R: SeedableRng,
226 Rsdr: Clone + Rng,
227{
228 fn clone(&self) -> Self {
229 let result: Result<_, convert::Infallible> =
230 Self::try_new(self.threshold, self.reseeder.clone());
231 result.unwrap()
232 }
233}
234
235impl<R, Rsdr> fmt::Debug for ReseedingRng<R, Rsdr>
236where
237 R: fmt::Debug,
238 Rsdr: fmt::Debug,
239{
240 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
241 f.debug_struct("ReseedingRng")
242 .field("inner", &self.inner)
243 .field("reseeder", &self.reseeder)
244 .field("threshold", &self.threshold)
245 .finish_non_exhaustive()
246 }
247}
248
249#[cfg(feature = "std_rng")]
250mod std_rng;
251
252#[cfg(feature = "std_rng")]
253pub use std_rng::StdReseedingRng;
254
255#[cfg(feature = "std_rng")]
256#[doc(no_inline)]
257pub use rand::RngExt; #[cfg(test)]
260mod mock;
261
262#[cfg(test)]
263mod tests {
264 use super::*;
265 use rand::rngs::{StdRng, SysRng};
266
267 #[test]
268 fn mirror_rand09_reseeding_rng() {
269 use rand_chacha09::{ChaCha12Core, ChaCha12Rng};
270 use rand09::{RngCore as _, SeedableRng as _};
271
272 use mock::Rand09Adapter as Adapter;
273
274 type OurImpl = ReseedingRng<Adapter, Adapter>;
275 type TheirImpl = rand09::rngs::ReseedingRng<ChaCha12Core, ChaCha12Rng>;
276
277 const N: usize = 1024 * 16 * 5 + 997;
278
279 let seed = rand09::random();
280 let mut o = OurImpl::try_new(1024 * 16, Adapter::from_seed(seed)).unwrap();
281 let mut t = TheirImpl::new(1024 * 16, ChaCha12Rng::from_seed(seed)).unwrap();
282
283 for _ in 0..(N / 4) {
284 assert_eq!(o.next_u32(), t.next_u32());
285 }
286
287 o.try_reseed().unwrap();
288 t.reseed().unwrap();
289
290 for _ in 0..(N / 8) {
291 assert_eq!(o.next_u64(), t.next_u64());
292 }
293
294 o.try_reseed().unwrap();
295 t.reseed().unwrap();
296
297 let mut buf_o = vec![0u8; 17 * 4];
298 let mut buf_t = vec![0u8; buf_o.len()];
299 for _ in 0..(N / buf_o.len()) {
300 o.fill_bytes(&mut buf_o[..]);
301 t.fill_bytes(&mut buf_t[..]);
302 assert_eq!(buf_o, buf_t);
303 }
304
305 o.try_reseed().unwrap();
306 t.reseed().unwrap();
307
308 buf_o.resize(1024 * 16 * 2 + 7 * 4, 0);
309 buf_t.resize(buf_o.len(), 0);
310 for _ in 0..(N / buf_o.len()) {
311 o.fill_bytes(&mut buf_o[..]);
312 t.fill_bytes(&mut buf_t[..]);
313 assert_eq!(buf_o, buf_t);
314 }
315 }
316
317 #[test]
318 fn reseed_after_threshold() {
319 let seed = rand::random();
320 let mut g1 = StdRng::from_rng(&mut StdRng::from_seed(seed));
321 let mut g2 =
322 ReseedingRng::<StdRng, _>::try_new(1024 * 64, StdRng::from_seed(seed)).unwrap();
323
324 for _ in 0..(64 * 1024 / (32 / 8 + 32 / 8 + 64 / 8)) {
325 assert_eq!(g1.next_u32(), g2.next_u32());
326 assert_eq!(g1.next_u32(), g2.next_u32());
327 assert_eq!(g1.next_u64(), g2.next_u64());
328 }
329
330 assert_ne!(g1.next_u32(), g2.next_u32());
331 assert_ne!(g1.next_u64(), g2.next_u64());
332 }
333
334 #[test]
335 fn reseed_after_clone() {
336 let mut g1 = ReseedingRng::<StdRng, _>::try_new(365, rand::rng()).unwrap();
337 assert_eq!(g1.threshold, 365);
338 assert_eq!(g1.bytes_consumed, 0);
339 g1.next_u32();
340 g1.next_u64();
341 assert_eq!(g1.bytes_consumed, 32 / 8 + 64 / 8);
342
343 let mut g2 = g1.clone();
344 assert_eq!(g2.threshold, 365);
345 assert_eq!(g2.bytes_consumed, 0);
346 assert_ne!(g1.next_u32(), g2.next_u32());
347 assert_ne!(g1.next_u64(), g2.next_u64());
348 assert_eq!(g1.bytes_consumed, (32 / 8 + 64 / 8) * 2);
349 assert_eq!(g2.bytes_consumed, 32 / 8 + 64 / 8);
350 }
351
352 #[test]
353 fn count_periodic_reseeds() {
354 use std::cell::Cell;
355
356 struct MockReseeder<'a> {
357 counter: &'a Cell<usize>,
358 }
359
360 impl TryRng for MockReseeder<'_> {
361 type Error = <SysRng as TryRng>::Error;
362
363 fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
364 self.counter.set(self.counter.get() + 1);
365 SysRng.try_next_u32()
366 }
367
368 fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
369 self.counter.set(self.counter.get() + 1);
370 SysRng.try_next_u64()
371 }
372
373 fn try_fill_bytes(&mut self, dst: &mut [u8]) -> Result<(), Self::Error> {
374 self.counter.set(self.counter.get() + 1);
375 SysRng.try_fill_bytes(dst)
376 }
377 }
378
379 let counter = Cell::new(0);
380 let reseeder = MockReseeder { counter: &counter };
381 let mut rng = ReseedingRng::<StdRng, _>::try_new(10, reseeder).unwrap();
382 assert_eq!(counter.get(), 1);
383
384 rng.fill_bytes(&mut [0u8; 10]);
385 assert_eq!(counter.get(), 1);
386 assert_eq!(rng.bytes_consumed, 10);
387 rng.fill_bytes(&mut [0u8; 1]);
388 assert_eq!(counter.get(), 2);
389 assert_eq!(rng.bytes_consumed, 1);
390 rng.fill_bytes(&mut [0u8; 9]);
391 assert_eq!(counter.get(), 2);
392 assert_eq!(rng.bytes_consumed, 10);
393 rng.fill_bytes(&mut [0u8; 1]);
394 assert_eq!(counter.get(), 3);
395 assert_eq!(rng.bytes_consumed, 1);
396 rng.fill_bytes(&mut [0u8; 25]);
397 assert_eq!(counter.get(), 5);
398 assert_eq!(rng.bytes_consumed, 6);
399
400 rng.next_u32();
401 assert_eq!(counter.get(), 5);
402 assert_eq!(rng.bytes_consumed, 10);
403 rng.next_u32();
404 assert_eq!(counter.get(), 6);
405 assert_eq!(rng.bytes_consumed, 4);
406 rng.next_u32();
407 assert_eq!(counter.get(), 6);
408 assert_eq!(rng.bytes_consumed, 8);
409 rng.next_u32(); assert_eq!(counter.get(), 7);
411 assert_eq!(rng.bytes_consumed, 4);
412
413 rng.fill_bytes(&mut [0u8; 7]);
414 assert_eq!(counter.get(), 8);
415 assert_eq!(rng.bytes_consumed, 1);
416 rng.next_u64();
417 assert_eq!(counter.get(), 8);
418 assert_eq!(rng.bytes_consumed, 9);
419 rng.next_u64(); assert_eq!(counter.get(), 9);
421 assert_eq!(rng.bytes_consumed, 8);
422 }
423
424 #[test]
425 #[should_panic]
426 fn panic_if_threshold_is_zero() {
427 let _ = ReseedingRng::<StdRng, _>::try_new(0, SysRng);
428 }
429
430 mod fallible {
432 use super::*;
433
434 const N: usize = 20 * 256;
435
436 #[test]
437 fn generate_random_numbers() {
438 let mut rng = ReseedingRng::<StdRng, _>::try_new(1024, SysRng).unwrap();
439
440 let arrays = (0..N)
441 .map(|_| rng.next_u32().to_le_bytes())
442 .collect::<Vec<_>>();
443 assert!(check_each_byte_for_randomness(&arrays));
444
445 let arrays = (0..N)
446 .map(|_| rng.next_u64().to_le_bytes())
447 .collect::<Vec<_>>();
448 assert!(check_each_byte_for_randomness(&arrays));
449
450 let mut buf = [0u8; 17];
451 let arrays = (0..N)
452 .map(|_| {
453 rng.fill_bytes(buf.as_mut());
454 buf
455 })
456 .collect::<Vec<_>>();
457 assert!(check_each_byte_for_randomness(&arrays));
458 }
459
460 #[test]
461 fn handle_corner_cases() {
462 let mut rng = ReseedingRng::<StdRng, _>::try_new(1, SysRng).unwrap();
463
464 let arrays = (0..N)
465 .map(|_| rng.next_u32().to_le_bytes())
466 .collect::<Vec<_>>();
467 assert!(check_each_byte_for_randomness(&arrays));
468
469 let arrays = (0..N)
470 .map(|_| rng.next_u64().to_le_bytes())
471 .collect::<Vec<_>>();
472 assert!(check_each_byte_for_randomness(&arrays));
473
474 let mut buf = [0u8; 5];
475 let arrays = (0..N)
476 .map(|_| {
477 rng.fill_bytes(buf.as_mut());
478 buf
479 })
480 .collect::<Vec<_>>();
481 assert!(check_each_byte_for_randomness(&arrays));
482
483 let mut buf = [0u8; 0];
484 for _ in 0..N {
485 rng.fill_bytes(buf.as_mut());
486 }
487 }
488 }
489
490 pub(crate) fn check_each_byte_for_randomness<const N: usize>(arrays: &[[u8; N]]) -> bool {
491 (0..N).all(|i| {
492 let mut freq = [0usize; 256];
493 for array in arrays {
494 freq[array[i] as usize] += 1; }
496
497 let expected = arrays.len() as f64 / 256.0;
498 let chi_squared = freq.iter().fold(0.0, |acc, &observed| {
499 let dev = observed as f64 - expected;
500 acc + dev * dev / expected
501 });
502
503 chi_squared < 330.52 })
505 }
506}