Skip to main content

enum_table_derive/
lib.rs

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        // SAFETY: `#variant_idents` covers every variant of `#ident` exactly once;
49        // unit variants and the absence of `#[repr(align(N))]` (checked above) rule
50        // out padding; and `sort_variants` sorts `VARIANTS` by unsigned bit-pattern.
51        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                            // SAFETY: see the `unsafe impl` block above.
61                            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}