Skip to main content

midnight_serialize_macros/
lib.rs

1// This file is part of midnight-ledger.
2// Copyright (C) 2025 Midnight Foundation
3// SPDX-License-Identifier: Apache-2.0
4// Licensed under the Apache License, Version 2.0 (the "License");
5// You may not use this file except in compliance with the License.
6// You may obtain a copy of the License at
7// http://www.apache.org/licenses/LICENSE-2.0
8// Unless required by applicable law or agreed to in writing, software
9// distributed under the License is distributed on an "AS IS" BASIS,
10// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
11// See the License for the specific language governing permissions and
12// limitations under the License.
13
14//! Derive macros for `midnight-serialize`.
15extern 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// Macro to implement Serializable for an object made up entirely of serializable
49// objects
50#[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            // Expands to an expression like
217            //      <A as Serializable>::serialize(a, writer)?;
218            //      <B as Serializable>::serialize(b, writer)?;
219            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    // Expands to
248    // let a = &self.a;
249    // let b = &self.b;
250    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    // Expands to
280    // (a, b)
281    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            // Expands to an expression like
343            //      0 + <A as Serializable>::serialized_size(a)
344            //      + <B as Serializable>::serialized_size(b)
345            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            // Expands to an expression like
406            //      a: <A as Deserializable>::deserialize(reader)?,
407            //      b: <B as Deserializable>::deserialize(reader)?,
408            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}