Skip to main content

ndn_tlv_derive/
lib.rs

1use quote::quote;
2use syn::{Data, Field, Fields, GenericParam};
3
4#[derive(deluxe::ParseMetaItem)]
5struct TlvAttrKW {
6    #[deluxe(default)]
7    internal: bool,
8}
9
10#[derive(deluxe::ExtractAttributes)]
11#[deluxe(attributes(tlv))]
12struct TlvAttr(#[deluxe(default)] usize, #[deluxe(flatten)] TlvAttrKW);
13
14#[derive(deluxe::ExtractAttributes)]
15struct TlvFieldAttr {
16    #[deluxe(default)]
17    default: bool,
18}
19
20fn decode_generics(
21    crate_name: &proc_macro2::TokenStream,
22    generics: syn::Generics,
23) -> (
24    proc_macro2::TokenStream,
25    proc_macro2::TokenStream,
26    proc_macro2::TokenStream,
27    proc_macro2::TokenStream,
28) {
29    let params: Vec<_> = generics
30        .params
31        .iter()
32        .filter_map(|x| {
33            if let GenericParam::Type(typ) = x {
34                let mut typ = typ.clone();
35                typ.eq_token = None;
36                typ.default = None;
37                Some(GenericParam::Type(typ))
38            } else {
39                None
40            }
41        })
42        .collect();
43    if params.len() > 0 {
44        (
45            quote! {
46                <#( #params ),*>
47            },
48            quote! {
49                where #(#params: #crate_name::TlvEncode, #params: #crate_name::TlvDecode),*
50            },
51            quote! {
52                where #(#params: #crate_name::TlvEncode),*
53            },
54            quote! {
55                where #(#params: #crate_name::TlvEncode),*
56            },
57        )
58    } else {
59        (quote! {}, quote! {}, quote! {}, quote! {})
60    }
61}
62
63fn derive_struct(
64    fields: Vec<Field>,
65    crate_name: proc_macro2::TokenStream,
66    named: bool,
67    derivee: proc_macro2::Ident,
68    generics: syn::Generics,
69    typ: usize,
70) -> proc_macro::TokenStream {
71    let (generic_args, decode_where, encode_where, tlv_where) =
72        decode_generics(&crate_name, generics);
73
74    let mut field_names = Vec::with_capacity(fields.len());
75
76    let impls = {
77        let mut initialisers = Vec::with_capacity(fields.len());
78        for (i, field) in fields.iter().enumerate() {
79            let ty = &field.ty;
80            if let Some(ref ident) = field.ident {
81                initialisers.push(quote! {
82                    #ident: <#ty as #crate_name::TlvDecode>::decode(&mut inner_data)?
83                });
84                field_names.push(quote!(#ident));
85            } else {
86                initialisers.push(quote! {
87                    <#ty as #crate_name::TlvDecode>::decode(&mut inner_data)?
88                });
89                let idx = syn::Index::from(i);
90                field_names.push(quote!(#idx));
91            }
92        }
93
94        let initialiser = if named {
95            quote! {Ok(Self { #(#initialisers,)* })}
96        } else {
97            quote! {
98                Ok(Self (#(#initialisers,)*))
99            }
100        };
101
102        let decode_impl = if typ == 0 {
103            quote! {
104                impl #generic_args #crate_name::TlvDecode for #derivee #generic_args #decode_where {
105                    fn decode(bytes: &mut #crate_name::bytes::Bytes) -> #crate_name::Result<Self> {
106                        let mut inner_data = bytes;
107                        #initialiser
108                    }
109                }
110            }
111        } else {
112            quote! {
113                impl #generic_args #crate_name::TlvDecode for #derivee #generic_args #decode_where {
114                    fn decode(bytes: &mut #crate_name::bytes::Bytes) -> #crate_name::Result<Self> {
115                        use #crate_name::bytes::Buf;
116                        #crate_name::find_tlv::<Self>(bytes, true)?;
117                        let _ = #crate_name::VarNum::decode(bytes)?;
118                        let length = #crate_name::VarNum::decode(bytes)?;
119                        if bytes.remaining() < length.into() {
120                            return Err(#crate_name::TlvError::UnexpectedEndOfStream);
121                        }
122                        let mut inner_data = bytes.split_to(length.into());
123
124                        #initialiser
125                    }
126                }
127            }
128        };
129
130        let encode_impl = {
131            let encode_header = if typ == 0 {
132                quote! {}
133            } else {
134                quote! {
135                    bytes.put(#crate_name::VarNum::from(Self::TYP).encode());
136                    bytes.put(#crate_name::VarNum::from(self.inner_size()).encode());
137                }
138            };
139
140            let size_header = if typ == 0 {
141                quote! {0}
142            } else {
143                quote! {
144                    #crate_name::VarNum::from(Self::TYP).size()
145                        + #crate_name::VarNum::from(self.inner_size()).size()
146                }
147            };
148            quote! {
149                impl #generic_args #crate_name::TlvEncode for #derivee #generic_args #encode_where {
150                    fn encode(&self) -> #crate_name::bytes::Bytes {
151                        use #crate_name::bytes::BufMut;
152                        let mut bytes = #crate_name::bytes::BytesMut::with_capacity(self.size());
153
154                        #encode_header
155                        #(
156                            bytes.put(self.#field_names.encode());
157                            )*
158
159                        bytes.freeze()
160                    }
161
162                    fn size(&self) -> usize {
163                        #size_header
164                            #(+ self.#field_names.size())*
165                    }
166                }
167            }
168        };
169
170        quote! {
171            #decode_impl
172            #encode_impl
173        }
174    };
175
176    let tlv_impl = if typ != 0 {
177        quote! {
178            impl #generic_args #crate_name::Tlv for #derivee #generic_args #tlv_where {
179                const TYP: usize = #typ;
180
181                fn inner_size(&self) -> usize {
182                    0 #(+ #crate_name::TlvEncode::size(&self.#field_names) )*
183                }
184            }
185        }
186    } else {
187        quote! {}
188    };
189
190    quote! {
191        #tlv_impl
192
193        #impls
194    }
195    .into()
196}
197
198#[proc_macro_derive(Tlv, attributes(tlv))]
199pub fn derive(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
200    let mut input = syn::parse2::<syn::DeriveInput>(input.into()).unwrap();
201
202    let TlvAttr(typ, kw) = deluxe::extract_attributes(&mut input).unwrap();
203
204    let derivee = input.ident;
205    let crate_name = if kw.internal {
206        quote! {crate}
207    } else {
208        quote! {::ndn_tlv}
209    };
210
211    match input.data {
212        Data::Union(_) => panic!("Deriving Tlv on Unions is not supported"),
213        Data::Struct(struct_data) => match struct_data.fields {
214            Fields::Unit => {
215                derive_struct(Vec::new(), crate_name, true, derivee, input.generics, typ)
216            }
217            Fields::Unnamed(unnamed_fields) => {
218                let mut fields = Vec::with_capacity(unnamed_fields.unnamed.len());
219                fields.extend(unnamed_fields.unnamed);
220                derive_struct(fields, crate_name, false, derivee, input.generics, typ)
221            }
222            Fields::Named(named_fields) => {
223                let mut fields = Vec::with_capacity(named_fields.named.len());
224                fields.extend(named_fields.named);
225                derive_struct(fields, crate_name, true, derivee, input.generics, typ)
226            }
227        },
228        Data::Enum(enm) => {
229            if typ != 0 {
230                panic!("Enums cannot have a TLV Type");
231            }
232            let mut variants = Vec::with_capacity(enm.variants.len());
233            let mut fields = Vec::with_capacity(enm.variants.len());
234            let mut default_variant = None;
235
236            let (generic_args, decode_where, encode_where, _tlv_where) =
237                decode_generics(&crate_name, input.generics);
238
239            for mut variant in enm.variants {
240                let attrs: TlvFieldAttr = deluxe::extract_attributes(&mut variant).unwrap();
241                if attrs.default {
242                    assert!(default_variant.is_none());
243                    default_variant = Some(variant.ident);
244                } else {
245                    variants.push(variant.ident);
246                }
247
248                if variant.fields.len() != 1 || !matches!(variant.fields, syn::Fields::Unnamed(_)) {
249                    panic!("Enum variants must have exactly 1 unnamed field");
250                }
251
252                if !attrs.default {
253                    fields.push(variant.fields.iter().next().unwrap().ty.clone());
254                }
255            }
256
257            let decode_default = {
258                if let Some(ref variant) = default_variant {
259                    quote! {
260                        _ => Ok(Self::#variant(#variant::decode(bytes)?)),
261                    }
262                } else {
263                    quote! {
264                        _ => Err(#crate_name::TlvError::TypeMismatch {
265                            expected: 0, // TODO
266                            found: typ.into(),
267                        }),
268                    }
269                }
270            };
271
272            let encode_default_encode = {
273                if let Some(ref variant) = default_variant {
274                    quote! {
275                        Self::#variant(x) => x.encode(),
276                    }
277                } else {
278                    quote! {}
279                }
280            };
281
282            let encode_default_size = {
283                if let Some(ref variant) = default_variant {
284                    quote! {
285                        Self::#variant(x) => x.size(),
286                    }
287                } else {
288                    quote! {}
289                }
290            };
291
292            quote! {
293                impl #generic_args #crate_name::TlvDecode for #derivee #generic_args #decode_where {
294                    fn decode(bytes: &mut #crate_name::bytes::Bytes) -> #crate_name::Result<Self> {
295                        let mut cur = bytes.clone();
296
297                        let typ = #crate_name::VarNum::decode(&mut cur)?;
298                        match typ.into() {
299                            #(
300                            <#fields>::TYP => Ok(Self::#variants(
301                                <#fields>::decode(bytes)?,
302                            )),
303                            )*
304                            #decode_default
305                        }
306                    }
307                }
308
309                impl #generic_args #crate_name::TlvEncode for #derivee #generic_args #encode_where {
310                    fn encode(&self) -> #crate_name::bytes::Bytes {
311                        match self {
312                            #(
313                            Self::#variants(x) => x.encode(),
314                            #encode_default_encode
315                            )*
316                        }
317                    }
318
319                    fn size(&self) -> usize {
320                        match self {
321                            #(
322                                Self::#variants(x) => x.size(),
323                                #encode_default_size
324                            )*
325                        }
326                    }
327                }
328            }
329            .into()
330        }
331    }
332}