moneta_fn 0.1.0

A set of macros to function profiling
Documentation
extern crate proc_macro;

use proc_macro::TokenStream;
use proc_macro2::{Ident, Span, TokenStream as TokenStream2};
use quote::quote;
use syn::{
    parse_macro_input, parse_quote, FnArg, Item, ItemFn, LitStr, Pat, PatType, ReturnType, Stmt,
    Type,
};

fn split_args(def: &mut ItemFn) -> Vec<(Box<Pat>, Ident, Box<Type>)> {
    let mut args = Vec::new();

    for (i, arg) in def.sig.inputs.iter_mut().enumerate() {
        let numbered = Ident::new(&format!("_arg{i}"), Span::call_site());
        match arg {
            FnArg::Typed(PatType { pat, ty, .. }) => {
                args.push((pat.clone(), numbered.clone(), ty.clone()));
                *pat = parse_quote!(mut #numbered);
            }
            FnArg::Receiver(_) => {
                todo!()
            }
        }
    }
    args
}

#[proc_macro]
pub fn count(func: TokenStream) -> TokenStream {
    let mut func: Vec<_> = parse_macro_input!(func as syn::Path)
        .segments
        .into_iter()
        .map(|seg| seg.ident)
        .collect();
    let fn_id = func.pop().unwrap();
    let id = Ident::new(&format!("__MONETA_FN_COUNT_{fn_id}"), Span::call_site());
    TokenStream::from(quote! { unsafe { #(#func::)* #id } })
}

#[proc_macro]
pub fn get_cache(func: TokenStream) -> TokenStream {
    let mut func: Vec<_> = parse_macro_input!(func as syn::Path)
        .segments
        .into_iter()
        .map(|seg| seg.ident)
        .collect();
    let fn_id = func.pop().unwrap();
    let id = Ident::new(&format!("__MONETA_FN_CACHE_{fn_id}"), Span::call_site());
    TokenStream::from(quote! { #id })
}

#[proc_macro_attribute]
pub fn moneta(meta: TokenStream, input: TokenStream) -> TokenStream {
    let mut outter = parse_macro_input!(input as ItemFn);
    let mut def_fn = outter.clone();
    let args: Vec<_> = split_args(&mut outter);

    let cache_id = Ident::new(
        &format!("__MONETA_FN_CACHE_{}", def_fn.sig.ident),
        Span::call_site(),
    );

    let counter_id = Ident::new(
        &format!("__MONETA_FN_COUNT_{}", def_fn.sig.ident),
        Span::call_site(),
    );

    let cache_ret = match def_fn.sig.output {
        ReturnType::Default => quote! { () },
        ReturnType::Type(_, ref ty) => {
            let ty = ty.clone();
            quote! { #ty }
        }
    };

    let name = outter.sig.ident.clone();
    let args_lit_name: Vec<_> = args
        .iter()
        .map(|(name, _, _)| LitStr::new(&format!("{:?}", quote! { #name }), Span::call_site()))
        .collect();
    let out_args: Vec<_> = args.iter().map(|(_, arg, _)| arg).collect();
    let def_args: Vec<_> = args.iter().map(|(name, _, _)| name).collect();
    let func_name = LitStr::new(&format!("{name}"), Span::call_site());
    let res_id = Ident::new("res", Span::call_site());
    let start_id = Ident::new("start", Span::call_site());

    let get_ret = if cfg!(feature = "cache") {
        quote! {
            let #res_id = #name (#(#out_args),*);
        }
    } else {
        quote! { let #res_id = #name (#(#out_args),*); }
    };

    let (trace_in, trace_out) = trace(&func_name, &start_id, &args_lit_name, &def_args);
    let cache_def = cache_def(&meta, &cache_id, &cache_ret);
    let (cache_get, cache_set) = cache(
        meta.to_string() != "no_cache",
        &def_args,
        &out_args,
        &cache_id,
        &res_id,
    );
    let (counter_def, counter_inc) = counter(&counter_id);

    let pre_injection = quote! {{
        #counter_inc
        #trace_in
        #cache_get
    }};
    let pre_injection = TokenStream::from(pre_injection);

    let post_injection = quote! {{
        let #start_id = std::time::Instant::now();
        #get_ret
        #trace_out
        #cache_set
        return #res_id;
    }};
    let post_injection = TokenStream::from(post_injection);
    def_fn
        .block
        .stmts
        .insert(0, parse_macro_input!(pre_injection as Stmt));

    outter.block.stmts = vec![
        Stmt::Item(Item::Fn(def_fn)),
        parse_macro_input!(post_injection as Stmt),
    ];

    let code = quote! {
        #counter_def
        #cache_def
        #outter
    };
    TokenStream::from(code)
}

fn cache_def(meta: &TokenStream, cache_id: &Ident, cache_ret: &TokenStream2) -> TokenStream2 {
    if meta.to_string() != "no_cache" && cfg!(feature = "cache") {
        quote! {
            lazy_static::lazy_static! {
                pub static ref #cache_id: std::sync::RwLock<hashbrown::HashMap<String, #cache_ret>> =
                    std::sync::RwLock::new(hashbrown::HashMap::new());
            }
        }
    } else {
        quote! {}
    }
}

fn trace(
    name_str: &LitStr,
    start_id: &Ident,
    out_args: &Vec<LitStr>,
    def_args: &Vec<&Box<Pat>>,
) -> (TokenStream2, TokenStream2) {
    let in_trace = if cfg!(feature = "trace") {
        quote! {{
            let args_fmt: String = [
                #(#out_args),*
            ].into_iter()
                .zip([#(format!("{:?}", #def_args)),*].into_iter())
                .map(|(n, v): (&str, String)| format!("\n\t{}: {}", n, v))
                .collect();
            println!("in {}: {:?}", #name_str, args_fmt);
        }}
    } else {
        quote! { ; }
    };

    let out_trace = if cfg!(feature = "time") {
        quote! {
            println!("out {}: {:?}", #name_str, #start_id.elapsed());
        }
    } else if cfg!(feature = "trace") {
        quote! {
            println!("out {}", #name_str);
        }
    } else {
        quote! { ; }
    };

    (in_trace, out_trace)
}

fn cache(
    enabled: bool,
    def_args: &Vec<&Box<Pat>>,
    out_args: &Vec<&Ident>,
    counter_id: &Ident,
    res_id: &Ident,
) -> (TokenStream2, TokenStream2) {
    let debug_fmt = LitStr::new(&"{:?}".repeat(def_args.len()), Span::call_site());

    let get_cache = if cfg!(feature = "cache") && enabled {
        quote! {{
            let values_fmt = format!(#debug_fmt, #(#def_args),*);
            if let Ok(reader) = #counter_id.read() {
                if let Some(val) = reader.get(&values_fmt) {
                    return val.clone();
                }
            }
        }}
    } else {
        quote! { ; }
    };

    let set_cache = if cfg!(feature = "cache") && enabled {
        quote! {{
            let values_fmt = format!(#debug_fmt, #(#out_args),*);
            if let Ok(mut writer) = #counter_id.write() {
                writer.entry(values_fmt).or_insert(#res_id.clone());
            }
        }}
    } else {
        quote! {{ ; }}
    };
    (get_cache, set_cache)
}

fn counter(counter_id: &Ident) -> (TokenStream2, TokenStream2) {
    let def = quote! {
        pub static mut #counter_id: usize = 0;
    };

    let inc = if cfg!(feature = "count") {
        quote! { unsafe { #counter_id += 1 }; }
    } else {
        quote! { ; }
    };

    (def, inc)
}