use std::collections::HashSet;
use pretty_assertions::assert_eq;
use rand::{SeedableRng, rngs::StdRng};
use crate::{
Evaluator, Pass, Passes,
support::{compile_valid, optimize, read_compilation_test_cases}
};
#[test]
fn test_common_subexpression_elimination()
{
let mut seen = HashSet::new();
for (index, (source, expected)) in read_compilation_test_cases(
include_str!("../../tests/test_common_subexpression_elimination.txt")
)
.iter()
.enumerate()
{
assert!(seen.insert(source), "duplicate test case: {}", source);
let function = compile_valid(source);
let optimized = optimize(
function.clone(),
Pass::CommonSubexpressionElimination.into()
);
let actual = format!("{}", optimized);
assert_eq!(
actual.trim(),
*expected,
"case {}: {}\nunoptimized:\n{}",
index + 1,
source,
function
);
}
}
#[test]
fn test_constant_commuting()
{
let mut seen = HashSet::new();
for (index, (source, expected)) in read_compilation_test_cases(
include_str!("../../tests/test_constant_commuting.txt")
)
.iter()
.enumerate()
{
assert!(seen.insert(source), "duplicate test case: {}", source);
let function = compile_valid(source);
let optimized =
optimize(function.clone(), Pass::ConstantCommuting.into());
let actual = format!("{}", optimized);
assert_eq!(
actual.trim(),
*expected,
"case {}: {}\nunoptimized:\n{}",
index + 1,
source,
function
);
}
}
#[test]
fn test_constant_folding()
{
let mut seen = HashSet::new();
for (index, (source, expected)) in read_compilation_test_cases(
include_str!("../../tests/test_constant_folding.txt")
)
.iter()
.enumerate()
{
assert!(seen.insert(source), "duplicate test case: {}", source);
let function = compile_valid(source);
let optimized =
optimize(function.clone(), Pass::ConstantFolding.into());
let actual = format!("{}", optimized);
assert_eq!(
actual.trim(),
*expected,
"case {}: {}\nunoptimized:\n{}",
index + 1,
source,
function
);
}
}
#[test]
fn test_strength_reduction()
{
let mut seen = HashSet::new();
for (index, (source, expected)) in read_compilation_test_cases(
include_str!("../../tests/test_strength_reduction.txt")
)
.iter()
.enumerate()
{
assert!(seen.insert(source), "duplicate test case: {}", source);
let function = compile_valid(source);
let optimized =
optimize(function.clone(), Pass::StrengthReduction.into());
let actual = format!("{}", optimized);
assert_eq!(
actual.trim(),
*expected,
"case {}: {}\nunoptimized:\n{}",
index + 1,
source,
function
);
}
}
#[test]
fn test_register_coalescence()
{
let mut seen = HashSet::new();
for (index, (source, expected)) in read_compilation_test_cases(
include_str!("../../tests/test_register_coalescence.txt")
)
.iter()
.enumerate()
{
assert!(seen.insert(source), "duplicate test case: {}", source);
let function = compile_valid(source);
let optimized =
optimize(function.clone(), Pass::RegisterCoalescing.into());
let actual = format!("{}", optimized);
assert_eq!(
actual.trim(),
*expected,
"case {}: {}\nunoptimized:\n{}",
index + 1,
source,
function
);
}
}
#[test]
fn test_full_optimization()
{
let mut seen = HashSet::new();
for (index, (source, expected)) in read_compilation_test_cases(
include_str!("../../tests/test_full_optimization.txt")
)
.iter()
.enumerate()
{
assert!(seen.insert(source), "duplicate test case: {}", source);
let function = compile_valid(source);
let optimized = optimize(function.clone(), Passes::all());
let actual = format!("{}", optimized);
assert_eq!(
actual.trim(),
*expected,
"case {}: {}\nunoptimized:\n{}",
index + 1,
source,
function
);
}
}
#[test]
fn test_optimize_max()
{
for (body, expected) in [
("@1 <- 3 max 5\n\t\treturn @1", "return 5"),
("@1 <- @0 max -2147483648\n\t\treturn @1", "return @0"),
(
"@1 <- 2147483647 max @0\n\t\treturn @1",
"return 2147483647"
),
("@1 <- @0 max @0\n\t\treturn @1", "return @0"),
(
"@1 <- @0 max 3\n\t\t@2 <- 5 max @1\n\t\treturn @2",
"@0 <- 5 max @0\n\t\treturn @0"
)
]
{
let registers = body.matches("<-").count() + 1;
let text = format!(
"Function({{x}}@0) r#{} âš…#0\n\textern[]\n\tbody:\n\t\t{}\n",
registers, body
);
let function = crate::Assembler::assemble(&text).unwrap();
let optimized = optimize(function, Passes::all()).to_string();
let actual = optimized.trim().split_once("body:\n\t\t").unwrap().1;
assert_eq!(actual, expected, "{}", body);
}
}
#[test]
fn test_coalescing_keeps_dead_writes_apart()
{
let function = crate::Assembler::assemble(
"\
Function({x}@0) r#4 âš…#0
\textern[]
\tbody:
\t\t@1 <- @0 + 1
\t\t@2 <- @0 + 2
\t\t@3 <- @1 + @0
\t\treturn @3
"
)
.unwrap();
let coalesced = optimize(function.clone(), Pass::RegisterCoalescing.into());
let evaluate = |function| {
Evaluator::new(function)
.evaluate([10], &mut StdRng::seed_from_u64(0))
.unwrap()
.result
};
assert_eq!(evaluate(coalesced.clone()), 21, "{}", coalesced);
assert_eq!(evaluate(function), 21);
}
const EDGES: [i32; 9] = [
i32::MIN,
i32::MIN + 1,
-65536,
-1,
0,
1,
65536,
i32::MAX - 1,
i32::MAX
];
const COMBINATIONS: usize = 10_000;
fn assert_exact(source: &str, passes: Passes)
{
let unoptimized = compile_valid(source);
let optimized = optimize(unoptimized.clone(), passes);
let arity = unoptimized.arity();
let variables = arity + unoptimized.externals.len();
let mut edges = EDGES.len();
while edges > 2 && edges.pow(variables as u32) > COMBINATIONS
{
edges -= 1;
}
let edges = &EDGES[EDGES.len() - edges..];
for mut combination in 0..edges.len().pow(variables as u32)
{
let values = (0..variables)
.map(|_| {
let value = edges[combination % edges.len()];
combination /= edges.len();
value
})
.collect::<Vec<_>>();
let evaluate = |function| {
let mut evaluator = Evaluator::new(function);
for (name, value) in
unoptimized.externals.iter().zip(&values[arity..])
{
evaluator.bind(name, *value).unwrap();
}
evaluator
.evaluate(
values[..arity].iter().copied(),
&mut StdRng::seed_from_u64(0)
)
.unwrap()
.result
};
assert_eq!(
evaluate(optimized.clone()),
evaluate(unoptimized.clone()),
"{} at {:?}\noptimized:\n{}",
source,
values,
optimized
);
}
}
#[test]
fn test_test_cases_are_exact()
{
let files = [
(
include_str!(
"../../tests/test_common_subexpression_elimination.txt"
),
Passes::from(Pass::CommonSubexpressionElimination)
),
(
include_str!("../../tests/test_constant_commuting.txt"),
Pass::ConstantCommuting.into()
),
(
include_str!("../../tests/test_constant_folding.txt"),
Pass::ConstantFolding.into()
),
(
include_str!("../../tests/test_strength_reduction.txt"),
Pass::StrengthReduction.into()
),
(
include_str!("../../tests/test_register_coalescence.txt"),
Pass::RegisterCoalescing.into()
),
(
include_str!("../../tests/test_full_optimization.txt"),
Passes::all()
)
];
for (file, passes) in files
{
for (source, _) in read_compilation_test_cases(file)
{
if compile_valid(source).rolling_record_count == 0
{
assert_exact(source, passes);
assert_exact(source, Passes::all());
}
}
}
}
#[test]
fn test_strength_reduction_is_exact()
{
for source in [
"{x}: {x} / {x}",
"{x}: --{x}",
"{x}: ---{x}",
"{x}: --{x} + 2",
"{x}, {y}: {a}@(-{x}) + {y} * 2 + -{a}",
"{x}: {x} * -1",
"{x}: {x} / -1",
"{x}: 0 - {x}",
"{x}: {x} * 2",
"{x}: {x} ^ 2",
"{x}: {x} % 1",
"{x}: {x} ^ 0",
"{x}: 1 ^ {x}"
]
{
assert_exact(source, Pass::StrengthReduction.into());
assert_exact(source, Passes::all());
}
}