use std::path::Path;
const CONTRACT: &str = include_str!("../../../contracts/kernel-fusion-v1.yaml");
const WINDOW_BEFORE: usize = 1;
const WINDOW_AFTER: usize = 3;
fn normalise(s: &str) -> String {
s.chars()
.filter(|c| *c != '_')
.flat_map(char::to_lowercase)
.collect()
}
fn check_call_site(root: &Path, fused: &str, call_site: &str) -> Result<(), String> {
let (path, line) = call_site
.rsplit_once(':')
.ok_or_else(|| format!("`{call_site}` has no `:line`"))?;
let line: usize = line
.parse()
.map_err(|_| format!("`{call_site}`: `{line}` is not a line number"))?;
let src = std::fs::read_to_string(root.join(path))
.map_err(|e| format!("`{path}` does not exist from the workspace root ({e})"))?;
let lines: Vec<&str> = src.lines().collect();
if line == 0 || line > lines.len() {
return Err(format!(
"`{path}` has {} lines, cited line {line}",
lines.len()
));
}
let kernel = fused.split_whitespace().next().unwrap_or(fused);
let stem = normalise(kernel.strip_suffix("Kernel").unwrap_or(kernel));
let lo = line.saturating_sub(1 + WINDOW_BEFORE);
let hi = (line + WINDOW_AFTER).min(lines.len());
let window = lines[lo..hi].join("\n");
if normalise(&window).contains(&stem) {
Ok(())
} else {
Err(format!(
"`{call_site}` does not call `{kernel}`; lines {}..={hi} read:\n{window}",
lo + 1
))
}
}
fn workspace_root() -> std::path::PathBuf {
Path::new(env!("CARGO_MANIFEST_DIR")).join("../..")
}
#[test]
fn every_active_fusion_call_site_is_a_live_call_of_its_kernel() {
let doc: serde_yaml_ng::Value = serde_yaml_ng::from_str(CONTRACT).expect("parse contract");
let decisions = doc["fusion_decisions"]
.as_mapping()
.expect("fusion_decisions mapping");
let root = workspace_root();
let mut checked = 0;
let mut broken = Vec::new();
for (key, d) in decisions {
if d["status"].as_str() != Some("ACTIVE") {
continue;
}
let id = d["id"].as_str().unwrap_or("?");
let fused = d["kernels"]["fused"].as_str().unwrap_or_default();
let Some(call_site) = d["call_site"].as_str() else {
broken.push(format!("\n - {id} ({key:?}): ACTIVE with no call_site"));
continue;
};
checked += 1;
if let Err(e) = check_call_site(&root, fused, call_site) {
broken.push(format!("\n - {id}: {e}"));
}
}
assert!(checked >= 7, "only {checked} ACTIVE call_sites checked");
assert!(
broken.is_empty(),
"ACTIVE fusion call_site(s) that are not a live call of their kernel:{}",
broken.concat()
);
}
#[test]
fn the_call_site_check_rejects_each_stale_shape() {
let root = workspace_root();
let live = "crates/aprender-serve/src/cuda/kernels_generate_gemm_cuda.rs";
let src = std::fs::read_to_string(root.join(live)).expect("read live generator");
let at = src
.lines()
.position(|l| l.contains("KernelType::FusedQKV {"))
.expect("FusedQKV arm")
+ 1;
let fused = "FusedQKVKernel (x)";
assert!(check_call_site(&root, fused, &format!("{live}:{at}")).is_ok());
for (stale, why) in [
(format!("realizar/src/cuda/{at}.rs:{at}"), "moved path"),
(format!("{live}:999999"), "line past end of file"),
(live.to_string(), "no line number"),
(format!("{live}:{}", at + 40), "drifted line"),
(
"crates/aprender-serve/src/cuda/generate.rs:267".to_string(),
"deleted orphan",
),
] {
assert!(
check_call_site(&root, fused, &stale).is_err(),
"{why}: `{stale}` was accepted"
);
}
}