use std::collections::{BTreeMap, HashSet};
#[cfg(feature = "parallel-histogram")]
use std::{
panic::{self, AssertUnwindSafe},
sync::atomic::{AtomicBool, Ordering}
};
use pretty_assertions::assert_eq;
#[cfg(feature = "parallel-histogram")]
use rayon::iter::ParallelIterator;
#[cfg(feature = "parallel-histogram")]
use crate::support::on_small_stack;
use crate::{
EvaluationError, Evaluator, HistogramBuilder, Passes, serial,
support::{compile_valid, optimize, read_histogram_test_cases}
};
pub fn histogram_test<I, T, B>(source: &'static str, builder: B)
where
I: 'static,
T: HistogramBuilder<'static, I>,
B: Fn(Evaluator) -> T
{
let mut seen = HashSet::new();
for (index, (source, args, externs, expected)) in
read_histogram_test_cases(source).iter().enumerate()
{
let key = (source, args.clone(), externs.clone());
let key = format!("{:?}", key);
assert!(seen.insert(key.clone()), "duplicate test case: {}", key);
let function = compile_valid(source);
let function = optimize(function, Passes::all());
let mut evaluator = Evaluator::new(function);
for (name, value) in externs.iter()
{
evaluator.bind(name, *value).unwrap();
}
let bounds = evaluator
.bounds_over(
args.iter().map(|arg| Some((*arg).into())),
externs.iter().map(|(name, value)| (*name, (*value).into()))
)
.unwrap();
assert!(
bounds.count.map(|c| c <= 50000).unwrap_or(true),
"case {}: {}: too many outcomes: {} > 50000",
index + 1,
key,
bounds.count.unwrap()
);
let expected = expected
.iter()
.map(|(key, value)| (*key, *value as u64))
.collect::<BTreeMap<_, _>>();
let builder = builder(evaluator);
let histogram = builder.build(args.iter().copied()).unwrap();
{
let histogram_map = histogram
.iter()
.map(|(outcome, count)| (*outcome, *count))
.collect::<BTreeMap<_, _>>();
assert_eq!(histogram_map, expected, "case {}: {}", index + 1, key);
}
let value_bounds = bounds.value;
assert_eq!(
*histogram.keys().min().unwrap(),
value_bounds.min,
"case {}: {}: min value bound mismatch",
index + 1,
key
);
assert_eq!(
*histogram.keys().max().unwrap(),
value_bounds.max,
"case {}: {}: max value bound mismatch",
index + 1,
key
);
if let Some(expected_outcomes) = bounds.count
{
let actual_outcomes: u128 =
histogram.values().map(|count| *count as u128).sum();
assert_eq!(
actual_outcomes,
expected_outcomes,
"case {}: {}: outcome count mismatch",
index + 1,
key
);
}
let total = histogram.total();
for (outcome, count) in histogram.iter()
{
let odds = histogram.odds(*outcome);
let expected_odds = (*count, total - *count);
assert_eq!(
odds,
expected_odds,
"case {}: {}: odds mismatch for outcome {}",
index + 1,
key,
outcome
);
let percent = histogram.percent_chance(*outcome);
let expected_percent = (*count as f64 / total as f64) * 100.0;
assert_eq!(
percent,
expected_percent,
"case {}: {}: percent mismatch for outcome {}",
index + 1,
key,
outcome
);
}
}
}
const HISTOGRAM_TEST_SOURCE: &str =
include_str!("../../tests/test_histograms.txt");
#[test]
fn test_serial_histogram_building()
{
histogram_test(HISTOGRAM_TEST_SOURCE, serial::HistogramBuilder::new)
}
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_histogram_building()
{
histogram_test(
HISTOGRAM_TEST_SOURCE,
crate::parallel::HistogramBuilder::new
)
}
#[cfg(feature = "parallel-histogram")]
const WORKER_STACK_SIZE: usize = 2 * 1024 * 1024;
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_histogram_many_dice_one_thread()
{
on_small_stack(|| build_many_dice(1))
}
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_histogram_many_dice_four_threads()
{
on_small_stack(|| build_many_dice(4))
}
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_histogram_many_faces_one_thread()
{
on_small_stack(|| build_many_faces(1))
}
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_histogram_many_faces_four_threads()
{
on_small_stack(|| build_many_faces(4))
}
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_histogram_consumer_panic()
{
on_small_stack(|| {
let builder = parallel_builder("6D6");
let panicked = AtomicBool::new(false);
let outcome = panic::catch_unwind(AssertUnwindSafe(|| {
in_pool(4, || {
builder.iter([]).unwrap().for_each(|state| {
if state.result == Some(21)
&& !panicked.swap(true, Ordering::Relaxed)
{
panic!("deliberate panic in the consumer")
}
})
})
}));
assert!(outcome.is_err());
})
}
#[cfg(feature = "parallel-histogram")]
fn build_many_dice(threads: usize)
{
const LIMIT: u64 = 1000;
let builder = parallel_builder("800D6");
let histogram =
in_pool(threads, || builder.build_with_limit([], LIMIT).unwrap());
assert_eq!(histogram.total(), LIMIT);
assert!(
histogram
.keys()
.all(|outcome| (800..=4800).contains(outcome)),
"implausible outcome: {:?}",
histogram.iter().collect::<BTreeMap<_, _>>()
);
}
#[cfg(feature = "parallel-histogram")]
fn build_many_faces(threads: usize)
{
const SOURCE: &str = "2D1000";
let builder = parallel_builder(SOURCE);
let parallel = in_pool(threads, || builder.build([]).unwrap());
let serial =
serial::HistogramBuilder::new(Evaluator::new(compile_valid(SOURCE)))
.build([])
.unwrap();
assert_eq!(parallel.total(), 1_000_000);
assert_eq!(
parallel.iter().collect::<BTreeMap<_, _>>(),
serial.iter().collect::<BTreeMap<_, _>>()
);
}
#[cfg(feature = "parallel-histogram")]
fn parallel_builder(source: &str) -> crate::parallel::HistogramBuilder
{
crate::parallel::HistogramBuilder::new(Evaluator::new(compile_valid(
source
)))
}
#[cfg(feature = "parallel-histogram")]
fn in_pool<R: Send>(threads: usize, f: impl FnOnce() -> R + Send) -> R
{
rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.stack_size(WORKER_STACK_SIZE)
.build()
.unwrap()
.install(f)
}
#[test]
fn test_histogram_max()
{
let function = crate::Assembler::assemble(
"\
Function() r#2 âš…#1
\textern[]
\tbody:
\t\tâš…0 <- roll range 1:6
\t\t@0 <- sum rolling record âš…0
\t\t@1 <- @0 max 3
\t\treturn @1
"
)
.unwrap();
let histogram = serial::HistogramBuilder::new(Evaluator::new(function))
.build([])
.unwrap();
assert_eq!(
histogram.iter().collect::<BTreeMap<_, _>>(),
BTreeMap::from([(&3, &3), (&4, &1), (&5, &1), (&6, &1)])
);
}
fn serial_builder(source: &str) -> serial::HistogramBuilder
{
serial::HistogramBuilder::new(Evaluator::new(compile_valid(source)))
}
fn least_budget(source: &str, args: &[i32]) -> u64
{
let builder = serial_builder(source);
let (mut low, mut high) = (0, 1u64 << 20);
assert!(builder.build_metered(args.iter().copied(), high).is_ok());
while low < high
{
let middle = low + (high - low) / 2;
match builder.build_metered(args.iter().copied(), middle)
{
Ok(_) => high = middle,
Err(EvaluationError::HistogramBudgetExhausted { .. }) =>
{
low = middle + 1
},
Err(e) => panic!("{}: {}", source, e)
}
}
low
}
#[test]
fn test_metered_charges()
{
for (source, args, branches) in [
("[1:6]", &[][..], 6),
("[3:1]", &[], 0),
("0D6", &[], 0),
("1D0", &[], 1),
("1D6", &[], 6),
("3D6", &[], 6 + 36 + 216),
("2D[1,2,3]", &[], 3 + 9),
("{n}: {n}D6", &[2], 42),
("[1:2] + [1:3]", &[], 2 + 2 * 3),
("(1D2)D2", &[], 2 + 2 + (2 + 2 * 2))
]
{
assert_eq!(least_budget(source, args), branches, "{}", source);
let builder = serial_builder(source);
assert_eq!(
builder
.build_metered(args.iter().copied(), branches)
.unwrap(),
builder.build(args.iter().copied()).unwrap(),
"{}",
source
);
}
}
#[test]
fn test_metered_refusal()
{
let builder = serial_builder("3D6");
assert_eq!(
builder.build_metered([], 257),
Err(EvaluationError::HistogramBudgetExhausted {
requested: 6,
remaining: 5,
consumed: 252
})
);
assert_eq!(
builder.build_metered([], 0),
Err(EvaluationError::HistogramBudgetExhausted {
requested: 6,
remaining: 0,
consumed: 0
})
);
}
#[test]
fn test_metered_vast_state_spaces()
{
let builder = serial_builder("{n}: {n}D6");
assert!(matches!(
builder.build_metered([i32::MAX], 10_000),
Err(EvaluationError::HistogramBudgetExhausted { .. })
));
let builder = serial_builder("{a}, {b}: [{a}:{b}]");
assert_eq!(
builder.build_metered([i32::MIN, i32::MAX], 10_000),
Err(EvaluationError::HistogramBudgetExhausted {
requested: 1 << 32,
remaining: 10_000,
consumed: 0
})
);
}
#[test]
fn test_widest_range()
{
let histogram = serial_builder("{a}, {b}: [{a}:{b}]")
.build_with_limit([i32::MIN, i32::MAX], 3)
.unwrap();
assert_eq!(
histogram.iter().collect::<BTreeMap<_, _>>(),
BTreeMap::from([
(&i32::MIN, &1),
(&(i32::MIN + 1), &1),
(&(i32::MIN + 2), &1)
])
);
}
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_metered_agrees()
{
for (source, args) in [
("3D6", &[][..]),
("(1D3)D3", &[]),
("{n}: {n}D4 drop lowest 1", &[5]),
("[1:6] * 2D[1,3,5] + 1D10", &[])
]
{
let least = least_budget(source, args);
let serial =
serial_builder(source).build(args.iter().copied()).unwrap();
let builder = parallel_builder(source);
for threads in [1, 4]
{
for _ in 0..8
{
let (within, beyond) = in_pool(threads, || {
(
builder.build_metered(args.iter().copied(), least),
builder.build_metered(args.iter().copied(), least - 1)
)
});
assert_eq!(within.as_ref(), Ok(&serial), "{}", source);
assert!(
matches!(
beyond,
Err(EvaluationError::HistogramBudgetExhausted { .. })
),
"{}: {:?}",
source,
beyond
);
}
}
}
}
#[cfg(feature = "parallel-histogram")]
#[test]
fn test_parallel_metered_vast_state_spaces()
{
let builder = parallel_builder("{n}: {n}D6");
for threads in [1, 4]
{
let refused =
in_pool(threads, || builder.build_metered([i32::MAX], 10_000));
assert!(
matches!(
refused,
Err(EvaluationError::HistogramBudgetExhausted { .. })
),
"{:?}",
refused
);
let histogram = in_pool(threads, || builder.build([2]).unwrap());
assert_eq!(histogram.total(), 36);
}
let builder = parallel_builder("{a}, {b}: [{a}:{b}]");
assert_eq!(
in_pool(4, || builder.build_metered([i32::MIN, i32::MAX], 10_000)),
Err(EvaluationError::HistogramBudgetExhausted {
requested: 1 << 32,
remaining: 10_000,
consumed: 0
})
);
}