use std::collections::HashSet;
use proc_macro2::TokenStream;
use quote::quote;
use syn::parse::Parser;
use syn::{Attribute, GenericParam, Meta};
pub struct AttributeParser {
metas: Vec<Meta>,
}
impl AttributeParser {
pub fn new(tokens: impl Into<TokenStream>) -> Result<Self, syn::Error> {
use syn::Token;
use syn::punctuated::Punctuated;
let parser = Punctuated::<Meta, Token![,]>::parse_terminated;
let tokens: TokenStream = tokens.into();
let metas = parser.parse2(tokens)?;
Ok(Self {
metas: metas.into_iter().collect(),
})
}
pub fn new_validated(
tokens: impl Into<TokenStream>,
allowed_keys: &[&str],
) -> Result<Self, TokenStream> {
let parser = Self::new(tokens).map_err(|e| e.to_compile_error())?;
parser.validate_keys(allowed_keys)?;
Ok(parser)
}
pub fn get_int<T>(&self, name: &str) -> Result<Option<T>, TokenStream>
where
T: std::str::FromStr,
T::Err: std::fmt::Display,
{
let Some(meta) = self
.metas
.iter()
.find(|meta| matches!(meta, Meta::NameValue(nv) if nv.path.is_ident(name)))
else {
return Ok(None);
};
let Meta::NameValue(nv) = meta else {
return Ok(None);
};
let syn::Expr::Lit(syn::ExprLit {
lit: syn::Lit::Int(lit),
..
}) = &nv.value
else {
return Err(syn::Error::new_spanned(
&nv.value,
format!("`{name}` must be an integer literal"),
)
.to_compile_error());
};
lit.base10_parse().map(Some).map_err(|err| {
syn::Error::new_spanned(lit, format!("invalid `{name}` value: {err}"))
.to_compile_error()
})
}
pub fn get_path_array(&self, name: &str) -> Result<Vec<syn::Path>, TokenStream> {
let Some(meta) = self
.metas
.iter()
.find(|meta| matches!(meta, Meta::NameValue(nv) if nv.path.is_ident(name)))
else {
return Ok(vec![]);
};
let Meta::NameValue(nv) = meta else {
return Ok(vec![]);
};
let syn::Expr::Array(arr) = &nv.value else {
return Err(syn::Error::new_spanned(
&nv.value,
format!("`{name}` must be an array of type paths, e.g. `[EventA, EventB]`"),
)
.to_compile_error());
};
let mut result = Vec::with_capacity(arr.elems.len());
for elem in &arr.elems {
if let syn::Expr::Path(path_expr) = elem {
result.push(path_expr.path.clone());
} else {
return Err(syn::Error::new_spanned(
elem,
format!("invalid `{name}` element: expected a type path"),
)
.to_compile_error());
}
}
Ok(result)
}
pub fn get_path(&self, name: &str) -> Option<syn::Path> {
self.metas.iter().find_map(|meta| {
if let Meta::NameValue(nv) = meta
&& nv.path.is_ident(name)
&& let syn::Expr::Path(p) = &nv.value
{
Some(p.path.clone())
} else {
None
}
})
}
pub fn get_expr_tokens(&self, name: &str) -> Option<TokenStream> {
self.metas.iter().find_map(|meta| {
if let Meta::NameValue(nv) = meta
&& nv.path.is_ident(name)
{
let expr = &nv.value;
Some(quote! { #expr })
} else {
None
}
})
}
pub fn validate_keys(&self, allowed: &[&str]) -> Result<(), TokenStream> {
for meta in &self.metas {
let Meta::NameValue(nv) = meta else {
return Err(syn::Error::new_spanned(
meta,
"invalid attribute syntax. Expected `key = value`",
)
.to_compile_error());
};
let Some(key_ident) = nv.path.get_ident() else {
return Err(syn::Error::new_spanned(
&nv.path,
"invalid attribute key. Expected a simple identifier",
)
.to_compile_error());
};
let key = key_ident.to_string();
if !allowed.contains(&key.as_str()) {
let allowed_list = allowed.join(", ");
return Err(syn::Error::new_spanned(
&nv.path,
format!("unknown attribute `{key}`. Expected one of: {allowed_list}"),
)
.to_compile_error());
}
}
Ok(())
}
}
pub fn deduplicate_type_generics(generics: &syn::Generics) -> TokenStream {
let mut seen = HashSet::new();
let mut unique_params = Vec::new();
for param in &generics.params {
let name = match param {
GenericParam::Type(t) => t.ident.to_string(),
GenericParam::Lifetime(l) => l.lifetime.to_string(),
GenericParam::Const(c) => c.ident.to_string(),
};
if seen.insert(name) {
match param {
GenericParam::Type(t) => {
let ident = &t.ident;
unique_params.push(quote! { #ident });
}
GenericParam::Lifetime(l) => {
let lifetime = &l.lifetime;
unique_params.push(quote! { #lifetime });
}
GenericParam::Const(c) => {
let ident = &c.ident;
unique_params.push(quote! { #ident });
}
}
}
}
if unique_params.is_empty() {
quote! {}
} else {
quote! { < #(#unique_params),* > }
}
}
pub fn has_derive(attrs: &[Attribute], derive_name: &str) -> bool {
use syn::punctuated::Punctuated;
use syn::{Path, Token};
attrs.iter().any(|attr| {
if !attr.path().is_ident("derive") {
return false;
}
let Meta::List(meta_list) = &attr.meta else {
return false;
};
let parser = Punctuated::<Path, Token![,]>::parse_terminated;
let Ok(paths) = parser.parse2(meta_list.tokens.clone()) else {
return false;
};
paths.iter().any(|path| {
path.segments
.last()
.map(|seg| seg.ident == derive_name)
.unwrap_or(false)
})
})
}
pub fn has_runnable_marker(attrs: &[Attribute]) -> bool {
attrs.iter().any(is_runnable_generated_attr)
}
pub fn is_runnable_generated_attr(attr: &Attribute) -> bool {
let path = attr.path();
path.is_ident("runnable_generated")
|| (path.segments.len() == 2
&& path.segments[0].ident == "rmk"
&& path.segments[1].ident == "runnable_generated")
|| (path.segments.len() == 3
&& path.segments[0].ident == "rmk"
&& path.segments[1].ident == "macros"
&& path.segments[2].ident == "runnable_generated")
}
pub fn attr_matches_name(attr: &Attribute, name: &str) -> bool {
attr.path()
.segments
.last()
.map(|seg| seg.ident == name)
.unwrap_or(false)
}