use proc_macro2::TokenStream;
use quote::quote;
use syn::spanned::Spanned;
use syn::visit_mut::VisitMut;
use crate::attrs::Scope;
use crate::metadata::{
extract_arc_inner_str, extract_inject_dyn_inner, extract_inject_inner,
extract_option_inject_dyn_inner, type_to_string,
};
pub fn expand_injectable_impl(attrs: TokenStream, item: TokenStream) -> syn::Result<TokenStream> {
let injectable_attrs = parse_impl_attrs(attrs)?;
let mut impl_block: syn::ItemImpl = syn::parse2(item)?;
let (type_name, self_ty, impl_generics) = extract_type_name(&impl_block)?;
let scan_result = scan_impl_methods(&impl_block)?;
AttrStripper.visit_item_impl_mut(&mut impl_block);
if let Some(constructor) = scan_result.constructor {
let provider_code = generate_provider(
&type_name,
&self_ty,
&impl_generics,
&constructor,
&scan_result.post_construct_hooks,
&scan_result.pre_destruct_hooks,
&injectable_attrs,
)?;
Ok(quote! { #impl_block #provider_code })
} else {
if scan_result.post_construct_hooks.is_empty() && scan_result.pre_destruct_hooks.is_empty()
{
return Err(syn::Error::new(
impl_block.self_ty.span(),
"#[injectable] without #[injectable(ctor)] requires at least one \
#[injectable(post_construct)] or #[injectable(pre_destruct)] method. \
For field injection without lifecycle hooks, use #[injectable] alone.",
));
}
let post_impl = generate_post_construct_impl(&type_name, &scan_result.post_construct_hooks);
let pre_impl = generate_pre_destruct_impl(&type_name, &scan_result.pre_destruct_hooks);
let hooks_submit = generate_hooks_entry_submit(
&type_name,
&scan_result.post_construct_hooks,
&scan_result.pre_destruct_hooks,
);
Ok(quote! { #impl_block #post_impl #pre_impl #hooks_submit })
}
}
struct InjectableImplAttrs {
scope: Scope,
}
impl Default for InjectableImplAttrs {
fn default() -> Self {
Self {
scope: Scope::Singleton,
}
}
}
fn parse_impl_attrs(attrs: TokenStream) -> syn::Result<InjectableImplAttrs> {
if attrs.is_empty() {
return Ok(InjectableImplAttrs::default());
}
let parsed: syn::punctuated::Punctuated<ImplArg, syn::Token![,]> =
syn::parse::Parser::parse2(syn::punctuated::Punctuated::parse_terminated, attrs)?;
let mut result = InjectableImplAttrs::default();
for arg in parsed {
match arg {
ImplArg::Scope(s) => {
result.scope = match s.as_str() {
"singleton" => Scope::Singleton,
"transient" => Scope::Transient,
"request" => Scope::Request,
other => Scope::Custom(other.to_string()),
};
}
}
}
Ok(result)
}
enum ImplArg {
Scope(String),
}
impl syn::parse::Parse for ImplArg {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
let ident: syn::Ident = input.parse()?;
if ident == "scope" {
input.parse::<syn::Token![=]>()?;
let lit: syn::LitStr = input.parse()?;
Ok(ImplArg::Scope(lit.value()))
} else {
Err(syn::Error::new(
ident.span(),
format!("unknown injectable_impl attribute: `{ident}`"),
))
}
}
}
struct ConstructorInfo {
method_name: syn::Ident,
is_async: bool,
params: Vec<ParamInfo>,
return_kind: ConstructorReturn,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ConstructorReturn {
SelfOwned,
ResultWrapped,
ResultInjectableError,
}
#[derive(Debug, Clone)]
enum FactoryFn {
Async(syn::Path),
Sync(syn::Path),
}
impl FactoryFn {
fn path(&self) -> &syn::Path {
match self {
FactoryFn::Async(p) | FactoryFn::Sync(p) => p,
}
}
fn is_async(&self) -> bool {
matches!(self, FactoryFn::Async(_))
}
}
struct ParamInfo {
name: syn::Ident,
ty: syn::Type,
ty_string: String,
factory_fn: Option<FactoryFn>,
}
struct HookInfo {
method_name: syn::Ident,
is_async: bool,
returns_result: bool,
}
struct ScanResult {
constructor: Option<ConstructorInfo>,
post_construct_hooks: Vec<HookInfo>,
pre_destruct_hooks: Vec<HookInfo>,
}
fn is_injectable_sub_arg(attr: &syn::Attribute, sub_arg: &str) -> bool {
if !attr.path().is_ident("injectable") {
return false;
}
attr.parse_args_with(|input: syn::parse::ParseStream| {
let ident: syn::Ident = input.parse()?;
Ok(ident == sub_arg)
})
.unwrap_or(false)
}
fn scan_impl_methods(impl_block: &syn::ItemImpl) -> syn::Result<ScanResult> {
let mut result = ScanResult {
constructor: None,
post_construct_hooks: Vec::new(),
pre_destruct_hooks: Vec::new(),
};
for item in &impl_block.items {
if let syn::ImplItem::Fn(method) = item {
let has_constructor = method
.attrs
.iter()
.any(|a| is_injectable_sub_arg(a, "ctor"));
let has_post_construct = method
.attrs
.iter()
.any(|a| is_injectable_sub_arg(a, "post_construct"));
let has_pre_destruct = method
.attrs
.iter()
.any(|a| is_injectable_sub_arg(a, "pre_destruct"));
if has_constructor {
if result.constructor.is_some() {
return Err(syn::Error::new(
method.sig.ident.span(),
"#[injectable] requires exactly one #[injectable(ctor)] method, but found multiple",
));
}
let params = extract_params(&method.sig)?;
result.constructor = Some(ConstructorInfo {
method_name: method.sig.ident.clone(),
is_async: method.sig.asyncness.is_some(),
return_kind: classify_constructor_return(&method.sig),
params,
});
}
if has_post_construct {
result.post_construct_hooks.push(HookInfo {
method_name: method.sig.ident.clone(),
is_async: method.sig.asyncness.is_some(),
returns_result: returns_result(&method.sig),
});
}
if has_pre_destruct {
result.pre_destruct_hooks.push(HookInfo {
method_name: method.sig.ident.clone(),
is_async: method.sig.asyncness.is_some(),
returns_result: returns_result(&method.sig),
});
}
}
}
Ok(result)
}
fn classify_constructor_return(sig: &syn::Signature) -> ConstructorReturn {
match &sig.output {
syn::ReturnType::Default => ConstructorReturn::SelfOwned,
syn::ReturnType::Type(_, ty) => {
let ty_str = type_to_string(ty);
if !ty_str.starts_with("Result") {
return ConstructorReturn::SelfOwned;
}
if ty_str.contains("InjectableError") {
ConstructorReturn::ResultInjectableError
} else {
ConstructorReturn::ResultWrapped
}
}
}
}
fn returns_result(sig: &syn::Signature) -> bool {
match &sig.output {
syn::ReturnType::Default => false, syn::ReturnType::Type(_, ty) => {
let ty_str = type_to_string(ty);
ty_str.starts_with("Result")
}
}
}
fn extract_params(sig: &syn::Signature) -> syn::Result<Vec<ParamInfo>> {
let mut params = Vec::new();
for input in &sig.inputs {
if let syn::FnArg::Typed(pat_type) = input {
let name = match &*pat_type.pat {
syn::Pat::Ident(pat_ident) => pat_ident.ident.clone(),
_ => {
return Err(syn::Error::new(
pat_type.pat.span(),
"constructor parameters must be named",
));
}
};
let ty = (*pat_type.ty).clone();
let ty_string = type_to_string(&ty);
let (has_inject, factory_fn) = parse_param_inject(&pat_type.attrs)?;
if extract_inject_inner(&ty).is_none() && !has_inject {
return Err(syn::Error::new(
ty.span(),
format!(
"parameter `{}: {}` is not auto-injectable; \
only `Inject<T>` parameters are injected automatically — \
annotate with `#[injectable(inject)]` to extract this from the container",
name, ty_string
),
));
}
params.push(ParamInfo {
name,
ty,
ty_string,
factory_fn,
});
}
}
Ok(params)
}
fn parse_param_inject(attrs: &[syn::Attribute]) -> syn::Result<(bool, Option<FactoryFn>)> {
for attr in attrs {
if attr.path().is_ident("injectable") {
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 on parameter: `{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(FactoryFn::Async(path)))
} else {
Ok(Some(FactoryFn::Sync(path)))
}
})?;
return Ok((true, factory));
}
}
Ok((false, None))
}
fn extract_type_name(
impl_block: &syn::ItemImpl,
) -> syn::Result<(syn::Ident, syn::Type, syn::Generics)> {
let ident = match &*impl_block.self_ty {
syn::Type::Path(type_path) => type_path
.path
.segments
.last()
.map(|s| s.ident.clone())
.ok_or_else(|| {
syn::Error::new(
impl_block.self_ty.span(),
"cannot determine type name from impl block",
)
})?,
_ => {
return Err(syn::Error::new(
impl_block.self_ty.span(),
"#[injectable] can only be used on impl blocks for named types",
));
}
};
Ok((
ident,
(*impl_block.self_ty).clone(),
impl_block.generics.clone(),
))
}
struct AttrStripper;
impl VisitMut for AttrStripper {
fn visit_impl_item_fn_mut(&mut self, node: &mut syn::ImplItemFn) {
node.attrs.retain(|a| !a.path().is_ident("injectable"));
for input in node.sig.inputs.iter_mut() {
if let syn::FnArg::Typed(pat_type) = input {
pat_type.attrs.retain(|a| !a.path().is_ident("injectable"));
}
}
syn::visit_mut::visit_impl_item_fn_mut(self, node);
}
}
fn generate_provider(
type_name: &syn::Ident,
self_ty: &syn::Type,
impl_generics: &syn::Generics,
constructor: &ConstructorInfo,
post_construct_hooks: &[HookInfo],
pre_destruct_hooks: &[HookInfo],
attrs: &InjectableImplAttrs,
) -> syn::Result<TokenStream> {
let provider_name = syn::Ident::new(
&format!("{}Provider", type_name),
proc_macro2::Span::call_site(),
);
let (gen_impl, ty_generics, where_clause) = impl_generics.split_for_impl();
let is_generic = !impl_generics.params.is_empty();
let (extract_statements, call_args, dep_strings) = generate_extraction_code(constructor)?;
let method_name = &constructor.method_name;
let await_token = if constructor.is_async {
quote! { .await }
} else {
quote! {}
};
let type_str = type_name.to_string();
let ctor_path = if is_generic {
quote! { <#self_ty>::#method_name }
} else {
quote! { #type_name::#method_name }
};
let construction = match constructor.return_kind {
ConstructorReturn::SelfOwned => quote! {
#ctor_path(#(#call_args),*) #await_token
},
ConstructorReturn::ResultWrapped => quote! {
#ctor_path(#(#call_args),*) #await_token
.map_err(|e| injectable_rs_runtime::InjectableError::ConstructionFailed {
type_name: #type_str,
reason: e.to_string(),
})?
},
ConstructorReturn::ResultInjectableError => quote! {
#ctor_path(#(#call_args),*) #await_token?
},
};
let post_construct_calls = generate_post_construct_calls(post_construct_hooks, &type_str);
let pre_destruct_impl = generate_pre_destruct_impl(type_name, pre_destruct_hooks);
let (pre_destruct_registration, return_instance) = if !pre_destruct_hooks.is_empty() {
(
quote! {
let __destructor_arc: std::sync::Arc<#self_ty> = std::sync::Arc::new(instance);
ctx.register_destructor_with_name(
#type_str,
std::sync::Arc::clone(&__destructor_arc) as std::sync::Arc<dyn injectable_rs_runtime::PreDestruct>,
);
let instance = std::sync::Arc::unwrap_or_clone(__destructor_arc);
},
quote! { Ok(instance) },
)
} else {
(quote! {}, quote! { Ok(instance) })
};
let scope_str = attrs.scope.as_str();
let graph_metadata = generate_graph_metadata(type_name, &dep_strings, scope_str);
let is_singleton: bool = attrs.scope != crate::attrs::Scope::Transient;
let arc_factory_submit = if is_generic {
quote! {}
} else {
crate::provider_gen::generate_arc_factory_submit(type_name, is_singleton)
};
let post_construct_impl = generate_post_construct_impl(type_name, post_construct_hooks);
let provider_struct = if is_generic {
let phantom = crate::provider_gen::phantom_for_generics(impl_generics);
quote! { pub struct #provider_name #ty_generics (#phantom); }
} else {
quote! { pub struct #provider_name; }
};
Ok(quote! {
#provider_struct
#[async_trait::async_trait]
impl #gen_impl injectable_rs_runtime::Provider<#self_ty>
for #provider_name #ty_generics
#where_clause
{
async fn provide(
ctx: &injectable_rs_runtime::ResolveContext,
) -> injectable_rs_runtime::InjectableResult<#self_ty> {
#(#extract_statements)*
let instance = #construction;
#post_construct_calls
#pre_destruct_registration
#return_instance
}
}
impl #gen_impl injectable_rs_runtime::Injectable for #self_ty
#where_clause
{
type Provider = #provider_name #ty_generics;
const IS_SINGLETON: bool = #is_singleton;
}
#post_construct_impl
#pre_destruct_impl
#graph_metadata
#arc_factory_submit
})
}
fn generate_extraction_code(
constructor: &ConstructorInfo,
) -> syn::Result<(Vec<TokenStream>, Vec<TokenStream>, Vec<String>)> {
let mut extract_statements = Vec::new();
let mut call_args = Vec::new();
let mut dep_strings = Vec::new();
for param in &constructor.params {
let name = ¶m.name;
let ty = ¶m.ty;
let ty_str = ¶m.ty_string;
if let Some(factory) = ¶m.factory_fn {
let path = factory.path();
if factory.is_async() {
extract_statements.push(quote! {
let #name: #ty = #path(ctx).await.map_err(|e|
injectable_rs_runtime::InjectableError::ConstructionFailed {
type_name: #ty_str,
reason: e.to_string(),
})?;
});
} else {
extract_statements.push(quote! {
let #name: #ty = #path(ctx);
});
}
call_args.push(quote! { #name });
continue;
}
if let Some(dyn_ty) = extract_inject_dyn_inner(ty) {
extract_statements.push(quote! {
let #name: #ty = {
let __arc = ctx.resolve_external::<::std::sync::Arc<#dyn_ty>>().await?;
injectable_rs_runtime::Inject::new(__arc)
};
});
call_args.push(quote! { #name });
continue;
}
if let Some(dyn_ty) = extract_option_inject_dyn_inner(ty) {
extract_statements.push(quote! {
let #name: #ty = match ctx.resolve_external::<::std::sync::Arc<#dyn_ty>>().await {
Ok(__arc) => Some(injectable_rs_runtime::Inject::new(__arc)),
Err(injectable_rs_runtime::InjectableError::MissingDependency { .. }) => None,
Err(__e) => return Err(__e),
};
});
call_args.push(quote! { #name });
continue;
}
extract_statements.push(quote! {
let #name: #ty =
<#ty as injectable_rs_runtime::Extract>::extract(ctx).await?;
});
call_args.push(quote! { #name });
if let Some(inner) = extract_inject_inner(ty) {
dep_strings.push(inner);
} else if let Some(inner) = extract_arc_inner_str(ty) {
dep_strings.push(inner);
} else {
dep_strings.push(ty_str.clone());
}
}
Ok((extract_statements, call_args, dep_strings))
}
fn generate_post_construct_calls(hooks: &[HookInfo], type_name_str: &str) -> TokenStream {
if hooks.is_empty() {
return quote! {};
}
let calls: Vec<TokenStream> = hooks
.iter()
.map(|hook| {
let hook_name = &hook.method_name;
let await_token = if hook.is_async {
quote! { .await }
} else {
quote! {}
};
if hook.returns_result {
quote! {
instance.#hook_name()#await_token.map_err(|e| injectable_rs_runtime::InjectableError::LifecycleHookFailed {
type_name: #type_name_str,
hook: "post_construct",
reason: e.to_string(),
})?;
}
} else {
quote! {
instance.#hook_name()#await_token;
}
}
})
.collect();
quote! { #(#calls)* }
}
fn generate_post_construct_impl(type_name: &syn::Ident, hooks: &[HookInfo]) -> TokenStream {
if hooks.is_empty() {
return quote! {};
}
let calls: Vec<TokenStream> = hooks
.iter()
.map(|hook| {
let hook_name = &hook.method_name;
let await_token = if hook.is_async {
quote! { .await }
} else {
quote! {}
};
if hook.returns_result {
quote! {
self.#hook_name()#await_token.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
}
} else {
quote! {
self.#hook_name()#await_token;
}
}
})
.collect();
quote! {
#[async_trait::async_trait]
impl injectable_rs_runtime::PostConstruct for #type_name {
async fn post_construct(&self) -> injectable_rs_runtime::HookResult {
#(#calls)*
Ok(())
}
}
}
}
fn generate_pre_destruct_impl(type_name: &syn::Ident, hooks: &[HookInfo]) -> TokenStream {
if hooks.is_empty() {
return quote! {};
}
let calls: Vec<TokenStream> = hooks
.iter()
.map(|hook| {
let hook_name = &hook.method_name;
let await_token = if hook.is_async {
quote! { .await }
} else {
quote! {}
};
if hook.returns_result {
quote! {
self.#hook_name()#await_token.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
}
} else {
quote! {
self.#hook_name()#await_token;
}
}
})
.collect();
quote! {
#[async_trait::async_trait]
impl injectable_rs_runtime::PreDestruct for #type_name {
async fn pre_destruct(&self) -> injectable_rs_runtime::HookResult {
#(#calls)*
Ok(())
}
}
}
}
fn generate_hooks_entry_submit(
type_name: &syn::Ident,
post_hooks: &[HookInfo],
pre_hooks: &[HookInfo],
) -> TokenStream {
if post_hooks.is_empty() && pre_hooks.is_empty() {
return quote! {};
}
let post_fn_name = syn::Ident::new(
&format!("__injectable_impl_post_{}", type_name),
proc_macro2::Span::call_site(),
);
let pre_fn_name = syn::Ident::new(
&format!("__injectable_impl_make_pre_{}", type_name),
proc_macro2::Span::call_site(),
);
let pre_adapter_name = syn::Ident::new(
&format!("__InjectableImplPreDestruct_{}", type_name),
proc_macro2::Span::call_site(),
);
let post_part = if !post_hooks.is_empty() {
let hook_calls: Vec<TokenStream> = post_hooks.iter().map(|hook| {
let method = &hook.method_name;
let await_tok = if hook.is_async { quote! { .await } } else { quote! {} };
if hook.returns_result {
quote! {
instance.#method()#await_tok.map_err(|e|
Box::new(e) as Box<dyn ::std::error::Error + ::std::marker::Send + ::std::marker::Sync>
)?;
}
} else {
quote! { instance.#method()#await_tok; }
}
}).collect();
quote! {
#[doc(hidden)]
#[allow(non_snake_case)]
fn #post_fn_name(
arc: ::std::sync::Arc<dyn ::std::any::Any + ::std::marker::Send + ::std::marker::Sync>,
) -> ::std::pin::Pin<Box<dyn ::std::future::Future<
Output = injectable_rs_runtime::HookResult
> + ::std::marker::Send + 'static>> {
let typed = ::std::sync::Arc::downcast::<#type_name>(arc)
.expect("InjectableHooksEntry TypeId guarantees correct type");
Box::pin(async move {
let instance: &_ = &*typed;
#(#hook_calls)*
Ok(())
})
}
}
} else {
quote! {}
};
let pre_part = if !pre_hooks.is_empty() {
let hook_calls: Vec<TokenStream> = pre_hooks.iter().map(|hook| {
let method = &hook.method_name;
let await_tok = if hook.is_async { quote! { .await } } else { quote! {} };
if hook.returns_result {
quote! {
instance.#method()#await_tok.map_err(|e|
Box::new(e) as Box<dyn ::std::error::Error + ::std::marker::Send + ::std::marker::Sync>
)?;
}
} else {
quote! { instance.#method()#await_tok; }
}
}).collect();
quote! {
#[doc(hidden)]
#[allow(non_camel_case_types)]
struct #pre_adapter_name(
::std::sync::Arc<dyn ::std::any::Any + ::std::marker::Send + ::std::marker::Sync>
);
#[async_trait::async_trait]
impl injectable_rs_runtime::PreDestruct for #pre_adapter_name {
async fn pre_destruct(&self) -> injectable_rs_runtime::HookResult {
let typed = ::std::sync::Arc::downcast::<#type_name>(
::std::sync::Arc::clone(&self.0)
).expect("InjectableHooksEntry TypeId guarantees correct type");
let instance: &_ = &*typed;
#(#hook_calls)*
Ok(())
}
}
#[doc(hidden)]
#[allow(non_snake_case)]
fn #pre_fn_name(
arc: ::std::sync::Arc<dyn ::std::any::Any + ::std::marker::Send + ::std::marker::Sync>,
) -> ::std::sync::Arc<dyn injectable_rs_runtime::PreDestruct> {
::std::sync::Arc::new(#pre_adapter_name(arc))
}
}
} else {
quote! {}
};
let post_fn_ref = if !post_hooks.is_empty() {
quote! { Some(#post_fn_name as injectable_rs_runtime::PostConstructFnPtr) }
} else {
quote! { None }
};
let pre_fn_ref = if !pre_hooks.is_empty() {
quote! { Some(#pre_fn_name as injectable_rs_runtime::MakePreDestructFnPtr) }
} else {
quote! { None }
};
quote! {
#post_part
#pre_part
injectable_rs_runtime::inventory::submit! {
injectable_rs_runtime::InjectableHooksEntry::new_const(
|| ::std::any::TypeId::of::<#type_name>(),
#post_fn_ref,
#pre_fn_ref,
)
}
}
}
fn generate_graph_metadata(
type_name: &syn::Ident,
dependencies: &[String],
scope: &str,
) -> TokenStream {
let type_str = type_name.to_string();
if dependencies.is_empty() {
quote! {
inventory::submit! {
injectable_rs_graph::GraphNode::leaf_with_scope(
#type_str,
#scope,
)
}
}
} else {
let dep_literals: Vec<_> = dependencies
.iter()
.map(|d| {
let d: &str = d;
quote! { #d }
})
.collect();
let dep_const_name = syn::Ident::new(
&format!(
"__INJECTABLE_GRAPH_DEPS_{}",
type_name.to_string().to_uppercase()
),
proc_macro2::Span::call_site(),
);
quote! {
#[allow(dead_code)]
const #dep_const_name: &[&str] = &[#(#dep_literals),*];
inventory::submit! {
injectable_rs_graph::GraphNode::with_scope(
#type_str,
#dep_const_name,
#scope,
)
}
}
}
}