abstract_getters_derive/
lib.rs1#![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#[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
108fn 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}