use proc_macro2::TokenStream;
use quote::quote;
use syn::spanned::Spanned;
use crate::metadata::{extract_inject_dyn_inner, extract_inject_inner, type_to_string};
pub(crate) fn expand_inject_fn(item: syn::ItemFn) -> syn::Result<TokenStream> {
for arg in &item.sig.inputs {
if matches!(arg, syn::FnArg::Receiver(_)) {
return Err(syn::Error::new(
arg.span(),
"#[injectable(factory)] cannot be applied to a method with `self`",
));
}
}
let fn_name = &item.sig.ident;
let vis = &item.vis;
let fn_name_str = fn_name.to_string();
let (inner_ty, is_result) = parse_return_inner(&item.sig.output)?;
let mut extract_stmts: Vec<TokenStream> = Vec::new();
for arg in &item.sig.inputs {
let syn::FnArg::Typed(pat_type) = arg else {
continue;
};
let name = match &*pat_type.pat {
syn::Pat::Ident(i) => i.ident.clone(),
_ => {
return Err(syn::Error::new(
pat_type.pat.span(),
"#[injectable(factory)] parameters must be named",
));
}
};
let ty = (*pat_type.ty).clone();
let ty_string = type_to_string(&ty);
let (has_inject, factory) = parse_inject_attr(&pat_type.attrs)?;
if extract_inject_inner(&ty).is_none() && !has_inject {
return Err(syn::Error::new(
ty.span(),
format!(
"parameter `{}: {}` requires `#[injectable(inject)]` annotation in \
`#[injectable(factory)]`; only `Inject<T>` is auto-injected",
name, ty_string
),
));
}
if let Some(factory_path) = factory {
extract_stmts.push(factory_path.gen_extract(name, &ty_string));
} else {
if let Some(dyn_ty) = extract_inject_dyn_inner(&ty) {
extract_stmts.push(quote! {
let #name: #ty = {
let __arc = __ctx.resolve_external::<::std::sync::Arc<#dyn_ty>>().await?;
injectable_rs_runtime::Inject::new(__arc)
};
});
} else {
extract_stmts.push(gen_standard_extract(name, &ty));
}
}
}
let body = &item.block;
let orig_ret = match &item.sig.output {
syn::ReturnType::Default => None,
syn::ReturnType::Type(_, ty) => Some(ty.as_ref().clone()),
};
let call = if is_result {
let ret_annotation = orig_ret.map(|ty| quote! { : #ty });
quote! {
let __injectable_result #ret_annotation = { #body };
__injectable_result.map_err(|e| injectable_rs_runtime::InjectableError::ConstructionFailed {
type_name: #fn_name_str,
reason: ::std::string::ToString::to_string(&e),
})
}
} else {
quote! { Ok({ #body }) }
};
let other_attrs: Vec<_> = item
.attrs
.iter()
.filter(|a| !a.path().is_ident("injectable"))
.collect();
Ok(quote! {
#(#other_attrs)*
#vis async fn #fn_name(
__ctx: &injectable_rs_runtime::ResolveContext,
) -> injectable_rs_runtime::InjectableResult<#inner_ty> {
#(#extract_stmts)*
#call
}
})
}
fn gen_standard_extract(name: syn::Ident, ty: &syn::Type) -> TokenStream {
quote! {
let #name: #ty =
<#ty as injectable_rs_runtime::Extract>::extract(__ctx).await?;
}
}
enum ParamFactory {
Async(syn::Path),
Sync(syn::Path),
}
impl ParamFactory {
fn gen_extract(self, name: syn::Ident, ty_str: &str) -> TokenStream {
match self {
ParamFactory::Async(path) => quote! {
let #name = #path(__ctx).await.map_err(|e|
injectable_rs_runtime::InjectableError::ConstructionFailed {
type_name: #ty_str,
reason: ::std::string::ToString::to_string(&e),
})?;
},
ParamFactory::Sync(path) => quote! {
let #name = #path(__ctx);
},
}
}
}
fn parse_inject_attr(attrs: &[syn::Attribute]) -> syn::Result<(bool, Option<ParamFactory>)> {
for attr in attrs {
if !attr.path().is_ident("injectable") {
continue;
}
let factory = attr.parse_args_with(|input: syn::parse::ParseStream| {
let kw: syn::Ident = input.parse()?;
if kw != "inject" {
return Err(syn::Error::new(
kw.span(),
format!(
"expected `inject` inside `#[injectable(...)]` on a parameter, \
found `{kw}`"
),
));
}
if input.is_empty() {
return Ok(None);
}
let content;
syn::parenthesized!(content in input);
let ident: syn::Ident = content.parse()?;
let is_async = if ident == "use_factory_async" || ident == "use_factory" {
true
} else if ident == "use_factory_sync" {
false
} else {
return Err(syn::Error::new(
ident.span(),
format!(
"unknown inject argument: `{ident}`; \
expected `use_factory_async = path` or `use_factory_sync = path`"
),
));
};
content.parse::<syn::Token![=]>()?;
let path: syn::Path = content.parse()?;
if is_async {
Ok(Some(ParamFactory::Async(path)))
} else {
Ok(Some(ParamFactory::Sync(path)))
}
})?;
return Ok((true, factory));
}
Ok((false, None))
}
#[allow(clippy::unwrap_used)]
fn parse_return_inner(output: &syn::ReturnType) -> syn::Result<(syn::Type, bool)> {
match output {
syn::ReturnType::Default => Ok((syn::parse_str("()").unwrap(), false)),
syn::ReturnType::Type(_, ty) => {
if let Some(inner) = extract_result_ok_ty(ty) {
Ok((inner, true))
} else {
Ok((*ty.clone(), false))
}
}
}
}
fn extract_result_ok_ty(ty: &syn::Type) -> Option<syn::Type> {
let syn::Type::Path(tp) = ty else { return None };
let seg = tp.path.segments.last()?;
if seg.ident != "Result" {
return None;
}
let syn::PathArguments::AngleBracketed(args) = &seg.arguments else {
return None;
};
let syn::GenericArgument::Type(inner) = args.args.first()? else {
return None;
};
Some(inner.clone())
}