enums.rsannotatedenums.rssource244 lines · 7.0 KB · raw
1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{Attribute, Data, DeriveInput, Fields};
4
5pub fn enum_display(input: TokenStream) -> TokenStream {
6    let ast: DeriveInput = syn::parse(input).expect("unable to parse input");
7
8    let name = &ast.ident;
9
10    let Data::Enum(data) = &ast.data else {
11        panic!("EnumDisplay derive macro can only be applied to enums; {name} is not an enum");
12    };
13
14    let match_arms: Vec<_> = data
15        .variants
16        .iter()
17        .map(|variant| {
18            let variant_name = &variant.ident;
19            assert!(variant.fields.is_empty(), "EnumDisplay macro only supports enums with only fieldless variants; {name}::{variant_name} has fields");
20
21            let variant_name_str = variant_name.to_string();
22            quote! {
23                Self::#variant_name => #variant_name_str
24            }
25        })
26        .collect();
27
28    let expanded = quote! {
29        impl #name {
30            pub fn to_str(&self) -> &'static str {
31                match self {
32                    #(#match_arms,)*
33                }
34            }
35        }
36
37        impl ::std::fmt::Display for #name {
38            fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result {
39                ::std::write!(f, "{}", self.to_str())
40            }
41        }
42    };
43
44    expanded.into()
45}
46
47pub fn enum_from_str(input: TokenStream) -> TokenStream {
48    let ast: DeriveInput = syn::parse(input).expect("unable to parse input");
49
50    let name = &ast.ident;
51
52    let Data::Enum(data) = &ast.data else {
53        panic!("EnumFromStr derive macro can only be applied to enums; {name} is not an enum");
54    };
55
56    let match_arms: Vec<_> = data
57        .variants
58        .iter()
59        .map(|variant| {
60            let variant_name = &variant.ident;
61            assert!(variant.fields.is_empty(), "EnumFromStr macro only supports enums with only fieldless variants; {name}::{variant_name} has fields");
62
63            let variant_name_lowercase = variant_name.to_string().to_ascii_lowercase();
64            quote! {
65                #variant_name_lowercase => ::std::result::Result::Ok(Self::#variant_name)
66            }
67        })
68        .collect();
69
70    let err_fmt_string = format!("invalid {name} string: '{{}}'");
71    let expanded = quote! {
72        impl ::std::str::FromStr for #name {
73            type Err = ::std::string::String;
74
75            fn from_str(s: &str) -> ::std::result::Result<Self, Self::Err> {
76                match s.to_ascii_lowercase().as_str() {
77                    #(#match_arms,)*
78                    _ => ::std::result::Result::Err(::std::format!(#err_fmt_string, s))
79                }
80            }
81        }
82    };
83
84    expanded.into()
85}
86
87pub fn enum_all(input: TokenStream) -> TokenStream {
88    let input: DeriveInput = syn::parse(input).expect("Unable to parse input");
89
90    let type_ident = &input.ident;
91    let Data::Enum(data) = &input.data else {
92        panic!("EnumAll only supports enums; {type_ident} is not an enum");
93    };
94
95    let variant_constructors: Vec<_> = data.variants.iter().map(|variant| {
96        let variant_ident = &variant.ident;
97        assert!(
98            matches!(variant.fields, Fields::Unit),
99            "EnumAll only supports enums with fieldless variants; {type_ident}::{variant_ident} is not a fieldless variant",
100        );
101
102        quote! {
103            Self::#variant_ident
104        }
105    }).collect();
106
107    let num_variants = data.variants.len();
108    let expanded = quote! {
109        impl #type_ident {
110            pub const ALL: [Self; #num_variants] = [#(#variant_constructors,)*];
111        }
112    };
113
114    expanded.into()
115}
116
117pub fn custom_value_enum(input: TokenStream) -> TokenStream {
118    let input: DeriveInput = syn::parse(input).expect("Unable to parse input");
119
120    let type_ident = &input.ident;
121
122    let Data::Enum(data) = &input.data else {
123        panic!("CustomValueEnum only supports enums");
124    };
125
126    let included_fields: Vec<_> = data
127        .variants
128        .iter()
129        .filter_map(|variant| {
130            let skip = value_enum_should_skip(&variant.attrs);
131            (!skip).then_some(&variant.ident)
132        })
133        .collect();
134
135    let expanded = quote! {
136        impl ::clap::ValueEnum for #type_ident {
137            fn value_variants<'a>() -> &'a [Self] {
138                const ALL: &[#type_ident] = &[
139                    #(#type_ident::#included_fields,)*
140                ];
141
142                ALL
143            }
144
145            fn to_possible_value(&self) -> ::std::option::Option<::clap::builder::PossibleValue> {
146                match self {
147                    #(
148                        Self::#included_fields => ::std::option::Option::Some(
149                            ::clap::builder::PossibleValue::new(self.to_str())
150                        ),
151                    )*
152                    _ => ::std::option::Option::None
153                }
154            }
155        }
156    };
157
158    expanded.into()
159}
160
161fn value_enum_should_skip(attrs: &[Attribute]) -> bool {
162    attrs.iter().any(|attr| {
163        if !attr.path().is_ident("value_enum") {
164            return false;
165        }
166
167        let mut skip = false;
168        attr.parse_nested_meta(|meta| {
169            if meta.path.is_ident("skip") {
170                skip = true;
171                Ok(())
172            } else {
173                Err(meta.error("Invalid value_enum meta"))
174            }
175        })
176        .expect("Failed to parse value_enum attribute");
177
178        skip
179    })
180}
181
182pub fn match_each_variant_macro(input: TokenStream) -> TokenStream {
183    let input: DeriveInput = syn::parse(input).expect("Unable to parse input");
184
185    let ident = &input.ident;
186
187    let Data::Enum(data) = &input.data else {
188        panic!("{ident} is not an enum");
189    };
190
191    let match_arms: Vec<_> = data
192        .variants
193        .iter()
194        .map(|variant| {
195            let variant_ident = &variant.ident;
196
197            let Fields::Unnamed(fields) = &variant.fields else {
198                panic!("{ident}::{variant_ident} should have unnamed fields");
199            };
200
201            assert_eq!(
202                fields.unnamed.len(),
203                1,
204                "{ident}::{variant_ident} has {} unnamed fields, expected 1",
205                fields.unnamed.len()
206            );
207
208            quote! {
209                #ident::#variant_ident($field) => $match_arm
210            }
211        })
212        .collect();
213
214    let variant_match_arms: Vec<_> = data
215        .variants
216        .iter()
217        .map(|variant| {
218            let variant_ident = &variant.ident;
219
220            // No need to re-validate the enum variants, that was done above
221
222            quote! {
223                #ident::#variant_ident($field) => #ident::#variant_ident($match_arm)
224            }
225        })
226        .collect();
227
228    let expanded = quote! {
229        macro_rules! match_each_variant {
230            ($value:expr, $field:ident => $match_arm:expr) => {
231                match $value {
232                    #(#match_arms,)*
233                }
234            };
235            ($value:expr, $field:ident => :variant($match_arm:expr)) => {
236                match $value {
237                    #(#variant_match_arms,)*
238                }
239            };
240        }
241    };
242
243    expanded.into()
244}