Просмотр исходного кода

util/serial: Implement SerialEncodable/SerialDecodable derive macros.

parazyd 4 лет назад
Родитель
Сommit
444531cc3d

+ 21 - 0
Cargo.lock

@@ -1559,6 +1559,8 @@ dependencies = [
  "bytes",
  "clap 3.1.6",
  "crypto_api_chachapoly",
+ "darkfi-derive",
+ "darkfi-derive-internal",
  "dirs 4.0.0",
  "drk-sdk",
  "fast-socks5",
@@ -1600,6 +1602,25 @@ dependencies = [
  "zeromq",
 ]
 
+[[package]]
+name = "darkfi-derive"
+version = "0.3.0"
+dependencies = [
+ "darkfi-derive-internal",
+ "proc-macro-crate 0.1.5",
+ "proc-macro2 1.0.36",
+ "syn 1.0.89",
+]
+
+[[package]]
+name = "darkfi-derive-internal"
+version = "0.3.0"
+dependencies = [
+ "proc-macro2 1.0.36",
+ "quote 1.0.17",
+ "syn 1.0.89",
+]
+
 [[package]]
 name = "darkfid"
 version = "0.3.0"

+ 7 - 0
Cargo.toml

@@ -27,6 +27,8 @@ members = [
     "bin/vanityaddr",
 
     "src/sdk",
+    "src/util/derive",
+    "src/util/derive-internal",
 ]
 
 [dependencies]
@@ -64,6 +66,8 @@ lazy_static = {version = "1.4.0", optional = true}
 fxhash = {version = "0.2.1", optional = true}
 indexmap = {version = "1.8.0", optional = true}
 itertools = {version = "0.10.3", optional = true}
+darkfi-derive = {path = "src/util/derive", optional = true}
+darkfi-derive-internal = {path = "src/util/derive-internal", optional = true}
 
 # Misc
 termion = {version = "1.5.6", optional = true}
@@ -157,6 +161,9 @@ util = [
     "dirs",
     "num-bigint",
     "fxhash",
+
+    "darkfi-derive",
+    "darkfi-derive-internal",
 ]
 
 rpc = [

+ 9 - 0
src/util/derive-internal/Cargo.toml

@@ -0,0 +1,9 @@
+[package]
+name = "darkfi-derive-internal"
+version = "0.3.0"
+edition = "2021"
+
+[dependencies]
+proc-macro2 = "1"
+quote = "1"
+syn = {version = "1", features = ["full", "fold"]}

+ 109 - 0
src/util/derive-internal/src/lib.rs

@@ -0,0 +1,109 @@
+//! Derive (de)serialization for structs, see src/util/derive
+use proc_macro2::TokenStream as TokenStream2;
+use quote::quote;
+use syn::{Fields, Ident, ItemStruct, WhereClause};
+
+pub fn struct_ser(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStream2> {
+    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 mut body = TokenStream2::new();
+    match &input.fields {
+        Fields::Named(fields) => {
+            let ln = quote! {
+                let mut len = 0;
+            };
+            body.extend(ln);
+
+            for field in &fields.named {
+                // TODO: Allow skip?
+
+                let field_name = field.ident.as_ref().unwrap();
+                let delta = quote! {
+                    len += self.#field_name.encode(&mut s)?;
+                };
+                body.extend(delta);
+
+                let field_type = &field.ty;
+                where_clause.predicates.push(
+                    syn::parse2(quote! {
+                        #field_type: #cratename::util::serial::Encodable
+                    })
+                    .unwrap(),
+                );
+            }
+
+            let ret = quote! {
+                Ok(len)
+            };
+            body.extend(ret)
+        }
+        Fields::Unnamed(_fields) => todo!(),
+        Fields::Unit => {}
+    }
+
+    Ok(quote! {
+        impl #cratename::util::serial::Encodable for #name #where_clause {
+            fn encode<S: std::io::Write>(&self, mut s: S) -> #cratename::Result<usize> {
+                #body
+            }
+        }
+    })
+}
+
+pub fn struct_de(input: &ItemStruct, cratename: Ident) -> syn::Result<TokenStream2> {
+    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 return_value = match &input.fields {
+        Fields::Named(fields) => {
+            let mut body = TokenStream2::new();
+
+            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)?,
+                    }
+                };
+
+                body.extend(delta);
+            }
+
+            quote! {
+                Self { #body }
+            }
+        }
+        Fields::Unnamed(_fields) => todo!(),
+        Fields::Unit => {
+            quote! {
+                Self {}
+            }
+        }
+    };
+
+    Ok(quote! {
+        impl #cratename::util::serial::Decodable for #name #where_clause {
+            fn decode<D: std::io::Read>(mut d: D) -> #cratename::Result<Self> {
+                Ok(#return_value)
+            }
+        }
+    })
+}

+ 14 - 0
src/util/derive/Cargo.toml

@@ -0,0 +1,14 @@
+[package]
+name = "darkfi-derive"
+version = "0.3.0"
+edition = "2021"
+
+[lib]
+proc-macro = true
+
+[dependencies]
+proc-macro-crate = "0.1.5"
+proc-macro2 = "1"
+syn = {version = "1", features = ["full", "fold"]}
+
+darkfi-derive-internal = {path = "../derive-internal"}

+ 47 - 0
src/util/derive/src/lib.rs

@@ -0,0 +1,47 @@
+extern crate proc_macro;
+use proc_macro::TokenStream;
+use proc_macro2::Span;
+use proc_macro_crate::crate_name;
+use syn::{Ident, ItemStruct};
+
+use darkfi_derive_internal::{struct_de, struct_ser};
+
+#[proc_macro_derive(SerialEncodable)]
+pub fn darkfi_serialize(input: TokenStream) -> TokenStream {
+    let cratename = Ident::new(
+        &crate_name("darkfi").unwrap_or_else(|_| "crate".to_string()),
+        Span::call_site(),
+    );
+
+    let res = if let Ok(input) = syn::parse::<ItemStruct>(input) {
+        struct_ser(&input, cratename)
+    } else {
+        // For now we only allow derive on structs
+        unreachable!()
+    };
+
+    TokenStream::from(match res {
+        Ok(res) => res,
+        Err(err) => err.to_compile_error(),
+    })
+}
+
+#[proc_macro_derive(SerialDecodable)]
+pub fn darkfi_deserialize(input: TokenStream) -> TokenStream {
+    let cratename = Ident::new(
+        &crate_name("darkfi").unwrap_or_else(|_| "crate".to_string()),
+        Span::call_site(),
+    );
+
+    let res = if let Ok(input) = syn::parse::<ItemStruct>(input) {
+        struct_de(&input, cratename)
+    } else {
+        // For now we only allow derive on structs
+        unreachable!()
+    };
+
+    TokenStream::from(match res {
+        Ok(res) => res,
+        Err(err) => err.to_compile_error(),
+    })
+}

+ 34 - 1
src/util/serial.rs

@@ -8,6 +8,8 @@ use std::{
 
 use num_bigint::BigUint;
 
+pub use darkfi_derive::{SerialDecodable, SerialEncodable};
+
 use super::endian;
 use crate::{Error, Result};
 
@@ -596,7 +598,7 @@ mod tests {
     use super::{
         deserialize, deserialize_partial,
         endian::{u16_to_array_le, u32_to_array_le, u64_to_array_le},
-        serialize, Encodable, Error, Result, VarInt,
+        serialize, Encodable, Error, Result, SerialDecodable, SerialEncodable, VarInt,
     };
     use std::{io, mem::discriminant};
 
@@ -834,4 +836,35 @@ mod tests {
             Some(::std::borrow::Cow::Borrowed("Andrew"))
         );
     }
+
+    #[derive(Clone, SerialEncodable, SerialDecodable)]
+    struct TestDerive0 {
+        foo: String,
+        bar: u64,
+    }
+
+    #[derive(Clone, SerialEncodable, SerialDecodable)]
+    struct TestDerive1 {
+        baz: TestDerive0,
+        meh: bool,
+    }
+
+    #[test]
+    fn serialize_deserialize_struct() {
+        let t0 = TestDerive0 { foo: String::from("Andrew"), bar: 42 };
+
+        let t1 = TestDerive1 { baz: t0.clone(), meh: false };
+
+        let t0_bytes = serialize(&t0);
+        let t1_bytes = serialize(&t1);
+
+        let t0_de: TestDerive0 = deserialize(&t0_bytes).unwrap();
+        let t1_de: TestDerive1 = deserialize(&t1_bytes).unwrap();
+
+        assert_eq!(t0.foo, t0_de.foo);
+        assert_eq!(t0.bar, t0_de.bar);
+        assert_eq!(t0.foo, t1_de.baz.foo);
+        assert_eq!(t0.bar, t1_de.baz.bar);
+        assert_eq!(t1.meh, t1_de.meh);
+    }
 }