Skip to main content

diode_macros/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3
4use syn::spanned::Spanned as _;
5use syn::{
6    Attribute, Data, DeriveInput, Error, FnArg, GenericArgument, ImplItem, ItemImpl, Pat,
7    PathArguments, Type,
8};
9
10fn extract_arc_type(ty: &Type) -> Option<Type> {
11    if let Type::Path(type_path) = ty
12        && let Some(segment) = type_path.path.segments.last()
13        && segment.ident == "Arc"
14        && let PathArguments::AngleBracketed(args) = &segment.arguments
15        && let Some(GenericArgument::Type(inner)) = args.args.first()
16    {
17        return Some(inner.clone());
18    }
19    None
20}
21
22fn extract_extract_type(attrs: &[Attribute]) -> Option<Type> {
23    for attr in attrs {
24        if attr.path().is_ident(EXTRACT_ATTR)
25            && let Ok(meta_list) = attr.meta.require_list()
26            && let Ok(ty) = syn::parse2::<Type>(meta_list.tokens.clone())
27        {
28            return Some(ty);
29        }
30    }
31    None
32}
33
34const EXTRACT_ATTR: &str = "inject";
35const FACTORY_ATTR: &str = "factory";
36
37/// Derive macro for Service trait
38#[proc_macro_derive(Service, attributes(inject))]
39pub fn derive_service(input: TokenStream) -> TokenStream {
40    let input = syn::parse_macro_input!(input as DeriveInput);
41    handle_derive_service(input)
42}
43
44/// Attribute macro for impl blocks with factory methods
45#[proc_macro_attribute]
46pub fn service(_attr: TokenStream, item: TokenStream) -> TokenStream {
47    if let Ok(item_impl) = syn::parse::<ItemImpl>(item) {
48        return handle_service_impl(item_impl);
49    }
50    TokenStream::from(
51        Error::new(
52            proc_macro2::Span::call_site(),
53            "#[service] can only be applied to impl blocks",
54        )
55        .to_compile_error(),
56    )
57}
58
59fn handle_derive_service(input: DeriveInput) -> TokenStream {
60    let name = &input.ident;
61    let fields = match &input.data {
62        Data::Struct(s) => &s.fields,
63        _ => {
64            return TokenStream::from(
65                Error::new(name.span(), "Only structs are supported").to_compile_error(),
66            );
67        }
68    };
69
70    let mut dependency_stmts = Vec::new();
71    let mut field_inits = Vec::new();
72    let mut field_lets = Vec::new();
73
74    match fields {
75        syn::Fields::Named(fields) => {
76            for field in &fields.named {
77                let field_ident = field.ident.as_ref().unwrap();
78                let field_ty = &field.ty;
79
80                if let Some(extract_type) = extract_extract_type(&field.attrs) {
81                    field_lets.push(quote! {
82                        let #field_ident = <#extract_type as ::diode::Extract<#field_ty>>::extract(ctx)?;
83                    });
84
85                    dependency_stmts.push(quote! {
86                        deps = deps.merge(<#extract_type as ::diode::Extract<#field_ty>>::dependencies());
87                    });
88
89                    field_inits.push(quote! { #field_ident: #field_ident });
90                } else if let Some(inner_type) = extract_arc_type(field_ty) {
91                    dependency_stmts.push(quote! {
92                        deps = deps.service::<#inner_type>();
93                    });
94
95                    field_lets.push(quote! {
96                        let #field_ident = ctx
97                            .get_component::<<#inner_type as ::diode::Service>::Handle>()
98                            .ok_or_else(|| {
99                                format!(
100                                    "Missing component: {}",
101                                    ::std::any::type_name::<<#inner_type as ::diode::Service>::Handle>()
102                                )
103                            })?;
104                    });
105
106                    field_inits.push(quote! { #field_ident: #field_ident });
107                } else {
108                    return TokenStream::from(
109                        Error::new(
110                            field_ty.span(),
111                            format!("Service dependencies must be of type Arc<T> or use #[{EXTRACT_ATTR}]",),
112                        )
113                        .to_compile_error(),
114                    );
115                }
116            }
117        }
118        syn::Fields::Unnamed(_) => {
119            return TokenStream::from(
120                Error::new(name.span(), "Tuple structs are not supported").to_compile_error(),
121            );
122        }
123        syn::Fields::Unit => {}
124    }
125
126    quote! {
127        impl ::diode::Service for #name {
128            type Handle = ::std::sync::Arc<Self>;
129
130            async fn build(
131                ctx: &::diode::AppContext
132            ) -> Result<Self::Handle, ::diode::StdError> {
133                #(#field_lets)*
134                Ok(::std::sync::Arc::new(Self {
135                    #(#field_inits,)*
136                }))
137            }
138
139            fn dependencies() -> ::diode::Dependencies {
140                use ::diode::ServiceDependencyExt as _;
141                let mut deps = ::diode::Dependencies::new();
142                #(#dependency_stmts)*
143                deps
144            }
145        }
146    }
147    .into()
148}
149
150fn handle_service_impl(input: ItemImpl) -> TokenStream {
151    if input.trait_.is_some() {
152        return TokenStream::from(
153            Error::new(input.span(), "Trait impls are not supported").to_compile_error(),
154        );
155    }
156
157    let self_ty = &input.self_ty;
158    let mut new_method = None;
159
160    for item in &input.items {
161        if let ImplItem::Fn(method) = item {
162            for attr in &method.attrs {
163                if attr.path().is_ident(FACTORY_ATTR) {
164                    if new_method.is_some() {
165                        return TokenStream::from(
166                            Error::new(attr.span(), "Only one constructor method allowed")
167                                .to_compile_error(),
168                        );
169                    }
170                    new_method = Some(method);
171                }
172            }
173        }
174    }
175
176    let method = match new_method {
177        Some(m) => m,
178        None => {
179            return TokenStream::from(
180                Error::new(input.span(), "No factory method found").to_compile_error(),
181            );
182        }
183    };
184
185    let method_name = &method.sig.ident;
186    let is_async = method.sig.asyncness.is_some();
187    let mut dependency_stmts = Vec::new();
188    let mut arg_inits = Vec::new();
189    let mut arg_names = Vec::new();
190
191    // Extract the actual return type to use as Handle
192    let return_type = match &method.sig.output {
193        syn::ReturnType::Default => {
194            return TokenStream::from(
195                Error::new(method.sig.span(), "Factory method must have a return type")
196                    .to_compile_error(),
197            );
198        }
199        syn::ReturnType::Type(_, ty) => ty.as_ref(),
200    };
201
202    // Determine the Handle type and whether the return type is Result
203    let (handle_type, is_result) = extract_handle_type(return_type);
204
205    // Create cleaned inputs without extract attributes
206    let mut cleaned_inputs = Vec::new();
207    let mut has_mut_ref = false;
208    let mut ref_count: usize = 0;
209
210    for fn_arg in &method.sig.inputs {
211        match fn_arg {
212            FnArg::Receiver(_) => {
213                return TokenStream::from(
214                    Error::new(
215                        fn_arg.span(),
216                        "Constructor method cannot have self parameter",
217                    )
218                    .to_compile_error(),
219                );
220            }
221            FnArg::Typed(pat_type) => {
222                let arg_ty = &pat_type.ty;
223
224                // Create cleaned parameter without extract attributes
225                let mut cleaned_pat_type = pat_type.clone();
226                cleaned_pat_type
227                    .attrs
228                    .retain(|attr| !attr.path().is_ident(EXTRACT_ATTR));
229                cleaned_inputs.push(FnArg::Typed(cleaned_pat_type));
230
231                if let Pat::Ident(pat_ident) = pat_type.pat.as_ref() {
232                    let arg_name = &pat_ident.ident;
233
234                    if let Some(extract_type) = extract_extract_type(&pat_type.attrs) {
235                        match arg_ty.as_ref() {
236                            Type::Reference(ref_ty) if ref_ty.mutability.is_some() => {
237                                has_mut_ref = true;
238                                ref_count += 1;
239                                arg_names.push(quote! { #arg_name.deref_mut() });
240                                let inner_ty = &ref_ty.elem;
241                                arg_inits.push(quote! {
242                                    let mut #arg_name = <#extract_type as ::diode::ExtractMut<#inner_ty>>::extract_mut(ctx)?;
243                                });
244                                dependency_stmts.push(quote! {
245                                    deps = deps.merge(<#extract_type as ::diode::ExtractRef<#inner_ty>>::dependencies());
246                                });
247                            }
248                            Type::Reference(ref_ty) => {
249                                ref_count += 1;
250                                arg_names.push(quote! { #arg_name.deref() });
251                                let inner_ty = &ref_ty.elem;
252                                arg_inits.push(quote! {
253                                    let #arg_name = <#extract_type as ::diode::ExtractRef<#inner_ty>>::extract_ref(ctx)?;
254                                });
255                                dependency_stmts.push(quote! {
256                                    deps = deps.merge(<#extract_type as ::diode::ExtractRef<#inner_ty>>::dependencies());
257                                });
258                            }
259                            _ => {
260                                arg_names.push(quote! { #arg_name });
261                                arg_inits.push(quote! {
262                                    let #arg_name = <#extract_type as ::diode::Extract<#arg_ty>>::extract(ctx)?;
263                                });
264                                dependency_stmts.push(quote! {
265                                    deps = deps.merge(<#extract_type as ::diode::Extract<#arg_ty>>::dependencies());
266                                });
267                            }
268                        };
269                    } else if let Some(inner_type) = extract_arc_type(arg_ty) {
270                        arg_names.push(quote! { #arg_name });
271                        dependency_stmts.push(quote! {
272                            deps = deps.service::<#inner_type>();
273                        });
274
275                        arg_inits.push(quote! {
276                            let #arg_name = ctx
277                                .get_component::<<#inner_type as ::diode::Service>::Handle>()
278                                .ok_or_else(|| {
279                                    format!(
280                                        "Missing component: {}",
281                                        ::std::any::type_name::<<#inner_type as ::diode::Service>::Handle>()
282                                    )
283                                })?;
284                        });
285                    } else {
286                        return TokenStream::from(
287                            Error::new(
288                                arg_ty.span(),
289                                format!(
290                                    "Arguments must be of type Arc<T> or use #[{EXTRACT_ATTR}]",
291                                ),
292                            )
293                            .to_compile_error(),
294                        );
295                    }
296                } else {
297                    return TokenStream::from(
298                        Error::new(pat_type.pat.span(), "Only simple bindings supported")
299                            .to_compile_error(),
300                    );
301                }
302            }
303        }
304    }
305
306    if has_mut_ref && ref_count > 1 {
307        return TokenStream::from(
308            Error::new(
309                method.sig.span(),
310                "Combining a `&mut` inject parameter with other `&` or `&mut` inject parameters \
311                 may cause a deadlock. Use `#[inject(AppContext)] ctx: &AppContext` and call \
312                 `get_component_ref`/`get_component_mut` manually, ensuring that guards do not \
313                 overlap.",
314            )
315            .to_compile_error(),
316        );
317    }
318
319    // Create cleaned input with extract and new attributes removed
320    let mut cleaned_input = input.clone();
321    for item in &mut cleaned_input.items {
322        if let ImplItem::Fn(method) = item
323            && method
324                .attrs
325                .iter()
326                .any(|attr| attr.path().is_ident(FACTORY_ATTR))
327        {
328            // Remove extract attributes from method parameters
329            method.sig.inputs = cleaned_inputs.into_iter().collect();
330            // Remove new attribute from method
331            method
332                .attrs
333                .retain(|attr| !attr.path().is_ident(FACTORY_ATTR));
334            break;
335        }
336    }
337
338    // Generate the method call based on whether it's async and returns Result
339    let method_call = if is_async {
340        quote! { Self::#method_name(#(#arg_names),*).await }
341    } else {
342        quote! { Self::#method_name(#(#arg_names),*) }
343    };
344
345    // Wrap the call based on whether the original method returns Result
346    let build_body = if is_result {
347        quote! {
348            #(#arg_inits)*
349            #method_call.map_err(|e| e.into())
350        }
351    } else {
352        quote! {
353            #(#arg_inits)*
354            Ok(#method_call)
355        }
356    };
357
358    quote! {
359        #cleaned_input
360
361        impl ::diode::Service for #self_ty {
362            type Handle = #handle_type;
363
364            async fn build(
365                ctx: &::diode::AppContext
366            ) -> Result<Self::Handle, ::diode::StdError> {
367                use ::std::ops::{Deref as _, DerefMut as _};
368                #build_body
369            }
370
371            fn dependencies() -> ::diode::Dependencies {
372                use ::diode::ServiceDependencyExt as _;
373                let mut deps = ::diode::Dependencies::new();
374                #(#dependency_stmts)*
375                deps
376            }
377        }
378    }
379    .into()
380}
381
382fn extract_handle_type(ty: &Type) -> (Type, bool) {
383    // Check if return type is Result<T, E>
384    if let Type::Path(type_path) = ty
385        && let Some(segment) = type_path.path.segments.last()
386        && segment.ident == "Result"
387        && let PathArguments::AngleBracketed(args) = &segment.arguments
388        && let Some(GenericArgument::Type(inner)) = args.args.first()
389    {
390        // Return the T from Result<T, E>
391        return (inner.clone(), true);
392    }
393    // Return the type as-is if not Result
394    (ty.clone(), false)
395}