1#[proc_macro_derive(Enumerable)]
2pub fn derive_enumerable(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
3 derive_enumerable_internal(syn::parse_macro_input!(input as syn::DeriveInput))
4 .unwrap_or_else(syn::Error::into_compile_error)
5 .into()
6}
7
8fn derive_enumerable_internal(input: syn::DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
9 let syn::Data::Enum(data_enum) = &input.data else {
10 return Err(syn::Error::new_spanned(
11 &input,
12 "Enumerable can only be derived for enums",
13 ));
14 };
15
16 if !input.generics.params.is_empty() {
17 return Err(syn::Error::new_spanned(
18 &input.generics,
19 "Enumerable cannot be derived for generic enums",
20 ));
21 }
22
23 if repr_align(&input.attrs)?.is_some() {
24 return Err(syn::Error::new_spanned(
25 &input,
26 "Enumerable cannot be derived for enums with `#[repr(align(N))]`: alignment \
27 padding could be read as uninitialized memory by this crate's byte-level \
28 comparisons",
29 ));
30 }
31
32 let variant_idents = data_enum
33 .variants
34 .iter()
35 .map(|v| {
36 if !matches!(v.fields, syn::Fields::Unit) {
37 return Err(syn::Error::new_spanned(
38 &v.fields,
39 "Enumerable can only be derived for unit variants",
40 ));
41 }
42 Ok(&v.ident)
43 })
44 .collect::<syn::Result<Vec<_>>>()?;
45
46 let ident = &input.ident;
47 let expanded = quote::quote! {
48 unsafe impl enum_table::Enumerable for #ident {
52 const VARIANTS: &'static [#ident] = &unsafe {
53 enum_table::__private::sort_variants([#(Self::#variant_idents),*])
54 };
55
56 fn variant_index(&self) -> usize {
57 match *self {
58 #(
59 Self::#variant_idents => const {
60 unsafe {
62 enum_table::__private::variant_index_of(&#ident::#variant_idents, <#ident as enum_table::Enumerable>::VARIANTS)
63 }
64 },
65 )*
66 }
67 }
68 }
69 };
70
71 Ok(expanded)
72}
73
74fn repr_align(attrs: &[syn::Attribute]) -> syn::Result<Option<usize>> {
75 let mut align = None::<usize>;
76 for attr in attrs {
77 if attr.path().is_ident("repr") {
78 attr.parse_nested_meta(|meta| {
79 if meta.path.is_ident("align") {
80 let content;
81 syn::parenthesized!(content in meta.input);
82 let lit: syn::LitInt = content.parse()?;
83 let n: usize = lit.base10_parse()?;
84 align = Some(n);
85 }
86 Ok(())
87 })?;
88 }
89 }
90 Ok(align)
91}