use proc_macro2::{Span, TokenStream};
use quote::quote;
use syn::{
parse::Parser, punctuated::Punctuated, spanned::Spanned, Attribute, Data,
DeriveInput, Expr, ExprLit, Fields, Lit, Meta, Path, Token,
};
pub fn derive_compact_repr(input: &DeriveInput) -> TokenStream {
let name = &input.ident;
let Data::Enum(data_enum) = &input.data else {
return err_span(
input.ident.span(),
"#[derive(CompactRepr)] only supports enums".to_string(),
);
};
let Some(_repr_int) = find_unsigned_repr(&input.attrs) else {
return err_span(
name.span(),
"#[derive(CompactRepr)] requires an unsigned integer repr \
(one of #[repr(u8)], #[repr(u16)], #[repr(u32)], #[repr(u64)] or \
#[repr(usize)])"
.to_string(),
);
};
let mut pairs: Vec<(Ident, u128)> = Vec::new();
let mut next: Option<u128> = Some(0);
for variant in &data_enum.variants {
if !matches!(variant.fields, Fields::Unit) {
return err_span(
variant.span(),
format!(
"#[derive(CompactRepr)] only supports fieldless (unit) \
variants; variant `{}` carries data",
variant.ident
),
);
}
let disc = match &variant.discriminant {
Some((_, expr)) => match eval_discriminant(expr) {
Ok(v) => v,
Err(msg) => return err_span(variant.span(), msg),
},
None => match next {
Some(n) => n,
None => {
return err_span(
name.span(),
format!("enum `{name}` has too many variants to assign discriminants"),
)
}
},
};
next = disc.checked_add(1);
pairs.push((variant.ident.clone(), disc));
}
let discriminants: Vec<u128> = pairs.iter().map(|(_, d)| *d).collect();
if discriminants.is_empty() {
return err_span(
name.span(),
"#[derive(CompactRepr)] does not support empty enums".to_string(),
);
}
let max = *discriminants.iter().max().expect("non-empty");
let bits: u32 = match max {
0 | 1 => 1,
2..=3 => 2,
4..=15 => 4,
_ => {
return err_span(
name.span(),
format!(
"#[derive(CompactRepr)] on `{name}`: the largest discriminant is \
{max}, which needs more than 4 bits. At 8 bits and above \
`Compact<{name}>` is the same size as a plain `Vec<{name}>` \
(use `#[repr(u8)]` or `#[repr(u16)]`) but adds encode/decode \
overhead, so compacting it is redundant. Drop `Compact` and \
store the field as a plain `{name}` (e.g. `flag: {name}` \
instead of `flag: Compact<{name}>`)."
),
);
}
};
let valid_values: Vec<usize> =
discriminants.iter().map(|d| *d as usize).collect();
let idents: Vec<Ident> = pairs.iter().map(|(id, _)| id.clone()).collect();
let discs: Vec<proc_macro2::Literal> = discriminants
.iter()
.map(|d| proc_macro2::Literal::usize_unsuffixed(*d as usize))
.collect();
let storage_ty = quote! { ::layout::bitpack::PackedArray<#bits> };
let first_ident = &idents[0];
quote! {
impl ::layout::CompactRepr for #name {
type Storage = #storage_ty;
const BITS: u32 = #bits;
#[inline]
fn encode(self) -> usize {
self as usize
}
#[inline]
fn decode(raw: usize) -> Self {
debug_assert!(
[#( #valid_values ),*].contains(&raw),
"invalid compact discriminant for {}",
stringify!(#name)
);
match raw {
#( #discs => Self::#idents, )*
_ => Self::#first_ident,
}
}
}
}
}
fn find_unsigned_repr(attrs: &[Attribute]) -> Option<Ident> {
for attr in attrs {
if !attr.path().is_ident("repr") {
continue;
}
let Meta::List(list) = &attr.meta else {
continue;
};
let parser = Punctuated::<Path, Token![,]>::parse_terminated;
let Ok(paths) = parser.parse2(list.tokens.clone()) else {
continue;
};
for path in paths {
if let Some(ident) = path.get_ident() {
if matches!(
ident.to_string().as_str(),
"u8" | "u16" | "u32" | "u64" | "usize"
) {
return Some(Ident::new(&ident.to_string(), ident.span()));
}
}
}
}
None
}
fn eval_discriminant(expr: &Expr) -> Result<u128, String> {
if let Expr::Lit(ExprLit {
lit: Lit::Int(li), ..
}) = expr
{
li.base10_parse::<u128>()
.map_err(|e| format!("#[derive(CompactRepr)]: could not parse discriminant literal: {e}"))
} else {
Err(
"#[derive(CompactRepr)] requires non-negative integer literal \
discriminants (custom const expressions are not supported)"
.to_string(),
)
}
}
use proc_macro2::Ident;
fn err_span(span: Span, msg: String) -> TokenStream {
syn::Error::new(span, msg).to_compile_error()
}