Support generics in serialize_with

This commit is contained in:
Sosthène Guédon
2025-06-05 10:42:07 +02:00
committed by sosthene-nitrokey
parent a632293e8e
commit 3f215c3fcb
2 changed files with 87 additions and 14 deletions
+62 -12
View File
@@ -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<proc_macro2::TokenStream> {
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<proc_macro2::TokenStream> {
fields
.iter()
.filter_map(|field| {
@@ -49,11 +60,12 @@ fn serialize_fields(fields: &[parse::Field], offset: usize) -> Vec<proc_macro2::
Some(f) => {
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<proc_macro2::
}
}
&__SerializeWith { value: &self.#member }
&__SerializeWith { value: &self.#member, phantom: ::core::marker::PhantomData::<#ident #ty_generics> }
})
}
};
@@ -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<proc_macro2::TokenStre
.collect()
}
fn match_fields(fields: &[parse::Field], offset: usize) -> Vec<proc_macro2::TokenStream> {
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<proc_macro2::TokenStream> {
fields
.iter()
.filter(|f| !f.skip_serializing_if.is_always())
@@ -186,11 +225,12 @@ fn match_fields(fields: &[parse::Field], offset: usize) -> Vec<proc_macro2::Toke
Some(f) => {
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<Self, __D::Error>
@@ -200,12 +240,13 @@ fn match_fields(fields: &[parse::Field], offset: usize) -> Vec<proc_macro2::Toke
Ok(__DeserializeWith {
value: #f(__deserializer)?,
phantom: ::core::marker::PhantomData,
lifetime: ::core::marker::PhantomData,
})
}
}
let __DeserializeWith { value, lifetime: _ } = map.next_value()?;
let __DeserializeWith { value, lifetime: _, phantom: _ } = map.next_value()?;
value
}
)
@@ -244,7 +285,6 @@ pub fn derive_deserialize(input: TokenStream) -> 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
+25 -2
View File
@@ -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<u8>,
}
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,
],
)
}
}