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#[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#[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 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 let (handle_type, is_result) = extract_handle_type(return_type);
204
205 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 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 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 method.sig.inputs = cleaned_inputs.into_iter().collect();
330 method
332 .attrs
333 .retain(|attr| !attr.path().is_ident(FACTORY_ATTR));
334 break;
335 }
336 }
337
338 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 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 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 (inner.clone(), true);
392 }
393 (ty.clone(), false)
395}