use proc_macro::TokenStream;
use proc_macro2::{Ident, Span};
use quote::quote;
use syn::{
parse::{Parse, ParseStream, Result},
parse_macro_input,
punctuated::Punctuated,
token::{Colon, Comma, Eq},
Attribute, Error, Expr, ItemStruct, LitStr, Meta, Type,
};
pub fn inject_impl(args: TokenStream, input: TokenStream) -> TokenStream {
let inject_args = parse_macro_input!(args as InjectArgs);
let mut item_struct = parse_macro_input!(input as ItemStruct);
match process_inject_attribute(&mut item_struct, inject_args) {
Ok(result) => result.into(),
Err(err) => err.to_compile_error().into(),
}
}
#[derive(Debug, Clone)]
struct InjectArgs {
services: Punctuated<ServiceDef, Comma>,
}
impl Parse for InjectArgs {
fn parse(input: ParseStream) -> Result<Self> {
Ok(InjectArgs {
services: input.parse_terminated(ServiceDef::parse, Comma)?,
})
}
}
#[derive(Debug, Clone)]
struct ServiceDef {
field_name: Ident,
field_type: Type,
#[allow(dead_code)]
service_name: Option<String>,
injection_type: InjectionType,
factory_expr: Option<Expr>,
attributes: Vec<Attribute>,
}
#[derive(Debug, Clone, PartialEq)]
enum InjectionType {
Regular,
Optional,
Named(String),
Scoped,
Factory,
Token,
}
impl Parse for ServiceDef {
fn parse(input: ParseStream) -> Result<Self> {
let attributes = input.call(Attribute::parse_outer)?;
let field_name: Ident = input.parse()?;
let _colon: Colon = input.parse()?;
let field_type: Type = input.parse()?;
let mut service_name = None;
let mut factory_expr = None;
let mut injection_type = InjectionType::Regular;
if input.peek(Eq) {
let _eq: Eq = input.parse()?;
if input.peek(LitStr) {
let lit: LitStr = input.parse()?;
service_name = Some(lit.value());
injection_type = InjectionType::Named(lit.value());
} else {
let expr: Expr = input.parse()?;
factory_expr = Some(expr);
injection_type = InjectionType::Factory;
}
}
for attr in &attributes {
if let Meta::Path(path) = &attr.meta {
if let Some(ident) = path.get_ident() {
match ident.to_string().as_str() {
"scoped" => injection_type = InjectionType::Scoped,
"factory" => injection_type = InjectionType::Factory,
_ => {}
}
}
}
}
if injection_type == InjectionType::Regular {
if is_option_type(&field_type) {
injection_type = InjectionType::Optional;
} else if is_token_reference(&field_type) {
injection_type = InjectionType::Token;
}
}
Ok(ServiceDef {
field_name,
field_type,
service_name,
injection_type,
factory_expr,
attributes,
})
}
}
fn is_option_type(ty: &Type) -> bool {
if let Type::Path(type_path) = ty {
if let Some(segment) = type_path.path.segments.last() {
return segment.ident == "Option";
}
}
false
}
fn is_token_reference(ty: &Type) -> bool {
if let Type::Reference(type_ref) = ty {
if let Type::Path(type_path) = type_ref.elem.as_ref() {
if let Some(segment) = type_path.path.segments.last() {
let type_name = segment.ident.to_string();
return type_name.ends_with("Token");
}
}
}
false
}
fn extract_token_type(ty: &Type) -> Result<&Type> {
if let Type::Reference(type_ref) = ty {
return Ok(type_ref.elem.as_ref());
}
Err(Error::new_spanned(
ty,
"Expected reference type (&TokenType)",
))
}
impl ServiceDef {
fn is_optional(&self) -> bool {
matches!(self.injection_type, InjectionType::Optional) || is_option_type(&self.field_type)
}
fn get_inner_type(&self) -> Result<&Type> {
if let Type::Path(type_path) = &self.field_type {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "Option" {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
if let Some(syn::GenericArgument::Type(inner_type)) = args.args.first() {
return Ok(inner_type);
}
}
return Err(Error::new_spanned(
&self.field_type,
"Failed to extract inner type from Option<T>",
));
}
}
}
Ok(&self.field_type)
}
fn get_service_type(&self) -> Result<&Type> {
if self.is_optional() {
self.get_inner_type()
} else {
Ok(&self.field_type)
}
}
fn get_token_type(&self) -> Result<&Type> {
if matches!(self.injection_type, InjectionType::Token) {
extract_token_type(&self.field_type)
} else {
Err(Error::new_spanned(
&self.field_type,
"Not a token-based service definition",
))
}
}
}
fn process_inject_attribute(
item_struct: &mut ItemStruct,
inject_args: InjectArgs,
) -> Result<proc_macro2::TokenStream> {
if inject_args.services.is_empty() {
return Err(Error::new(
Span::call_site(),
"#[inject] requires at least one service definition",
));
}
let service_fields = generate_service_fields(&inject_args.services)?;
match &mut item_struct.fields {
syn::Fields::Named(fields) => {
for field in service_fields {
fields.named.push(field);
}
}
syn::Fields::Unnamed(_) => {
return Err(Error::new_spanned(
item_struct,
"#[inject] can only be applied to structs with named fields",
));
}
syn::Fields::Unit => {
item_struct.fields = syn::Fields::Named(syn::FieldsNamed {
brace_token: Default::default(),
named: service_fields.into_iter().collect(),
});
}
}
let from_ioc_container_impl =
generate_from_ioc_container_method(&item_struct.ident, &inject_args.services)?;
Ok(quote! {
#item_struct
#from_ioc_container_impl
})
}
fn generate_service_fields(services: &Punctuated<ServiceDef, Comma>) -> Result<Vec<syn::Field>> {
let mut fields = Vec::new();
for service in services {
let field_name = &service.field_name;
let field_core_type = match service.injection_type {
InjectionType::Token => {
let token_type = service.get_token_type()?;
quote! { std::sync::Arc<<#token_type as elif_core::container::ServiceToken>::Service> }
}
_ => {
let service_type = service.get_service_type()?;
quote! { std::sync::Arc<#service_type> }
}
};
let wrapped_type = if service.is_optional() {
quote! { Option<#field_core_type> }
} else {
field_core_type
};
let field = syn::Field {
attrs: service.attributes.clone(),
vis: syn::Visibility::Inherited, mutability: syn::FieldMutability::None,
ident: Some(field_name.clone()),
colon_token: Some(Default::default()),
ty: syn::parse2(wrapped_type)?,
};
fields.push(field);
}
Ok(fields)
}
fn generate_from_ioc_container_method(
struct_name: &Ident,
services: &Punctuated<ServiceDef, Comma>,
) -> Result<proc_macro2::TokenStream> {
let mut field_initializers = Vec::new();
for service in services {
let field_name = &service.field_name;
let service_type = service.get_service_type()?;
let initializer = if service.is_optional() {
match &service.injection_type {
InjectionType::Regular | InjectionType::Optional => {
quote! {
#field_name: container.try_resolve::<#service_type>()
}
}
InjectionType::Named(name) => {
quote! {
#field_name: container.try_resolve_named::<#service_type>(#name)
}
}
InjectionType::Scoped => {
quote! {
#field_name: container.try_resolve_scoped::<#service_type>(&scope_id)
}
}
InjectionType::Factory => {
if let Some(factory_expr) = &service.factory_expr {
quote! {
#field_name: {
let factory = #factory_expr;
Some(Arc::new(factory(container)?))
}
}
} else {
quote! {
#field_name: container.try_resolve::<#service_type>()
}
}
}
InjectionType::Token => {
let token_type = service.get_token_type()?;
quote! {
#field_name: container.try_resolve_by_token::<#token_type>()
}
}
}
} else {
match &service.injection_type {
InjectionType::Regular => {
quote! {
#field_name: container.resolve::<#service_type>()
.map_err(|e| format!("Failed to inject service {}: {}", stringify!(#service_type), e))?
}
}
InjectionType::Named(name) => {
quote! {
#field_name: container.resolve_named::<#service_type>(#name)
.map_err(|e| format!("Failed to inject named service {}({}): {}", stringify!(#service_type), #name, e))?
}
}
InjectionType::Scoped => {
quote! {
#field_name: container.resolve_scoped::<#service_type>(&scope_id)
.map_err(|e| format!("Failed to inject scoped service {}: {}", stringify!(#service_type), e))?
}
}
InjectionType::Factory => {
if let Some(factory_expr) = &service.factory_expr {
quote! {
#field_name: {
let factory = #factory_expr;
Arc::new(factory(container)?)
}
}
} else {
quote! {
#field_name: container.resolve::<#service_type>()
.map_err(|e| format!("Failed to inject factory service {}: {}", stringify!(#service_type), e))?
}
}
}
InjectionType::Optional => {
quote! {
#field_name: container.resolve::<#service_type>()
.map_err(|e| format!("Failed to inject service {}: {}", stringify!(#service_type), e))?
}
}
InjectionType::Token => {
let token_type = service.get_token_type()?;
quote! {
#field_name: container.resolve_by_token::<#token_type>()
.map_err(|e| format!("Failed to inject token service {}: {}", stringify!(#token_type), e))?
}
}
}
};
field_initializers.push(initializer);
}
Ok(quote! {
impl #struct_name {
pub fn from_ioc_container(
container: &elif_core::container::IocContainer,
scope: Option<&elif_core::container::ScopeId>
) -> Result<Self, String> {
let scope_id = match scope {
Some(s) => s.clone(),
None => container.create_scope()
.map_err(|e| format!("Failed to create scope: {}", e))?
};
Ok(Self {
#(#field_initializers),*
})
}
}
})
}
#[cfg(test)]
mod tests {
use super::*;
use syn::parse_quote;
#[test]
fn test_is_token_reference_with_token_suffix() {
let email_token: Type = parse_quote!(&EmailNotificationToken);
assert!(is_token_reference(&email_token));
let db_token: Type = parse_quote!(&DatabaseToken);
assert!(is_token_reference(&db_token));
let user_token: Type = parse_quote!(&UserServiceToken);
assert!(is_token_reference(&user_token));
}
#[test]
fn test_is_token_reference_without_token_suffix() {
let config: Type = parse_quote!(&Config);
assert!(!is_token_reference(&config));
let connection: Type = parse_quote!(&DatabaseConnection);
assert!(!is_token_reference(&connection));
let str_ref: Type = parse_quote!(&str);
assert!(!is_token_reference(&str_ref));
let service: Type = parse_quote!(&UserService);
assert!(!is_token_reference(&service));
}
#[test]
fn test_is_token_reference_non_references() {
let owned_token: Type = parse_quote!(EmailNotificationToken);
assert!(!is_token_reference(&owned_token));
let owned_service: Type = parse_quote!(UserService);
assert!(!is_token_reference(&owned_service));
let arc_service: Type = parse_quote!(Arc<UserService>);
assert!(!is_token_reference(&arc_service));
}
#[test]
fn test_is_token_reference_with_paths() {
let module_token: Type = parse_quote!(&crate::tokens::EmailNotificationToken);
assert!(is_token_reference(&module_token));
let module_non_token: Type = parse_quote!(&crate::config::DatabaseConfig);
assert!(!is_token_reference(&module_non_token));
let std_ref: Type = parse_quote!(&std::collections::HashMap<String, String>);
assert!(!is_token_reference(&std_ref));
}
}