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