midnight_serialize_macros/
lib.rs1extern crate proc_macro;
16use proc_macro2::{Ident, Span, TokenStream};
17use quote::{quote, quote_spanned};
18use syn::parse::Parser;
19use syn::punctuated::Punctuated;
20use syn::spanned::Spanned;
21use syn::{
22 Data, DeriveInput, Fields, GenericParam, Generics, Index, Meta, Token, parse_macro_input,
23 parse_quote,
24};
25
26fn deserializable_add_trait_bounds(mut generics: Generics, phantom: &[Ident]) -> Generics {
27 for param in &mut generics.params {
28 if let GenericParam::Type(ref mut type_param) = *param
29 && !phantom.contains(&type_param.ident)
30 {
31 type_param.bounds.push(parse_quote!(Deserializable));
32 }
33 }
34 generics
35}
36
37fn tagged_add_trait_bounds(mut generics: Generics, phantom: &[Ident]) -> Generics {
38 for param in &mut generics.params {
39 if let GenericParam::Type(ref mut type_param) = *param
40 && !phantom.contains(&type_param.ident)
41 {
42 type_param.bounds.push(parse_quote!(Tagged));
43 }
44 }
45 generics
46}
47
48#[proc_macro_derive(Serializable, attributes(tag, phantom))]
51pub fn derive_serializable(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
52 let input = parse_macro_input!(input as DeriveInput);
53
54 let name = input.ident;
55
56 let phantom_generics = input
57 .attrs
58 .iter()
59 .find_map(|attr| match &attr.meta {
60 Meta::List(l) if l.path.is_ident("phantom") => {
61 let parser = Punctuated::<Ident, Token![,]>::parse_separated_nonempty;
62 parser
63 .parse2(l.tokens.clone())
64 .ok()
65 .map(|punct| punct.iter().cloned().collect::<Vec<_>>())
66 }
67 _ => None,
68 })
69 .unwrap_or(vec![]);
70
71 let generics = serializable_add_trait_bounds(input.generics.clone(), &phantom_generics);
72 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
73
74 let tag = input.attrs.iter().find_map(|attr| match &attr.meta {
75 Meta::NameValue(nv) if nv.path.is_ident("tag") => Some(&nv.value),
76 _ => None,
77 });
78
79 let de_generics = deserializable_add_trait_bounds(input.generics.clone(), &phantom_generics);
80 let (de_impl_generics, de_ty_generics, de_where_clause) = de_generics.split_for_impl();
81
82 let serialize = serialize(&input.data);
83 let deserialize = deserialize(&input.data);
84 let size = size(&input.data);
85
86 let mut expanded = quote! {
87 impl #impl_generics Serializable for #name #ty_generics #where_clause {
88 fn serialize(&self, writer: &mut impl ::std::io::Write) -> Result<(), ::std::io::Error> {
89 #serialize
90 Ok(())
91 }
92
93 fn serialized_size(&self) -> usize {
94 #size
95 }
96 }
97
98 impl #de_impl_generics Deserializable for #name #de_ty_generics #de_where_clause {
99 fn deserialize(reader: &mut impl ::std::io::Read, recursion_depth: u32) -> Result<Self, ::std::io::Error> {
100 #deserialize
101 }
102 }
103 };
104
105 if let Some(tag) = tag {
106 let tag_generics = tagged_add_trait_bounds(input.generics, &phantom_generics);
107 let (tag_impl_generics, tag_ty_generics, tag_where_clause) = tag_generics.split_for_impl();
108
109 let generics = tag_generics
110 .params
111 .iter()
112 .filter_map(|param| match param {
113 GenericParam::Type(ty) if !phantom_generics.contains(&ty.ident) => Some(&ty.ident),
114 _ => None,
115 })
116 .collect::<Vec<_>>();
117
118 let tag_expand = if generics.is_empty() {
119 quote! { ::std::borrow::Cow::Borrowed(#tag) }
120 } else {
121 let mut fstring = String::new();
122 fstring.push_str("{}(");
123 for i in 0..generics.len() {
124 if i > 0 {
125 fstring.push(',');
126 }
127 fstring.push_str("{}");
128 }
129 fstring.push(')');
130 quote! { ::std::borrow::Cow::Owned(::std::format!(#fstring, #tag, #( <#generics as Tagged>::tag() ),*)) }
131 };
132 let tag_factor_expand = tag_factors(&input.data);
133
134 expanded.extend(quote! {
135 impl #tag_impl_generics Tagged for #name #tag_ty_generics #tag_where_clause {
136 fn tag() -> ::std::borrow::Cow<'static, ::core::primitive::str> {
137 #tag_expand
138 }
139 fn tag_unique_factor() -> String {
140 #tag_factor_expand
141 }
142 }
143 });
144 }
145
146 proc_macro::TokenStream::from(expanded)
147}
148
149fn tag_factors_fields_fmt_str(fields: &Fields) -> String {
150 let nfields = fields.iter().count();
151 let mut res = String::new();
152 res.push('(');
153 for i in 0..nfields {
154 if i != 0 {
155 res.push(',');
156 }
157 res.push_str("{}");
158 }
159 res.push(')');
160 res
161}
162
163fn tag_factors_fields_fmt_args(fields: &Fields) -> impl Iterator<Item = TokenStream> {
164 fields.iter().map(|field| &field.ty).map(|ty| {
165 quote! {
166 <#ty>::tag()
167 }
168 })
169}
170
171fn tag_factors(data: &Data) -> TokenStream {
172 let fmt_str = match data {
173 Data::Struct(data) => format!("({})", tag_factors_fields_fmt_str(&data.fields)),
174 Data::Enum(data) => {
175 let mut res = String::new();
176 res.push('[');
177 for (i, variant) in data.variants.iter().enumerate() {
178 if i != 0 {
179 res.push(',');
180 }
181 res.push_str(&tag_factors_fields_fmt_str(&variant.fields));
182 }
183 res.push(']');
184 res
185 }
186 Data::Union(_) => unimplemented!(),
187 };
188 let fmt_args: Box<dyn Iterator<Item = TokenStream>> = match data {
189 Data::Struct(data) => Box::new(tag_factors_fields_fmt_args(&data.fields)),
190 Data::Enum(data) => Box::new(
191 data.variants
192 .iter()
193 .flat_map(|var| tag_factors_fields_fmt_args(&var.fields)),
194 ),
195 Data::Union(_) => unimplemented!(),
196 };
197 quote! {
198 format!(#fmt_str, #(#fmt_args),*)
199 }
200}
201
202fn serializable_add_trait_bounds(mut generics: Generics, phantom: &[Ident]) -> Generics {
203 for param in &mut generics.params {
204 if let GenericParam::Type(ref mut type_param) = *param
205 && !phantom.contains(&type_param.ident)
206 {
207 type_param.bounds.push(parse_quote!(Serializable));
208 }
209 }
210 generics
211}
212
213fn serialize_fields(fields: &Fields) -> TokenStream {
214 match fields {
215 Fields::Named(fields) => {
216 let recurse = fields.named.iter().map(|f| {
220 let name = &f.ident;
221 let ty = &f.ty;
222 quote_spanned! {f.span()=>
223 <#ty as Serializable>::serialize(#name, writer)?;
224 }
225 });
226 quote! {
227 #(#recurse)*
228 }
229 }
230 Fields::Unnamed(fields) => {
231 let recurse = fields.unnamed.iter().enumerate().map(|(i, f)| {
232 let name = Ident::new(&format!("var_{}", i), Span::call_site());
233 let ty = &f.ty;
234 quote_spanned! {f.span()=>
235 <#ty as Serializable>::serialize(#name, writer)?;
236 }
237 });
238 quote! {
239 #(#recurse)*
240 }
241 }
242 Fields::Unit => TokenStream::new(),
243 }
244}
245
246fn unpack_struct(fields: &Fields) -> TokenStream {
247 match fields {
251 Fields::Named(fields) => {
252 let recurse = fields.named.iter().map(|var| {
253 let name = &var.ident;
254 quote_spanned!(var.span()=>
255 let #name = &self.#name;
256 )
257 });
258 quote! {
259 #(#recurse)*
260 }
261 }
262 Fields::Unnamed(fields) => {
263 let recurse = fields.unnamed.iter().enumerate().map(|(i, var)| {
264 let name = Ident::new(&format!("var_{}", i), Span::call_site());
265 let index = Index::from(i);
266 quote_spanned!(var.span()=>
267 let #name = &self.#index;
268 )
269 });
270 quote! {
271 #(#recurse)*
272 }
273 }
274 Fields::Unit => TokenStream::new(),
275 }
276}
277
278fn unpack_enum(fields: &Fields) -> TokenStream {
279 match fields {
282 Fields::Named(fields) => {
283 let recurse = fields.named.iter().map(|var| {
284 let name = &var.ident;
285 quote_spanned!(var.span()=>
286 #name,
287 )
288 });
289 quote! {
290 {#(#recurse)*}
291 }
292 }
293 Fields::Unnamed(fields) => {
294 let recurse = fields.unnamed.iter().enumerate().map(|(i, var)| {
295 let name = Ident::new(&format!("var_{}", i), Span::call_site());
296 quote_spanned!(var.span()=>
297 #name,
298 )
299 });
300 quote! {
301 (#(#recurse)*)
302 }
303 }
304 Fields::Unit => TokenStream::new(),
305 }
306}
307
308fn serialize(data: &Data) -> TokenStream {
309 match *data {
310 Data::Struct(ref data) => {
311 let unpack = unpack_struct(&data.fields);
312 let fields = serialize_fields(&data.fields);
313 quote! {
314 #unpack #fields
315 }
316 }
317 Data::Enum(ref data) => {
318 let recurse = data.variants.iter().enumerate().map(|(i, var)| {
319 let fields = serialize_fields(&var.fields);
320 let unpack = unpack_enum(&var.fields);
321 let ty = &var.ident;
322 quote_spanned! {var.span()=>
323 Self::#ty #unpack => {
324 <u8 as Serializable>::serialize(&(#i as u8), writer)?;
325 #fields
326 },
327 }
328 });
329 quote! {
330 match self {
331 #(#recurse)*
332 }
333 }
334 }
335 Data::Union(_) => TokenStream::new(),
336 }
337}
338
339fn size_fields(fields: &Fields) -> TokenStream {
340 match fields {
341 Fields::Named(fields) => {
342 let recurse = fields.named.iter().map(|f| {
346 let name = &f.ident;
347 let ty = &f.ty;
348 quote_spanned! {f.span()=>
349 + <#ty as Serializable>::serialized_size(#name)
350 }
351 });
352 quote! {
353 0 #(#recurse)*
354 }
355 }
356 Fields::Unnamed(fields) => {
357 let recurse = fields.unnamed.iter().enumerate().map(|(i, f)| {
358 let name = Ident::new(&format!("var_{}", i), Span::call_site());
359 let ty = &f.ty;
360 quote_spanned! {f.span()=>
361 + <#ty as Serializable>::serialized_size(#name)
362 }
363 });
364 quote! {
365 0 #(#recurse)*
366 }
367 }
368 Fields::Unit => quote! { 0 },
369 }
370}
371
372fn size(data: &Data) -> TokenStream {
373 match *data {
374 Data::Struct(ref data) => {
375 let unpack = unpack_struct(&data.fields);
376 let fields = size_fields(&data.fields);
377 quote! {
378 #unpack #fields
379 }
380 }
381 Data::Enum(ref data) => {
382 let recurse = data.variants.iter().map(|var| {
383 let unpack = unpack_enum(&var.fields);
384 let fields = size_fields(&var.fields);
385 let ty = &var.ident;
386 quote_spanned! {var.span()=>
387 Self::#ty #unpack => {
388 1 + #fields
389 }
390 }
391 });
392 quote! {
393 match self {
394 #(#recurse)*
395 }
396 }
397 }
398 Data::Union(_) => unimplemented!(),
399 }
400}
401
402fn deserialize_fields(fields: &Fields) -> TokenStream {
403 match fields {
404 Fields::Named(fields) => {
405 let recurse = fields.named.iter().map(|f| {
409 let name = &f.ident;
410 let ty = &f.ty;
411 quote_spanned! {f.span()=>
412 #name: <#ty as Deserializable>::deserialize(reader, recursion_depth)?,
413 }
414 });
415 quote! {
416 {#(#recurse)*}
417 }
418 }
419 Fields::Unnamed(fields) => {
420 let recurse = fields.unnamed.iter().map(|f| {
421 let ty = &f.ty;
422 quote_spanned! {f.span()=>
423 <#ty as Deserializable>::deserialize(reader, recursion_depth)?,
424 }
425 });
426 quote! {
427 (#(#recurse)*)
428 }
429 }
430 Fields::Unit => quote! {},
431 }
432}
433
434fn deserialize(data: &Data) -> TokenStream {
435 match *data {
436 Data::Struct(ref data) => {
437 let fields = deserialize_fields(&data.fields);
438 quote! {
439 Ok(Self #fields)
440 }
441 }
442 Data::Enum(ref data) => {
443 let recurse = data.variants.iter().enumerate().map(|(i, var)| {
444 let i = i as u8;
445 let fields = deserialize_fields(&var.fields);
446 let name = &var.ident;
447 quote_spanned! {var.span()=>
448 #i => Ok(Self::#name #fields),
449 }
450 });
451 quote! {
452 let discriminant = <u8 as Deserializable>::deserialize(reader, recursion_depth)?;
453 match discriminant {
454 #(#recurse)*
455 _ => Err(::std::io::Error::new(::std::io::ErrorKind::InvalidData, "unrecognised discriminant"))
456 }
457 }
458 }
459 Data::Union(_) => unimplemented!(),
460 }
461}