use std::num::NonZeroU64;
use pretty_assertions::assert_eq;
use rand::{Rng as _, SeedableRng, rngs::StdRng};
use crate::{
Budget, Distribution, EvaluationError, Evaluator, Probability, Sampled,
Unobserved, Weight, support::compile_valid
};
fn samples(n: u64) -> NonZeroU64 { NonZeroU64::new(n).unwrap() }
fn confidence(numerator: u8, denominator: u8) -> Probability
{
Probability::new(Weight::from(numerator), Weight::from(denominator))
.unwrap()
}
fn exact(evaluator: &Evaluator) -> Distribution
{
evaluator
.plan_distribution([])
.unwrap()
.build(Budget::UNLIMITED, &Unobserved)
.unwrap()
}
fn deviation(sampled: &Sampled, exact: &Distribution) -> f64
{
exact
.iter()
.map(|(outcome, _)| {
(sampled.counts().cdf(outcome).to_f64()
- exact.cdf(outcome).to_f64())
.abs()
})
.fold(0.0, f64::max)
}
#[test]
fn test_error_bound()
{
let mut evaluator = Evaluator::new(compile_valid("1"));
let mut rng = StdRng::seed_from_u64(0);
let sampled = evaluator.sample([], &mut rng, samples(50_000), 0).unwrap();
let epsilon = sampled.error_bound(&confidence(19, 20));
assert!((epsilon - (40f64.ln() / 100_000.0).sqrt()).abs() < 1e-15);
assert!((epsilon - 0.006_073_6).abs() < 1e-7, "{epsilon}");
assert_eq!(format!("{epsilon:.4}"), "0.0061");
let quadrupled =
evaluator.sample([], &mut rng, samples(200_000), 0).unwrap();
let halved = quadrupled.error_bound(&confidence(19, 20));
assert!((halved - epsilon / 2.0).abs() < 1e-15);
assert_eq!(
sampled.error_bound(&Probability::ZERO),
(2f64.ln() / 100_000.0).sqrt()
);
assert_eq!(sampled.error_bound(&Probability::ONE), f64::INFINITY);
assert_eq!(
sampled.error_bound(&confidence(38, 40)),
sampled.error_bound(&confidence(19, 20))
);
}
#[test]
fn test_sample_within_error_bound()
{
for source in [
"3D6",
"4D6 drop lowest 1",
"(1D4)D6",
"[1:6] + 1D20 drop highest",
"{x}@(1D6) + {x} * 2"
]
{
let mut evaluator = Evaluator::new(compile_valid(source));
let exact = exact(&evaluator);
let mut rng = StdRng::seed_from_u64(7093814572093487);
let n = samples(50_000);
let sampled = evaluator.sample([], &mut rng, n, u64::MAX).unwrap();
assert_eq!(sampled.samples(), n, "{source}");
assert_eq!(
sampled.counts().total(),
&Weight::from(n.get()),
"{source}"
);
assert!(
sampled.counts().min() >= exact.min()
&& sampled.counts().max() <= exact.max(),
"{source}"
);
let epsilon = sampled.error_bound(&confidence(19, 20));
let deviation = deviation(&sampled, &exact);
assert!(deviation <= epsilon, "{source}: {deviation} > {epsilon}");
}
}
#[test]
fn test_sample_confidence()
{
let mut evaluator = Evaluator::new(compile_valid("2D6"));
let exact = exact(&evaluator);
let confidence = confidence(9, 10);
let strays = (0..200u64)
.filter(|&seed| {
let mut rng = StdRng::seed_from_u64(seed);
let sampled = evaluator
.sample([], &mut rng, samples(500), u64::MAX)
.unwrap();
deviation(&sampled, &exact) > sampled.error_bound(&confidence)
})
.count();
assert!(strays <= 20, "{strays} of 200 samples strayed");
}
#[test]
fn test_sample_refuses_before_rolling()
{
let mut evaluator = Evaluator::new(compile_valid("{x}: {x}D6"));
let seed = 3409857120983475;
let mut rng = StdRng::seed_from_u64(seed);
let mut untouched = StdRng::seed_from_u64(seed);
assert_eq!(
evaluator.sample([i32::MAX], &mut rng, samples(50_000), 1_000_000),
Err(EvaluationError::DiceBudgetExhausted {
requested: i32::MAX as u64,
remaining: 1_000_000,
consumed: 0
})
);
assert_eq!(
evaluator.sample([], &mut rng, samples(50_000), 1_000_000),
Err(EvaluationError::BadArity {
expected: 1,
given: 0
})
);
assert_eq!(rng.next_u64(), untouched.next_u64());
}
#[test]
fn test_sample_dice_budget_spans_samples()
{
let mut evaluator = Evaluator::new(compile_valid("{x}: ({x}D1)D6"));
let mut rng = StdRng::seed_from_u64(8471029384710293);
let sampled = evaluator.sample([1], &mut rng, samples(3), 6).unwrap();
assert_eq!(sampled.counts().total(), &Weight::from(3u8));
assert_eq!(
evaluator.sample([1], &mut rng, samples(3), 5),
Err(EvaluationError::DiceBudgetExhausted {
requested: 1,
remaining: 0,
consumed: 5
})
);
}
#[test]
fn test_sample_display()
{
let mut evaluator = Evaluator::new(compile_valid("1D1 + 2"));
let mut rng = StdRng::seed_from_u64(0);
let sampled = evaluator.sample([], &mut rng, samples(3), 3).unwrap();
assert_eq!(sampled.to_string(), "sampled (n = 3):\n3: 3\n");
}
#[cfg(feature = "serde")]
#[test]
fn test_sample_serde()
{
let mut evaluator = Evaluator::new(compile_valid("1D1 + 2"));
let mut rng = StdRng::seed_from_u64(0);
let sampled = evaluator.sample([], &mut rng, samples(3), 3).unwrap();
let json = serde_json::to_string(&sampled).unwrap();
assert_eq!(
json,
r#"{"samples":3,"counts":{"weights":{"3":"3"},"total":"3"}}"#
);
assert_eq!(serde_json::from_str::<Sampled>(&json).unwrap(), sampled);
for (invalid, reason) in [
(
r#"{"samples":4,"counts":{"weights":{"3":"3"},"total":"3"}}"#,
"total 3, not its 4 samples"
),
(
r#"{"samples":0,"counts":{"weights":{"3":"3"},"total":"3"}}"#,
"nonzero"
),
(
r#"{"samples":3,"counts":{"weights":{"3":"2"},"total":"3"}}"#,
"not the sum"
),
(
r#"{"samples":3,"counts":{"weights":{"3":"3"},"total":"3"},"x":1}"#,
"unknown field"
)
]
{
let error = serde_json::from_str::<Sampled>(invalid).unwrap_err();
assert!(error.to_string().contains(reason), "{invalid}: {error}");
}
}