use proc_macro::TokenStream;
use quote::{format_ident, quote};
use std::sync::atomic::{AtomicUsize, Ordering};
use syn::{parse_macro_input, FnArg, ItemFn, Pat, ReturnType};
use crate::utils::ferro;
static MEMO_COUNTER: AtomicUsize = AtomicUsize::new(0);
pub fn memoize_impl(attr: TokenStream, input: TokenStream) -> TokenStream {
let _ = attr;
let input_fn = parse_macro_input!(input as ItemFn);
let ferro = ferro();
if input_fn.sig.asyncness.is_none() {
return syn::Error::new_spanned(
&input_fn.sig,
"#[memoize] can only be applied to `async fn`",
)
.to_compile_error()
.into();
}
let n = MEMO_COUNTER.fetch_add(1, Ordering::Relaxed);
let marker_name = format_ident!("__FerroMemoMarker{n}");
let fn_vis = &input_fn.vis;
let fn_name = &input_fn.sig.ident;
let fn_generics = &input_fn.sig.generics;
let fn_block = &input_fn.block;
let fn_attrs = &input_fn.attrs;
let fn_output = &input_fn.sig.output;
let all_inputs: Vec<_> = input_fn.sig.inputs.iter().collect();
let value_inputs: Vec<_> = input_fn
.sig
.inputs
.iter()
.filter(|a| !matches!(a, FnArg::Receiver(_)))
.collect();
let mut value_arg_names: Vec<proc_macro2::Ident> = Vec::new();
let mut value_arg_types: Vec<&syn::Type> = Vec::new();
for arg in &value_inputs {
match arg {
FnArg::Typed(pat_type) => {
let ty = &*pat_type.ty;
match &*pat_type.pat {
Pat::Ident(pat_ident) => {
value_arg_names.push(pat_ident.ident.clone());
value_arg_types.push(ty);
}
other => {
return syn::Error::new_spanned(
other,
"#[memoize] arguments must be simple identifiers in v17.0; \
got a destructuring pattern",
)
.to_compile_error()
.into();
}
}
}
FnArg::Receiver(_) => {
}
}
}
let return_ty: proc_macro2::TokenStream = match fn_output {
ReturnType::Default => quote! { () },
ReturnType::Type(_, ty) => quote! { #ty },
};
let output = quote! {
#(#fn_attrs)*
#fn_vis async fn #fn_name #fn_generics(#(#all_inputs),*) #fn_output
where
#( #value_arg_types: ::std::hash::Hash, )*
#return_ty: ::std::clone::Clone + ::std::marker::Send + ::std::marker::Sync + 'static,
{
struct #marker_name;
let __ferro_memo_key = #ferro::memo::MemoKey::new::<#marker_name, _>(
&( #( &#value_arg_names, )* ),
);
if let ::std::option::Option::Some(__ferro_store) =
#ferro::memo::current_memo_store()
{
let __ferro_slot = __ferro_store.get_or_insert(
__ferro_memo_key,
move || {
::std::boxed::Box::pin(async move {
let __ferro_result: #return_ty = { #fn_block };
::std::sync::Arc::new(__ferro_result)
as ::std::sync::Arc<
dyn ::std::any::Any + ::std::marker::Send + ::std::marker::Sync,
>
})
},
);
let __ferro_arc = __ferro_slot.await;
return ::std::clone::Clone::clone(
__ferro_arc
.downcast_ref::<#return_ty>()
.expect("MemoStore type invariant violated"),
);
}
{ #fn_block }
}
};
output.into()
}