diff --git a/src/lib.rs b/src/lib.rs index 8a90962..6d0117d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -34,11 +34,22 @@ use parse::Skip; use proc_macro::TokenStream; use proc_macro2::{Ident, Span}; use quote::{format_ident, quote}; -use syn::{parse_macro_input, Lifetime, LifetimeParam, TypeParamBound}; +use syn::{ + parse_macro_input, ImplGenerics, Lifetime, LifetimeParam, TypeGenerics, TypeParamBound, + WhereClause, +}; use crate::parse::Input; -fn serialize_fields(fields: &[parse::Field], offset: usize) -> Vec { +fn serialize_fields( + fields: &[parse::Field], + offset: usize, + impl_generics_serialize: ImplGenerics<'_>, + ty_generics_serialize: TypeGenerics<'_>, + ty_generics: &TypeGenerics<'_>, + where_clause: Option<&WhereClause>, + ident: &Ident, +) -> Vec { fields .iter() .filter_map(|field| { @@ -49,11 +60,12 @@ fn serialize_fields(fields: &[parse::Field], offset: usize) -> Vec { let ty = &field.ty; quote!({ - struct __SerializeWith<'__lifetime> { + struct __SerializeWith #impl_generics_serialize { value: &'__lifetime #ty, + phantom: ::core::marker::PhantomData<#ident #ty_generics>, } - impl<'__lifetime> serde::Serialize for __SerializeWith<'__lifetime> { + impl #impl_generics_serialize serde::Serialize for __SerializeWith #ty_generics_serialize #where_clause { fn serialize<__S>( &self, __s: __S, @@ -65,7 +77,7 @@ fn serialize_fields(fields: &[parse::Field], offset: usize) -> Vec } }) } }; @@ -111,7 +123,6 @@ pub fn derive_serialize(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as Input); let ident = input.ident; let num_fields = count_serialized_fields(&input.fields); - let serialize_fields = serialize_fields(&input.fields, input.attrs.offset); let (_, ty_generics, where_clause) = input.generics.split_for_impl(); let mut generics_cl = input.generics.clone(); generics_cl.type_params_mut().for_each(|t| { @@ -120,6 +131,26 @@ pub fn derive_serialize(input: TokenStream) -> TokenStream { }); let (impl_generics, _, _) = generics_cl.split_for_impl(); + let mut generics_cl2 = generics_cl.clone(); + + generics_cl2 + .params + .push(syn::GenericParam::Lifetime(LifetimeParam::new( + Lifetime::new("'__lifetime", Span::call_site()), + ))); + + let (impl_generics_serialize, ty_generics_serialize, _) = generics_cl2.split_for_impl(); + + let serialize_fields = serialize_fields( + &input.fields, + input.attrs.offset, + impl_generics_serialize, + ty_generics_serialize, + &ty_generics, + where_clause, + &ident, + ); + TokenStream::from(quote! { #[automatically_derived] impl #impl_generics serde::Serialize for #ident #ty_generics #where_clause { @@ -173,7 +204,15 @@ fn unwrap_expected_fields(fields: &[parse::Field]) -> Vec Vec { +fn match_fields( + fields: &[parse::Field], + offset: usize, + impl_generics_with_de: &ImplGenerics<'_>, + ty_generics: &TypeGenerics<'_>, + ty_generics_with_de: &TypeGenerics<'_>, + where_clause: Option<&WhereClause>, + struct_ident: &Ident, +) -> Vec { fields .iter() .filter(|f| !f.skip_serializing_if.is_always()) @@ -186,11 +225,12 @@ fn match_fields(fields: &[parse::Field], offset: usize) -> Vec { let ty = &field.ty; quote!({ - struct __DeserializeWith< 'de> { + struct __DeserializeWith #impl_generics_with_de { value: #ty, + phantom: ::core::marker::PhantomData<#struct_ident #ty_generics>, lifetime: ::core::marker::PhantomData<&'de ()>, } - impl<'de> serde::Deserialize<'de> for __DeserializeWith<'de> { + impl #impl_generics_with_de serde::Deserialize<'de> for __DeserializeWith #ty_generics_with_de #where_clause { fn deserialize<__D>( __deserializer: __D, ) -> Result @@ -200,12 +240,13 @@ fn match_fields(fields: &[parse::Field], offset: usize) -> Vec TokenStream { let ident = input.ident; let none_fields = none_fields(&input.fields); 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 (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); @@ -271,7 +311,17 @@ pub fn derive_deserialize(input: TokenStream) -> TokenStream { .push_value(TypeParamBound::Verbatim(quote!(serde::Deserialize<'de>))); }); - let (impl_generics_with_de, _, _) = generics_cl.split_for_impl(); + let (impl_generics_with_de, ty_generics_with_de, _) = generics_cl.split_for_impl(); + + let match_fields = match_fields( + &input.fields, + input.attrs.offset, + &impl_generics_with_de, + &ty_generics, + &ty_generics_with_de, + where_clause, + &ident, + ); let the_loop = if !input.fields.is_empty() { // NB: In the previous "none_fields", we use the actual struct's diff --git a/tests/basics.rs b/tests/basics.rs index 3063e15..e00e503 100644 --- a/tests/basics.rs +++ b/tests/basics.rs @@ -548,12 +548,12 @@ mod generics { #[test] fn deserialize_with() { #[derive(serde_indexed::DeserializeIndexed, PartialEq, Eq, Debug)] - struct SerializeWith { + struct DeserializeWith { #[serde(deserialize_with = "serde_bytes::deserialize")] data: Vec, } - let value = SerializeWith { data: vec![0; 128] }; + let value = DeserializeWith { data: vec![0; 128] }; assert_de_tokens( &value, @@ -588,4 +588,27 @@ mod generics { ], ) } + + #[test] + fn with_lifetime() { + #[derive( + serde_indexed::SerializeIndexed, serde_indexed::DeserializeIndexed, PartialEq, Eq, Debug, + )] + struct SerializeWith<'a> { + #[serde(with = "serde_bytes")] + data: &'a [u8], + } + + let value = SerializeWith { data: &[0; 128] }; + + assert_tokens( + &value, + &[ + Token::Map { len: Some(1) }, + Token::U64(0), + Token::BorrowedBytes(&[0; 128]), + Token::MapEnd, + ], + ) + } }