Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
91 changes: 71 additions & 20 deletions wacore/derive/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -545,7 +545,7 @@ fn is_option_type(ty: &syn::Type) -> bool {
// variant; optional #[wire_alias = "..."] adds parser-side aliases;
// #[wire(skip)] on a field excludes it from JSON; #[wire_fallback] with
// { tag: String } catches unknown tags.
// Emits: wire_tag(), impl Serialize (SerializeMap), and a sibling
// Emits: wire_tag(), impl Serialize (SerializeStruct), and a sibling
// <Name>Tag unit enum (unit-string WireEnum) for parser dispatch.
// Adding `content = "data"` selects Serde's adjacent representation,
// supports tuple payloads, and emits both Serialize and Deserialize.
Expand Down Expand Up @@ -1598,15 +1598,43 @@ fn expand_wire_enum_tagged(
}
}
} else {
let struct_name_lit = name.to_string();
// Struct-field path, not map keys: serializers that intern struct keys
// cannot intern a name that arrives as a value.
let serialize_arms: Vec<_> = infos
.iter()
.map(|info| {
let id = &info.ident;
let emit_arm = |pattern: proc_macro2::TokenStream,
len_prelude: proc_macro2::TokenStream,
entries: Vec<proc_macro2::TokenStream>| {
quote! {
#pattern => {
#len_prelude
let mut __state = ::serde::Serializer::serialize_struct(
serializer, #struct_name_lit, __len
)?;
::serde::ser::SerializeStruct::serialize_field(
&mut __state, #discriminator_lit, self.wire_tag()
)?;
#(#entries)*
::serde::ser::SerializeStruct::end(__state)
}
}
};
if info.is_fallback {
quote! { #name::#id { tag: _ } => {} }
emit_arm(
quote! { #name::#id { tag: _ } },
quote! { let __len = 1usize; },
Vec::new(),
)
} else {
match &info.fields {
Fields::Unit => quote! { #name::#id => {} },
Fields::Unit => emit_arm(
quote! { #name::#id },
quote! { let __len = 1usize; },
Vec::new(),
),
Fields::Named(named) => {
let bindings: Vec<proc_macro2::TokenStream> = named
.named
Expand All @@ -1620,38 +1648,67 @@ fn expand_wire_enum_tagged(
}
})
.collect();
let entries: Vec<proc_macro2::TokenStream> = named
let serialized: Vec<&syn::Field> = named
.named
.iter()
.filter(|field| !field_has_wire_skip(&field.attrs))
.collect();
let optional: Vec<&syn::Field> = serialized
.iter()
.copied()
.filter(|field| is_option_type(&field.ty))
.collect();
// Exact, not an upper bound: length-prefixed formats
// encode this count.
let constant_count = serialized.len() - optional.len() + 1;
let len_prelude = if optional.is_empty() {
quote! { let __len = #constant_count; }
} else {
let increments = optional.iter().map(|field| {
let id = field.ident.as_ref().unwrap();
quote! {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
if ::core::option::Option::is_some(#id) {
__len += 1;
}
}
});
quote! {
let mut __len = #constant_count;
#(#increments)*
}
};
let entries: Vec<proc_macro2::TokenStream> = serialized
.iter()
.map(|field| {
let id = field.ident.as_ref().unwrap();
let key = id.to_string();
if is_option_type(&field.ty) {
quote! {
if let ::core::option::Option::Some(__v) = #id {
::serde::ser::SerializeMap::serialize_entry(
&mut map, #key, __v
::serde::ser::SerializeStruct::serialize_field(
&mut __state, #key, __v
)?;
}
}
} else {
quote! {
::serde::ser::SerializeMap::serialize_entry(
&mut map, #key, #id
::serde::ser::SerializeStruct::serialize_field(
&mut __state, #key, #id
)?;
}
}
})
.collect();
quote! {
#name::#id { #(#bindings),* } => {
#(#entries)*
}
}
emit_arm(
quote! { #name::#id { #(#bindings),* } },
len_prelude,
entries,
)
}
// Scoped to the offending variant so it cannot shadow
// the arms that follow it.
Fields::Unnamed(_) => quote! {
compile_error!("tagged WireEnum tuple variants require #[wire(content = \"...\")]");
#name::#id(..) => compile_error!("tagged WireEnum tuple variants require #[wire(content = \"...\")]")
},
}
}
Expand All @@ -1664,15 +1721,9 @@ fn expand_wire_enum_tagged(
&self,
serializer: S,
) -> ::core::result::Result<S::Ok, S::Error> {
use ::serde::ser::SerializeMap;
let mut map = serializer.serialize_map(None)?;
::serde::ser::SerializeMap::serialize_entry(
&mut map, #discriminator_lit, self.wire_tag()
)?;
match self {
#(#serialize_arms,)*
}
::serde::ser::SerializeMap::end(map)
}
}
}
Expand Down
Loading
Loading