use proc_macro::{
Delimiter, Group as Paired, Ident as Id, Span, TokenStream,
TokenTree::{self, *},
};
use std::iter;
use pm2_types::EnumRepr;
use cg2::{Generator, f};
use heck::ToShoutySnakeCase;
use crate::{get_repr, seek_and_collect_ident};
pub fn run(attr: TokenStream, item: TokenStream) -> TokenStream {
let item_tokens: Vec<TokenTree> = item.clone().into_iter().collect();
let attr_tokens: Vec<TokenTree> = attr.clone().into_iter().collect();
let mut item_stream = item.clone().into_iter();
let repr = if let Some(r) = get_repr(&item_tokens) {
match EnumRepr::try_from(r.as_str()) {
Ok(r) => Some(r),
Err(r) => panic!("unknown repr: {r}"),
}
} else {
None
};
let mut tokens_before = vec![];
if !seek_and_collect_ident(&mut item_stream, &mut tokens_before, "enum") {
panic!("missing enum keyword");
}
let Ident(i) = item_stream.next().expect("missing name") else {
panic!("name not an Ident");
};
let name = i.to_string();
let name_ident = i;
let Group(body) = item_stream.next().expect("missing body") else {
panic!("generic params and where clauses are not allowed");
};
let body_tokens: Vec<TokenTree> = body.stream().into_iter().collect();
let variants = get_enum_variants(&body_tokens);
if variants.is_empty() {
panic!("must have atleast 1 variant");
}
let opts = get_opts(&attr_tokens);
let mut g = Generator::new();
g.implement()
.space()
.ident(&name)
.block(|g: &mut Generator| {
cg2::bind!($g, declare_const);
declare_const!(pub REPR: Option<&str>, |g: &mut Generator| {
if let Some(r) = repr {
g.out(f!("Some({:?})", r.as_str()));
} else {
g.out("None");
}
});
declare_const!(pub NUM_VARIANTS: usize = variants.len());
declare_const!(pub FIRST_VARIANT: Self = f!("Self::{}", variants.first().unwrap().name));
declare_const!(pub LAST_VARIANT: Self = f!("Self::{}", variants.last().unwrap().name));
if let Some(r) = repr
&& r.is_primitive()
{
declare_const!(pub MAX_DISCRIMINANT: r.as_str(), "Self::LAST_VARIANT as _");
}
});
for o in opts {
match o {
Opt::DerefDiscriminant => {
let Some(r) = repr.filter(|r| r.is_primitive()) else {
panic!("deref_discriminant: must have a primitive repr");
};
g.implement()
.space()
.out(f!("::core::ops::Deref for {}", &name))
.block(|g: &mut Generator| {
g.out(f!("type Target = {};", r.as_str()));
g.out("fn deref(&self) -> &Self::Target")
.block(|g: &mut Generator| {
g.unsafe_block("::core::mem::transmute(self)");
});
});
}
Opt::Bitflags { mut name_format } => {
let Some(r) = repr.filter(|r| r.is_primitive()) else {
panic!("bitflags: must have a primitive repr");
};
use EnumRepr::*;
let bits = match r {
U8 => 1,
U16 => 2,
U32 => 4,
U64 => 8,
U128 => 16,
_ => panic!("bitflags: repr must be an unsigned integer apart from usize"),
} * 8;
if variants.len() > bits {
panic!(
"bitflags: too many variants for repr. only {bits} variants allowed for {}",
r.as_str(),
);
}
match name_format.matches("{}").count() {
0 => panic!("bitflags: name_format missing placeholder"),
1 => (),
2.. => panic!("bitflags: name_format currently only supports a single placeholder"),
};
name_format = name_format[1..name_format.len() - 1].trim().to_string();
if name_format == "{}" {
panic!("bitflags: name_format missing pattern");
}
g.implement()
.space()
.ident(&name)
.block(|g: &mut Generator| {
for (i, v) in variants.iter().enumerate() {
g.public().space().declare_const(
name_format.replace("{}", &v.name.to_shouty_snake_case()),
r.as_str(),
1u128 << i,
);
}
});
}
}
}
let mut main_stream: TokenStream = g.code().parse().expect("failed to parse new code");
for t in tokens_before {
main_stream.extend(iter::once(t));
}
main_stream.extend(iter::once(Id::new("enum", Span::call_site())));
main_stream.extend(iter::once(name_ident));
let mut inner_stream = TokenStream::new();
for v in variants {
for t in v.tokens_before {
inner_stream.extend(iter::once(t.clone()));
}
inner_stream.extend(iter::once(v.name_ident));
}
main_stream.extend(iter::once(Paired::new(Delimiter::Brace, inner_stream)));
main_stream
}
#[derive(Debug)]
enum Opt {
DerefDiscriminant,
Bitflags { name_format: String },
}
fn get_opts(tt: &[TokenTree]) -> Vec<Opt> {
let mut opts = vec![];
for (i, t) in tt.iter().enumerate() {
if let Ident(i) = t
&& i.to_string() == "deref_discriminant"
{
opts.push(Opt::DerefDiscriminant);
continue;
}
if let Ident(id) = t
&& id.to_string() == "bitflags"
&& i < tt.len() - 1
&& let Punct(p) = &tt[i + 1]
&& *p == '='
&& (i + 1) < tt.len() - 1
&& let Literal(l) = &tt[i + 2]
{
opts.push(Opt::Bitflags {
name_format: l.to_string(),
});
continue;
}
}
opts
}
#[derive(Debug)]
struct EnumVariant<'a> {
pub name: String,
pub name_ident: Id,
pub _attrs: Vec<TokenStream>,
pub tokens_before: Vec<&'a TokenTree>,
pub _tokens_after: Vec<&'a TokenTree>,
}
fn get_enum_variants(tt: &[TokenTree]) -> Vec<EnumVariant<'_>> {
let mut variants = vec![];
let mut idx = 0;
let mut attrs = Vec::new();
let mut tokens_before = Vec::new();
while idx < tt.len() {
let t = &tt[idx];
if let Ident(i) = t {
let variant = EnumVariant {
name: i.to_string(),
name_ident: i.clone(),
_attrs: attrs.clone(),
tokens_before: tokens_before.clone(),
_tokens_after: vec![],
};
attrs.clear();
tokens_before.clear();
variants.push(variant);
idx += 1;
continue;
}
if let Punct(p) = t
&& *p == '#'
&& idx < tt.len() - 1
&& let Group(g) = &tt[idx + 1]
&& g.delimiter() == Delimiter::Bracket
&& let mut s = g.stream().into_iter()
&& let Some(Ident(id)) = s.next()
&& id.to_string() == "pm2"
{
attrs.push(g.stream());
idx += 2;
continue;
}
tokens_before.push(t);
idx += 1;
}
variants
}