diff options
Diffstat (limited to 'engine-macros/src/reflection/enum_impl.rs')
| -rw-r--r-- | engine-macros/src/reflection/enum_impl.rs | 334 |
1 files changed, 178 insertions, 156 deletions
diff --git a/engine-macros/src/reflection/enum_impl.rs b/engine-macros/src/reflection/enum_impl.rs index 14d91a7..96efd71 100644 --- a/engine-macros/src/reflection/enum_impl.rs +++ b/engine-macros/src/reflection/enum_impl.rs @@ -1,10 +1,12 @@ +use proc_macro2::TokenStream; use quote::{format_ident, quote}; use crate::reflection::default_value::gen_get_default_value_fn; use crate::reflection::options_attr::OptionsAttr; use crate::util::find_engine_crate_path; -pub fn generate(input: syn::ItemEnum, options: OptionsAttr) -> proc_macro2::TokenStream +pub fn generate(input: &syn::ItemEnum, options: &OptionsAttr) + -> proc_macro2::TokenStream { let engine_crate_path = find_engine_crate_path().unwrap(); @@ -35,7 +37,7 @@ pub fn generate(input: syn::ItemEnum, options: OptionsAttr) -> proc_macro2::Toke let impls: &mut dyn Iterator<Item = proc_macro2::TokenStream> = if input.generics.params.is_empty() { &mut [generate_impls( - &input, + input, None, is_unit_only, &variant_lookup_match_arms, @@ -45,7 +47,7 @@ pub fn generate(input: syn::ItemEnum, options: OptionsAttr) -> proc_macro2::Toke } else { &mut options.impl_with_generics.iter().map(|generic_args| { generate_impls( - &input, + input, Some(generic_args), is_unit_only, &variant_lookup_match_arms, @@ -69,7 +71,7 @@ fn generate_impls( { let get_default_value_fn = gen_get_default_value_fn(&input.ident, generic_args); - let variants = generate_variants(&input.variants, &engine_crate_path); + let variants = generate_variants(&input.variants, engine_crate_path); let generics_type_aliases = input .generics @@ -77,9 +79,8 @@ fn generate_impls( .zip( generic_args .iter() - .map(|generic_args| &generic_args.args) - .flatten() - .flat_map(|generic_arg| match generic_arg { + .flat_map(|generic_args| &generic_args.args) + .filter_map(|generic_arg| match generic_arg { syn::GenericArgument::Type(ty) => Some(ty), _ => None, }), @@ -165,43 +166,7 @@ fn generate_variants<'a>( let variant_name = syn::LitStr::new(&variant.ident.to_string(), variant.ident.span()); - let fields = match &variant.fields { - syn::Fields::Unit => quote! { None }, - syn::Fields::Named(named_fields) => { - let fields = named_fields.named.iter().enumerate().map( - |(variant_field_index, variant_field)| { - generate_enum_variant_field( - variant_field, - variant_field_index, - &engine_crate_path, - ) - }, - ); - - quote! { - Some(#engine_crate_path::reflection::EnumVariantFields::Named { - fields: &[#(#fields),*] - }) - } - } - syn::Fields::Unnamed(unnamed_fields) => { - let fields = unnamed_fields.unnamed.iter().enumerate().map( - |(variant_field_index, variant_field)| { - generate_enum_variant_field( - variant_field, - variant_field_index, - &engine_crate_path, - ) - }, - ); - - quote! { - Some(#engine_crate_path::reflection::EnumVariantFields::Unnamed { - fields: &[#(#fields),*] - }) - } - } - }; + let fields = gen_variant_fields(variant, engine_crate_path); let variant_field_vars = variant.fields.iter().enumerate().map( |(variant_field_index, variant_field)| { @@ -227,9 +192,7 @@ fn generate_variants<'a>( syn::Fields::Unit => { let variant_ident = &variant.ident; - quote! { - Self::#variant_ident - } + quote! { Self::#variant_ident } } syn::Fields::Named(named_fields) => { let variant_ident = &variant.ident; @@ -241,9 +204,7 @@ fn generate_variants<'a>( let field_var_ident = format_ident!("field_{field_ident}"); - quote! { - #field_ident: *#field_var_ident - } + quote! { #field_ident: *#field_var_ident } }); quote! { @@ -259,74 +220,20 @@ fn generate_variants<'a>( |(field_index, _)| { let field_var_ident = format_ident!("field_{field_index}"); - quote! { - *#field_var_ident - } + quote! { *#field_var_ident } }, ); quote! { - Self::#variant_ident( - #(#fields),* - ) + Self::#variant_ident(#(#fields),*) } } }; - let variant_ident = &variant.ident; - - let variant_pattern = match &variant.fields { - syn::Fields::Unit => quote! { Self::#variant_ident }, - syn::Fields::Named(syn::FieldsNamed { named: named_fields, .. }) => { - let field_idents = named_fields.iter().map(|field| { - let Some(field_ident) = &field.ident else { - unreachable!(); - }; - - field_ident - }); - - quote! { Self::#variant_ident { #(#field_idents),* } } - } - syn::Fields::Unnamed(syn::FieldsUnnamed { - unnamed: unnamed_fields, .. - }) => { - let field_idents = unnamed_fields - .iter() - .enumerate() - .map(|(field_index, _)| format_ident!("field_{field_index}")); - - quote! { Self::#variant_ident(#(#field_idents),*) } - } - }; - - let field_index_match_arms = match &variant.fields { - syn::Fields::Unit => quote! {}, - syn::Fields::Named(syn::FieldsNamed { named: named_fields, .. }) => { - let match_arms = - named_fields.iter().enumerate().map(|(field_index, field)| { - let Some(field_ident) = &field.ident else { - unreachable!(); - }; - - quote! { #field_index => { #field_ident } } - }); - - quote! { #(#match_arms)* } - } - syn::Fields::Unnamed(syn::FieldsUnnamed { - unnamed: unnamed_fields, .. - }) => { - let match_arms = - unnamed_fields.iter().enumerate().map(|(field_index, _)| { - let field_ident = format_ident!("field_{field_index}"); - - quote! { #field_index => { #field_ident } } - }); - - quote! { #(#match_arms)* } - } - }; + let VariantTryGetFieldFns { + immutable: try_get_field_immutable_fn, + mutable: try_get_field_mutable_fn, + } = gen_variant_try_get_field_functions(variant, engine_crate_path); quote! { #engine_crate_path::reflection::EnumVariant { @@ -352,52 +259,8 @@ fn generate_variants<'a>( Ok(()) }, - try_get_field: |target, field_index| { { - #![allow(unreachable_code)] - - use std::any::Any; - use #engine_crate_path::reflection::GetError; - - let target = target - .downcast_ref::<Self>() - .ok_or(GetError::WrongTargetType)?; - - let #variant_pattern = target else { - return Err(GetError::WrongTargetEnumVariant); - }; - - let field: &dyn Any = match field_index { - #field_index_match_arms - _ => { - return Err(GetError::IndexOutOfBounds); - } - }; - - Ok(field) - } }, - try_get_field_mut: |target, field_index| { { - #![allow(unreachable_code)] - - use std::any::Any; - use #engine_crate_path::reflection::GetError; - - let target = target - .downcast_mut::<Self>() - .ok_or(GetError::WrongTargetType)?; - - let #variant_pattern = target else { - return Err(GetError::WrongTargetEnumVariant); - }; - - let field: &mut dyn Any = match field_index { - #field_index_match_arms - _ => { - return Err(GetError::IndexOutOfBounds); - } - }; - - Ok(field) - } }, + try_get_field: #try_get_field_immutable_fn, + try_get_field_mut: #try_get_field_mutable_fn, } } }) @@ -483,3 +346,162 @@ fn gen_get_optional_type_reflection( .field_type_reflection() } } + +fn gen_variant_fields( + variant: &syn::Variant, + engine_crate_path: &syn::Path, +) -> TokenStream +{ + match &variant.fields { + syn::Fields::Unit => quote! { None }, + syn::Fields::Named(named_fields) => { + let fields = named_fields.named.iter().enumerate().map( + |(variant_field_index, variant_field)| { + generate_enum_variant_field( + variant_field, + variant_field_index, + engine_crate_path, + ) + }, + ); + + quote! { + Some(#engine_crate_path::reflection::EnumVariantFields::Named { + fields: &[#(#fields),*] + }) + } + } + syn::Fields::Unnamed(unnamed_fields) => { + let fields = unnamed_fields.unnamed.iter().enumerate().map( + |(variant_field_index, variant_field)| { + generate_enum_variant_field( + variant_field, + variant_field_index, + engine_crate_path, + ) + }, + ); + + quote! { + Some(#engine_crate_path::reflection::EnumVariantFields::Unnamed { + fields: &[#(#fields),*] + }) + } + } + } +} + +struct VariantTryGetFieldFns +{ + immutable: TokenStream, + mutable: TokenStream, +} + +fn gen_variant_try_get_field_functions( + variant: &syn::Variant, + engine_crate_path: &syn::Path, +) -> VariantTryGetFieldFns +{ + let field_index_match_arms = match &variant.fields { + syn::Fields::Unit => quote! {}, + syn::Fields::Named(syn::FieldsNamed { named: named_fields, .. }) => { + let match_arms = + named_fields.iter().enumerate().map(|(field_index, field)| { + let Some(field_ident) = &field.ident else { + unreachable!(); + }; + + quote! { #field_index => { #field_ident } } + }); + + quote! { #(#match_arms)* } + } + syn::Fields::Unnamed(syn::FieldsUnnamed { unnamed: unnamed_fields, .. }) => { + let match_arms = unnamed_fields.iter().enumerate().map(|(field_index, _)| { + let field_ident = format_ident!("field_{field_index}"); + + quote! { #field_index => { #field_ident } } + }); + + quote! { #(#match_arms)* } + } + }; + + let variant_ident = &variant.ident; + + let variant_pattern = match &variant.fields { + syn::Fields::Unit => quote! { Self::#variant_ident }, + syn::Fields::Named(syn::FieldsNamed { named: named_fields, .. }) => { + let field_idents = named_fields.iter().map(|field| { + let Some(field_ident) = &field.ident else { + unreachable!(); + }; + + field_ident + }); + + quote! { Self::#variant_ident { #(#field_idents),* } } + } + syn::Fields::Unnamed(syn::FieldsUnnamed { unnamed: unnamed_fields, .. }) => { + let field_idents = unnamed_fields + .iter() + .enumerate() + .map(|(field_index, _)| format_ident!("field_{field_index}")); + + quote! { Self::#variant_ident(#(#field_idents),*) } + } + }; + + VariantTryGetFieldFns { + immutable: quote! { + |target, field_index| { { + #![allow(unreachable_code)] + + use std::any::Any; + use #engine_crate_path::reflection::GetError; + + let target = target + .downcast_ref::<Self>() + .ok_or(GetError::WrongTargetType)?; + + let #variant_pattern = target else { + return Err(GetError::WrongTargetEnumVariant); + }; + + let field: &dyn Any = match field_index { + #field_index_match_arms + _ => { + return Err(GetError::IndexOutOfBounds); + } + }; + + Ok(field) + } } + }, + mutable: quote! { + |target, field_index| { { + #![allow(unreachable_code)] + + use std::any::Any; + use #engine_crate_path::reflection::GetError; + + let target = target + .downcast_mut::<Self>() + .ok_or(GetError::WrongTargetType)?; + + let #variant_pattern = target else { + return Err(GetError::WrongTargetEnumVariant); + }; + + let field: &mut dyn Any = match field_index { + #field_index_match_arms + _ => { + return Err(GetError::IndexOutOfBounds); + } + }; + + Ok(field) + } } + }, + } +} |
