1#![forbid(unsafe_code)]
2#![warn(missing_docs)]
3
4use 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)]
37pub 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))]
50pub 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)]
63pub 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))]
76pub 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}