use darling::FromMeta;
use proc_macro2::TokenStream as TokenStream2;
use syn::{ItemFn, ItemMod};
#[derive(Debug, Clone, Copy, FromMeta)]
pub enum Overwritten {
Usb,
ChipConfig,
ChipInit,
Entry,
BindInterrupt,
}
pub(crate) fn find_overwritten(item_fn: &ItemFn) -> Option<darling::Result<Overwritten>> {
let marker = item_fn
.attrs
.iter()
.find(|attr| attr.path().is_ident("Override") || attr.path().is_ident("Overwritten"))?;
if item_fn
.attrs
.iter()
.any(|attr| attr.path().is_ident("cfg") || attr.path().is_ident("cfg_attr"))
{
return None;
}
Some(Overwritten::from_meta(&marker.meta))
}
pub(crate) fn validate_overwritten_attrs(item_mod: &ItemMod) -> Option<TokenStream2> {
let mut errors: Vec<darling::Error> = Vec::new();
if let Some((_, items)) = &item_mod.content {
for item in items {
if let syn::Item::Fn(item_fn) = item
&& let Some(Err(e)) = find_overwritten(item_fn)
{
errors.push(e);
}
}
}
if errors.is_empty() {
None
} else {
Some(darling::Error::multiple(errors).write_errors())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn parse_fn(src: &str) -> ItemFn {
syn::parse_str(src).expect("test fn should parse")
}
#[test]
fn no_overwritten_attribute_is_none() {
assert!(find_overwritten(&parse_fn("fn f() {}")).is_none());
assert!(find_overwritten(&parse_fn("#[inline]\nfn f() {}")).is_none());
}
#[test]
fn valid_variant_parses() {
let res = find_overwritten(&parse_fn("#[Overwritten(entry)]\nfn f() {}"));
assert!(matches!(res, Some(Ok(Overwritten::Entry))));
}
#[test]
fn override_spelling_is_accepted() {
let res = find_overwritten(&parse_fn("#[Override(chip_config)]\nfn f() {}"));
assert!(matches!(res, Some(Ok(Overwritten::ChipConfig))));
}
#[test]
fn bind_interrupt_passes_validation() {
let res = find_overwritten(&parse_fn("#[Override(bind_interrupt)]\nfn f() {}"));
assert!(matches!(res, Some(Ok(Overwritten::BindInterrupt))));
}
#[test]
fn doc_comment_does_not_disable_override() {
let res = find_overwritten(&parse_fn("/// my entry\n#[Overwritten(entry)]\nfn f() {}"));
assert!(matches!(res, Some(Ok(Overwritten::Entry))));
}
#[test]
fn cfg_gated_override_is_ignored() {
assert!(
find_overwritten(&parse_fn(
"#[cfg(feature = \"x\")]\n#[Override(chip_config)]\nfn f() {}"
))
.is_none()
);
assert!(
find_overwritten(&parse_fn(
"#[cfg(feature = \"x\")]\n#[Override(Entry)]\nfn f() {}"
))
.is_none()
);
assert!(
find_overwritten(&parse_fn(
"#[cfg_attr(test, inline)]\n#[Override(entry)]\nfn f() {}"
))
.is_none()
);
}
#[test]
fn inert_attribute_does_not_disable_override() {
let res = find_overwritten(&parse_fn(
"#[allow(dead_code)]\n#[Override(entry)]\nfn f() {}",
));
assert!(matches!(res, Some(Ok(Overwritten::Entry))));
let res = find_overwritten(&parse_fn(
"#[allow(dead_code)]\n#[Override(Entry)]\nfn f() {}",
));
assert!(matches!(res, Some(Err(_))));
}
#[test]
fn cfg_gated_override_produces_no_module_error() {
let item_mod: ItemMod =
syn::parse_str("mod kb { #[cfg(feature = \"x\")] #[Overwritten(Entry)] fn run() {} }")
.unwrap();
assert!(validate_overwritten_attrs(&item_mod).is_none());
}
#[test]
fn miscased_variant_is_an_error() {
let res = find_overwritten(&parse_fn("#[Overwritten(Entry)]\nfn f() {}"));
assert!(matches!(res, Some(Err(_))));
let res = find_overwritten(&parse_fn("#[Override(Entry)]\nfn f() {}"));
let err = match res {
Some(Err(e)) => e,
other => panic!("expected Some(Err(_)), got {other:?}"),
};
assert!(
err.to_string().contains("entry"),
"unexpected message: {err}"
);
}
#[test]
fn invalid_attr_in_module_becomes_compile_error() {
let item_mod: ItemMod =
syn::parse_str("mod kb { #[Overwritten(Entry)] fn run() {} }").unwrap();
let tokens = validate_overwritten_attrs(&item_mod).expect("should produce errors");
assert!(tokens.to_string().contains("compile_error"));
}
#[test]
fn valid_module_produces_no_errors() {
let item_mod: ItemMod =
syn::parse_str("mod kb { #[Overwritten(entry)] fn run() {} fn helper() {} }").unwrap();
assert!(validate_overwritten_attrs(&item_mod).is_none());
}
}