use polydat::ast::PortType;
use polydat::dsl::compile_polydat;
use polydat::kernel::{InputKind, PolydatProgram};
fn compile(src: &str) -> polydat::kernel::PolydatKernel {
compile_polydat(src).unwrap_or_else(|e| panic!("compile failed: {e}\n{src}"))
}
fn elements(p: &PolydatProgram, idx: usize) -> Vec<(String, PortType)> {
p.traversals()[idx].elements.clone()
}
#[test]
fn element_types_follow_the_source_table() {
let k = compile(
"input cycle: u64\n\
for a in 1..4, b in 10,20,30, c in 1.5,2.5, d in load,verify, e in true,false, p in partitions(\"*/2\", 100) {\n\
x := hash(a)\n\
}\n",
);
assert_eq!(
elements(k.program(), 0),
vec![
("a".to_string(), PortType::U64),
("b".to_string(), PortType::U64),
("c".to_string(), PortType::F64),
("d".to_string(), PortType::Str),
("e".to_string(), PortType::Bool),
("p".to_string(), PortType::Ext),
]
);
}
#[test]
fn int_among_floats_widens_and_mixed_types_are_an_error() {
let k = compile("input cycle: u64\nfor m in 1,2.5,3 {\n x := m\n}\n");
assert_eq!(elements(k.program(), 0)[0].1, PortType::F64);
let err = compile_polydat("input cycle: u64\nfor m in 1,load {\n x := m\n}\n").unwrap_err();
assert!(err.contains("mixes"), "{err}");
}
#[test]
fn generator_call_sources_take_the_node_return_type() {
let k = compile("input cycle: u64\nfor g in hash_range(cycle, 10) {\n x := g\n}\n");
assert_eq!(elements(k.program(), 0)[0].1, PortType::U64);
}
#[test]
fn body_is_a_child_program_with_iteration_externs() {
let k = compile(
"input cycle: u64\nfor k in 1..4, limit in 10,20,30 {\n f := hash(k)\n g := u64_add(limit, k)\n}\n",
);
let parent = k.program();
assert_eq!(parent.traversals().len(), 1);
let t = &parent.traversals()[0];
assert_eq!(t.span.line, 2);
let child = &t.program;
for name in ["k", "limit"] {
let idx = child
.find_input(name)
.unwrap_or_else(|| panic!("child lacks {name}"));
assert_eq!(child.input_kind(idx), Some(InputKind::IterationExtern));
assert_eq!(child.input_port_type(name), Some(PortType::U64));
}
assert_eq!(child.coord_count(), 1);
assert_eq!(child.input_name_by_idx(0), Some("cycle"));
assert!(child.output_names().contains(&"f"));
assert!(child.output_names().contains(&"g"));
assert!(!parent.output_names().contains(&"f"));
}
#[test]
fn outer_wires_cascade_with_the_parents_types() {
let k = compile(
"input cycle: u64\nextern scale: f64 = 2.0\nbase := hash(cycle)\nlabel := \"run\"\n\
for k in 1..4 {\n f := u64_add(base, k)\n s := \"{label}-{k}\"\n z := f64_mul(scale, 1.5)\n}\n",
);
let t = &k.program().traversals()[0];
let mut cascade = t.cascade.clone();
cascade.sort_by(|a, b| a.0.cmp(&b.0));
assert_eq!(
cascade,
vec![
("base".to_string(), PortType::U64),
("label".to_string(), PortType::Str),
("scale".to_string(), PortType::F64),
]
);
for (name, ty) in &cascade {
assert_eq!(t.program.input_port_type(name), Some(*ty));
}
}
#[test]
fn traversal_over_a_producer_resolves_its_comprehension() {
let k = compile(
"input cycle: u64\nsweep := for k in 1..4, limit in 10,20,30 order halton/5\nfor sweep {\n f := hash(k)\n g := u64_add(limit, k)\n}\n",
);
let p = k.program();
assert_eq!(p.producers().len(), 1);
assert_eq!(p.producers()[0].name, "sweep");
assert_eq!(
elements(p, 0),
vec![
("k".to_string(), PortType::U64),
("limit".to_string(), PortType::U64)
]
);
let err =
compile_polydat("input cycle: u64\nfor nowhere {\n f := hash(cycle)\n}\n").unwrap_err();
assert!(err.contains("no producer named 'nowhere'"), "{err}");
}
#[test]
fn nested_bodies_compile_to_nested_programs_once_each() {
let k = compile(
"input cycle: u64\nfor p in partitions(\"*/4\", 1000) {\n outer := cardinality(p)\n\
for t in 0..20 {\n mid := u64_add(outer, t)\n for d in 0..50 {\n leaf := u64_add(mid, d)\n }\n }\n}\n",
);
let top = k.program();
assert_eq!(top.traversals().len(), 1);
let level1 = &top.traversals()[0].program;
assert_eq!(level1.traversals().len(), 1);
let level2 = &level1.traversals()[0].program;
assert_eq!(level2.traversals().len(), 1);
let level3 = &level2.traversals()[0].program;
assert!(level3.traversals().is_empty());
assert_eq!(
level2.traversals()[0].cascade,
vec![("mid".to_string(), PortType::U64)]
);
assert_eq!(
level1.traversals()[0].cascade,
vec![("outer".to_string(), PortType::U64)]
);
assert!(level3.output_names().contains(&"leaf"));
}
#[test]
fn body_type_errors_are_reported_at_compile_time() {
let err = compile_polydat(
"input cycle: u64\nfor name in load,verify {\n x := u64_add(name, 1)\n}\n",
)
.unwrap_err();
assert!(err.contains("for name in load,verify"), "{err}");
assert!(err.contains("body failed to compile"), "{err}");
}
#[test]
fn unknown_outer_name_is_an_error_in_the_body() {
let err = compile_polydat("input cycle: u64\nfor k in 1..4 {\n x := hash(missing)\n}\n")
.unwrap_err();
assert!(err.contains("missing"), "{err}");
}
#[test]
fn body_may_not_declare_another_coordinate() {
let err = compile_polydat(
"input cycle: u64\nfor k in 1..4 {\n input other: u64\n x := hash(other)\n}\n",
)
.unwrap_err();
assert!(err.contains("cannot declare input 'other'"), "{err}");
}
#[test]
fn a_body_with_a_cursor_over_an_element_compiles() {
let k = compile(
"input cycle: u64\nfor p in partitions(\"*/4\", 1000) {\n cursor rows = range(0, 1000) over p\n row := mod_in(cycle, rows.cursor)\n}\n",
);
let child = &k.program().traversals()[0].program;
assert_eq!(child.cursor_schemas().len(), 1);
assert!(child.output_names().contains(&"row"));
}