enumeric 0.1.2

numeric range enum variant generation
Documentation
//! Procedural macro to automatically expand enum variants based on a numeric range.
//!
//! # Example
//!
//! ```rust
//! use enumeric::range_enum;
//!
//! #[range_enum]
//! enum MyEnum {
//!     /// Expands to: Item0, Item1, Item2
//!     #[range(0..3)]
//!     Item,
//!
//!     /// This remains unchanged.
//!     Other,
//!
//!     /// Expands to: Data10(u8), Data11(u8)
//!     #[range(10..12)]
//!     Data(u8),
//!
//!     /// Expands to: Code20 { id: u16 }, Code21 { id: u16 }, ..., Code22 { id: u16 }
//!     #[range(20..=22)]
//!     Code { id: u16 },
//! }
//! ```
//!
//! Currently supports `#[range(start..end)]` and `#[range(start..=end)]` syntax.
//! Only integer literals are accepted in the range.
//!
//! Struct, tuple, and unit variants are all supported.

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") {
                // parse range of discriminant
                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(),
                };

                // parse base variant name
                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;
            }
        }

        // keep original variant if not range
        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()
}

// FIXME: use bigger type instead of `u64`
#[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()
    }
}