Skip to main content

abstract_getters_derive/
lib.rs

1#![allow(clippy::needless_doctest_main)]
2#![doc = include_str!(concat!(env!("CARGO_MANIFEST_DIR"), "/README.md"))]
3#![warn(missing_docs)]
4
5use convert_case::Casing;
6use proc_macro::TokenStream;
7use quote::{ToTokens, quote};
8use syn::{
9    Data, DeriveInput, Fields, Generics, Ident, Index, Type, parse_macro_input, spanned::Spanned,
10};
11
12/// Derives the [Getters](abstract_getters::Getters) trait for a struct.
13#[proc_macro_derive(Getters)]
14pub fn derive_getters(input: TokenStream) -> TokenStream {
15    let input = parse_macro_input!(input as DeriveInput);
16    let name = input.ident;
17    let generics = input.generics;
18    let struct_mod_ident = Ident::new(
19        &name.to_string().to_case(convert_case::Case::Snake),
20        name.span(),
21    );
22
23    let field_impls = match input.data {
24        Data::Struct(data_struct) => {
25            generate_for_fields(&name, &generics, data_struct.fields, None)
26        }
27        Data::Enum(data_enum) => {
28            let variants = data_enum.variants.into_iter().map(|variant| {
29                let variant_module = Ident::new(
30                    &variant.ident.to_string().to_case(convert_case::Case::Snake),
31                    variant.ident.span(),
32                );
33                let field_impls =
34                    generate_for_fields(&name, &generics, variant.fields, Some(&variant.ident));
35                quote! {
36                    pub mod #variant_module {
37                        use super::*;
38                        #field_impls
39                    }
40                }
41            });
42            quote! {
43                #(#variants)*
44            }
45        }
46        Data::Union(union_data) => syn::Error::new(
47            union_data.union_token.span(),
48            "Getters cannot be derived for unions",
49        )
50        .to_compile_error(),
51    };
52
53    let expanded = quote! {
54        pub mod #struct_mod_ident {
55            use super::*;
56            #field_impls
57        }
58    };
59
60    TokenStream::from(expanded)
61}
62
63fn generate_for_fields(
64    name: &Ident,
65    generics: &Generics,
66    fields: Fields,
67    enum_variant_name: Option<&Ident>,
68) -> proc_macro2::TokenStream {
69    let field_impls_iter =
70        match fields {
71            Fields::Named(fields_named) => &mut fields_named.named.into_iter().map(|field| {
72                let field_ident = field.ident.expect("A named field");
73                generate_for_field(
74                    field_ident.clone(),
75                    field_ident,
76                    name,
77                    field.ty,
78                    generics,
79                    enum_variant_name,
80                )
81            }) as &mut dyn Iterator<Item = _>,
82
83            Fields::Unnamed(fields_unnamed) => &mut fields_unnamed
84                .unnamed
85                .into_iter()
86                .enumerate()
87                .map(|(index, field)| {
88                    let field_struct = Ident::new(&format!("_{index}"), field.span());
89                    let field_index = Index::from(index);
90                    generate_for_field(
91                        field_struct,
92                        field_index,
93                        name,
94                        field.ty,
95                        generics,
96                        enum_variant_name,
97                    )
98                }) as &mut dyn Iterator<Item = _>,
99
100            _ => &mut std::iter::empty() as &mut dyn Iterator<Item = _>,
101        };
102
103    quote! {
104        #(#field_impls_iter)*
105    }
106}
107
108/// Generate an owned, mutable and referential impl for a field by generation a struct with the field's name
109/// and implementing the [Field](abstract_getters::Field) trait for it.
110fn generate_for_field<N: ToTokens>(
111    field_struct: Ident,
112    field_name: N,
113    struct_name: &Ident,
114    ty: Type,
115    generics: &Generics,
116    enum_variant_name: Option<&Ident>,
117) -> proc_macro2::TokenStream {
118    let struct_params = &generics.params;
119    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
120
121    if let Some(variant_name) = enum_variant_name {
122        quote! {
123            #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
124            #[allow(non_camel_case_types)]
125            pub struct #field_struct;
126            impl #impl_generics abstract_getters::Field<#field_struct> for #struct_name #ty_generics #where_clause {
127                type Type = Option<#ty>;
128                fn field(self) -> <Self as abstract_getters::Field<#field_struct>>::Type {
129                    match self {
130                        Self::#variant_name{#field_name: __get_field, ..} => Some(__get_field),
131                        _ => None,
132                    }
133                }
134            }
135
136            impl <'__top_level, #struct_params> abstract_getters::Field<#field_struct> for &'__top_level #struct_name #ty_generics #where_clause {
137                type Type = Option<&'__top_level #ty>;
138                fn field(self) -> <Self as abstract_getters::Field<#field_struct>>::Type {
139                    match self {
140                        #struct_name::#variant_name{#field_name: __get_field, ..} => Some(__get_field),
141                        _ => None,
142                    }
143                }
144            }
145            impl <'__top_level, #struct_params> abstract_getters::Field<#field_struct> for &'__top_level mut #struct_name #ty_generics #where_clause {
146                type Type = Option<&'__top_level mut #ty>;
147                fn field(self) -> <Self as abstract_getters::Field<#field_struct>>::Type {
148                    match self {
149                        #struct_name::#variant_name{#field_name: __get_field, ..} => Some(__get_field),
150                        _ => None,
151                    }
152                }
153            }
154        }
155    } else {
156        quote! {
157            #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
158            #[allow(non_camel_case_types)]
159            pub struct #field_struct;
160            impl #impl_generics abstract_getters::Field<#field_struct> for #struct_name #ty_generics #where_clause {
161                type Type = #ty;
162                fn field(self) -> <Self as abstract_getters::Field<#field_struct>>::Type { self.#field_name }
163            }
164            impl<'__top_level, #struct_params> abstract_getters::Field<#field_struct> for &'__top_level #struct_name #ty_generics #where_clause {
165                type Type = &'__top_level #ty;
166                fn field(self) -> <Self as abstract_getters::Field<#field_struct>>::Type { &self.#field_name }
167            }
168            impl<'__top_level, #struct_params> abstract_getters::Field<#field_struct> for &'__top_level mut #struct_name #ty_generics #where_clause {
169                type Type = &'__top_level mut #ty;
170                fn field(self) -> <Self as abstract_getters::Field<#field_struct>>::Type { &mut self.#field_name }
171            }
172        }
173    }
174}