Parcourir la source

util/derive: allow to skip fields with attributes

ghassmo il y a 4 ans
Parent
commit
7cecc88a3d
3 fichiers modifiés avec 51 ajouts et 17 suppressions
  1. 37 15
      src/util/derive-internal/src/lib.rs
  2. 2 2
      src/util/derive/src/lib.rs
  3. 12 0
      src/util/serial.rs

+ 37 - 15
src/util/derive-internal/src/lib.rs

@@ -20,7 +20,11 @@ pub fn struct_ser(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStre
             body.extend(ln);
 
             for field in &fields.named {
-                // TODO: Allow skip?
+                if let Some(attr) = field.attrs.iter().next() {
+                    if attr.path.is_ident("skip_serialize") {
+                        continue
+                    }
+                }
 
                 let field_name = field.ident.as_ref().unwrap();
                 let delta = quote! {
@@ -90,20 +94,38 @@ pub fn struct_de(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStrea
 
             for field in &fields.named {
                 let field_name = field.ident.as_ref().unwrap();
-                // TODO: Allow skip?
-                let delta = {
-                    let field_type = &field.ty;
-                    where_clause.predicates.push(
-                        syn::parse2(quote! {
-                            #field_type: #cratename::util::serial::Decodable
-                        })
-                        .unwrap(),
-                    );
-
-                    quote! {
-                        #field_name: #cratename::util::serial::Decodable::decode(&mut d)?,
-                    }
-                };
+                let delta: TokenStream2;
+                let attr = field.attrs.iter().next();
+
+                if attr.is_some() && attr.unwrap().path.is_ident("skip_serialize") {
+                    delta = {
+                        let field_type = &field.ty;
+                        where_clause.predicates.push(
+                            syn::parse2(quote! {
+                                #field_type: core::default::Default
+                            })
+                            .unwrap(),
+                        );
+
+                        quote! {
+                            #field_name: #field_type::default(),
+                        }
+                    };
+                } else {
+                    delta = {
+                        let field_type = &field.ty;
+                        where_clause.predicates.push(
+                            syn::parse2(quote! {
+                                #field_type: #cratename::util::serial::Decodable
+                            })
+                            .unwrap(),
+                        );
+
+                        quote! {
+                            #field_name: #cratename::util::serial::Decodable::decode(&mut d)?,
+                        }
+                    };
+                }
 
                 body.extend(delta);
             }

+ 2 - 2
src/util/derive/src/lib.rs

@@ -6,7 +6,7 @@ use syn::{Ident, ItemStruct};
 
 use darkfi_derive_internal::{struct_de, struct_ser};
 
-#[proc_macro_derive(SerialEncodable)]
+#[proc_macro_derive(SerialEncodable, attributes(skip_serialize))]
 pub fn darkfi_serialize(input: TokenStream) -> TokenStream {
     let found_crate = crate_name("darkfi").expect("darkfi is found in Cargo.toml");
 
@@ -30,7 +30,7 @@ pub fn darkfi_serialize(input: TokenStream) -> TokenStream {
     })
 }
 
-#[proc_macro_derive(SerialDecodable)]
+#[proc_macro_derive(SerialDecodable, attributes(skip_serialize))]
 pub fn darkfi_deserialize(input: TokenStream) -> TokenStream {
     let found_crate = crate_name("darkfi").expect("darkfi is found in Cargo.toml");
 

+ 12 - 0
src/util/serial.rs

@@ -852,22 +852,34 @@ mod tests {
     #[derive(Debug, PartialEq, Clone, SerialEncodable, SerialDecodable)]
     struct TestDerive2(u64);
 
+    #[derive(Debug, PartialEq, Clone, SerialEncodable, SerialDecodable)]
+    struct TestDerive3 {
+        foo: u64,
+        #[skip_serialize]
+        bar: u64,
+        meh: u64,
+    }
+
     #[test]
     fn serialize_deserialize_struct() {
         let t0 = TestDerive0 { foo: String::from("Andrew"), bar: 42 };
         let t1 = TestDerive1 { baz: t0.clone(), meh: false };
         let t2 = TestDerive2(u64::MAX);
+        let t3 = TestDerive3 { foo: 30, bar: 20, meh: 44 };
 
         let t0_bytes = serialize(&t0);
         let t1_bytes = serialize(&t1);
         let t2_bytes = serialize(&t2);
+        let t3_bytes = serialize(&t3);
 
         let t0_de: TestDerive0 = deserialize(&t0_bytes).unwrap();
         let t1_de: TestDerive1 = deserialize(&t1_bytes).unwrap();
         let t2_de: TestDerive2 = deserialize(&t2_bytes).unwrap();
+        let t3_de: TestDerive3 = deserialize(&t3_bytes).unwrap();
 
         assert_eq!(t0, t0_de);
         assert_eq!(t1, t1_de);
         assert_eq!(t2, t2_de);
+        assert_eq!(t3_de, TestDerive3 { foo: 30, bar: 0, meh: 44 });
     }
 }