luanti_protocol_derive/
lib.rs1#![expect(
4 missing_docs,
5 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 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 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
87fn 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
238fn 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}