pm2 0.1.4

Useful proc macros.
Documentation
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();

  // constants
  g.implement()
    .space()
    .ident(&name)
    .block(|g: &mut Generator| {
      cg2::bind!($g, declare_const);

      // REPR
      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");
        }
      });

      // NUM_VARIANTS
      declare_const!(pub NUM_VARIANTS: usize = variants.len());

      // FIRST_VARIANT
      // LAST_VARIANT
      declare_const!(pub FIRST_VARIANT: Self = f!("Self::{}", variants.first().unwrap().name));
      declare_const!(pub LAST_VARIANT: Self = f!("Self::{}", variants.last().unwrap().name));

      // MAX_DISCRIMINANT
      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 {
  // previously Seal
  DerefDiscriminant,
  Bitflags { name_format: String },
}

fn get_opts(tt: &[TokenTree]) -> Vec<Opt> {
  let mut opts = vec![];

  for (i, t) in tt.iter().enumerate() {
    // deref_discriminant
    if let Ident(i) = t
      && i.to_string() == "deref_discriminant"
    {
      opts.push(Opt::DerefDiscriminant);
      continue;
    }

    // bitflags = "..."
    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
}