Skip to main content

comprehensive_macros/
lib.rs

1//! Macros in support of [`comprehensive`]. It is not necessary to depend on this crate directly.
2//!
3//! [`comprehensive`]: https://docs.rs/comprehensive/latest/comprehensive/
4
5// Would impose a requirement for rustc 1.88
6// https://github.com/rust-lang/rust/pull/132833
7#![allow(clippy::collapsible_if)]
8
9extern crate proc_macro;
10use convert_case::{Case, Casing};
11use proc_macro2::{Span, TokenStream};
12use quote::{ToTokens, format_ident, quote, quote_spanned};
13use syn::punctuated::Punctuated;
14use syn::spanned::Spanned;
15use syn::{
16    Attribute, Data, DeriveInput, Fields, GenericArgument, Generics, Ident, Lit, LitBool, LitStr,
17    Path, PathArguments, Type, Visibility, parse_macro_input,
18};
19
20enum DependencyType<'a> {
21    Concrete(&'a Type, bool),
22    Trait(&'a Type, bool),
23    Weak(&'a Type),
24    NewStyle(&'a Type),
25}
26
27enum DependencyLabel {
28    Arc,
29    Option,
30    Vec,
31    PhantomData,
32}
33
34impl DependencyLabel {
35    fn expect_arc_inside(&self) -> bool {
36        match *self {
37            Self::Arc => false,
38            Self::Option => true,
39            Self::PhantomData => false,
40            Self::Vec => true,
41        }
42    }
43
44    fn into_dependency_type<'a>(
45        self,
46        ty: &'a Type,
47        may_fail: bool,
48    ) -> Result<DependencyType<'a>, Span> {
49        let (expect_trait, result) = match self {
50            Self::Arc => (false, DependencyType::Concrete(ty, false)),
51            Self::Option => (false, DependencyType::Concrete(ty, true)),
52            Self::Vec => (true, DependencyType::Trait(ty, may_fail)),
53            Self::PhantomData => (false, DependencyType::Weak(ty)),
54        };
55        if expect_trait == matches!(ty, Type::TraitObject(_)) {
56            Ok(result)
57        } else {
58            Err(ty.span())
59        }
60    }
61}
62
63fn find_path_with_1_generic_type(ty: &Type) -> Result<(DependencyLabel, &Type), Span> {
64    let Type::Path(path) = ty else {
65        return Err(ty.span());
66    };
67    let last_segment = path.path.segments.last().ok_or_else(|| path.span())?;
68    let dep_type = if last_segment.ident == "Arc" {
69        DependencyLabel::Arc
70    } else if last_segment.ident == "Vec" {
71        DependencyLabel::Vec
72    } else if last_segment.ident == "Option" {
73        DependencyLabel::Option
74    } else if last_segment.ident == "PhantomData" {
75        DependencyLabel::PhantomData
76    } else {
77        return Err(path.span());
78    };
79    // a = <T> or <Arc<dyn Tr>> or <Arc<T>>
80    let PathArguments::AngleBracketed(ref generics) = last_segment.arguments else {
81        return Err(last_segment.arguments.span());
82    };
83    if generics.args.len() != 1 {
84        return Err(generics.span());
85    };
86    let generic = generics.args.first().unwrap();
87    let GenericArgument::Type(ty) = generic else {
88        return Err(generic.span());
89    };
90    Ok((dep_type, ty))
91}
92
93fn find_dependency_type<'a>(
94    orig_ty: &'a Type,
95    attrs: &[Attribute],
96) -> Result<DependencyType<'a>, Span> {
97    let may_fail = attrs.iter().any(|a| a.path().is_ident("may_fail"));
98    let old_style = attrs.iter().any(|a| a.path().is_ident("old_style"));
99    if !old_style && !may_fail {
100        return Ok(DependencyType::NewStyle(orig_ty));
101    }
102    // We accept:
103    // Arc<T>
104    // Option<Arc<T>>
105    // Vec<Arc<dyn Tr>>
106    // PhantomData<T>
107    let (dep_type, mut ty) = find_path_with_1_generic_type(orig_ty)?;
108    if dep_type.expect_arc_inside() {
109        let (inner_dep_type, inner_ty) = find_path_with_1_generic_type(ty)?;
110        if !matches!(inner_dep_type, DependencyLabel::Arc) {
111            return Err(orig_ty.span());
112        }
113        ty = inner_ty;
114    }
115    dep_type.into_dependency_type(ty, may_fail)
116}
117
118fn produce_concrete(
119    dep_types: &Vec<Result<DependencyType<'_>, TokenStream>>,
120    for_optional: bool,
121) -> impl Iterator<Item = TokenStream> {
122    dep_types
123        .iter()
124        .enumerate()
125        .filter_map(move |(i, r)| match r {
126            Ok(DependencyType::Concrete(ty, optional)) if *optional == for_optional => {
127                let temp = format_ident!("dep_{}", i);
128                Some(quote! {
129                    let #temp = ::comprehensive::assembly::Registrar::< #ty >::produce(cx);
130                })
131            }
132            Ok(DependencyType::Trait(ty, optional)) if *optional == for_optional => {
133                let temp = format_ident!("dep_{}", i);
134                Some(if for_optional {
135                    quote! { let #temp = cx.produce_trait::< #ty >(); }
136                } else {
137                    quote! { let #temp = cx.produce_trait_fallible::< #ty >(); }
138                })
139            }
140            Ok(DependencyType::NewStyle(ty)) => {
141                let temp = format_ident!("dep_{}", i);
142                Some(if for_optional {
143                    quote! { let #temp = < #ty as ::comprehensive::dependencies::ResourceDependency >::produce_late(cx, #temp ); }
144                } else {
145                    quote! { let #temp = < #ty as ::comprehensive::dependencies::ResourceDependency >::produce_early(cx); }
146                })
147            }
148            _ => None,
149        })
150}
151
152fn derive_r_d_struct(name: &Ident, generics: &Generics, fields: &Fields) -> TokenStream {
153    const NO_FIELDS: &Punctuated<syn::Field, syn::token::Comma> = &Punctuated::new();
154    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
155    let dep_types = match fields {
156        Fields::Named(f) => &f.named,
157        Fields::Unnamed(f) => &f.unnamed,
158        Fields::Unit => NO_FIELDS,
159    }
160    .iter()
161    .map(|f| match find_dependency_type(&f.ty, &f.attrs) {
162        Ok(dty) => Ok(dty),
163        Err(span) => Err(quote_spanned! {
164            span => compile_error!("each field of a ResourceDependencies struct must have a type matching one of: Arc<T>, Option<Arc<T>>, Vec<Arc<dyn Tr>>, PhantomData<T>");
165        }),
166    })
167    .collect::<Vec<_>>();
168
169    let registrations = dep_types.iter().map(|r| match r {
170        Ok(DependencyType::Concrete(ty, _)) => quote! {
171            ::comprehensive::assembly::Registrar::< #ty >::register(cx);
172        },
173        Ok(DependencyType::Trait(ty, _)) => quote! {
174            cx.require_trait::< #ty >();
175        },
176        Ok(DependencyType::Weak(ty)) => quote! {
177            ::comprehensive::assembly::Registrar::< #ty >::register_without_dependency(cx);
178        },
179        Ok(DependencyType::NewStyle(ty)) => quote! {
180            < #ty as ::comprehensive::dependencies::ResourceDependency >::register(cx);
181        },
182        Err(ts) => ts.clone(),
183    });
184    // Produce all of the required dependencies, collecting the errors.
185    let productions1 = produce_concrete(&dep_types, false);
186    // Return if any failed.
187    let productions2 = dep_types.iter().enumerate().filter_map(|(i, r)| match r {
188        Ok(DependencyType::Concrete(_, false)) => {
189            let temp = format_ident!("dep_{}", i);
190            Some(quote! { let #temp = #temp ?; })
191        }
192        Ok(DependencyType::Trait(_, false)) => {
193            let temp = format_ident!("dep_{}", i);
194            Some(quote! { let #temp = #temp ?; })
195        }
196        Ok(DependencyType::NewStyle(_)) => {
197            let temp = format_ident!("dep_{}", i);
198            Some(quote! { let #temp = #temp ?; })
199        }
200        _ => None,
201    });
202    // Produce all of the optional dependencies.
203    let productions3 = produce_concrete(&dep_types, true);
204    let definition = match fields {
205        Fields::Named(f) => {
206            let elements =
207                f.named
208                    .iter()
209                    .zip(dep_types.iter())
210                    .enumerate()
211                    .map(|(i, (field, dt))| {
212                        let name = field.ident.as_ref().unwrap();
213                        match dt {
214                            Ok(DependencyType::Concrete(_, false)) => {
215                                let temp = format_ident!("dep_{}", i);
216                                quote! { #name: #temp , }
217                            }
218                            Ok(DependencyType::Concrete(_, true)) => {
219                                let temp = format_ident!("dep_{}", i);
220                                quote! { #name: #temp .ok(), }
221                            }
222                            Ok(DependencyType::Trait(_, _)) => {
223                                let temp = format_ident!("dep_{}", i);
224                                quote! { #name: #temp , }
225                            }
226                            Ok(DependencyType::Weak(_)) => {
227                                quote! { #name: ::std::marker::PhantomData, }
228                            }
229                            Ok(DependencyType::NewStyle(_)) => {
230                                let temp = format_ident!("dep_{}", i);
231                                quote! { #name: #temp ?, }
232                            }
233                            Err(ts) => ts.clone(),
234                        }
235                    });
236            quote! {
237                ::std::result::Result::Ok(Self { #( #elements )* })
238            }
239        }
240        Fields::Unnamed(_) => {
241            let elements = dep_types.iter().enumerate().map(|(i, dt)| match dt {
242                Ok(DependencyType::Concrete(_, false)) => {
243                    let temp = format_ident!("dep_{}", i);
244                    quote! { #temp , }
245                }
246                Ok(DependencyType::Concrete(_, true)) => {
247                    let temp = format_ident!("dep_{}", i);
248                    quote! { #temp .ok(), }
249                }
250                Ok(DependencyType::Trait(_, _)) => {
251                    let temp = format_ident!("dep_{}", i);
252                    quote! { #temp , }
253                }
254                Ok(DependencyType::Weak(_)) => {
255                    quote! { ::std::marker::PhantomData, }
256                }
257                Ok(DependencyType::NewStyle(_)) => {
258                    let temp = format_ident!("dep_{}", i);
259                    quote! { #temp ?, }
260                }
261                Err(ts) => ts.clone(),
262            });
263            quote! {
264                ::std::result::Result::Ok(Self ( #( #elements )* ))
265            }
266        }
267        Fields::Unit => quote! { ::std::result::Result::Ok(Self) },
268    };
269
270    quote! {
271        #[automatically_derived]
272        impl #impl_generics ::comprehensive::ResourceDependencies for #name #ty_generics #where_clause {
273            fn register(cx: &mut ::comprehensive::assembly::RegisterContext) {
274                #( #registrations )*
275            }
276
277            fn produce(cx: &mut ::comprehensive::assembly::ProduceContext) -> ::std::result::Result<Self, ::std::boxed::Box<dyn ::std::error::Error>> {
278                #( #productions1 )*
279                #( #productions2 )*
280                #( #productions3 )*
281                #definition
282            }
283        }
284    }
285}
286
287/// This macro should be used to derive the
288/// [`ResourceDependencies`](https://docs.rs/comprehensive/latest/comprehensive/assembly/trait.ResourceDependencies.html)
289/// trait for expressing dependencies between resources.
290///
291/// It takes a struct as input. The types of the fields of the struct should all
292/// match one of:
293///
294/// - [`Arc<T>`](std::sync::Arc) where T is a Resource. That resource will be a
295///   required dependency.
296/// - [`Option<Arc<T>>`](std::option::Option) where T is a Resource. That resource
297///   will be an optional dependency with the value being set no [`None`] if
298///   the dependency fails initialisation.
299/// - [`Vec<Arc<dyn T>>`](Vec) where T is a trait that might be implemented by
300///   some resources. All of the resources that exist in the graph and
301///   implement that trait and declare that they do so in their
302///   [`Resource`](https://docs.rs/comprehensive/latest/comprehensive/v1/trait.Resource.html)
303///   definition (using
304///   [`v1::resource`](https://docs.rs/comprehensive/latest/comprehensive/v1/attr.resource.html))
305///   will be collected here.
306///
307///   By default, if some resources matching the requested trait exist but
308///   fail initialisation, this will be considered an error: this set of
309///   dependencies will fail to construct and the error will be bubbled up.
310///   This mode is suitable for resources that wish to collect all available
311///   dependencies of a given type and not silently ignore a failing subset.
312///
313///   If a struct field of type [`Vec`] is annotated with `#[may_fail]`,
314///   then resources matching the requested trait exist but which fail
315///   initialisation will instead be dropped and a vector containing only
316///   the successful ones will be produced. This mode is suitable for
317///   resources that degrade well if some dependencies are not available
318///   (especially ones seeking just one working dependency resource from
319///   a set of possible ones.
320///
321///   Note that in order to be selected for this, a Resource has to exist
322///   somewhere in the graph already, which means that somewhere it must
323///   be declared as a dependency under its concrete name. For this, a
324///   common expected pattern is that the Assembly's top-level dependencies
325///   request it under its concrete type but otherwise do not make any use
326///   of it, enabling one or more resources elsewhere in the graph to
327///   discover it under its trait interface.
328/// - [`PhantomData<T>`](std::marker::PhantomData) where T is a Resource.
329///   The resource
330///   will be made available to the assembly, but no dependency on it is
331///   introduced in the graph. The resource will be actually included in
332///   the assembly only if something depends on it in some other way.
333///
334///   This is only useful if another resource depends upon `T` via a trait
335///   that it exposes. In that case, the
336///   [`PhantomData`](std::marker::PhantomData) dependency serves to import
337///   `T` so that it can be discovered.
338///
339/// See
340/// [`ResourceDependencies`](https://docs.rs/comprehensive/latest/comprehensive/assembly/trait.ResourceDependencies.html)
341/// for usage information.
342#[proc_macro_derive(ResourceDependencies, attributes(may_fail, old_style))]
343pub fn derive_resource_dependencies(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
344    let input: DeriveInput = parse_macro_input!(item);
345    match input.data {
346        Data::Struct(ref s) => derive_r_d_struct(&input.ident, &input.generics, &s.fields),
347        _ => quote_spanned! {
348            input.span() => compile_error!("`#[derive(ResourceDependencies)]` requires a struct");
349        },
350    }
351    .into()
352}
353
354fn derive_grpc_service_internal(
355    name: &Ident,
356    generics: &Generics,
357    attrs: &[Attribute],
358) -> Result<TokenStream, syn::Error> {
359    let mut implementation: Option<syn::Type> = None;
360    let mut service: Option<syn::Type> = None;
361    let mut descriptor: Option<syn::Expr> = None;
362    for attr in attrs {
363        if attr.path().is_ident("implementation") {
364            implementation = Some(attr.parse_args()?);
365        } else if attr.path().is_ident("service") {
366            service = Some(attr.parse_args()?);
367        } else if attr.path().is_ident("descriptor") {
368            descriptor = Some(attr.parse_args()?);
369        }
370    }
371    let Some(implementation) = implementation else {
372        return Ok(quote! {
373            compile_error!("`[#implementation(T)]` is required");
374        });
375    };
376    let Some(service) = service else {
377        return Ok(quote! {
378            compile_error!("`[#service(T)]` is required");
379        });
380    };
381    let descriptor_registration = match descriptor {
382        Some(d) => quote! {
383            d.server.register_encoded_file_descriptor_set( #d );
384        },
385        None => quote! {},
386    };
387
388    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
389    Ok(quote! {
390        #[automatically_derived]
391        impl #impl_generics ::comprehensive::Resource for #name #ty_generics #where_clause {
392            type Args = ::comprehensive::NoArgs;
393            type Dependencies = ::comprehensive_grpc::GrpcServiceDependencies< #implementation >;
394            const NAME: &str = ::comprehensive_grpc::const_format::concatcp!(
395                < #service ::< #implementation > as ::tonic::server::NamedService>::NAME,
396                " gRPC service"
397            );
398
399            fn new(
400                d: ::comprehensive_grpc::GrpcServiceDependencies< #implementation >,
401                _: ::comprehensive::NoArgs,
402            ) -> ::std::result::Result<Self, std::boxed::Box<dyn ::std::error::Error>> {
403                #descriptor_registration
404                d.server.add_service( #service ::from_arc(d.implementation))?;
405                Ok(Self)
406            }
407        }
408
409        #[automatically_derived]
410        impl #impl_generics ::comprehensive_grpc::GrpcService for #name #ty_generics #where_clause {}
411    })
412}
413
414/// This macro is obsolete after GrpcService was converted to expect
415/// [`Resource`](https://docs.rs/comprehensive/latest/comprehensive/v1/trait.Resource.html).
416#[proc_macro_derive(GrpcService, attributes(implementation, service, descriptor))]
417pub fn derive_grpc_service(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
418    let input: DeriveInput = parse_macro_input!(item);
419    derive_grpc_service_internal(&input.ident, &input.generics, &input.attrs)
420        .unwrap_or_else(|e| {
421            let e = e.to_compile_error();
422            quote! { #e }
423        })
424        .into()
425}
426
427fn get_str_lit_val(v: &syn::Expr) -> Result<LitStr, Span> {
428    let syn::Expr::Lit(exprlit) = v else {
429        return Err(v.span());
430    };
431    let syn::Lit::Str(ref litstr) = exprlit.lit else {
432        return Err(exprlit.lit.span());
433    };
434    Ok(litstr.clone())
435}
436
437fn is_router(f: &syn::Field) -> bool {
438    f.attrs.iter().any(|a| a.path().is_ident("router"))
439}
440
441fn derive_h_s_i(
442    name: &Ident,
443    data: &Data,
444    generics: &Generics,
445    attrs: &[Attribute],
446) -> Result<TokenStream, syn::Error> {
447    let mut flag_prefix: Option<LitStr> = None;
448    for attr in attrs {
449        if attr.path().is_ident("flag_prefix") {
450            flag_prefix = match get_str_lit_val(&attr.meta.require_name_value()?.value) {
451                Ok(prefix) => Some(prefix),
452                Err(span) => {
453                    return Ok(quote_spanned! {
454                        span => compile_error!("flag_prefix argument must be str literal");
455                    });
456                }
457            };
458        }
459    }
460    let Some(flag_prefix) = flag_prefix else {
461        return Ok(quote! {
462            compile_error!("`[#flag_prefix = \"foo_\"]` is required");
463        });
464    };
465    let Data::Struct(st) = data else {
466        return Ok(quote! {
467            compile_error!("`#[derive(HttpServingInstance)]` requires a struct");
468        });
469    };
470    let router_members: Vec<syn::Member> = match st.fields {
471        Fields::Named(ref f) => f
472            .named
473            .iter()
474            .filter_map(|field| {
475                if is_router(field) {
476                    Some(syn::Member::Named(field.ident.clone().unwrap()))
477                } else {
478                    None
479                }
480            })
481            .take(2)
482            .collect(),
483        Fields::Unnamed(ref f) => f
484            .unnamed
485            .iter()
486            .enumerate()
487            .filter_map(|(i, field)| {
488                if is_router(field) {
489                    Some(syn::Member::Unnamed(syn::Index {
490                        index: i as u32,
491                        span: field.span(),
492                    }))
493                } else {
494                    None
495                }
496            })
497            .take(2)
498            .collect(),
499        Fields::Unit => Vec::new(),
500    };
501    if router_members.len() != 1 {
502        return Ok(quote! {
503            compile_error!("exactly 1 struct field must be annotated with #[router]");
504        });
505    }
506    let router_member = router_members.first().unwrap();
507
508    let http_port_flag_name = format!("{}http-port", flag_prefix.value());
509    let http_port_flag_name_lit = LitStr::new(&http_port_flag_name, flag_prefix.span());
510    let http_bind_addr_flag_name = format!("{}http-bind-addr", flag_prefix.value());
511    let http_bind_addr_flag_name_lit = LitStr::new(&http_bind_addr_flag_name, flag_prefix.span());
512    let https_port_flag_name = format!("{}https-port", flag_prefix.value());
513    let https_port_flag_name_lit = LitStr::new(&https_port_flag_name, flag_prefix.span());
514    let https_bind_addr_flag_name = format!("{}https-bind-addr", flag_prefix.value());
515    let https_bind_addr_flag_name_lit = LitStr::new(&https_bind_addr_flag_name, flag_prefix.span());
516
517    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
518    Ok(quote! {
519        #[automatically_derived]
520        impl #impl_generics ::comprehensive_http::HttpServingInstance for #name #ty_generics #where_clause {
521            const HTTP_PORT_FLAG_NAME: &str = #http_port_flag_name_lit ;
522            const HTTP_BIND_ADDR_FLAG_NAME: &str = #http_bind_addr_flag_name_lit ;
523            const HTTPS_PORT_FLAG_NAME: &str = #https_port_flag_name_lit ;
524            const HTTPS_BIND_ADDR_FLAG_NAME: &str = #https_bind_addr_flag_name_lit ;
525
526            fn get_router(&self) -> ::axum::Router {
527                self. #router_member .clone()
528            }
529        }
530    })
531}
532
533#[proc_macro_derive(HttpServingInstance, attributes(flag_prefix, router))]
534pub fn derive_http_serving_instance(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
535    let input: DeriveInput = parse_macro_input!(item);
536    derive_h_s_i(&input.ident, &input.data, &input.generics, &input.attrs)
537        .unwrap_or_else(|e| {
538            let e = e.to_compile_error();
539            quote! { #e }
540        })
541        .into()
542}
543
544fn path_and_single_generic_type(ty: &Type) -> Result<(&Path, &Type), Span> {
545    let Type::Path(path) = ty else {
546        return Err(ty.span());
547    };
548    let Some(last) = path.path.segments.last() else {
549        return Err(path.path.segments.span());
550    };
551    let PathArguments::AngleBracketed(ref generics) = last.arguments else {
552        return Err(last.arguments.span());
553    };
554    if generics.args.len() != 1 {
555        return Err(generics.span());
556    }
557    let GenericArgument::Type(gty) = generics.args.first().unwrap() else {
558        return Err(generics.span());
559    };
560    Ok((&path.path, gty))
561}
562
563fn client_type(ty: &Type) -> Result<(bool, &Path), Span> {
564    let (path1, inner) = path_and_single_generic_type(ty)?;
565    let seg = &path1.segments;
566    if seg.len() == 1 {
567        let seg1 = &seg.first().unwrap().ident;
568        if *seg1 == Ident::new("Option", seg1.span()) {
569            let (path2, _) = path_and_single_generic_type(inner)?;
570            return Ok((true, path2));
571        }
572    }
573    Ok((false, path1))
574}
575
576fn derive_grpc_client_struct(
577    vis: &Visibility,
578    name: &Ident,
579    generics: &Generics,
580    fields: &Fields,
581    attrs: &[Attribute],
582) -> TokenStream {
583    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
584    let mut fields_it = fields.iter();
585    let client_field = fields_it.next().unwrap();
586    let cts = client_field.ty.span();
587
588    let mut propagate_health = true;
589    let mut deps = quote! { GRPCClientDependencies };
590    let mut defaults = quote! {};
591    for attr in attrs {
592        if attr.path().is_ident("no_propagate_health") {
593            propagate_health = false;
594        }
595        if attr.path().is_ident("no_tls") {
596            deps = quote! { GRPCClientDependenciesNoTls };
597        }
598        if let syn::Meta::List(l) = &attr.meta {
599            if l.path.is_ident("defaults") {
600                let tokens = &l.tokens;
601                defaults = quote! {
602                    fn instance_defaults() -> ::comprehensive_grpc::client::GrpcClientResourceDefaults {
603                        #tokens
604                    }
605                };
606            }
607        }
608        if attr.path().is_ident("defaults") {
609            deps = quote! { GRPCClientDependenciesNoTls };
610        }
611    }
612
613    let (is_option, client_type) = match client_type(&client_field.ty) {
614        Ok(v) => v,
615        Err(span) => {
616            return quote_spanned! {
617                span => compile_error!("First field of struct must be pb::client::Type<_> or Option<pb::client::Type<_>>");
618            };
619        }
620    };
621
622    let mut builder = client_type.clone();
623    if let Some(last) = builder.segments.last_mut() {
624        last.arguments = PathArguments::None;
625    }
626    builder
627        .segments
628        .push(Ident::new("with_origin", builder.segments.last().span()).into());
629
630    let name_str = name.to_string();
631    let name_lit = Lit::Str(LitStr::new(&name_str, name.span()));
632    let label = Lit::Str(LitStr::new(&name_str.to_case(Case::Snake), name.span()));
633    let required = Lit::Bool(LitBool::new(!is_option, client_type.span()));
634    let flag_prefix = format!("{}-", name_str.to_case(Case::Kebab));
635    let flag_prefix_span = name.span();
636    let flag_prefix = Lit::Str(LitStr::new(&flag_prefix, flag_prefix_span));
637
638    let (producer, cloner, client_return_type) = if is_option {
639        (
640            quote_spanned! { cts => param.map(|(stack, uri)| #builder (stack, uri)) },
641            quote! { as_ref().map(|c| c.clone()) },
642            quote_spanned! { cts => Option < #client_type > },
643        )
644    } else {
645        (
646            // unwrap okay because we made it a required arg in Clap
647            quote_spanned! { cts => { let (stack, uri) = param.unwrap(); #builder (stack, uri) } },
648            quote! { clone() },
649            client_type.to_token_stream(),
650        )
651    };
652    let worker_field = fields_it.next();
653    let (builder, get0, maybe_get1) = match client_field.ident {
654        None => (
655            if worker_field.is_some() {
656                quote_spanned! { fields.span() => Self( #producer , worker ) }
657            } else {
658                quote_spanned! { fields.span() => Self( #producer ) }
659            },
660            quote! { self.0 },
661            worker_field.map(|_| quote! { self.1 }),
662        ),
663        Some(ref client_field_name) => {
664            if let Some(f) = worker_field {
665                let worker_field_name = f.ident.as_ref().unwrap();
666                (
667                    quote_spanned! {
668                        fields.span() => Self {
669                            #client_field_name : #producer ,
670                            #worker_field_name : worker,
671                        }
672                    },
673                    quote! { self. #client_field_name },
674                    Some(quote! { self. #worker_field_name }),
675                )
676            } else {
677                (
678                    quote_spanned! {
679                        fields.span() => Self {
680                            #client_field_name : #producer ,
681                        }
682                    },
683                    quote! { self. #client_field_name },
684                    None,
685                )
686            }
687        }
688    };
689
690    let resource = if let Some(get1) = maybe_get1 {
691        quote! {
692            impl #impl_generics ::comprehensive::v0::Resource for #name #ty_generics #where_clause {
693                type Args = ::comprehensive_grpc::client::GrpcClientArgs<Self>;
694                type Dependencies = ::comprehensive_grpc::client:: #deps ;
695                const NAME: &'static str = #name_lit ;
696
697                fn new(d: ::comprehensive_grpc::client:: #deps , a: ::comprehensive_grpc::client::GrpcClientArgs<Self>) -> ::std::result::Result<Self, ::std::boxed::Box<dyn ::std::error::Error>> {
698                    let (param, worker) = ::comprehensive_grpc::client::new(a, #label , #propagate_health , d)?;
699                    Ok( #builder )
700                }
701
702                async fn run(&self) -> ::std::result::Result<(), ::std::boxed::Box<dyn ::std::error::Error>> {
703                    #get1 .go().await;
704                    Ok(())
705                }
706            }
707
708            impl #impl_generics ::comprehensive::AnyResource for #name #ty_generics #where_clause {
709                type Target = ::comprehensive::v0::ResourceProvider< #name #ty_generics >;
710            }
711        }
712    } else {
713        quote! {
714            impl #impl_generics ::comprehensive::v1::Resource for #name #ty_generics #where_clause {
715                type Args = ::comprehensive_grpc::client::GrpcClientArgs<Self>;
716                type Dependencies = ::comprehensive_grpc::client:: #deps ;
717                type CreationError = ::std::boxed::Box<dyn ::std::error::Error>;
718                const NAME: &'static str = #name_lit ;
719
720                fn new(
721                    d: ::comprehensive_grpc::client:: #deps ,
722                    a: ::comprehensive_grpc::client::GrpcClientArgs<Self>,
723                    api: &mut ::comprehensive::v1::AssemblyRuntime<'_>,
724                ) -> ::std::result::Result<::std::sync::Arc<Self>, ::std::boxed::Box<dyn ::std::error::Error>> {
725                    let (param, worker) = ::comprehensive_grpc::client::new(a, #label , #propagate_health , d)?;
726                    api.set_task(async move { worker.go().await; Ok(()) });
727                    Ok(::std::sync::Arc::new( #builder ))
728                }
729            }
730
731            impl #impl_generics ::comprehensive::AnyResource for #name #ty_generics #where_clause {
732                type Target = ::comprehensive::v1::ResourceProvider< #name #ty_generics >;
733            }
734        }
735    };
736
737    quote! {
738        #[automatically_derived]
739        impl #impl_generics ::comprehensive_grpc::client::InstanceDescriptor for #name #ty_generics #where_clause {
740            const REQUIRED: bool = #required ;
741            ::comprehensive_grpc::declare_client_flag_name_constants!( #flag_prefix );
742            #defaults
743        }
744
745        #[automatically_derived]
746        #resource
747
748        #[automatically_derived]
749        impl #impl_generics #name #ty_generics #where_clause {
750            #vis fn client(&self) -> #client_return_type {
751                #get0 . #cloner
752            }
753        }
754    }
755}
756
757/// Declare a resource for a gRPC client using a particular gRPC service to
758/// a particular backend.
759///
760/// Use this derive macro on a struct with a single field:
761/// a [`tonic`] gRPC client type, parameterised with [`Channel`].
762///
763/// The single field may be wrapped in an [`Option`].
764///   - If it is, then the client is considered optional and will
765///     be [`Some`] only if a URI for it is given on the command line.
766///   - If it is not, then the client is considered required and
767///     the program will fail at startup unless a URI for it is given.
768///
769/// ```
770/// # mod pb {
771/// #     pub mod test_client {
772/// #         #[derive(Clone)]
773/// #         pub struct TestClient<T>(std::marker::PhantomData<T>);
774/// #         impl<T> TestClient<T> {
775/// #             pub fn with_origin<U>(_: T, _: U) -> Self {
776/// #                 Self(std::marker::PhantomData)
777/// #             }
778/// #         }
779/// #     }
780/// # }
781/// use comprehensive_grpc::GrpcClient;
782/// use comprehensive_grpc::client::Channel;
783///
784/// #[derive(GrpcClient)]
785/// struct MyClientResource(
786///     pb::test_client::TestClient<Channel>,
787/// );
788/// ```
789///
790/// Normally, the health of the gRPC client will count toward the health of
791/// the [`Assembly`] as a whole. To prevent that, add `#[no_propagate_health]`.
792///
793/// The attribute `#[no_tls]` may be used to prevent the channel from
794/// supporting TLS even when the `tls` feature is enabled. That attribute
795/// is not expected to be widely useful and exists only to prevent a
796/// circular dependency. In this case, the struct field should referncen
797/// `ChannelNoTls` instead of `Channel`.
798///
799/// To customise the default values of the command line flags used to set the gRPC
800/// client channels' default parameters, add `#[defaults(foo)]` where `foo` is a
801/// block of code that evaluates to [`GrpcClientResourceDefaults`].
802///
803/// [`tonic`]: https://docs.rs/tonic/latest/tonic/
804/// [`Channel`]: https://docs.rs/comprehensive_grpc/latest/comprehensive_grpc/client/type.Channel.html
805/// [`GrpcClientResourceDefaults`]: https://docs.rs/comprehensive_grpc/latest/comprehensive_grpc/client/type.GrpcClientResourceDefaults.html
806/// [`Assembly`]: https://docs.rs/comprehensive/latest/comprehensive/assembly/struct.Assembly.html
807#[proc_macro_derive(GrpcClient, attributes(defaults, no_propagate_health, no_tls))]
808pub fn derive_grpc_client(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
809    let input: DeriveInput = parse_macro_input!(item);
810    match input.data {
811        Data::Struct(ref s) if s.fields.len() <= 2 => derive_grpc_client_struct(&input.vis, &input.ident, &input.generics, &s.fields, &input.attrs),
812        _ => quote_spanned! {
813            input.span() => compile_error!("`#[derive(GrpcClient)]` requires a struct with exactly 1 field (or 2, for backward compatibility");
814        },
815    }
816    .into()
817}
818
819fn type_unless_self_colon_colon(ty: &Type) -> Option<Type> {
820    let Type::Path(typ) = ty else {
821        return Some(ty.clone());
822    };
823    if typ
824        .path
825        .segments
826        .first()
827        .map(|s| s.ident == "Self")
828        .unwrap_or(false)
829    {
830        // The function signature is referencing the associated
831        // type so we cannot infer the associated type from the
832        // function signature.
833        None
834    } else {
835        Some(Type::Path(typ.clone()))
836    }
837}
838
839fn get_fnarg_type(arg: &syn::FnArg) -> Result<Option<Type>, TokenStream> {
840    match arg {
841        syn::FnArg::Typed(pat) => Ok(type_unless_self_colon_colon(&pat.ty)),
842        _ => Err(quote_spanned! {
843            arg.span() => compile_error!("expected a typed argument");
844        }),
845    }
846}
847
848fn bad_return_type(ty: &Type) -> TokenStream {
849    quote_spanned! {
850        ty.span() => compile_error!("expected return type Result<_, _>");
851    }
852}
853
854fn parse_v1resource_error_return_type(ty: &Type) -> Result<Option<Type>, TokenStream> {
855    let Type::Path(typ) = ty else {
856        return Err(bad_return_type(ty));
857    };
858    let Some(result) = typ.path.segments.last() else {
859        return Err(bad_return_type(ty));
860    };
861    let syn::PathArguments::AngleBracketed(ref args) = result.arguments else {
862        return Err(bad_return_type(ty));
863    };
864    if args.args.len() != 2 {
865        return Err(bad_return_type(ty));
866    }
867    let syn::GenericArgument::Type(ref err_ty) = args.args[1] else {
868        return Err(bad_return_type(ty));
869    };
870    Ok(type_unless_self_colon_colon(err_ty))
871}
872
873enum ExportType<A, B, C, D> {
874    General(A),
875    Grpc(B),
876    ProtoDescriptor(C),
877    NotOurs(D),
878}
879
880impl<A, B, C, D> ExportType<A, B, C, D> {
881    fn ours(&self) -> bool {
882        !matches!(self, Self::NotOurs(_))
883    }
884}
885
886#[proc_macro_attribute]
887pub fn v1resource(
888    _attr: proc_macro::TokenStream,
889    item: proc_macro::TokenStream,
890) -> proc_macro::TokenStream {
891    let mut block: syn::ItemImpl = parse_macro_input!(item);
892    let mut errors = Vec::new();
893
894    let mut dependencies = None;
895    let mut args = None;
896    let mut creation_error = None;
897
898    let mut name_already_specified = false;
899    let mut dependencies_already_specified = false;
900    let mut args_already_specified = false;
901    let mut creation_error_already_specified = false;
902
903    for item in &block.items {
904        match item {
905            syn::ImplItem::Fn(f) => {
906                if f.sig.ident == "new" {
907                    if f.sig.inputs.len() == 3 {
908                        match get_fnarg_type(&f.sig.inputs[0]) {
909                            Ok(Some(ty)) => {
910                                dependencies = Some(ty);
911                            }
912                            Ok(None) => (),
913                            Err(e) => {
914                                errors.push(e);
915                            }
916                        }
917                        match get_fnarg_type(&f.sig.inputs[1]) {
918                            Ok(Some(ty)) => {
919                                args = Some(ty);
920                            }
921                            Ok(None) => (),
922                            Err(e) => {
923                                errors.push(e);
924                            }
925                        }
926                    } else {
927                        errors.push(quote_spanned! {
928                            f.sig.inputs.span() => compile_error!("expected Resource::new to take exactly 3 arguments");
929                        });
930                    }
931                    match f.sig.output {
932                        syn::ReturnType::Type(_, ref ty) => {
933                            // Result<Arc<Self>, Self::CreationError>
934                            match parse_v1resource_error_return_type(ty) {
935                                Ok(maybe_error) => {
936                                    creation_error = maybe_error;
937                                }
938                                Err(e) => {
939                                    errors.push(e);
940                                }
941                            }
942                        }
943                        _ => {
944                            errors.push(quote_spanned! {
945                                f.sig.output.span() => compile_error!("expected a return type");
946                            });
947                        }
948                    }
949                }
950            }
951            syn::ImplItem::Const(ico) => {
952                if ico.ident == "NAME" {
953                    name_already_specified = true;
954                }
955            }
956            syn::ImplItem::Type(ity) => {
957                if ity.ident == "Dependencies" {
958                    dependencies_already_specified = true;
959                }
960                if ity.ident == "Args" {
961                    args_already_specified = true;
962                }
963                if ity.ident == "CreationError" {
964                    creation_error_already_specified = true;
965                }
966            }
967            _ => (),
968        }
969    }
970    if !name_already_specified {
971        let name = LitStr::new(
972            &block.self_ty.to_token_stream().to_string(),
973            block.self_ty.span(),
974        );
975        block.items.push(syn::ImplItem::Verbatim(quote! {
976            const NAME: &str = #name ;
977        }));
978    }
979    if !dependencies_already_specified {
980        if let Some(d) = dependencies {
981            block.items.push(syn::ImplItem::Verbatim(quote_spanned! {
982                d.span() => type Dependencies = #d ;
983            }));
984        }
985    }
986    if !args_already_specified {
987        if let Some(a) = args {
988            block.items.push(syn::ImplItem::Verbatim(quote_spanned! {
989                a.span() => type Args = #a ;
990            }));
991        }
992    }
993    if !creation_error_already_specified {
994        if let Some(e) = creation_error {
995            block.items.push(syn::ImplItem::Verbatim(quote_spanned! {
996                e.span() => type CreationError = #e ;
997            }));
998        }
999    }
1000    let (ours, not_ours): (Vec<_>, Vec<_>) = block
1001        .attrs
1002        .into_iter()
1003        .map(|a| {
1004            if matches!(a.style, syn::AttrStyle::Outer) {
1005                match a.meta {
1006                    syn::Meta::List(ref l) => {
1007                        if l.path.is_ident("export") {
1008                            ExportType::General(l.parse_args::<Type>())
1009                        } else if l.path.is_ident("export_grpc") {
1010                            ExportType::Grpc(l.parse_args::<Path>())
1011                        } else if l.path.is_ident("proto_descriptor") {
1012                            ExportType::ProtoDescriptor(l.parse_args::<syn::Expr>())
1013                        } else {
1014                            ExportType::NotOurs(a)
1015                        }
1016                    }
1017                    _ => ExportType::NotOurs(a),
1018                }
1019            } else {
1020                ExportType::NotOurs(a)
1021            }
1022        })
1023        .partition(|ono| ono.ours());
1024    block.attrs = not_ours
1025        .into_iter()
1026        .filter_map(|ono| match ono {
1027            ExportType::NotOurs(v) => Some(v),
1028            _ => None,
1029        })
1030        .collect();
1031    let mut grpc_exports = ours
1032        .iter()
1033        .filter_map(|ono| match ono {
1034            ExportType::Grpc(Ok(pa)) => Some(quote_spanned! {
1035                pa.span() => server.add_service( #pa ::from_arc(self))?;
1036            }),
1037            _ => None,
1038        })
1039        .peekable();
1040    let mut grpc_descriptors = ours
1041        .iter()
1042        .filter_map(|ono| match ono {
1043            ExportType::ProtoDescriptor(Ok(ex)) => Some(quote_spanned! {
1044                ex.span() => server.register_encoded_file_descriptor_set( #ex );
1045            }),
1046            _ => None,
1047        })
1048        .peekable();
1049    let (impl_generics, _, where_clause) = block.generics.split_for_impl();
1050    let self_ty = &block.self_ty;
1051    let grpc_derive = if grpc_exports.peek().is_some() || grpc_descriptors.peek().is_some() {
1052        quote! {
1053            #[automatically_derived]
1054            impl #impl_generics ::comprehensive_grpc::GrpcService for #self_ty #where_clause {
1055                fn add_to_server(
1056                    self: Arc<Self>,
1057                    server: &mut ::comprehensive_grpc::server::GrpcServiceAdder,
1058                ) -> Result<(), ::comprehensive_grpc::ComprehensiveGrpcError> {
1059                    #( #grpc_descriptors )*
1060                    #( #grpc_exports )*
1061                    Ok(())
1062                }
1063            }
1064        }
1065    } else {
1066        quote! {}
1067    };
1068    let mut exports = ours.into_iter().filter_map(|ono| match ono {
1069        ExportType::General(Ok(ty)) => Some(quote_spanned! {
1070            ty.span() => installer.offer(|s| ::std::sync::Arc::clone(s) as ::std::sync::Arc< #ty >);
1071        }),
1072        ExportType::General(Err(e)) => Some(e.to_compile_error()),
1073        ExportType::Grpc(Ok(pa)) => Some(quote_spanned! {
1074            pa.span() => installer.offer(|s| ::std::sync::Arc::clone(s) as ::std::sync::Arc<dyn ::comprehensive_grpc::GrpcService>);
1075        }),
1076        ExportType::Grpc(Err(e)) => Some(e.to_compile_error()),
1077        ExportType::ProtoDescriptor(Ok(_)) => None,
1078        ExportType::ProtoDescriptor(Err(e)) => Some(e.to_compile_error()),
1079        ExportType::NotOurs(_) => None,
1080    }).peekable();
1081    if exports.peek().is_some() {
1082        block.items.push(syn::ImplItem::Verbatim(quote! {
1083            fn provide_as_trait<'provide_as_trait>(installer: &'provide_as_trait mut ::comprehensive::v1::TraitInstaller<'_, 'provide_as_trait, '_, Self>) {
1084                #( #exports )*
1085            }
1086        }));
1087    }
1088    for e in errors {
1089        block.items.push(syn::ImplItem::Verbatim(e));
1090    }
1091    quote! {
1092        #block
1093
1094        #[automatically_derived]
1095        impl #impl_generics ::comprehensive::AnyResource for #self_ty #where_clause {
1096            type Target = ::comprehensive::v1::ResourceProvider< #self_ty >;
1097        }
1098
1099        #grpc_derive
1100    }
1101    .into()
1102}