1#![forbid(unsafe_code)]
2#![warn(missing_docs)]
3
4use 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)]
19pub 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))]
32pub 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)]
45pub 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))]
58pub 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}