extern crate proc_macro;
use proc_macro::TokenStream;
use quote::quote;
use syn::{
parse_macro_input, ItemFn, Type, ReturnType, GenericArgument, PathArguments,
parse::Parse, parse::ParseStream, Error, Result as SynResult,
visit_mut::{self, VisitMut}, Expr, Ident, Lit,
spanned::Spanned,
};
fn to_pascal_case(s: &str) -> String {
let mut pascal = String::new();
let mut capitalize = true;
for c in s.chars() {
if c == '_' {
capitalize = true;
} else if capitalize {
pascal.push(c.to_ascii_uppercase());
capitalize = false;
} else {
pascal.push(c);
}
}
pascal
}
struct PurityCheckVisitor {
errors: Vec<Error>,
}
impl VisitMut for PurityCheckVisitor {
fn visit_expr_mut(&mut self, i: &mut Expr) {
match i {
Expr::Unsafe(e) => {
self.errors.push(Error::new(
e.span(),
"impure `unsafe` block found in function marked as `pure`",
));
}
Expr::Macro(e) => {
if e.mac.path.is_ident("asm") {
self.errors.push(Error::new(
e.span(),
"impure inline assembly found in function marked as `pure`",
));
}
}
Expr::MethodCall(e) => {
self.errors.push(Error::new(
e.span(), "method calls are not supported in pure functions"
));
}
Expr::Call(call_expr) => {
if let Expr::Path(expr_path) = &*call_expr.func {
let path = &expr_path.path;
if let Some(segment) = path.segments.last() {
if segment.ident == "Ok" || segment.ident == "Err" {
visit_mut::visit_expr_call_mut(self, call_expr);
return;
}
}
let mut zst_path = path.clone();
if let Some(last_segment) = zst_path.segments.last_mut() {
let ident_str = last_segment.ident.to_string();
let pascal_case_ident = to_pascal_case(&ident_str);
last_segment.ident = Ident::new(&pascal_case_ident, last_segment.ident.span());
let _fn_name = path.segments.last().unwrap().ident.to_string();
visit_mut::visit_expr_call_mut(self, call_expr);
let new_node = syn::parse_quote!({
{
let _ = || {
fn _assert_pure_function<T: crate::traits::IsPure>(_: T) {}
_assert_pure_function(#zst_path);
};
#call_expr
}
});
*i = new_node;
return; }
} else {
self.errors.push(Error::new_spanned(&call_expr.func, "closures and other complex function call expressions are not supported in pure functions"));
}
}
_ => {}
}
visit_mut::visit_expr_mut(self, i);
}
}
#[proc_macro_attribute]
pub fn pure(_args: TokenStream, item: TokenStream) -> TokenStream {
let mut input_fn = parse_macro_input!(item as ItemFn);
let mut visitor = PurityCheckVisitor { errors: vec![] };
let mut new_body_box = input_fn.block.clone();
visitor.visit_block_mut(&mut new_body_box);
if !visitor.errors.is_empty() {
let combined_errors = visitor.errors.into_iter().reduce(|mut a, b| {
a.combine(b);
a
});
if let Some(errors) = combined_errors {
return errors.to_compile_error().into();
}
}
input_fn.block = new_body_box;
let fn_name_str = input_fn.sig.ident.to_string();
let zst_name = Ident::new(&to_pascal_case(&fn_name_str), input_fn.sig.ident.span());
let expanded = quote! {
#input_fn
#[doc(hidden)]
struct #zst_name;
#[doc(hidden)]
impl crate::traits::IsPure for #zst_name {}
};
TokenStream::from(expanded)
}
struct AttributeArgs {
strategy_type: Type,
}
impl Parse for AttributeArgs {
fn parse(input: ParseStream) -> SynResult<Self> {
let strategy_type: Type = input.parse()?;
Ok(AttributeArgs { strategy_type })
}
}
fn extract_result_types(return_type: &Type) -> SynResult<(Type, Type)> {
if let Type::Path(type_path) = return_type {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "Result" {
if let PathArguments::AngleBracketed(args) = &segment.arguments {
if args.args.len() == 2 {
if let (
GenericArgument::Type(ok_type),
GenericArgument::Type(err_type)
) = (&args.args[0], &args.args[1]) {
return Ok((ok_type.clone(), err_type.clone()));
}
}
}
}
}
}
Err(Error::new_spanned(
return_type,
"Expected function to return Result<T, E>"
))
}
#[proc_macro_attribute]
pub fn error_strategy(args: TokenStream, item: TokenStream) -> TokenStream {
let input_fn = parse_macro_input!(item as ItemFn);
let args = parse_macro_input!(args as AttributeArgs);
let strategy_type = args.strategy_type;
let fn_name = &input_fn.sig.ident;
let fn_vis = &input_fn.vis;
let fn_inputs = &input_fn.sig.inputs;
let fn_body = &input_fn.block;
let fn_asyncness = &input_fn.sig.asyncness;
let fn_generics = &input_fn.sig.generics;
let where_clause = &input_fn.sig.generics.where_clause;
let (ok_type, err_type) = match &input_fn.sig.output {
ReturnType::Type(_, ty) => {
match extract_result_types(ty) {
Ok(types) => types,
Err(e) => return e.to_compile_error().into(),
}
}
ReturnType::Default => {
return Error::new_spanned(
&input_fn.sig,
"Function must return Result<T, E>"
).to_compile_error().into();
}
};
let original_impl_name = syn::Ident::new(
&format!("{}_original_impl", fn_name),
fn_name.span()
);
let strategy_name = quote!(#strategy_type).to_string();
let param_names: Vec<_> = input_fn.sig.inputs.iter().filter_map(|arg| {
if let syn::FnArg::Typed(pat_type) = arg {
if let syn::Pat::Ident(pat_ident) = &*pat_type.pat {
Some(&pat_ident.ident)
} else {
None
}
} else {
None
}
}).collect();
let function_call = if fn_asyncness.is_some() {
quote! { #original_impl_name(#(#param_names),*).await }
} else {
quote! { #original_impl_name(#(#param_names),*) }
};
let expanded = quote! {
#[doc(hidden)]
#fn_asyncness fn #original_impl_name #fn_generics (#fn_inputs) -> Result<#ok_type, #err_type> #where_clause
#fn_body
#fn_vis #fn_asyncness fn #fn_name #fn_generics (#fn_inputs) -> crate::PipexResult<#ok_type, #err_type> #where_clause {
let result = #function_call;
crate::PipexResult::new(result, #strategy_name)
}
};
TokenStream::from(expanded)
}
struct MemoizedArgs {
capacity: Option<usize>,
}
impl Parse for MemoizedArgs {
fn parse(input: ParseStream) -> SynResult<Self> {
let mut capacity = None;
while !input.is_empty() {
let lookahead = input.lookahead1();
if lookahead.peek(syn::Ident) {
let ident: Ident = input.parse()?;
if ident == "capacity" {
input.parse::<syn::Token![=]>()?;
let lit: Lit = input.parse()?;
if let Lit::Int(lit_int) = lit {
capacity = Some(lit_int.base10_parse()?);
} else {
return Err(Error::new_spanned(lit, "capacity must be an integer"));
}
} else {
return Err(Error::new_spanned(ident, "unknown attribute argument"));
}
if input.peek(syn::Token![,]) {
input.parse::<syn::Token![,]>()?;
}
} else {
return Err(lookahead.error());
}
}
Ok(MemoizedArgs { capacity })
}
}
impl Default for MemoizedArgs {
fn default() -> Self {
Self { capacity: Some(1000) } }
}
#[proc_macro_attribute]
pub fn memoized(args: TokenStream, item: TokenStream) -> TokenStream {
let args = if args.is_empty() {
MemoizedArgs::default()
} else {
parse_macro_input!(args as MemoizedArgs)
};
let input_fn = parse_macro_input!(item as ItemFn);
let fn_name = &input_fn.sig.ident;
let fn_vis = &input_fn.vis;
let fn_inputs = &input_fn.sig.inputs;
let fn_output = &input_fn.sig.output;
let fn_generics = &input_fn.sig.generics;
let where_clause = &input_fn.sig.generics.where_clause;
let fn_asyncness = &input_fn.sig.asyncness;
let cache_name = Ident::new(&format!("{}_CACHE", fn_name.to_string().to_uppercase()), fn_name.span());
let param_names: Vec<_> = input_fn.sig.inputs.iter().filter_map(|arg| {
if let syn::FnArg::Typed(pat_type) = arg {
if let syn::Pat::Ident(pat_ident) = &*pat_type.pat {
Some(&pat_ident.ident)
} else {
None
}
} else {
None
}
}).collect();
let original_fn_name = Ident::new(&format!("{}_original", fn_name), fn_name.span());
let capacity = args.capacity.unwrap_or(1000);
let return_type = match &input_fn.sig.output {
ReturnType::Default => quote! { () },
ReturnType::Type(_, ty) => quote! { #ty },
};
let key_type = if param_names.is_empty() {
quote! { () }
} else {
let param_types: Vec<_> = input_fn.sig.inputs.iter().filter_map(|arg| {
if let syn::FnArg::Typed(pat_type) = arg {
Some(&pat_type.ty)
} else {
None
}
}).collect();
if param_types.len() == 1 {
quote! { #(#param_types)* }
} else {
quote! { (#(#param_types),*) }
}
};
let key_creation = if param_names.is_empty() {
quote! { () }
} else if param_names.len() == 1 {
let param = ¶m_names[0];
quote! { #param.clone() }
} else {
quote! { (#(#param_names.clone()),*) }
};
let fn_call = if fn_asyncness.is_some() {
quote! { #original_fn_name(#(#param_names),*).await }
} else {
quote! { #original_fn_name(#(#param_names),*) }
};
let fn_body = &input_fn.block;
let expanded = quote! {
#fn_asyncness fn #original_fn_name #fn_generics (#fn_inputs) #fn_output #where_clause
#fn_body
#fn_vis #fn_asyncness fn #fn_name #fn_generics (#fn_inputs) #fn_output #where_clause {
#[cfg(feature = "memoization")]
{
use std::sync::Arc;
static #cache_name: crate::once_cell::sync::Lazy<crate::dashmap::DashMap<#key_type, #return_type>> = crate::once_cell::sync::Lazy::new(|| {
crate::dashmap::DashMap::with_capacity(#capacity)
});
let cache = &#cache_name;
let key = #key_creation;
if let Some(cached_result) = cache.get(&key) {
return cached_result.clone();
}
let result = #fn_call;
if cache.len() < #capacity {
cache.insert(key, result.clone());
}
result
}
#[cfg(not(feature = "memoization"))]
{
#fn_call
}
}
};
TokenStream::from(expanded)
}