use proc_macro2::TokenStream;
use quote::quote;
use crate::metadata::{
extract_arc_inner, extract_arc_inner_str, extract_inject_dyn_inner, extract_inject_inner,
extract_option_inject_dyn_inner,
};
#[derive(Debug, Clone)]
pub enum FieldInjectKind {
Inject,
External,
Factory(syn::Path),
Provider(syn::Path),
}
#[derive(Debug, Clone)]
pub struct FieldInfo {
pub name: Option<syn::Ident>,
pub ty: syn::Type,
pub ty_string: String,
pub inject_kind: FieldInjectKind,
}
pub fn generate_field_injection_provider(
type_name: &syn::Ident,
generics: &syn::Generics,
fields: &[FieldInfo],
scope: &str,
has_post_construct: bool,
) -> TokenStream {
let provider_name = provider_ident(type_name);
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let is_generic = !generics.params.is_empty();
let field_statements: Vec<TokenStream> = fields
.iter()
.enumerate()
.map(|(i, field)| {
let field_ty = &field.ty;
match &field.inject_kind {
FieldInjectKind::Inject => {
let var_expr = if let Some(dyn_ty) = extract_inject_dyn_inner(field_ty) {
quote! {
{
let __arc = ctx.resolve_external::<::std::sync::Arc<#dyn_ty>>().await?;
injectable_rs_runtime::Inject::new(__arc)
}
}
} else if let Some(dyn_ty) = extract_option_inject_dyn_inner(field_ty) {
quote! {
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),
}
}
} else {
quote! {
<#field_ty as injectable_rs_runtime::Extract>::extract(ctx).await?
}
};
if let Some(name) = &field.name {
quote! { let #name = #var_expr; }
} else {
let temp_name = syn::Ident::new(
&format!("__field_{}", i),
proc_macro2::Span::call_site(),
);
quote! { let #temp_name = #var_expr; }
}
}
FieldInjectKind::External => {
let var_expr = quote! {
ctx.resolve_external::<#field_ty>().await?
};
if let Some(name) = &field.name {
quote! { let #name = #var_expr; }
} else {
let temp_name = syn::Ident::new(
&format!("__field_{}", i),
proc_macro2::Span::call_site(),
);
quote! { let #temp_name = #var_expr; }
}
}
FieldInjectKind::Factory(factory_path) => {
let ty_str = &field.ty_string;
let wrap = factory_wrap_for_field_type(&field.ty);
let var_expr = quote! {
{
let __v = #factory_path(ctx).await.map_err(|e|
injectable_rs_runtime::InjectableError::ConstructionFailed {
type_name: #ty_str,
reason: e.to_string(),
})?;
#wrap
}
};
if let Some(name) = &field.name {
quote! { let #name = #var_expr; }
} else {
let temp_name = syn::Ident::new(
&format!("__field_{}", i),
proc_macro2::Span::call_site(),
);
quote! { let #temp_name = #var_expr; }
}
}
FieldInjectKind::Provider(factory_path) => {
let var_expr = quote! { #factory_path(ctx) };
if let Some(name) = &field.name {
quote! { let #name = #var_expr; }
} else {
let temp_name = syn::Ident::new(
&format!("__field_{}", i),
proc_macro2::Span::call_site(),
);
quote! { let #temp_name = #var_expr; }
}
}
}
})
.collect();
let construction = if fields.is_empty() {
quote! { #type_name }
} else if fields.first().is_some_and(|f| f.name.is_some()) {
let field_names: Vec<_> = fields
.iter()
.filter_map(|f| f.name.as_ref())
.cloned()
.collect();
quote! { #type_name { #(#field_names),* } }
} else {
let field_refs: Vec<_> = fields
.iter()
.enumerate()
.map(|(i, _)| {
syn::Ident::new(&format!("__field_{}", i), proc_macro2::Span::call_site())
})
.collect();
quote! { #type_name(#(#field_refs),*) }
};
#[allow(clippy::unnecessary_filter_map)]
let dep_strings: Vec<String> = fields
.iter()
.filter(|f| matches!(f.inject_kind, FieldInjectKind::Inject))
.filter_map(|f| {
let ty_str = &f.ty_string;
if extract_inject_dyn_inner(&f.ty).is_some()
|| extract_option_inject_dyn_inner(&f.ty).is_some()
{
None
} else if let Some(inner) = extract_inject_inner(&f.ty) {
Some(inner)
} else if let Some(inner) = extract_arc_inner_str(&f.ty) {
Some(inner)
} else {
Some(ty_str.clone())
}
})
.collect();
let graph_metadata = if is_generic {
quote! {}
} else {
generate_graph_metadata_from_strings(type_name, &dep_strings, scope)
};
let is_singleton: bool = scope != "transient";
let arc_factory_submit = if is_generic {
quote! {}
} else {
generate_arc_factory_submit(type_name, is_singleton)
};
let hooks_dispatch = if is_generic {
quote! { Ok(instance) }
} else {
generate_inventory_hooks_dispatch(type_name)
};
let legacy_hooks_submit = if is_generic {
quote! {}
} else {
generate_legacy_hooks_submit(type_name, has_post_construct, false)
};
let provider_struct = if is_generic {
let phantom = phantom_for_generics(generics);
quote! { pub struct #provider_name #impl_generics (#phantom); }
} else {
quote! { pub struct #provider_name; }
};
quote! {
#provider_struct
#[async_trait::async_trait]
impl #impl_generics injectable_rs_runtime::Provider<#type_name #ty_generics>
for #provider_name #ty_generics
#where_clause
{
async fn provide(
ctx: &injectable_rs_runtime::ResolveContext,
) -> injectable_rs_runtime::InjectableResult<#type_name #ty_generics> {
#(#field_statements)*
let instance = #construction;
#hooks_dispatch
}
}
impl #impl_generics injectable_rs_runtime::Injectable for #type_name #ty_generics
#where_clause
{
type Provider = #provider_name #ty_generics;
const IS_SINGLETON: bool = #is_singleton;
}
#graph_metadata
#arc_factory_submit
#legacy_hooks_submit
}
}
pub(crate) fn phantom_for_generics(generics: &syn::Generics) -> proc_macro2::TokenStream {
let types: Vec<_> = generics
.type_params()
.map(|tp| {
let id = &tp.ident;
quote! { #id }
})
.collect();
let lifetimes: Vec<_> = generics
.lifetimes()
.map(|lp| {
let lt = &lp.lifetime;
quote! { &#lt () }
})
.collect();
quote! { ::std::marker::PhantomData<(#(#types,)* #(#lifetimes,)*)> }
}
pub(crate) fn generate_arc_factory_submit(
type_name: &syn::Ident,
is_singleton: bool,
) -> TokenStream {
let type_id_fn_name = syn::Ident::new(
&format!("__injectable_type_id_{}", type_name),
proc_macro2::Span::call_site(),
);
let provide_fn_name = syn::Ident::new(
&format!("__injectable_provide_{}", type_name),
proc_macro2::Span::call_site(),
);
let resolve_expr = if is_singleton {
quote! {
<::std::sync::Arc<#type_name> as injectable_rs_runtime::Extract>::extract(&ctx).await
}
} else {
quote! {
<#type_name as injectable_rs_runtime::Injectable>::Provider::provide(&ctx)
.await
.map(::std::sync::Arc::new)
}
};
quote! {
#[doc(hidden)]
#[allow(non_snake_case)]
fn #type_id_fn_name() -> ::std::any::TypeId {
::std::any::TypeId::of::<::std::sync::Arc<#type_name>>()
}
#[doc(hidden)]
#[allow(non_snake_case)]
fn #provide_fn_name(
ctx: ::std::sync::Arc<injectable_rs_runtime::ResolveContext>,
) -> ::std::pin::Pin<Box<dyn ::std::future::Future<
Output = injectable_rs_runtime::InjectableResult<Box<dyn ::std::any::Any + Send>>
> + Send + 'static>> {
Box::pin(async move {
#resolve_expr
.map(|arc| -> Box<dyn ::std::any::Any + Send> { Box::new(arc) })
})
}
injectable_rs_runtime::inventory::submit! {
injectable_rs_runtime::InjectableArcFactory::new_const(
stringify!(#type_name),
#type_id_fn_name,
#provide_fn_name,
)
}
}
}
fn generate_graph_metadata_from_strings(
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,
)
}
}
}
}
fn factory_wrap_for_field_type(field_ty: &syn::Type) -> TokenStream {
if extract_inject_inner(field_ty).is_some() {
quote! { injectable_rs_runtime::Inject::new(::std::sync::Arc::new(__v)) }
} else if extract_arc_inner(field_ty).is_some() {
quote! { ::std::sync::Arc::new(__v) }
} else {
quote! { __v }
}
}
fn provider_ident(type_name: &syn::Ident) -> syn::Ident {
syn::Ident::new(
&format!("{}Provider", type_name),
proc_macro2::Span::call_site(),
)
}
pub(crate) fn generate_inventory_hooks_dispatch(type_name: &syn::Ident) -> TokenStream {
let type_str = type_name.to_string();
quote! {
let mut __post_fn_opt: Option<injectable_rs_runtime::PostConstructFnPtr> = None;
let __type_id = ::std::any::TypeId::of::<#type_name>();
for __h in injectable_rs_runtime::inventory::iter::<injectable_rs_runtime::InjectableHooksEntry>() {
if __h.type_id() == __type_id {
__post_fn_opt = __h.post_construct_fn();
break;
}
}
if let Some(__post_fn) = __post_fn_opt {
let __post_arc: ::std::sync::Arc<#type_name> = ::std::sync::Arc::new(instance);
let __post_arc_cloned: ::std::sync::Arc<#type_name> = ::std::sync::Arc::clone(&__post_arc);
let __arc_any: ::std::sync::Arc<dyn ::std::any::Any + ::std::marker::Send + ::std::marker::Sync>
= __post_arc_cloned;
__post_fn(__arc_any).await.map_err(|e|
injectable_rs_runtime::InjectableError::LifecycleHookFailed {
type_name: #type_str,
hook: "post_construct",
reason: e.to_string(),
}
)?;
let instance = match ::std::sync::Arc::try_unwrap(__post_arc) {
Ok(v) => v,
Err(_) => panic!(
"post_construct hook must not retain the Arc past its completion"
),
};
Ok(instance)
} else {
Ok(instance)
}
}
}
pub(crate) fn generate_legacy_hooks_submit(
type_name: &syn::Ident,
has_post_construct: bool,
has_pre_destruct: bool,
) -> TokenStream {
if !has_post_construct && !has_pre_destruct {
return quote! {};
}
let post_fn_name = syn::Ident::new(
&format!("__injectable_legacy_post_{}", type_name),
proc_macro2::Span::call_site(),
);
let pre_fn_name = syn::Ident::new(
&format!("__injectable_legacy_make_pre_{}", type_name),
proc_macro2::Span::call_site(),
);
let pre_adapter_name = syn::Ident::new(
&format!("__InjectableLegacyPreDestruct_{}", type_name),
proc_macro2::Span::call_site(),
);
let post_part = if has_post_construct {
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 {
injectable_rs_runtime::PostConstruct::post_construct(&*typed).await
})
}
}
} else {
quote! {}
};
let pre_part = if has_pre_destruct {
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");
injectable_rs_runtime::PreDestruct::pre_destruct(&*typed).await
}
}
#[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 has_post_construct {
quote! { Some(#post_fn_name as injectable_rs_runtime::PostConstructFnPtr) }
} else {
quote! { None }
};
let pre_fn_ref = if has_pre_destruct {
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,
)
}
}
}