From 056374115f8e625f3ad91cd0426a6ed8966dde31 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Sosth=C3=A8ne=20Gu=C3=A9don?= Date: Tue, 1 Aug 2023 09:28:20 +0200 Subject: [PATCH] Add support for lifetimes --- src/lib.rs | 42 +++++++++++++++++++++++++++++++++++------- src/parse.rs | 14 +++++++++++++- 2 files changed, 48 insertions(+), 8 deletions(-) diff --git a/src/lib.rs b/src/lib.rs index 2fe3bf6..da41932 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -30,6 +30,8 @@ extern crate proc_macro; mod parse; +use std::iter; + use proc_macro::TokenStream; use quote::{format_ident, quote}; use syn::parse_macro_input; @@ -85,9 +87,10 @@ pub fn derive_serialize(input: TokenStream) -> TokenStream { let ident = input.ident; let num_fields = count_serialized_fields(&input.fields); let serialize_fields = serialize_fields(&input.fields, input.attrs.offset); + let (_, lifetimes) = lifetimes(&input.lifetimes); TokenStream::from(quote! { - impl serde::Serialize for #ident { + impl<#(#lifetimes),*> serde::Serialize for #ident<#(#lifetimes),*> { fn serialize(&self, serializer: S) -> core::result::Result where S: serde::Serializer @@ -167,6 +170,21 @@ fn all_fields(fields: &[parse::Field]) -> Vec { .collect() } +fn lifetimes( + lifetimes: &[syn::Lifetime], +) -> (proc_macro2::TokenStream, Vec) { + let lifetimes: Vec<_> = lifetimes + .into_iter() + .map(|l| { + quote! {#l} + }) + .collect(); + let de_lifetime = quote! { + 'de: #(#lifetimes)+* + }; + (de_lifetime, lifetimes) +} + #[proc_macro_derive(DeserializeIndexed, attributes(serde, serde_indexed))] pub fn derive_deserialize(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as Input); @@ -175,6 +193,9 @@ pub fn derive_deserialize(input: TokenStream) -> TokenStream { let unwrap_expected_fields = unwrap_expected_fields(&input.fields); let match_fields = match_fields(&input.fields, input.attrs.offset); let all_fields = all_fields(&input.fields); + let (de_lifetime, lifetimes) = lifetimes(&input.lifetimes); + let impl_lifetimes = iter::once(&de_lifetime).chain(&lifetimes); + let impl_lifetimes2 = impl_lifetimes.clone(); let the_loop = if !input.fields.is_empty() { // NB: In the previous "none_fields", we use the actual struct's @@ -195,22 +216,29 @@ pub fn derive_deserialize(input: TokenStream) -> TokenStream { quote! {} }; + let phantom_datas_ty = lifetimes + .iter() + .map(|l| quote!(core::marker::PhantomData<&#l ()>)); + let phantom_datas_values = lifetimes + .iter() + .map(|_l| quote!(core::marker::PhantomData::default())); + TokenStream::from(quote! { - impl<'de> serde::Deserialize<'de> for #ident { + impl<#(#impl_lifetimes),*> serde::Deserialize<'de> for #ident<#(#lifetimes),*> { fn deserialize(deserializer: D) -> core::result::Result where D: serde::Deserializer<'de>, { - struct IndexedVisitor; + struct IndexedVisitor<#(#lifetimes),*>(#(#phantom_datas_ty),*); - impl<'de> serde::de::Visitor<'de> for IndexedVisitor { - type Value = #ident; + impl<#(#impl_lifetimes2),*> serde::de::Visitor<'de> for IndexedVisitor<#(#lifetimes),*> { + type Value = #ident<#(#lifetimes),*>; fn expecting(&self, formatter: &mut core::fmt::Formatter) -> core::fmt::Result { formatter.write_str(stringify!(#ident)) } - fn visit_map(self, mut map: V) -> core::result::Result<#ident, V::Error> + fn visit_map(self, mut map: V) -> core::result::Result where V: serde::de::MapAccess<'de>, { @@ -224,7 +252,7 @@ pub fn derive_deserialize(input: TokenStream) -> TokenStream { } } - deserializer.deserialize_map(IndexedVisitor {}) + deserializer.deserialize_map(IndexedVisitor(#(#phantom_datas_values),*)) } } }) diff --git a/src/parse.rs b/src/parse.rs index ad5287b..c5435fe 100644 --- a/src/parse.rs +++ b/src/parse.rs @@ -1,11 +1,12 @@ use proc_macro2::Span; use syn::parse::{Error, Parse, ParseStream, Result}; -use syn::{Data, DeriveInput, Fields, Ident, Token}; +use syn::{Data, DeriveInput, Fields, Ident, Lifetime, Token}; pub struct Input { pub ident: Ident, pub attrs: StructAttrs, pub fields: Vec, + pub lifetimes: Vec, } #[derive(Default)] @@ -68,6 +69,14 @@ fn parse_attrs(attrs: &Vec) -> Result { Ok(struct_attrs) } +fn lifetimes(generics: &syn::Generics) -> Vec { + generics + .lifetimes() + .into_iter() + .map(|l| l.lifetime.clone()) + .collect() +} + impl Parse for Input { fn parse(input: ParseStream) -> Result { let call_site = Span::call_site(); @@ -91,12 +100,15 @@ impl Parse for Input { let fields = fields_from_ast(&syn_fields.named); + let lifetimes = lifetimes(&derive_input.generics); + //serde::internals::ast calls `fields_from_ast(cx, &fields.named, attrs, container_default)` Ok(Input { ident: derive_input.ident, attrs, fields, + lifetimes, }) } }