#![allow(unused_variables)]
use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
use std::time::Duration;
use syn::{
FnArg, ItemFn, LitStr, Pat, PatType, Result, ReturnType, Token, Type, parse::Parse,
parse::ParseStream, parse_macro_input,
};
#[derive(Default)]
struct ProviderArgs {
interval: Option<Duration>,
cache_expiration: Option<Duration>,
stale_time: Option<Duration>,
inject: Vec<syn::Type>, }
impl Parse for ProviderArgs {
fn parse(input: ParseStream) -> Result<Self> {
let mut args = ProviderArgs::default();
while !input.is_empty() {
let ident: syn::Ident = input.parse()?;
input.parse::<Token![=]>()?;
match ident.to_string().as_str() {
"interval" => {
let lit: LitStr = input.parse()?;
let duration_str = lit.value();
let duration = humantime::parse_duration(&duration_str).map_err(|e| {
syn::Error::new_spanned(lit, format!("Invalid duration format: {}", e))
})?;
args.interval = Some(duration);
}
"cache_expiration" => {
let lit: LitStr = input.parse()?;
let duration_str = lit.value();
let duration = humantime::parse_duration(&duration_str).map_err(|e| {
syn::Error::new_spanned(lit, format!("Invalid duration format: {}", e))
})?;
args.cache_expiration = Some(duration);
}
"stale_time" => {
let lit: LitStr = input.parse()?;
let duration_str = lit.value();
let duration = humantime::parse_duration(&duration_str).map_err(|e| {
syn::Error::new_spanned(lit, format!("Invalid duration format: {}", e))
})?;
args.stale_time = Some(duration);
}
"inject" => {
let content;
syn::bracketed!(content in input);
let types = content.parse_terminated(syn::Type::parse, Token![,])?;
args.inject = types.into_iter().collect();
}
_ => return Err(syn::Error::new_spanned(ident, "Unknown argument")),
}
if input.peek(Token![,]) {
input.parse::<Token![,]>()?;
}
}
Ok(args)
}
}
#[proc_macro_attribute]
pub fn provider(args: TokenStream, input: TokenStream) -> TokenStream {
let provider_args = if args.is_empty() {
ProviderArgs::default()
} else {
match syn::parse(args) {
Ok(args) => args,
Err(err) => return err.to_compile_error().into(),
}
};
let input_fn = parse_macro_input!(input as ItemFn);
let result = generate_provider(input_fn, provider_args);
match result {
Ok(tokens) => tokens.into(),
Err(err) => err.to_compile_error().into(),
}
}
fn generate_provider(input_fn: ItemFn, provider_args: ProviderArgs) -> Result<TokenStream2> {
let info = extract_provider_info(&input_fn)?;
let ProviderInfo {
fn_vis,
fn_block,
output_type,
error_type,
struct_name,
..
} = &info;
let enhanced_fn_block = generate_dependency_injection(&provider_args.inject, fn_block);
let interval_impl = generate_interval_impl(&provider_args);
let cache_expiration_impl = generate_cache_expiration_impl(&provider_args);
let stale_time_impl = generate_stale_time_impl(&provider_args);
let common_struct = generate_common_struct_and_const(&info);
if input_fn.sig.inputs.is_empty() {
Ok(quote! {
#common_struct
impl #struct_name {
#fn_vis async fn call() -> Result<#output_type, #error_type> {
#enhanced_fn_block
}
}
impl ::dioxus_provider::hooks::Provider<()> for #struct_name {
type Output = #output_type;
type Error = #error_type;
fn run(&self, _param: ()) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send {
Self::call()
}
#interval_impl
#cache_expiration_impl
#stale_time_impl
}
})
} else {
let params = extract_all_params(&input_fn)?;
if params.len() == 1 {
let param = ¶ms[0];
let param_name = ¶m.name;
let param_type = ¶m.ty;
Ok(quote! {
#common_struct
impl #struct_name {
#fn_vis async fn call(#param_name: #param_type) -> Result<#output_type, #error_type> {
#enhanced_fn_block
}
}
impl ::dioxus_provider::hooks::Provider<#param_type> for #struct_name {
type Output = #output_type;
type Error = #error_type;
fn run(&self, #param_name: #param_type) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send {
Self::call(#param_name)
}
#interval_impl
#cache_expiration_impl
#stale_time_impl
}
})
} else {
let param_names: Vec<_> = params.iter().map(|p| &p.name).collect();
let param_types: Vec<_> = params.iter().map(|p| &p.ty).collect();
let tuple_type = quote! { (#(#param_types,)*) };
Ok(quote! {
#common_struct
impl #struct_name {
#fn_vis async fn call(#(#param_names: #param_types,)*) -> Result<#output_type, #error_type> {
#enhanced_fn_block
}
}
impl ::dioxus_provider::hooks::Provider<#tuple_type> for #struct_name {
type Output = #output_type;
type Error = #error_type;
fn run(&self, params: #tuple_type) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send {
let (#(#param_names,)*) = params;
Self::call(#(#param_names,)*)
}
#interval_impl
#cache_expiration_impl
#stale_time_impl
}
})
}
}
}
fn generate_duration_impl(method_name: &str, duration: Option<Duration>) -> TokenStream2 {
if let Some(duration) = duration {
let duration_secs = duration.as_secs();
let method_ident = syn::Ident::new(method_name, proc_macro2::Span::call_site());
quote! {
fn #method_ident(&self) -> Option<::std::time::Duration> {
Some(::std::time::Duration::from_secs(#duration_secs))
}
}
} else {
quote! {}
}
}
fn generate_interval_impl(provider_args: &ProviderArgs) -> TokenStream2 {
generate_duration_impl("interval", provider_args.interval)
}
fn generate_cache_expiration_impl(provider_args: &ProviderArgs) -> TokenStream2 {
generate_duration_impl("cache_expiration", provider_args.cache_expiration)
}
fn generate_stale_time_impl(provider_args: &ProviderArgs) -> TokenStream2 {
generate_duration_impl("stale_time", provider_args.stale_time)
}
struct ProviderInfo {
fn_vis: syn::Visibility,
fn_attrs: Vec<syn::Attribute>,
fn_block: Box<syn::Block>,
output_type: Type,
error_type: Type,
struct_name: syn::Ident,
fn_name: syn::Ident,
}
struct ParamInfo {
name: syn::Ident,
ty: Type,
}
fn extract_provider_info(input_fn: &ItemFn) -> Result<ProviderInfo> {
let fn_name = input_fn.sig.ident.clone();
let fn_vis = input_fn.vis.clone();
let fn_attrs = input_fn.attrs.clone();
let fn_block = input_fn.block.clone();
let (output_type, error_type) = extract_result_types(&input_fn.sig.output)?;
let struct_name = syn::Ident::new(
&to_pascal_case(&fn_name.to_string()),
proc_macro2::Span::call_site(),
);
Ok(ProviderInfo {
fn_vis,
fn_attrs,
fn_block,
output_type,
error_type,
struct_name,
fn_name,
})
}
fn generate_common_struct_and_const(info: &ProviderInfo) -> TokenStream2 {
let struct_name = &info.struct_name;
let fn_attrs = &info.fn_attrs;
let fn_name = &info.fn_name;
quote! {
#[derive(Clone, PartialEq)]
#(#fn_attrs)*
pub struct #struct_name;
impl Default for #struct_name {
fn default() -> Self {
Self
}
}
pub fn #fn_name() -> #struct_name {
#struct_name
}
}
}
fn extract_all_params(input_fn: &ItemFn) -> Result<Vec<ParamInfo>> {
let mut params = Vec::new();
for input in &input_fn.sig.inputs {
match input {
FnArg::Typed(PatType { pat, ty, .. }) => {
if let Pat::Ident(pat_ident) = &**pat {
params.push(ParamInfo {
name: pat_ident.ident.clone(),
ty: (**ty).clone(),
});
} else {
return Err(syn::Error::new_spanned(
pat,
"Only simple parameter names are supported",
));
}
}
FnArg::Receiver(_) => {
return Err(syn::Error::new_spanned(
input,
"Methods with self parameter are not supported",
));
}
}
}
Ok(params)
}
fn extract_result_types(return_type: &ReturnType) -> Result<(Type, Type)> {
match return_type {
ReturnType::Default => Err(syn::Error::new_spanned(
return_type,
"Provider functions must return Result<T, E>",
)),
ReturnType::Type(_, ty) => {
if let Type::Path(type_path) = &**ty {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "Result" {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
if args.args.len() == 2 {
let mut args_iter = args.args.iter();
let output_type = match args_iter.next().unwrap() {
syn::GenericArgument::Type(ty) => ty.clone(),
_ => {
return Err(syn::Error::new_spanned(
args,
"Result must have type arguments",
));
}
};
let error_type = match args_iter.next().unwrap() {
syn::GenericArgument::Type(ty) => ty.clone(),
_ => {
return Err(syn::Error::new_spanned(
args,
"Result must have type arguments",
));
}
};
return Ok((output_type, error_type));
}
}
}
}
}
Err(syn::Error::new_spanned(
return_type,
"Provider functions must return Result<T, E>",
))
}
}
}
fn to_pascal_case(s: &str) -> String {
let mut result = String::new();
let mut capitalize_next = true;
for c in s.chars() {
if c == '_' {
capitalize_next = true;
} else if capitalize_next {
result.push(c.to_ascii_uppercase());
capitalize_next = false;
} else {
result.push(c);
}
}
result
}
fn generate_dependency_injection(inject_types: &[syn::Type], original_block: &syn::Block) -> syn::Block {
if inject_types.is_empty() {
return original_block.clone();
}
let injection_stmts: Vec<_> = inject_types
.iter()
.map(|ty| {
let var_name = syn::Ident::new(
&format!("injected_{}", to_pascal_case("e!(#ty).to_string().to_lowercase())),
proc_macro2::Span::call_site(),
);
syn::parse_quote! {
let #var_name = ::dioxus_provider::injection::inject::<#ty>()
.map_err(|e| format!("Dependency injection failed for {}: {}", stringify!(#ty), e))?;
}
})
.collect();
let mut new_block = original_block.clone();
new_block.stmts.splice(0..0, injection_stmts);
new_block
}