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