use std::{collections::HashSet, ops::RangeInclusive};
use pretty_assertions::assert_eq;
use rand::{Rng as _, SeedableRng, rngs::StdRng};
use crate::{
EvaluationError, Evaluator, HistogramBuilder as _, Passes,
RollingRecordKind,
support::{
compile_valid, on_small_stack, optimize, read_evaluation_test_cases
}
};
#[test]
fn test_evaluation()
{
let mut seen = HashSet::new();
for (index, (source, args, externs, expected)) in
read_evaluation_test_cases(include_str!(
"../../tests/test_evaluation.txt"
))
.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_eq!(
bounds.to_string(),
*expected,
"case {}: {}",
index + 1,
key
);
let mut rng = StdRng::seed_from_u64(24987829587102357);
for _ in 0..1000
{
let result =
evaluator.evaluate(args.iter().copied(), &mut rng).unwrap();
assert!(
result.dice <= bounds.dice,
"case {}: {}: too many dice: {} > {}",
index + 1,
key,
result.dice,
bounds.dice
);
assert_eq!(
result.dice,
result
.records
.iter()
.map(|record| record.results.len() as u64)
.sum::<u64>(),
"case {}: {}: misreported dice",
index + 1,
key
);
let bounds: RangeInclusive<i32> = bounds.value.into();
assert!(
bounds.contains(&result.result),
"case {}: {}: result out of bounds: {} ∉ {}..={}: rolls: {}",
index + 1,
key,
result.result,
bounds.start(),
bounds.end(),
result
.records
.iter()
.map(|record| record.results[0].to_string())
.collect::<Vec<_>>()
.join(", ")
);
for record in result.records
{
match record.kind
{
RollingRecordKind::Uninitialized => unreachable!(),
RollingRecordKind::Range { start, end } =>
{
assert_eq!(
record.results.len(),
1,
"case {}: {}: wrong number of results",
index + 1,
key
);
if end < start
{
assert_eq!(
record.results[0],
0,
"case {}: {}: roll out of bounds: {} ∉ {}..={}",
index + 1,
key,
record.results[0],
start,
end
);
}
else
{
assert!(
(start..=end).contains(&record.results[0]),
"case {}: {}: roll out of bounds: {} ∉ {}..={}",
index + 1,
key,
record.results[0],
start,
end
);
}
},
RollingRecordKind::Standard { count, faces } =>
{
assert_eq!(
record.results.len(),
count.max(0) as usize,
"case {}: {}: wrong number of results",
index + 1,
key
);
for result in record.results
{
match faces <= 0
{
false => assert!(
(1..=faces).contains(&result),
"case {}: {}: roll out of bounds: {} ∉ 1..={}",
index + 1,
key,
result,
faces
),
true => assert_eq!(
result,
0,
"case {}: {}: roll out of bounds: {} ≠0",
index + 1,
key,
result
)
}
}
},
RollingRecordKind::Custom { count, faces } =>
{
assert_eq!(
record.results.len(),
count as usize,
"case {}: {}: wrong number of results",
index + 1,
key
);
for result in record.results
{
assert!(
faces.contains(&result),
"case {}: {}: roll out of bounds: {} ∉ {:?}",
index + 1,
key,
result,
faces
);
}
}
}
}
}
}
}
#[test]
fn test_bad_arity()
{
let function = compile_valid("{x}: {x}");
let mut evaluator = Evaluator::new(function);
let mut rng = StdRng::seed_from_u64(409568093489576902);
assert_eq!(
evaluator.evaluate([], &mut rng),
Err(EvaluationError::BadArity {
expected: 1,
given: 0
})
);
assert_eq!(
evaluator.evaluate([1, 2].iter().copied(), &mut rng),
Err(EvaluationError::BadArity {
expected: 1,
given: 2
})
);
assert_eq!(
evaluator.bounds_over([], []),
Err(EvaluationError::BadArity {
expected: 1,
given: 0
})
);
assert_eq!(
evaluator.bounds_over([Some(1.into()), Some(2.into())], []),
Err(EvaluationError::BadArity {
expected: 1,
given: 2
})
);
}
#[test]
fn test_unrecognized_external()
{
let function = compile_valid("{x}: {x}");
let mut evaluator = Evaluator::new(function);
assert_eq!(
evaluator.bind("y", 1),
Err(EvaluationError::UnrecognizedExternal("y"))
);
}
#[test]
fn test_bind_canonical_external()
{
let function = compile_valid("{a\n b} + 1");
assert_eq!(function.externals, vec!["a b".to_string()]);
let mut evaluator = Evaluator::new(function);
assert_eq!(evaluator.bind("a b", 1), Ok(()));
assert_eq!(
evaluator.bind("a\n b", 1),
Err(EvaluationError::UnrecognizedExternal("a\n b"))
);
}
#[test]
fn test_rolling_record_count()
{
assert_eq!(RollingRecordKind::<i32>::Uninitialized.count(), None);
let mut rng = StdRng::seed_from_u64(69873748728957892);
for (src, expected) in [("[3:8]", 1), ("3D6", 3), ("8D[-1, -1, -2, 5]", 8)]
{
let function = compile_valid(src);
let mut evaluator = Evaluator::new(function);
let result = evaluator.evaluate([], &mut rng).unwrap();
assert_eq!(result.records.len(), 1);
assert_eq!(result.records[0].kind.count(), Some(expected));
}
}
#[test]
fn test_dice_budget_refuses_before_rolling()
{
on_small_stack(|| {
for (source, arg) in [
("{x}: {x}D6", i32::MAX),
("{x}: {x}D[1, 2, 3]", i32::MAX),
("{x}: ({x} * {x})D6", 46_341)
]
{
let mut evaluator = Evaluator::new(compile_valid(source));
let seed = 5829175027591875;
let mut rng = StdRng::seed_from_u64(seed);
let mut untouched = StdRng::seed_from_u64(seed);
assert_eq!(
evaluator.evaluate_metered([arg], &mut rng, 100),
Err(EvaluationError::DiceBudgetExhausted {
requested: i32::MAX as u64,
remaining: 100,
consumed: 0
}),
"{}",
source
);
assert_eq!(rng.next_u64(), untouched.next_u64(), "{}", source);
}
});
}
#[test]
fn test_dice_budget_exact()
{
let mut evaluator = Evaluator::new(compile_valid("{x}: ({x}D1)D6"));
let mut rng = StdRng::seed_from_u64(2098357109857129);
let evaluation = evaluator.evaluate_metered([7], &mut rng, 14).unwrap();
assert_eq!(evaluation.dice, 14);
assert_eq!(
evaluator.evaluate_metered([7], &mut rng, 13),
Err(EvaluationError::DiceBudgetExhausted {
requested: 7,
remaining: 6,
consumed: 7
})
);
assert_eq!(
evaluator.evaluate_metered([8], &mut rng, 14),
Err(EvaluationError::DiceBudgetExhausted {
requested: 8,
remaining: 6,
consumed: 8
})
);
}
#[test]
fn test_dice_budget_costs()
{
let mut rng = StdRng::seed_from_u64(7120985710298375);
let mut evaluator = Evaluator::new(compile_valid("{x}: {x}D6"));
for count in [0, -1, i32::MIN]
{
let evaluation =
evaluator.evaluate_metered([count], &mut rng, 0).unwrap();
assert_eq!(evaluation.dice, 0, "{}", count);
}
for source in ["{x}: [1:{x}]", "{x}: [{x}:1]"]
{
let mut evaluator = Evaluator::new(compile_valid(source));
assert_eq!(
evaluator.evaluate_metered([6], &mut rng, 0),
Err(EvaluationError::DiceBudgetExhausted {
requested: 1,
remaining: 0,
consumed: 0
}),
"{}",
source
);
let evaluation = evaluator.evaluate_metered([6], &mut rng, 1).unwrap();
assert_eq!(evaluation.dice, 1, "{}", source);
}
}
#[test]
fn test_dice_budget_agrees_with_unmetered()
{
let mut evaluator = Evaluator::new(compile_valid(
"4D6 drop lowest + [1:20] + 3D[-1, 0, 1] + (1D4)D8"
));
let seed = 9812750918273509;
let bounds = evaluator.bounds_over([], []).unwrap();
assert_eq!(bounds.dice, 13);
for i in 0..100
{
let unmetered = evaluator
.evaluate([], &mut StdRng::seed_from_u64(seed + i))
.unwrap();
let metered = evaluator
.evaluate_metered([], &mut StdRng::seed_from_u64(seed + i), 13)
.unwrap();
assert_eq!(metered, unmetered);
}
}
#[test]
fn test_max()
{
let mut evaluator = Evaluator::new(
crate::Assembler::assemble(
"\
Function({x}@0) r#2 âš…#0
\textern[]
\tbody:
\t\t@1 <- @0 max 0
\t\treturn @1
"
)
.unwrap()
);
for (x, expected) in
[(i32::MIN, 0), (-3, 0), (0, 0), (5, 5), (i32::MAX, i32::MAX)]
{
let evaluation = evaluator
.evaluate([x], &mut StdRng::seed_from_u64(0))
.unwrap();
assert_eq!(evaluation.result, expected, "{}", x);
assert_eq!(evaluation.dice, 0, "{}", x);
}
for ((min, max), expected) in
[((-3, 5), (0, 5)), ((-7, -2), (0, 0)), ((2, 9), (2, 9))]
{
let bounds = evaluator
.bounds_over([Some((min, max).into())], [])
.unwrap();
assert_eq!(
(bounds.value.min, bounds.value.max),
expected,
"[{}, {}]",
min,
max
);
assert_eq!(bounds.dice, 0);
}
}
#[test]
fn test_negative_drop_count_drops_nothing()
{
for source in [
"(1D1 drop lowest 2 drop lowest -1)D6",
"(1D1 drop highest 2 drop highest -1)D6",
"(1D1 drop lowest -1 drop lowest 2)D6"
]
{
let mut evaluator = Evaluator::new(compile_valid(source));
let bounds = evaluator.bounds_over([], []).unwrap();
assert_eq!(bounds.dice, 1, "{}", source);
let mut rng = StdRng::seed_from_u64(6120957120985710);
let evaluation = evaluator.evaluate_metered([], &mut rng, 1).unwrap();
assert_eq!(evaluation.dice, 1, "{}", source);
assert_eq!(evaluation.result, 0, "{}", source);
}
}
#[test]
fn test_drop_order_is_irrelevant()
{
for (source, reversed) in [
(
"{x}: 3D6 drop lowest {x} drop lowest -1",
"{x}: 3D6 drop lowest -1 drop lowest {x}"
),
(
"{x}: 3D6 drop highest {x} drop highest -2",
"{x}: 3D6 drop highest -2 drop highest {x}"
)
]
{
let functions = [source, reversed].map(|source| {
let function = crate::compile_unoptimized(source).unwrap();
[function.clone(), optimize(function, Passes::all())]
});
for (i, x) in [-1, 0, 1, 2, 3, 4].into_iter().enumerate()
{
let seed = 1098275019827350 + i as u64;
let results = functions
.iter()
.flatten()
.map(|function| {
Evaluator::new(function.clone())
.evaluate([x], &mut StdRng::seed_from_u64(seed))
.unwrap()
.result
})
.collect::<Vec<_>>();
assert!(
results.iter().all(|&result| result == results[0]),
"{} with {}: {:?}",
source,
x,
results
);
}
}
}
#[test]
fn test_optimizer_respects_clamping()
{
for source in [
"{x}: {x}D1",
"{x}: {x}D[5]",
"{x}: {x}D1 drop lowest 1",
"{x}: {x}D1 drop lowest 5",
"{x}: {x}D[5] drop highest 1",
"{x}: 3D1 drop lowest {x}",
"{x}: 3D[5] drop lowest {x}",
"{x}: 1D1 drop lowest {x}",
"{x}: 1D6 drop highest {x}",
"{x}: 1D[2, 3, 4] drop lowest {x}",
"{x}: 3D1 drop lowest -1 + {x}",
"{x}: 3D6 drop lowest 2 drop lowest -1 + {x}",
"{x}: 3D6 drop lowest {x} drop lowest -1",
"{x}: 3D6 drop lowest {x} drop lowest 1 drop lowest 1",
"{x}: 3D6 drop highest -2 drop highest {x} drop lowest 1",
"{x}: ({x}D1 drop lowest 2 drop lowest -1)D6"
]
{
let function = crate::compile_unoptimized(source).unwrap();
let optimized = optimize(function.clone(), Passes::all());
for x in [i32::MIN, -3, -1, 0, 1, 2, 3, 5]
{
let [unoptimized, optimized] =
[&function, &optimized].map(|function| {
crate::serial::HistogramBuilder::new(Evaluator::new(
function.clone()
))
.build([x])
.unwrap()
.iter()
.map(|(outcome, count)| (*outcome, *count))
.collect::<std::collections::BTreeMap<_, _>>()
});
assert_eq!(optimized, unoptimized, "{} with {}", source, x);
}
}
}