use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, Error, Expr, ExprLit, ExprRange, ItemEnum, Lit, LitInt, Meta};
#[proc_macro_attribute]
pub fn range_enum(_: TokenStream, item: TokenStream) -> TokenStream {
let input = parse_macro_input!(item as ItemEnum);
let mut generated_variants = Vec::default();
let vis = &input.vis;
let generics = &input.generics;
let enum_ident = &input.ident;
let attrs = &input.attrs;
for variant in input.variants.iter() {
let mut range_variant = false;
for attr in &variant.attrs {
if attr.path().is_ident("range") {
let Meta::List(meta_list) = attr.meta.clone() else {
continue;
};
let Expr::Range(range) = syn::parse2::<Expr>(meta_list.tokens.clone()).unwrap()
else {
continue;
};
let range = match ParsedRange::try_new(range) {
Ok(r) => r,
Err(err) => return err.to_compile_error().into(),
};
let base = &variant.ident;
let (start, end) = match (range.start, range.end) {
(Some(start), Some(end)) => (start, end),
_ => unimplemented!("Currently only x..y and x..=y supported."),
};
for i in start..end {
let variant_name = syn::Ident::new(&format!("{}{}", base, i), base.span());
let fields = &variant.fields;
let discriminant = variant
.discriminant
.as_ref()
.map(|(_, expr)| quote! { = #expr });
generated_variants.push(quote! {
#variant_name #fields #discriminant,
});
}
range_variant = true;
break;
}
}
if !range_variant {
let variant_name = &variant.ident;
let fields = &variant.fields;
let discriminant = variant
.discriminant
.as_ref()
.map(|(_, expr)| quote! { = #expr });
generated_variants.push(quote! {
#variant_name #fields #discriminant,
});
}
}
let output = quote! {
#(#attrs)*
#vis enum #enum_ident #generics {
#(#generated_variants)*
}
};
output.into()
}
#[derive(Copy, Clone, Debug)]
struct ParsedRange {
start: Option<u64>,
end: Option<u64>,
}
impl ParsedRange {
fn try_new(range: ExprRange) -> Result<ParsedRange, Error> {
let start = match range.start.as_deref() {
Some(Expr::Lit(ExprLit {
lit: Lit::Int(i), ..
})) => Some(parse_litint_auto(i)),
Some(expr) => {
return Err(Error::new_spanned(
expr,
"Expected integer literal for range start.",
))
}
_ => None,
};
let end_raw = match range.end.as_deref() {
Some(Expr::Lit(ExprLit {
lit: Lit::Int(i), ..
})) => Some(parse_litint_auto(i)),
Some(expr) => {
return Err(Error::new_spanned(
expr,
"Expected integer literal for range end.",
))
}
_ => None,
};
let end = if let Some(end) = end_raw {
Some(match range.limits {
syn::RangeLimits::Closed(_) => end + 1,
syn::RangeLimits::HalfOpen(_) => end,
})
} else {
None
};
Ok(ParsedRange { start, end })
}
}
fn parse_litint_auto(lit: &LitInt) -> u64 {
let s = lit.to_string();
if let Some(hex) = s.strip_prefix("0x") {
u64::from_str_radix(hex, 16).unwrap()
} else if let Some(oct) = s.strip_prefix("0o") {
u64::from_str_radix(oct, 8).unwrap()
} else if let Some(bin) = s.strip_prefix("0b") {
u64::from_str_radix(bin, 2).unwrap()
} else {
s.parse::<u64>().unwrap()
}
}