Parcourir la source

serial/derive-internal: Port to syn 2

parazyd il y a 3 ans
Parent
commit
835935eaac

+ 2 - 2
src/serial/derive-internal/Cargo.toml

@@ -9,6 +9,6 @@ license = "AGPL-3.0-only"
 edition = "2021"
 
 [dependencies]
-proc-macro2 = "1.0.63"
+proc-macro2 = "1.0.64"
 quote = "1.0.29"
-syn = {version = "1.0.109", features = ["full", "fold"]}
+syn = {version = "2.0.25", features = ["full", "fold"]}

+ 12 - 28
src/serial/derive-internal/src/helpers.rs

@@ -16,39 +16,23 @@
  * along with this program.  If not, see <https://www.gnu.org/licenses/>.
  */
 
-use quote::ToTokens;
-use syn::{Attribute, Meta};
-//use syn::{spanned::Spanned, Attribute, Error, Meta, NestedMeta, Path};
+use syn::{Attribute, Path};
 
 pub fn contains_skip(attrs: &[Attribute]) -> bool {
-    for attr in attrs.iter() {
-        if let Ok(Meta::Path(path)) = attr.parse_meta() {
-            if path.to_token_stream().to_string().as_str() == "skip_serialize" {
-                return true
-            }
-        }
-    }
-    false
+    attrs.iter().any(|attr| attr.path().is_ident("skip_serialize"))
 }
 
-/*
-pub fn contains_initialize_with(attrs: &[Attribute]) -> syn::Result<Option<Path>> {
+pub fn contains_initialize_with(attrs: &[Attribute]) -> Option<Path> {
     for attr in attrs.iter() {
-        if let Ok(Meta::List(meta_list)) = attr.parse_meta() {
-            if meta_list.path.to_token_stream().to_string().as_str() == "init_serialize" {
-                if meta_list.nested.len() != 1 {
-                    return Err(Error::new(
-                        meta_list.span(),
-                        "init_serialize requires exactly one initialization method.",
-                    ))
-                }
-                let nested_meta = meta_list.nested.iter().next().unwrap();
-                if let NestedMeta::Meta(Meta::Path(path)) = nested_meta {
-                    return Ok(Some(path.clone()))
-                }
-            }
+        if attr.path().is_ident("init_serialize") {
+            let mut res = None;
+            let _ = attr.parse_nested_meta(|meta| {
+                res = Some(meta.path);
+                Ok(())
+            });
+            return res
         }
     }
-    Ok(None)
+
+    None
 }
-*/

+ 208 - 113
src/serial/derive-internal/src/lib.rs

@@ -16,139 +16,219 @@
  * along with this program.  If not, see <https://www.gnu.org/licenses/>.
  */
 
-//! Derive (de)serialization for structs, see src/serial/derive
-use proc_macro2::{Span, TokenStream as TokenStream2};
+//! Derive (de)serialization for enums and structs, see src/serial/derive
+use std::collections::HashMap;
+
+use proc_macro2::{Ident, Span, TokenStream};
 use quote::quote;
-use syn::{Fields, Ident, Index, ItemEnum, ItemStruct, WhereClause};
+use syn::{
+    punctuated::Punctuated, token::Comma, Fields, FieldsNamed, FieldsUnnamed, Index, ItemEnum,
+    ItemStruct, Variant, WhereClause, WherePredicate,
+};
 
 mod helpers;
-use helpers::contains_skip;
+use helpers::{contains_initialize_with, contains_skip};
 
-pub fn enum_ser(input: &ItemEnum, cratename: Ident) -> syn::Result<TokenStream2> {
-    let name = &input.ident;
+struct VariantParts {
+    where_predicates: Vec<WherePredicate>,
+    variant_header: TokenStream,
+    variant_body: TokenStream,
+    variant_idx_body: TokenStream,
+}
+
+/// Calculates the discriminant that will be assigned by the compiler.
+/// See: https://doc.rust-lang.org/reference/items/enumerations.html#assigning-discriminant-values
+fn discriminant_map(variants: &Punctuated<Variant, Comma>) -> HashMap<Ident, TokenStream> {
+    let mut map = HashMap::new();
+
+    let mut next_discriminant_if_not_specified = quote! {0};
+
+    for variant in variants {
+        let this_discriminant = variant
+            .discriminant
+            .clone()
+            .map_or_else(|| quote! { #next_discriminant_if_not_specified }, |(_, e)| quote! { #e });
+
+        next_discriminant_if_not_specified = quote! { #this_discriminant + 1 };
+        map.insert(variant.ident.clone(), this_discriminant);
+    }
+
+    map
+}
+
+fn named_fields(
+    cratename: &Ident,
+    enum_ident: &Ident,
+    variant_ident: &Ident,
+    discriminant_value: &TokenStream,
+    fields: &FieldsNamed,
+) -> syn::Result<VariantParts> {
+    let mut where_predicates: Vec<WherePredicate> = vec![];
+    let mut variant_header = TokenStream::new();
+    let mut variant_body = TokenStream::new();
+
+    for field in &fields.named {
+        if !contains_skip(&field.attrs) {
+            let field_ident = field.ident.clone().unwrap();
+
+            variant_header.extend(quote! { #field_ident, });
+
+            let field_type = &field.ty;
+            where_predicates.push(
+                syn::parse2(quote! {
+                    #field_type: #cratename::Encodable
+                })
+                .unwrap(),
+            );
+
+            variant_body.extend(quote! {
+                #field_ident.encode(writer)?;
+            })
+        }
+    }
+
+    // `..` pattern matching works even if all fields were specified
+    variant_header = quote! { { #variant_header .. }};
+    let variant_idx_body = quote!(
+        #enum_ident::#variant_ident { .. } => #discriminant_value,
+    );
+
+    Ok(VariantParts { where_predicates, variant_header, variant_body, variant_idx_body })
+}
+
+fn unnamed_fields(
+    cratename: &Ident,
+    enum_ident: &Ident,
+    variant_ident: &Ident,
+    discriminant_value: &TokenStream,
+    fields: &FieldsUnnamed,
+) -> syn::Result<VariantParts> {
+    let mut where_predicates: Vec<WherePredicate> = vec![];
+    let mut variant_header = TokenStream::new();
+    let mut variant_body = TokenStream::new();
+
+    for (field_idx, field) in fields.unnamed.iter().enumerate() {
+        let field_idx = u32::try_from(field_idx).expect("up to 2^32 fields are supported");
+        if contains_skip(&field.attrs) {
+            let field_ident = Ident::new(format!("_id{}", field_idx).as_str(), Span::mixed_site());
+            variant_header.extend(quote! { #field_ident, });
+        } else {
+            let field_ident = Ident::new(format!("id{}", field_idx).as_str(), Span::mixed_site());
+            variant_header.extend(quote! { #field_ident, });
+
+            let field_type = &field.ty;
+            where_predicates.push(
+                syn::parse2(quote! {
+                    #field_type: #cratename::Encodable
+                })
+                .unwrap(),
+            );
+
+            variant_body.extend(quote! {
+                #field_ident.encode(writer)?;
+            })
+        }
+    }
+
+    variant_header = quote! { ( #variant_header )};
+    let variant_idx_body = quote!(
+        #enum_ident::#variant_ident(..) => #discriminant_value,
+    );
+
+    Ok(VariantParts { where_predicates, variant_header, variant_body, variant_idx_body })
+}
+
+pub fn enum_ser(input: &ItemEnum, cratename: Ident) -> syn::Result<TokenStream> {
+    let enum_ident = &input.ident;
     let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
     let mut where_clause = where_clause.map_or_else(
         || WhereClause { where_token: Default::default(), predicates: Default::default() },
         Clone::clone,
     );
+    let mut all_variants_idx_body = TokenStream::new();
+    let mut fields_body = TokenStream::new();
+    let discriminants = discriminant_map(&input.variants);
 
-    let mut variant_idx_body = TokenStream2::new();
-    let mut fields_body = TokenStream2::new();
-    for (variant_idx, variant) in input.variants.iter().enumerate() {
-        let variant_idx = u8::try_from(variant_idx).expect("up to 256 enum variants are supported");
+    for variant in input.variants.iter() {
         let variant_ident = &variant.ident;
-        let mut variant_header = TokenStream2::new();
-        let mut variant_body = TokenStream2::new();
-        match &variant.fields {
-            Fields::Named(fields) => {
-                for field in &fields.named {
-                    let field_name = field.ident.as_ref().unwrap();
-                    if contains_skip(&field.attrs) {
-                        variant_header.extend(quote! { _ #field_name, }); // TODO: Test this
-                        continue
-                    } else {
-                        let field_type = &field.ty;
-                        where_clause.predicates.push(
-                            syn::parse2(quote! {
-                                #field_type: #cratename::Encodable
-                            })
-                            .unwrap(),
-                        );
-                        variant_header.extend(quote! { #field_name, });
-                    }
-                    variant_body.extend(quote! {
-                        len += self.#field_name.encode(&mut s);
-                    })
+        let discriminant_value = discriminants.get(variant_ident).unwrap();
+        let VariantParts { where_predicates, variant_header, variant_body, variant_idx_body } =
+            match &variant.fields {
+                Fields::Named(fields) => {
+                    named_fields(&cratename, enum_ident, variant_ident, discriminant_value, fields)?
                 }
-                variant_header = quote! { { #variant_header } };
-                variant_idx_body.extend(quote!(
-                    #name::#variant_ident { .. } => #variant_idx,
-                ));
-            }
-            Fields::Unnamed(fields) => {
-                for (field_idx, field) in fields.unnamed.iter().enumerate() {
-                    let field_idx =
-                        u32::try_from(field_idx).expect("up to 2^32 fields are supported");
-                    if contains_skip(&field.attrs) {
-                        let field_ident =
-                            Ident::new(format!("_id{}", field_idx).as_str(), Span::call_site());
-                        variant_header.extend(quote! { #field_ident, });
-                        continue
-                    } else {
-                        let field_type = &field.ty;
-                        where_clause.predicates.push(
-                            syn::parse2(quote! {
-                                #field_type: #cratename::Encodable
-                            })
-                            .unwrap(),
-                        );
-
-                        let field_ident =
-                            Ident::new(format!("id{}", field_idx).as_str(), Span::call_site());
-                        variant_header.extend(quote! { #field_ident, });
-                        variant_body.extend(quote! {
-                            len += self.#field_ident.encode(&mut s)?;
-                        })
+                Fields::Unnamed(fields) => unnamed_fields(
+                    &cratename,
+                    enum_ident,
+                    variant_ident,
+                    discriminant_value,
+                    fields,
+                )?,
+                Fields::Unit => {
+                    let variant_idx_body = quote!(
+                        #enum_ident::#variant_ident => #discriminant_value,
+                    );
+                    VariantParts {
+                        where_predicates: vec![],
+                        variant_header: TokenStream::new(),
+                        variant_body: TokenStream::new(),
+                        variant_idx_body,
                     }
                 }
-                variant_header = quote! { ( #variant_header )};
-                variant_idx_body.extend(quote!(
-                    #name::#variant_ident(..) => #variant_idx,
-                ));
-            }
-            Fields::Unit => {
-                variant_idx_body.extend(quote!(
-                    #name::#variant_ident => #variant_idx,
-                ));
-            }
-        }
+            };
+        where_predicates.into_iter().for_each(|predicate| where_clause.predicates.push(predicate));
+        all_variants_idx_body.extend(variant_idx_body);
         fields_body.extend(quote!(
-            #name::#variant_ident #variant_header => {
+            #enum_ident::#variant_ident #variant_header => {
                 #variant_body
             }
         ))
     }
 
     Ok(quote! {
-        impl #impl_generics #cratename::Encodable for #name #ty_generics #where_clause {
+        impl #impl_generics #cratename::Encodable for #enum_ident #ty_generics #where_clause {
             fn encode<S: std::io::Write>(&self, mut s: S) -> ::core::result::Result<usize, std::io::Error> {
                 let variant_idx: u8 = match self {
-                    #variant_idx_body
+                    #all_variants_idx_body
                 };
 
-                s.write_all(&variant_idx.to_le_bytes())?;
-                let mut len = 1;
+                let bytes = variant_idx.to_le_bytes();
+
+                writer.write_all(&bytes)?;
 
                 match self {
                     #fields_body
                 }
 
-                Ok(len)
+                Ok(bytes.len())
+
             }
         }
     })
 }
 
-pub fn enum_de(input: &ItemEnum, cratename: Ident) -> syn::Result<TokenStream2> {
+pub fn enum_de(input: &ItemEnum, cratename: Ident) -> syn::Result<TokenStream> {
     let name = &input.ident;
     let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
     let mut where_clause = where_clause.map_or_else(
         || WhereClause { where_token: Default::default(), predicates: Default::default() },
         Clone::clone,
     );
+    let init_method = contains_initialize_with(&input.attrs);
+    let mut variant_arms = TokenStream::new();
+    let discriminants = discriminant_map(&input.variants);
 
-    let mut variant_arms = TokenStream2::new();
-    for (variant_idx, variant) in input.variants.iter().enumerate() {
-        let variant_idx = u8::try_from(variant_idx).expect("up to 256 enum variants are supported");
+    for variant in input.variants.iter() {
         let variant_ident = &variant.ident;
-        let mut variant_header = TokenStream2::new();
+        let discriminant = discriminants.get(variant_ident).unwrap();
+        let mut variant_header = TokenStream::new();
         match &variant.fields {
             Fields::Named(fields) => {
                 for field in &fields.named {
                     let field_name = field.ident.as_ref().unwrap();
                     if contains_skip(&field.attrs) {
                         variant_header.extend(quote! {
-                                #field_name: Default::default(),
+                            #field_name: Default::default(),
                         });
                     } else {
                         let field_type = &field.ty;
@@ -164,7 +244,7 @@ pub fn enum_de(input: &ItemEnum, cratename: Ident) -> syn::Result<TokenStream2>
                         });
                     }
                 }
-                variant_header = quote! { { #variant_header } };
+                variant_header = quote! { { #variant_header }};
             }
             Fields::Unnamed(fields) => {
                 for field in fields.unnamed.iter() {
@@ -179,42 +259,44 @@ pub fn enum_de(input: &ItemEnum, cratename: Ident) -> syn::Result<TokenStream2>
                             .unwrap(),
                         );
 
-                        variant_header.extend(quote! { #cratename::Decodable::decode(&mut d)?, });
+                        variant_header.extend(quote! {
+                            #cratename::Decodable::decode(&mut d)?,
+                        });
                     }
                 }
-                variant_header = quote! { ( #variant_header ) };
+                variant_header = quote! { ( #variant_header )};
             }
             Fields::Unit => {}
         }
-
         variant_arms.extend(quote! {
-            #variant_idx => #name::#variant_ident #variant_header ,
+            if variant_tag == #discriminant { #name::#variant_ident #variant_header } else
         });
     }
 
-    let variant_idx = quote! {
-        let variant_idx: u8 = #cratename::Decodable::decode(&mut d)?;
+    let init = if let Some(method_ident) = init_method {
+        quote! {
+            return_value.#method_ident();
+        }
+    } else {
+        quote! {}
     };
 
     Ok(quote! {
-        impl #impl_generics #cratename::Decodable for #name #ty_generics #where_clause {
-            fn decode<D: std::io::Read>(mut d: D) -> ::core::result::Result<Self, std::io::Error> {
-                #variant_idx
-
-                let return_value = match variant_idx {
-                    #variant_arms
-                    _ => {
-                        let msg = format!("Unexpected variant index: {:?}", variant_idx);
-                        return Err(std::io::Error::new(std::io::ErrorKind::InvalidInput, msg));
-                    }
-                };
-                Ok(return_value)
-            }
+    impl #impl_generics #cratename::Decodable for #name #ty_generics #where_clause {
+        fn decode<D: std::io::Read>(mut d: D) -> ::core::result::Result<Self, std::io::Error> {
+            let mut return_value =
+                #variant_arms {
+                    let msg = format!("Unexpected variant index: {:?}", variant_idx);
+                    return Err(std::io::Error::new(std::io::ErrorKind::InvalidInput, msg));
+                }
+            };
+            #init
+            Ok(return_value)
         }
     })
 }
 
-pub fn struct_ser(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStream2> {
+pub fn struct_ser(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStream> {
     let name = &input.ident;
     let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
     let mut where_clause = where_clause.map_or_else(
@@ -222,7 +304,7 @@ pub fn struct_ser(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStre
         Clone::clone,
     );
 
-    let mut body = TokenStream2::new();
+    let mut body = TokenStream::new();
 
     match &input.fields {
         Fields::Named(fields) => {
@@ -272,7 +354,7 @@ pub fn struct_ser(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStre
     })
 }
 
-pub fn struct_de(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStream2> {
+pub fn struct_de(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStream> {
     let name = &input.ident;
     let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
     let mut where_clause = where_clause.map_or_else(
@@ -280,13 +362,14 @@ pub fn struct_de(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStrea
         Clone::clone,
     );
 
+    let init_method = contains_initialize_with(&input.attrs);
     let return_value = match &input.fields {
         Fields::Named(fields) => {
-            let mut body = TokenStream2::new();
+            let mut body = TokenStream::new();
             for field in &fields.named {
                 let field_name = field.ident.as_ref().unwrap();
 
-                let delta: TokenStream2 = if contains_skip(&field.attrs) {
+                let delta = if contains_skip(&field.attrs) {
                     quote! {
                         #field_name: Default::default(),
                     }
@@ -310,7 +393,7 @@ pub fn struct_de(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStrea
             }
         }
         Fields::Unnamed(fields) => {
-            let mut body = TokenStream2::new();
+            let mut body = TokenStream::new();
             for _ in 0..fields.unnamed.len() {
                 let delta = quote! {
                     #cratename::Decodable::decode(&mut d)?,
@@ -328,11 +411,23 @@ pub fn struct_de(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStrea
         }
     };
 
-    Ok(quote! {
+    if let Some(method_ident) = init_method {
+        Ok(quote! {
         impl #impl_generics #cratename::Decodable for #name #ty_generics #where_clause {
             fn decode<D: std::io::Read>(mut d: D) -> ::core::result::Result<Self, std::io::Error> {
-                Ok(#return_value)
+                let mut return_value = #return_value;
+                return_value.#method_ident();
+                Ok(return_value)
             }
         }
-    })
+        })
+    } else {
+        Ok(quote! {
+            impl #impl_generics #cratename::Decodable for #name #ty_generics #where_clause {
+                fn decode<D: std::io::Read>(mut d: D) -> ::core::result::Result<Self, std::io::Error> {
+                    Ok(#return_value)
+                }
+            }
+        })
+    }
 }