1#![forbid(unsafe_code)]
2#![warn(missing_docs)]
3
4use 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)]
34pub 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))]
47pub 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)]
60pub 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))]
73pub 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}