extern crate proc_macro;
use proc_macro::TokenStream;
use quote::quote;
use syn::{ext::IdentExt, fold::Fold, punctuated::Punctuated, spanned::Spanned, visit::Visit, *};
struct RemoveMut;
impl Fold for RemoveMut {
fn fold_pat_ident(&mut self, mut i: PatIdent) -> PatIdent {
i.mutability = None;
i
}
fn fold_arg_self(&mut self, mut i: ArgSelf) -> ArgSelf {
i.mutability = None;
i
}
}
struct HasSelfType(bool);
impl<'ast> Visit<'ast> for HasSelfType {
fn visit_ident(&mut self, i: &'ast Ident) {
if i == "Self" {
self.0 = true;
}
}
fn visit_item(&mut self, _: &'ast Item) {
}
}
enum Kind {
UnsafeFn,
SafeBody,
}
struct FnOrMethod {
attrs: Vec<Attribute>,
vis: Visibility,
constness: Option<token::Const>,
asyncness: Option<token::Async>,
unsafety: Option<token::Unsafe>,
abi: Option<Abi>,
ident: Ident,
decl: FnDecl,
block: Option<Block>,
semi_token: Option<token::Semi>,
}
impl From<ItemFn> for FnOrMethod {
fn from(itemfn: ItemFn) -> FnOrMethod {
FnOrMethod {
attrs: itemfn.attrs,
vis: itemfn.vis,
constness: itemfn.constness,
asyncness: itemfn.asyncness,
unsafety: itemfn.unsafety,
abi: itemfn.abi,
ident: itemfn.ident,
decl: *itemfn.decl,
block: Some(*itemfn.block),
semi_token: None,
}
}
}
impl From<TraitItemMethod> for FnOrMethod {
fn from(m: TraitItemMethod) -> FnOrMethod {
FnOrMethod {
attrs: m.attrs,
vis: Visibility::Inherited,
constness: m.sig.constness,
asyncness: None,
unsafety: m.sig.unsafety,
abi: m.sig.abi,
ident: m.sig.ident,
decl: m.sig.decl,
block: m.default,
semi_token: m.semi_token,
}
}
}
#[proc_macro_attribute]
pub fn unsafe_fn(_attr: TokenStream, item: TokenStream) -> TokenStream {
if let Ok(m) = parse::<TraitItemMethod>(item.clone()) {
return unsafe_fn_impl(m.into(), Kind::UnsafeFn);
}
let item = parse_macro_input!(item as Item);
match item {
Item::Fn(f) => unsafe_fn_impl(f.into(), Kind::UnsafeFn),
Item::Trait(t) => quote!(unsafe #t).into(),
_ => Error::new(
item.span(),
"#[unsafe_fn] can only be applied to functions or traits",
)
.to_compile_error()
.into(),
}
}
#[proc_macro_attribute]
pub fn safe_body(_attr: TokenStream, item: TokenStream) -> TokenStream {
if let Ok(m) = parse::<TraitItemMethod>(item.clone()) {
return unsafe_fn_impl(m.into(), Kind::SafeBody);
}
let item = parse_macro_input!(item as ItemFn);
unsafe_fn_impl(item.into(), Kind::SafeBody)
}
fn unsafe_fn_impl(
FnOrMethod {
attrs,
vis,
constness,
asyncness,
unsafety,
abi,
ident,
decl,
block,
semi_token,
}: FnOrMethod,
k: Kind,
) -> TokenStream {
let unsafety = match (k, unsafety) {
(Kind::UnsafeFn, None) => <Token![unsafe]>::default(),
(Kind::SafeBody, Some(u)) => u,
(Kind::UnsafeFn, Some(u)) => {
return Error::new(u.span(), "#[unsafe_fn] already marked unsafe")
.to_compile_error()
.into()
}
(Kind::SafeBody, None) => {
return Error::new(
proc_macro::Span::call_site().into(),
"#[safe_body] function must be marked as unsafe",
)
.to_compile_error()
.into()
}
};
let FnDecl {
fn_token,
generics,
paren_token: _paren_token,
inputs,
variadic,
output,
} = &decl;
let (impl_generics, _, where_clause) = generics.split_for_impl();
let unsafe_fn_name = Ident::new(
&format!("__unsafe_fn_{}", ident.unraw().to_string()),
ident.span(),
);
let block = match block {
None => {
let inner_where = match &where_clause {
Some(w) => quote!(#w, Self:Sized),
None => quote!(where Self:Sized),
};
return quote!(
#(#attrs)* #vis #constness #asyncness #unsafety #abi
#fn_token #ident #impl_generics (#inputs #variadic) #output #where_clause
#semi_token
#[doc(hide)]
#[inline]
#constness #asyncness
#fn_token #unsafe_fn_name #impl_generics (#inputs #variadic) #output #inner_where
{ ::std::panic!("Not to be called"); }
)
.into();
}
Some(block) => block,
};
let mut main_param = Punctuated::<FnArg, Token!(,)>::new();
let mut sub_param = Punctuated::<FnArg, Token!(,)>::new();
let mut sub_args = Punctuated::<Ident, Token!(,)>::new();
let mut wrap_self = false;
for it in inputs.iter() {
match it {
FnArg::SelfRef(_) | FnArg::SelfValue(_) => {
sub_param.push(it.clone());
main_param.push(RemoveMut.fold_fn_arg(it.clone()));
wrap_self = true;
}
FnArg::Captured(ArgCaptured {
pat,
colon_token,
ty,
}) => {
if let Pat::Ident(i) = pat {
main_param.push(RemoveMut.fold_fn_arg(it.clone()));
sub_param.push(it.clone());
if i.ident == "self" {
wrap_self = true;
} else {
sub_args.push(i.ident.clone());
}
} else {
let name = Ident::new(&format!("__unsafe_fn_arg{}", sub_args.len()), it.span());
main_param.push(parse(quote!(#name #colon_token #ty).into()).unwrap());
sub_param.push(it.clone());
sub_args.push(name);
}
}
FnArg::Inferred(_) => {
unimplemented!();
}
FnArg::Ignored(_) => {
main_param.push(it.clone());
}
}
}
let fun = quote! {
#[doc(hide)]
#[inline]
#constness #asyncness #fn_token #unsafe_fn_name #impl_generics (#sub_param #variadic) #output #where_clause {
#block
}
};
let fdecl = quote! {
#(#attrs)* #vis #constness #asyncness #unsafety #abi
#fn_token #ident #impl_generics (#main_param #variadic) #output #where_clause
};
let type_params: Vec<_> = generics.type_params().map(|x| &x.ident).collect();
let turbo = if type_params.is_empty() {
quote!()
} else {
quote!(::< #(#type_params),* >)
};
let r = if wrap_self {
quote! {
#fun
#fdecl {
self.#unsafe_fn_name #turbo (#sub_args)
}
}
} else if {
let mut has_self = HasSelfType(false);
has_self.visit_fn_decl(&decl);
has_self.visit_block(&block);
has_self.0
} {
quote! {
#fun
#fdecl {
Self::#unsafe_fn_name #turbo (#sub_args)
}
}
} else {
quote!(
#fdecl {
#fun
#unsafe_fn_name #turbo (#sub_args)
}
)
};
r.into()
}