Skip to main content

rustbinary_derive/
lib.rs

1#![forbid(unsafe_code)]
2#![warn(missing_docs)]
3
4//! Procedural derives used by `rustbinary`.
5//!
6//! Applications normally enable the matching `rustbinary` feature and use the
7//! derives re-exported by the main crate. Generated paths intentionally refer
8//! to `::rustbinary`, keeping the runtime traits and wire implementation owned
9//! by one crate.
10
11use proc_macro::TokenStream;
12use quote::{quote, ToTokens};
13use syn::{
14    parse_macro_input, parse_quote, Data, DataEnum, DataStruct, DeriveInput, Expr, Fields,
15    Generics, Lit, Meta, Type,
16};
17
18#[proc_macro_derive(Fingerprint)]
19/// Derives `rustbinary::Fingerprint` from structural type metadata.
20///
21/// Struct field names/types and enum variant names/order participate in the
22/// identifier. Unions are rejected. Generic type parameters receive the
23/// corresponding `Fingerprint` bound.
24pub fn derive_fingerprint(input: TokenStream) -> TokenStream {
25    let input = parse_macro_input!(input as DeriveInput);
26    fingerprint_impl(&input)
27        .unwrap_or_else(syn::Error::into_compile_error)
28        .into()
29}
30
31#[proc_macro_derive(StaticSize, attributes(bits))]
32/// Derives compile-time normal and bit-packed size bounds.
33///
34/// Fields must implement `rustbinary::StaticSize`. An optional `#[bits = N]`
35/// attribute contributes the explicit packed width. Dynamic collections do not
36/// provide a finite `StaticSize` implementation.
37pub fn derive_static_size(input: TokenStream) -> TokenStream {
38    let input = parse_macro_input!(input as DeriveInput);
39    static_size_impl(&input)
40        .unwrap_or_else(syn::Error::into_compile_error)
41        .into()
42}
43
44#[proc_macro_derive(Reflect)]
45/// Derives allocation-free structural reflection metadata.
46///
47/// Generated metadata contains the declared type name, fields, field type
48/// tokens, declaration indexes, and enum variants. No registry or runtime
49/// initialization is generated.
50pub fn derive_reflect(input: TokenStream) -> TokenStream {
51    let input = parse_macro_input!(input as DeriveInput);
52    reflect_impl(&input)
53        .unwrap_or_else(syn::Error::into_compile_error)
54        .into()
55}
56
57#[proc_macro_derive(BitPacked, attributes(bits))]
58/// Derives `rustbinary::BitPack` for structs and enums.
59///
60/// Fields with `#[bits = N]` use `BitValue` range validation. Other fields
61/// recursively use `BitPack`. Enum tags use the minimum bit width and unknown
62/// decoded tags are rejected.
63pub fn derive_bit_packed(input: TokenStream) -> TokenStream {
64    let input = parse_macro_input!(input as DeriveInput);
65    bit_packed_impl(&input)
66        .unwrap_or_else(syn::Error::into_compile_error)
67        .into()
68}
69
70fn add_bound(mut generics: Generics, bound: syn::Path) -> Generics {
71    for parameter in generics.type_params_mut() {
72        parameter.bounds.push(parse_quote!(#bound));
73    }
74    generics
75}
76
77fn field_name(index: usize, field: &syn::Field) -> String {
78    field
79        .ident
80        .as_ref()
81        .map_or_else(|| index.to_string(), ToString::to_string)
82}
83
84fn hash_field(index: usize, field: &syn::Field) -> proc_macro2::TokenStream {
85    let name = field_name(index, field);
86    let ty = &field.ty;
87    quote! {
88        hash = ::rustbinary::schema::hash_bytes(hash, #name.as_bytes());
89        hash = ::rustbinary::schema::hash_u64(hash, <#ty as ::rustbinary::Fingerprint>::TYPE_FINGERPRINT);
90    }
91}
92
93fn fingerprint_impl(input: &DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
94    let name = &input.ident;
95    let generics = add_bound(
96        input.generics.clone(),
97        parse_quote!(::rustbinary::Fingerprint),
98    );
99    let (impl_generics, type_generics, where_clause) = generics.split_for_impl();
100    let body = match &input.data {
101        Data::Struct(data) => {
102            let fields = data
103                .fields
104                .iter()
105                .enumerate()
106                .map(|(index, field)| hash_field(index, field));
107            quote! {
108                let mut hash = ::rustbinary::schema::hash_bytes(
109                    ::rustbinary::schema::FNV_OFFSET,
110                    concat!(module_path!(), "::", stringify!(#name), "|struct").as_bytes(),
111                );
112                #(#fields)*
113                hash
114            }
115        }
116        Data::Enum(data) => fingerprint_enum(name, data),
117        Data::Union(_) => {
118            return Err(syn::Error::new_spanned(
119                input,
120                "Fingerprint cannot be derived for unions",
121            ))
122        }
123    };
124    Ok(quote! {
125        impl #impl_generics ::rustbinary::Fingerprint for #name #type_generics #where_clause {
126            const TYPE_FINGERPRINT: u64 = { #body };
127        }
128    })
129}
130
131fn fingerprint_enum(name: &syn::Ident, data: &DataEnum) -> proc_macro2::TokenStream {
132    let variants = data
133        .variants
134        .iter()
135        .enumerate()
136        .map(|(variant_index, variant)| {
137            let variant_name = variant.ident.to_string();
138            let index = variant_index as u64;
139            let fields = variant
140                .fields
141                .iter()
142                .enumerate()
143                .map(|(field_index, field)| hash_field(field_index, field));
144            quote! {
145                hash = ::rustbinary::schema::hash_u64(hash, #index);
146                hash = ::rustbinary::schema::hash_bytes(hash, #variant_name.as_bytes());
147                #(#fields)*
148            }
149        });
150    quote! {
151        let mut hash = ::rustbinary::schema::hash_bytes(
152            ::rustbinary::schema::FNV_OFFSET,
153            concat!(module_path!(), "::", stringify!(#name), "|enum").as_bytes(),
154        );
155        #(#variants)*
156        hash
157    }
158}
159
160fn static_field_size(field: &syn::Field, packed: bool) -> proc_macro2::TokenStream {
161    let ty = &field.ty;
162    if packed {
163        quote!(<#ty as ::rustbinary::StaticSize>::PACKED_MAX_SIZE)
164    } else {
165        quote!(<#ty as ::rustbinary::StaticSize>::MAX_SIZE)
166    }
167}
168
169fn static_field_bits(field: &syn::Field) -> syn::Result<proc_macro2::TokenStream> {
170    let ty = &field.ty;
171    Ok(match declared_bits(field)? {
172        Some(width) => quote!(#width),
173        None => quote!(<#ty as ::rustbinary::StaticSize>::PACKED_MAX_BITS),
174    })
175}
176
177fn sum_field_bits(fields: &Fields) -> syn::Result<proc_macro2::TokenStream> {
178    let mut sum = quote!(0usize);
179    for field in fields {
180        let bits = static_field_bits(field)?;
181        sum = quote!(::rustbinary::static_size::saturating_add(#sum, #bits));
182    }
183    Ok(sum)
184}
185
186fn sum_fields(fields: &Fields, packed: bool) -> proc_macro2::TokenStream {
187    fields.iter().fold(quote!(0usize), |sum, field| {
188        let size = static_field_size(field, packed);
189        quote!(::rustbinary::static_size::saturating_add(#sum, #size))
190    })
191}
192
193fn max_variants(data: &DataEnum, packed: bool) -> proc_macro2::TokenStream {
194    data.variants
195        .iter()
196        .fold(quote!(0usize), |maximum, variant| {
197            let size = sum_fields(&variant.fields, packed);
198            quote!(::rustbinary::static_size::max(#maximum, #size))
199        })
200}
201
202fn static_size_impl(input: &DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
203    let name = &input.ident;
204    let generics = add_bound(
205        input.generics.clone(),
206        parse_quote!(::rustbinary::StaticSize),
207    );
208    let (impl_generics, type_generics, where_clause) = generics.split_for_impl();
209    let (maximum, packed_bits) = match &input.data {
210        Data::Struct(data) => (
211            sum_fields(&data.fields, false),
212            sum_field_bits(&data.fields)?,
213        ),
214        Data::Enum(data) => {
215            let maximum = max_variants(data, false);
216            let tag_bits = if data.variants.len() <= 1 {
217                0usize
218            } else {
219                (usize::BITS - (data.variants.len() - 1).leading_zeros()) as usize
220            };
221            let mut packed_payload = quote!(0usize);
222            for variant in &data.variants {
223                let bits = sum_field_bits(&variant.fields)?;
224                packed_payload = quote!(::rustbinary::static_size::max(#packed_payload, #bits));
225            }
226            (
227                quote!(::rustbinary::static_size::saturating_add(5, #maximum)),
228                quote!(::rustbinary::static_size::saturating_add(#tag_bits, #packed_payload)),
229            )
230        }
231        Data::Union(_) => {
232            return Err(syn::Error::new_spanned(
233                input,
234                "StaticSize cannot be derived for unions",
235            ))
236        }
237    };
238    Ok(quote! {
239        impl #impl_generics ::rustbinary::StaticSize for #name #type_generics #where_clause {
240            const MAX_SIZE: usize = #maximum;
241            const PACKED_MAX_BITS: usize = #packed_bits;
242            const PACKED_MAX_SIZE: usize = ::rustbinary::static_size::bytes_for_bits(#packed_bits);
243        }
244    })
245}
246
247fn type_name(ty: &Type) -> String {
248    ty.to_token_stream().to_string().replace(' ', "")
249}
250
251fn reflect_fields(fields: &Fields) -> proc_macro2::TokenStream {
252    let descriptors = fields.iter().enumerate().map(|(index, field)| {
253        let name = field_name(index, field);
254        let ty = type_name(&field.ty);
255        quote!(::rustbinary::FieldInfo { name: #name, type_name: #ty, index: #index })
256    });
257    quote!(&[#(#descriptors),*])
258}
259
260fn reflect_impl(input: &DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
261    let name = &input.ident;
262    let generics = input.generics.clone();
263    let (impl_generics, type_generics, where_clause) = generics.split_for_impl();
264    let shape = match &input.data {
265        Data::Struct(DataStruct { fields, .. }) => {
266            let fields = reflect_fields(fields);
267            quote!(::rustbinary::TypeShape::Struct(#fields))
268        }
269        Data::Enum(data) => {
270            let variants = data.variants.iter().enumerate().map(|(index, variant)| {
271                let variant_name = variant.ident.to_string();
272                let fields = reflect_fields(&variant.fields);
273                quote!(::rustbinary::VariantInfo { name: #variant_name, index: #index, fields: #fields })
274            });
275            quote!(::rustbinary::TypeShape::Enum(&[#(#variants),*]))
276        }
277        Data::Union(_) => {
278            return Err(syn::Error::new_spanned(
279                input,
280                "Reflect cannot be derived for unions",
281            ))
282        }
283    };
284    Ok(quote! {
285        impl #impl_generics ::rustbinary::Reflect for #name #type_generics #where_clause {
286            const TYPE_NAME: &'static str = concat!(module_path!(), "::", stringify!(#name));
287            const SHAPE: ::rustbinary::TypeShape = #shape;
288        }
289    })
290}
291
292fn declared_bits(field: &syn::Field) -> syn::Result<Option<usize>> {
293    let Some(attribute) = field
294        .attrs
295        .iter()
296        .find(|attribute| attribute.path().is_ident("bits"))
297    else {
298        return Ok(None);
299    };
300    let value = match &attribute.meta {
301        Meta::NameValue(name_value) => match &name_value.value {
302            Expr::Lit(expression) => match &expression.lit {
303                Lit::Int(value) => value.base10_parse()?,
304                _ => {
305                    return Err(syn::Error::new_spanned(
306                        expression,
307                        "bits must be an integer",
308                    ))
309                }
310            },
311            expression => {
312                return Err(syn::Error::new_spanned(
313                    expression,
314                    "bits must be an integer",
315                ))
316            }
317        },
318        Meta::List(_) => attribute.parse_args::<syn::LitInt>()?.base10_parse()?,
319        Meta::Path(_) => return Err(syn::Error::new_spanned(attribute, "use #[bits = N]")),
320    };
321    if value == 0 || value > 128 {
322        return Err(syn::Error::new_spanned(
323            attribute,
324            "bit width must be between 1 and 128",
325        ));
326    }
327    Ok(Some(value))
328}
329
330fn add_bit_bounds(
331    mut generics: Generics,
332    fields: impl Iterator<Item = syn::Field>,
333) -> syn::Result<Generics> {
334    let where_clause = generics.make_where_clause();
335    for field in fields {
336        let has_declared_bits = declared_bits(&field)?.is_some();
337        let ty = field.ty;
338        if has_declared_bits {
339            where_clause
340                .predicates
341                .push(parse_quote!(#ty: ::rustbinary::BitValue));
342        } else {
343            where_clause
344                .predicates
345                .push(parse_quote!(#ty: ::rustbinary::BitPack));
346        }
347    }
348    Ok(generics)
349}
350
351fn bit_count(fields: &Fields) -> syn::Result<proc_macro2::TokenStream> {
352    let mut total = quote!(0usize);
353    for field in fields {
354        let ty = &field.ty;
355        let bits = match declared_bits(field)? {
356            Some(width) => quote!(#width),
357            None => quote!(<#ty as ::rustbinary::BitPack>::MAX_BITS),
358        };
359        total = quote!(#total.saturating_add(#bits));
360    }
361    Ok(total)
362}
363
364fn pack_statement(
365    field: &syn::Field,
366    value: proc_macro2::TokenStream,
367    borrowed: bool,
368) -> syn::Result<proc_macro2::TokenStream> {
369    let ty = &field.ty;
370    Ok(match declared_bits(field)? {
371        Some(width) => {
372            let value = if borrowed { quote!(*#value) } else { value };
373            quote! {
374                writer.write(<#ty as ::rustbinary::BitValue>::encode_bits(#value, #width)?, #width)?;
375            }
376        }
377        None => quote! {
378            <#ty as ::rustbinary::BitPack>::pack(#value, writer)?;
379        },
380    })
381}
382
383fn unpack_expression(field: &syn::Field) -> syn::Result<proc_macro2::TokenStream> {
384    let ty = &field.ty;
385    Ok(match declared_bits(field)? {
386        Some(width) => quote! {
387            <#ty as ::rustbinary::BitValue>::decode_bits(reader.read(#width)?, #width)?
388        },
389        None => quote! {
390            <#ty as ::rustbinary::BitPack>::unpack(reader)?
391        },
392    })
393}
394
395fn bit_packed_impl(input: &DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
396    let name = &input.ident;
397    let fields = match &input.data {
398        Data::Struct(data) => data.fields.iter().cloned().collect::<Vec<_>>(),
399        Data::Enum(data) => data
400            .variants
401            .iter()
402            .flat_map(|variant| variant.fields.iter().cloned())
403            .collect(),
404        Data::Union(_) => {
405            return Err(syn::Error::new_spanned(
406                input,
407                "BitPacked cannot be derived for unions",
408            ))
409        }
410    };
411    let generics = add_bit_bounds(input.generics.clone(), fields.into_iter())?;
412    let (impl_generics, type_generics, where_clause) = generics.split_for_impl();
413    let (maximum, pack, unpack) = match &input.data {
414        Data::Struct(data) => bit_packed_struct(name, data)?,
415        Data::Enum(data) => bit_packed_enum(name, data)?,
416        Data::Union(_) => unreachable!(),
417    };
418    Ok(quote! {
419        impl #impl_generics ::rustbinary::BitPack for #name #type_generics #where_clause {
420            const MAX_BITS: usize = #maximum;
421            fn pack(&self, writer: &mut ::rustbinary::BitWriter<'_>) -> ::rustbinary::Result<()> {
422                #pack
423                Ok(())
424            }
425            fn unpack(reader: &mut ::rustbinary::BitReader<'_>) -> ::rustbinary::Result<Self> {
426                #unpack
427            }
428        }
429    })
430}
431
432fn bit_packed_struct(
433    name: &syn::Ident,
434    data: &DataStruct,
435) -> syn::Result<(
436    proc_macro2::TokenStream,
437    proc_macro2::TokenStream,
438    proc_macro2::TokenStream,
439)> {
440    let maximum = bit_count(&data.fields)?;
441    let mut packs = Vec::new();
442    let mut values = Vec::new();
443    for (index, field) in data.fields.iter().enumerate() {
444        let member = field
445            .ident
446            .clone()
447            .map(syn::Member::Named)
448            .unwrap_or_else(|| syn::Member::Unnamed(syn::Index::from(index)));
449        packs.push(pack_statement(field, quote!(&self.#member), true)?);
450        values.push(unpack_expression(field)?);
451    }
452    let construct = match &data.fields {
453        Fields::Named(fields) => {
454            let names = fields
455                .named
456                .iter()
457                .map(|field| field.ident.as_ref().expect("named"));
458            quote!(#name { #(#names: #values),* })
459        }
460        Fields::Unnamed(_) => quote!(#name(#(#values),*)),
461        Fields::Unit => quote!(#name),
462    };
463    Ok((maximum, quote!(#(#packs)*), quote!(Ok(#construct))))
464}
465
466fn bit_packed_enum(
467    name: &syn::Ident,
468    data: &DataEnum,
469) -> syn::Result<(
470    proc_macro2::TokenStream,
471    proc_macro2::TokenStream,
472    proc_macro2::TokenStream,
473)> {
474    if data.variants.is_empty() {
475        return Err(syn::Error::new_spanned(
476            name,
477            "empty enums cannot be bit-packed",
478        ));
479    }
480    let tag_bits = if data.variants.len() <= 1 {
481        0usize
482    } else {
483        (usize::BITS - (data.variants.len() - 1).leading_zeros()) as usize
484    };
485    let mut maximum = quote!(0usize);
486    let mut pack_arms = Vec::new();
487    let mut unpack_arms = Vec::new();
488    for (variant_index, variant) in data.variants.iter().enumerate() {
489        let variant_name = &variant.ident;
490        let payload_bits = bit_count(&variant.fields)?;
491        maximum = quote!(::rustbinary::__bitpack_max(#maximum, #payload_bits));
492        let bindings = (0..variant.fields.len())
493            .map(|index| syn::Ident::new(&format!("field_{index}"), variant.ident.span()))
494            .collect::<Vec<_>>();
495        let pattern = match &variant.fields {
496            Fields::Named(fields) => {
497                let names = fields
498                    .named
499                    .iter()
500                    .map(|field| field.ident.as_ref().expect("named"));
501                quote!(Self::#variant_name { #(#names: #bindings),* })
502            }
503            Fields::Unnamed(_) => quote!(Self::#variant_name(#(#bindings),*)),
504            Fields::Unit => quote!(Self::#variant_name),
505        };
506        let packs = variant
507            .fields
508            .iter()
509            .zip(&bindings)
510            .map(|(field, binding)| pack_statement(field, quote!(#binding), true))
511            .collect::<syn::Result<Vec<_>>>()?;
512        pack_arms.push(quote! {
513            #pattern => {
514                writer.write(#variant_index as u128, #tag_bits)?;
515                #(#packs)*
516            }
517        });
518        let values = variant
519            .fields
520            .iter()
521            .map(unpack_expression)
522            .collect::<syn::Result<Vec<_>>>()?;
523        let construct = match &variant.fields {
524            Fields::Named(fields) => {
525                let names = fields
526                    .named
527                    .iter()
528                    .map(|field| field.ident.as_ref().expect("named"));
529                quote!(Self::#variant_name { #(#names: #values),* })
530            }
531            Fields::Unnamed(_) => quote!(Self::#variant_name(#(#values),*)),
532            Fields::Unit => quote!(Self::#variant_name),
533        };
534        unpack_arms.push(quote!(#variant_index => Ok(#construct)));
535    }
536    Ok((
537        quote!((#tag_bits).saturating_add(#maximum)),
538        quote!(match self { #(#pack_arms),* }),
539        quote! {
540            match reader.read(#tag_bits)? as usize {
541                #(#unpack_arms,)*
542                _ => Err(::rustbinary::Error::BitPacking("unknown packed enum variant")),
543            }
544        },
545    ))
546}