use std::fs;
use std::path::{Path, PathBuf};
use syn::visit::Visit;
const GUARD_ALLOW: &[&str] = &[
"type_ref",
];
const EXCLUDED: &[&str] = &["raw.rs", "macros.rs", "claim.rs"];
fn sources() -> Vec<(PathBuf, String)> {
let mut files = Vec::new();
collect_rs(
&Path::new(env!("CARGO_MANIFEST_DIR")).join("src"),
&mut files,
);
files.sort();
files
.into_iter()
.map(|p| {
let text =
fs::read_to_string(&p).unwrap_or_else(|e| panic!("read {}: {e}", p.display()));
(p, text)
})
.collect()
}
fn collect_rs(dir: &Path, out: &mut Vec<PathBuf>) {
for entry in fs::read_dir(dir).expect("read_dir src") {
let path = entry.expect("dir entry").path();
if path.is_dir() {
collect_rs(&path, out);
} else if path.extension().is_some_and(|e| e == "rs")
&& !path
.file_name()
.and_then(|n| n.to_str())
.is_some_and(|n| EXCLUDED.contains(&n))
{
out.push(path);
}
}
}
fn is_database(ty: &syn::Type) -> bool {
matches!(ty, syn::Type::Path(p) if p.path.segments.last().is_some_and(|s| s.ident == "Database"))
}
fn call_segments(call: &syn::ExprCall) -> Option<Vec<String>> {
let syn::Expr::Path(p) = &*call.func else {
return None;
};
Some(
p.path
.segments
.iter()
.map(|s| s.ident.to_string())
.collect(),
)
}
#[derive(Default)]
struct BodyScan {
reaches_kernel: bool,
has_guard: bool,
}
impl<'ast> Visit<'ast> for BodyScan {
fn visit_expr_call(&mut self, node: &'ast syn::ExprCall) {
if let Some(segs) = call_segments(node) {
if segs.len() == 2
&& matches!(segs[0].as_str(), "sys" | "idakit_sys")
&& segs[1].starts_with(|c: char| c.is_ascii_lowercase())
{
self.reaches_kernel = true;
}
if segs.last().map(String::as_str) == Some("ensure_kernel_thread") {
self.has_guard = true;
}
}
syn::visit::visit_expr_call(self, node);
}
}
struct GuardCheck<'a> {
path: &'a Path,
on_database: bool,
violations: &'a mut Vec<String>,
}
impl<'ast> Visit<'ast> for GuardCheck<'_> {
fn visit_item_impl(&mut self, node: &'ast syn::ItemImpl) {
let outer = self.on_database;
self.on_database = is_database(&node.self_ty);
syn::visit::visit_item_impl(self, node);
self.on_database = outer;
}
fn visit_impl_item_fn(&mut self, node: &'ast syn::ImplItemFn) {
if !self.on_database {
return;
}
let mut scan = BodyScan::default();
scan.visit_block(&node.block);
let name = node.sig.ident.to_string();
if scan.reaches_kernel && !scan.has_guard && !GUARD_ALLOW.contains(&name.as_str()) {
self.violations.push(format!(
"{}: Database::{name} calls idakit_sys directly but never runs \
ensure_kernel_thread(); a migrated database would touch the kernel on the wrong \
thread",
self.path.display()
));
}
}
}
#[test]
fn database_kernel_methods_are_guarded() {
let mut violations = Vec::new();
for (path, src) in sources() {
let file =
syn::parse_file(&src).unwrap_or_else(|e| panic!("parse {}: {e}", path.display()));
GuardCheck {
path: &path,
on_database: false,
violations: &mut violations,
}
.visit_file(&file);
}
assert!(
violations.is_empty(),
"every Database method that calls idakit_sys must first run \
claim::ensure_kernel_thread(); Database is Send, so it may have migrated threads. Add the \
guard, or if a guarded Database method it calls first already claims the thread, add the \
method to GUARD_ALLOW:\n{}",
violations.join("\n")
);
}