use proc_macro::TokenStream;
use quote::quote;
use syn::{parse::Parser, parse_macro_input, punctuated::Punctuated, Expr, ItemFn, Lit, Meta, Token};
#[proc_macro_attribute]
pub fn cached(args: TokenStream, item: TokenStream) -> TokenStream {
let parser = Punctuated::<Meta, Token![,]>::parse_terminated;
let args = parser.parse(args).expect("Failed to parse arguments");
let input = parse_macro_input!(item as ItemFn);
let mut service_name = "default".to_string();
let mut ttl = quote! { None };
let mut key_pattern = None;
let mut key_prefix = None;
for arg in args {
if let Meta::NameValue(nv) = arg {
if nv.path.is_ident("service") {
if let Expr::Lit(expr_lit) = nv.value {
if let Lit::Str(lit) = expr_lit.lit {
service_name = lit.value();
}
}
} else if nv.path.is_ident("ttl") {
if let Expr::Lit(expr_lit) = nv.value {
if let Lit::Int(lit) = expr_lit.lit {
let val = lit.base10_parse::<u64>().unwrap();
ttl = quote! { Some(#val) };
}
}
} else if nv.path.is_ident("key") {
if let Expr::Lit(expr_lit) = nv.value {
if let Lit::Str(lit) = expr_lit.lit {
key_pattern = Some(lit.value());
}
}
} else if nv.path.is_ident("key_prefix") {
if let Expr::Lit(expr_lit) = nv.value {
if let Lit::Str(lit) = expr_lit.lit {
key_prefix = Some(lit.value());
}
}
}
}
}
let fn_name = &input.sig.ident;
let fn_args = &input.sig.inputs;
let fn_output = &input.sig.output;
let fn_block = &input.block;
let vis = &input.vis;
let return_type = match fn_output {
syn::ReturnType::Default => quote! { () },
syn::ReturnType::Type(_, ty) => {
if let syn::Type::Path(path) = &**ty {
if let Some(seg) = path.path.segments.last() {
if seg.ident == "Result" {
if let syn::PathArguments::AngleBracketed(args) = &seg.arguments {
if let Some(first_arg) = args.args.first() {
quote! { #first_arg }
} else {
quote! { #ty }
}
} else {
quote! { #ty }
}
} else {
quote! { #ty }
}
} else {
quote! { #ty }
}
} else {
quote! { #ty }
}
}
};
let arg_names: Vec<_> = fn_args
.iter()
.filter_map(|arg| {
if let syn::FnArg::Typed(pat_type) = arg {
if let syn::Pat::Ident(pat_ident) = &*pat_type.pat {
return Some(&pat_ident.ident);
}
}
None
})
.collect();
let arg_names_cloned: Vec<_> = arg_names
.iter()
.map(|name| {
quote! { (#name).clone() }
})
.collect();
let key_gen_with_cloned_args = if let Some(pattern) = key_pattern {
quote! {
format!(#pattern)
}
} else if let Some(prefix) = key_prefix {
if arg_names.is_empty() {
quote! { format!("{}:{}:{}", #service_name, #prefix, stringify!(#fn_name)) }
} else {
quote! {
format!("{}:{}:{}:{:?}", #service_name, #prefix, stringify!(#fn_name), (#(#arg_names_cloned),*))
}
}
} else {
if arg_names.is_empty() {
quote! { format!("{}:{}", #service_name, stringify!(#fn_name)) }
} else {
quote! {
format!("{}:{}:{:?}", #service_name, stringify!(#fn_name), (#(#arg_names_cloned),*))
}
}
};
let output = quote! {
#vis async fn #fn_name(#fn_args) #fn_output {
let cache_key = #key_gen_with_cloned_args;
let cache = match ::oxcache::__internal_get_cache(#service_name) {
Some(c) => c,
None => return async { #fn_block }.await,
};
if let Ok(Some(bytes)) = cache.get_bytes(&cache_key).await {
if let Ok(val) = cache.unified_serializer().deserialize::<#return_type>(&bytes) {
return ::std::result::Result::Ok(val);
}
}
let result = async { #fn_block }.await;
if let Ok(ref val) = result {
if let Ok(bytes) = cache.unified_serializer().serialize(val) {
let _ = cache.set_bytes(&cache_key, bytes, #ttl).await;
}
}
result
}
};
output.into()
}