Skip to main content

luanti_protocol_derive/
lib.rs

1//! Derive macros for luanti-protocol
2
3#![expect(
4    missing_docs,
5    // clippy::missing_panics_doc,
6    // clippy::missing_errors_doc,
7    clippy::expect_used,
8    clippy::unwrap_used,
9    clippy::unimplemented,
10    reason = "//TODO add documentation and improve error handling"
11)]
12
13use proc_macro2::Ident;
14use proc_macro2::Literal;
15use proc_macro2::TokenStream;
16use quote::ToTokens;
17use quote::quote;
18use quote::quote_spanned;
19use syn::Data;
20use syn::DeriveInput;
21use syn::Field;
22use syn::Generics;
23use syn::Index;
24use syn::Type;
25use syn::TypeParam;
26use syn::parse_macro_input;
27use syn::punctuated::Punctuated;
28use syn::spanned::Spanned;
29
30#[proc_macro_derive(LuantiSerialize, attributes(wrap))]
31pub fn luanti_serialize(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
32    let input = parse_macro_input!(input as DeriveInput);
33    let name = input.ident;
34    let serialize_body = make_serialize_body(&name, &input.data);
35
36    // The struct must include Serialize in the bounds of any type
37    // that need to be serializable.
38    let impl_generic = input.generics.to_token_stream();
39    let name_generic = strip_generic_bounds(&input.generics).to_token_stream();
40    let where_generic = input.generics.where_clause;
41
42    let expanded = quote! {
43        impl #impl_generic Serialize for #name #name_generic #where_generic {
44            type Input = Self;
45            fn serialize<S: Serializer>(value: &Self::Input, ser: &mut S) -> SerializeResult {
46                #serialize_body
47                Ok(())
48            }
49        }
50    };
51    proc_macro::TokenStream::from(expanded)
52}
53
54#[proc_macro_derive(LuantiDeserialize, attributes(wrap))]
55pub fn luanti_deserialize(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
56    let input = parse_macro_input!(input as DeriveInput);
57    let name = input.ident;
58    let deserialize_body = make_deserialize_body(&name, &input.data);
59
60    // The struct must include Deserialize in the bounds of any type
61    // that need to be serializable.
62    let impl_generic = input.generics.to_token_stream();
63    let name_generic = strip_generic_bounds(&input.generics).to_token_stream();
64    let where_generic = input.generics.where_clause;
65
66    let expanded = quote! {
67        impl #impl_generic Deserialize for #name #name_generic #where_generic {
68            type Output = Self;
69            fn deserialize(deser: &mut Deserializer) -> DeserializeResult<Self> {
70                #deserialize_body
71            }
72        }
73    };
74    proc_macro::TokenStream::from(expanded)
75}
76
77fn get_wrapped_type(field: &Field) -> Type {
78    let mut ty = field.ty.clone();
79    for attr in &field.attrs {
80        if attr.path().is_ident("wrap") {
81            ty = attr.parse_args::<Type>().unwrap();
82        }
83    }
84    ty
85}
86
87/// For struct, fields are serialized/deserialized in order.
88/// For enum, tags are assumed u8, consecutive, starting with 0.
89fn make_serialize_body(input_name: &Ident, data: &Data) -> TokenStream {
90    match *data {
91        Data::Struct(ref data) => match data.fields {
92            syn::Fields::Named(ref fields) => {
93                let recurse = fields.named.iter().map(|field| {
94                    let name = &field.ident;
95                    let ty = get_wrapped_type(field);
96                    quote_spanned! {field.span() =>
97                        <#ty as Serialize>::serialize(&value.#name, ser)?;
98                    }
99                });
100                quote! {
101                    #(#recurse)*
102                }
103            }
104            syn::Fields::Unnamed(ref fields) => {
105                let recurse = fields.unnamed.iter().enumerate().map(|(index, field)| {
106                    let index = Index::from(index);
107                    let ty = get_wrapped_type(field);
108                    quote_spanned! {field.span() =>
109                        <#ty as Serialize>::serialize(&value.#index, ser)?;
110                    }
111                });
112                quote! {
113                    #(#recurse)*
114                }
115            }
116            syn::Fields::Unit => {
117                quote! {}
118            }
119        },
120        Data::Enum(ref body) => {
121            let recurse = body.variants.iter().enumerate().map(|(index, variant)| {
122                if !variant.fields.is_empty() {
123                    quote_spanned! {variant.span() =>
124                        compile_error!("Cannot handle fields yet");
125                    }
126                } else if variant.discriminant.is_some() {
127                    quote_spanned! {variant.span() =>
128                        compile_error!("Cannot handle discriminant yet");
129                    }
130                } else {
131                    let id = &variant.ident;
132                    let i = Literal::u8_unsuffixed(
133                        u8::try_from(index).expect("variant index exceeds range of u8"),
134                    );
135                    quote_spanned! {variant.span() =>
136                        #id => #i,
137                    }
138                }
139            });
140            quote! {
141                    use #input_name::*;
142                    let tag = match value {
143                        #(#recurse)*
144                    };
145                    u8::serialize(&tag, ser)?;
146            }
147        }
148        Data::Union(_) => unimplemented!(),
149    }
150}
151
152fn make_deserialize_body(input_name: &Ident, data: &Data) -> TokenStream {
153    match *data {
154        Data::Struct(ref data) => match data.fields {
155            syn::Fields::Named(ref fields) => {
156                let assignments = fields.named.iter().map(|field| {
157                    let name = &field.ident;
158                    let ty = get_wrapped_type(field);
159                    quote_spanned! {field.span() =>
160                        log::trace!(stringify!("deserializing field", #input_name, #name));
161                        #[allow(unused_qualifications)]
162                        let #name = anyhow::Context::context(<#ty as Deserialize>::deserialize(deser), stringify!("failed to deserialize field", #input_name, #name))?;
163
164                        log::trace!("result: {:?} - {} bytes left", #name, deser.remaining());
165                    }
166                });
167                let fields = fields.named.iter().map(|field| {
168                    let name = &field.ident;
169                    quote_spanned! { field.span() => #name, }
170                });
171                quote! {
172                    #(#assignments)*
173                    Ok(Self { #(#fields)* })
174                }
175            }
176            syn::Fields::Unnamed(ref fields) => {
177                let recurse = fields.unnamed.iter().enumerate().map(|(index, field)| {
178                    let index = Index::from(index);
179                    let ty = get_wrapped_type(field);
180                    quote_spanned! {field.span() =>
181                        #index: <#ty as Deserialize>::deserialize(deser)?,
182                    }
183                });
184                let inner = quote! {
185                    #(#recurse)*
186                };
187                quote! {
188                    Ok(Self {
189                        #inner
190                    })
191                }
192            }
193            syn::Fields::Unit => {
194                let inner = quote! {};
195                quote! {
196                    Ok(Self {
197                        #inner
198                    })
199                }
200            }
201        },
202        Data::Enum(ref body) => {
203            let recurse = body.variants.iter().enumerate().map(|(index, variant)| {
204                if !variant.fields.is_empty() {
205                    quote_spanned! {variant.span() =>
206                        compile_error!("Cannot handle fields yet");
207                    }
208                } else if variant.discriminant.is_some() {
209                    quote_spanned! {variant.span() =>
210                        compile_error!("Cannot handle discriminant yet");
211                    }
212                } else {
213                    let id = &variant.ident;
214                    let i = Literal::u8_unsuffixed(
215                        u8::try_from(index).expect("variant index exceeds range of u8"),
216                    );
217                    quote_spanned! {variant.span() =>
218                        #i => #id,
219
220                    }
221                }
222            });
223
224            let input_name_str = Literal::string(&input_name.to_string());
225            quote! {
226                    use #input_name::*;
227                    let tag = u8::deserialize(deser)?;
228                    Ok(match tag {
229                        #(#recurse)*
230                        _ => bail!("Invalid {} tag: {}", #input_name_str, tag),
231                    })
232            }
233        }
234        Data::Union(_) => unimplemented!(),
235    }
236}
237
238/// Converts <T: Trait, S: Trait2> into <T, S>
239fn strip_generic_bounds(input: &Generics) -> Generics {
240    let input = input.clone();
241    Generics {
242        lt_token: input.lt_token,
243        params: {
244            let mut params = input.params.clone();
245            params.iter_mut().for_each(|param| {
246                *param = match param.clone() {
247                    syn::GenericParam::Type(param) => syn::GenericParam::Type(TypeParam {
248                        attrs: Vec::new(),
249                        ident: param.ident.clone(),
250                        colon_token: None,
251                        bounds: Punctuated::new(),
252                        eq_token: None,
253                        default: None,
254                    }),
255                    any => any,
256                }
257            });
258            params
259        },
260        gt_token: input.gt_token,
261        where_clause: None,
262    }
263}