clean up procedural macro

This commit is contained in:
steven-omaha
2023-01-12 23:27:32 +01:00
parent 28f95a42c7
commit e37e43c3d1
+79 -37
View File
@@ -1,39 +1,74 @@
use proc_macro::TokenStream; use proc_macro::TokenStream;
use quote::quote; use quote::quote;
use syn::DeriveInput; use syn::DeriveInput;
use syn::__private::TokenStream2;
#[proc_macro_derive(Register)] #[proc_macro_derive(Register)]
pub fn register(input: TokenStream) -> TokenStream { pub fn register(input: TokenStream) -> TokenStream {
let input = syn::parse::<DeriveInput>(input).unwrap(); let input = syn::parse::<DeriveInput>(input).unwrap();
let name = &input.ident; let name = &input.ident;
let enum_data = if let syn::Data::Enum(enum_data) = &input.data { let enum_data = if let syn::Data::Enum(enum_data) = &input.data {
enum_data enum_data
} else { } else {
panic!("`Register` can only be used on enums"); panic!("`Register` can only be used on enums");
}; };
let first_variant = &enum_data.variants[0].ident; let first_variant = &enum_data.variants[0].ident;
let variant_matches_backend = enum_data.variants.iter().map(|variant| { let variant_matches_backend = generate_variant_backend(enum_data);
let variant_name = &variant.ident; let variant_matches_next = generate_variant_matches_next(enum_data);
quote! { let variant_imports = generate_variant_imports(enum_data);
Self::#variant_name => Box::new(#variant_name::new()),
}
});
let variant_matches_next = enum_data.variants.iter().enumerate().map(|(i, variant)| { let expanded = compile_output(
let variant_name = &variant.ident; name,
if i == enum_data.variants.len() - 1 { first_variant,
quote! { variant_matches_backend,
Self::#variant_name => None, variant_matches_next,
} variant_imports,
} else { );
let next_variant = &enum_data.variants[i + 1].ident;
quote! {
Self::#variant_name => Some(Self::#next_variant),
}
}
});
TokenStream::from(expanded)
}
fn compile_output<T, U, V>(
name: &syn::Ident,
first_variant: &syn::Ident,
variant_backend: T,
variant_next: U,
variant_imports: V,
) -> TokenStream2
where
T: Iterator<Item = TokenStream2>,
U: Iterator<Item = TokenStream2>,
V: Iterator<Item = TokenStream2>,
{
let expanded = quote! {
#(#variant_imports)*
impl #name {
pub fn iter() -> BackendIter {
BackendIter {
next: Some(Self::#first_variant),
}
}
fn get_backend(&self) -> Box<dyn Backend> {
match self {
#(#variant_backend)*
}
}
fn next(&self) -> Option<Self> {
match self {
#(#variant_next)*
}
}
}
};
expanded
}
fn generate_variant_imports(enum_data: &syn::DataEnum) -> impl Iterator<Item = TokenStream2> + '_ {
let variant_imports = enum_data.variants.iter().map(|variant| { let variant_imports = enum_data.variants.iter().map(|variant| {
let variant_name = &variant.ident; let variant_name = &variant.ident;
// let variant_module = variant_name.clone(); // let variant_module = variant_name.clone();
@@ -48,27 +83,34 @@ pub fn register(input: TokenStream) -> TokenStream {
} }
}); });
variant_imports
}
let expanded = quote! { fn generate_variant_matches_next(
#(#variant_imports)* enum_data: &syn::DataEnum,
) -> impl Iterator<Item = TokenStream2> + '_ {
impl #name { let variant_matches_next = enum_data.variants.iter().enumerate().map(|(i, variant)| {
pub fn iter() -> BackendIter { let variant_name = &variant.ident;
BackendIter { if i == enum_data.variants.len() - 1 {
next: Some(Self::#first_variant), quote! {
Self::#variant_name => None,
} }
} else {
let next_variant = &enum_data.variants[i + 1].ident;
quote! {
Self::#variant_name => Some(Self::#next_variant),
} }
fn get_backend(&self) -> Box<dyn Backend> {
match self {
#(#variant_matches_backend)*
} }
} });
fn next(&self) -> Option<Self> { variant_matches_next
match self { }
#(#variant_matches_next)*
} fn generate_variant_backend(enum_data: &syn::DataEnum) -> impl Iterator<Item = TokenStream2> + '_ {
} let variant_matches_backend = enum_data.variants.iter().map(|variant| {
let variant_name = &variant.ident;
quote! {
Self::#variant_name => Box::new(#variant_name::new()),
} }
}; });
TokenStream::from(expanded) variant_matches_backend
} }