use rand_chacha::{ChaCha12Rng, ChaCha20Rng, ChaCha8Rng};
use crate::error::CheckpointError;
pub trait SnapshotRng: Sized {
fn capture(&self) -> Result<Vec<u8>, CheckpointError>;
fn restore(bytes: &[u8]) -> Result<Self, CheckpointError>;
}
macro_rules! impl_snapshot_rng {
($rng:ty) => {
impl SnapshotRng for $rng {
fn capture(&self) -> Result<Vec<u8>, CheckpointError> {
bincode::serialize(self).map_err(|e| CheckpointError::SerializeError(Box::new(e)))
}
fn restore(bytes: &[u8]) -> Result<Self, CheckpointError> {
bincode::deserialize(bytes)
.map_err(|e| CheckpointError::DeserializeError(Box::new(e)))
}
}
};
}
impl_snapshot_rng!(ChaCha8Rng);
impl_snapshot_rng!(ChaCha12Rng);
impl_snapshot_rng!(ChaCha20Rng);
#[cfg(test)]
mod tests {
use super::*;
use rand::{Rng, SeedableRng};
#[test]
fn test_snapshot_rng_round_trip_is_bit_identical() {
let mut rng = ChaCha8Rng::seed_from_u64(42);
for _ in 0..37 {
let _: u64 = rng.gen();
}
let bytes = rng.capture().unwrap();
let mut restored = ChaCha8Rng::restore(&bytes).unwrap();
for _ in 0..1000 {
let a: u64 = rng.gen();
let b: u64 = restored.gen();
assert_eq!(a, b, "restored RNG diverged from the original");
}
}
#[test]
fn test_snapshot_rng_variants() {
let mut c12 = ChaCha12Rng::seed_from_u64(7);
let _: u32 = c12.gen();
let bytes = c12.capture().unwrap();
let mut r12 = ChaCha12Rng::restore(&bytes).unwrap();
assert_eq!(c12.gen::<u64>(), r12.gen::<u64>());
let mut c20 = ChaCha20Rng::seed_from_u64(9);
let _: u32 = c20.gen();
let bytes = c20.capture().unwrap();
let mut r20 = ChaCha20Rng::restore(&bytes).unwrap();
assert_eq!(c20.gen::<u64>(), r20.gen::<u64>());
}
}