use proc_macro::TokenStream;
use proc_macro2::{Ident, TokenStream as TokenStream2, TokenTree};
use quote::quote;
use syn::{parse_macro_input, ItemTrait};
mod types;
mod apply;
mod parse;
mod codegen;
mod preprocess;
use codegen::generate_impl;
use parse::{parse_item, Cursor};
use types::{reset_fresh_counter, Op};
#[proc_macro_attribute]
pub fn batch_impl(attr: TokenStream, item: TokenStream) -> TokenStream {
expand_attr_macro(attr, item, true)
}
#[proc_macro_attribute]
pub fn batch_impl_only(attr: TokenStream, item: TokenStream) -> TokenStream {
expand_attr_macro(attr, item, false)
}
fn expand_attr_macro(attr: TokenStream, item: TokenStream, include_trait: bool) -> TokenStream {
reset_fresh_counter();
let trait_item = parse_macro_input!(item as ItemTrait);
let trait_name = trait_item.ident.clone();
let attr_vec = TokenStream2::from(attr).into_iter().collect::<Vec<_>>();
let mut cursor = Cursor::new(&attr_vec);
let expanded = match preprocess::expand_tokens(&mut cursor, &trait_item) {
Ok(tokens) => tokens,
Err(err) => return err.into(),
};
cursor = Cursor::new(&expanded);
let trait_name_ts = quote![#trait_name];
let is_unsafe = trait_item.unsafety.is_some();
let start_trait = if include_trait { Some(trait_item) } else { None };
let impls = parse_batch_trait_entry(
&mut cursor, Op::Comma, &trait_name_ts, &trait_name,
is_unsafe, start_trait,
);
impls.into()
}
#[proc_macro]
pub fn batch_trait(input: TokenStream) -> TokenStream {
reset_fresh_counter();
let tokens = TokenStream2::from(input).into_iter().collect::<Vec<_>>();
let mut cursor = Cursor::new(&tokens);
let mut result = quote![];
loop {
while cursor.is_punct(';') {
cursor.bump();
}
if cursor.at_end() {
break;
}
let is_unsafe = if matches!(cursor.peek(), Some(TokenTree::Ident(id)) if *id == "unsafe") {
cursor.bump();
true
} else {
false
};
let path_start = cursor.pos();
let mut depth = 0i32;
while let Some(token) = cursor.peek() {
match token {
TokenTree::Punct(p) if p.as_char() == '<' => {
depth += 1;
cursor.bump();
}
TokenTree::Punct(p) if p.as_char() == '>' => {
depth -= 1;
cursor.bump();
}
TokenTree::Punct(p) if p.as_char() == ':' && depth == 0 => {
if matches!(cursor.peek_at(1), Some(TokenTree::Punct(p2)) if p2.as_char() == ':') {
cursor.bump();
cursor.bump();
} else {
break;
}
}
_ => cursor.bump(),
}
}
let trait_path = cursor.slice_since(path_start);
if trait_path.is_empty() {
result.extend(generate_compile_error("batch_trait! 中期望 trait 名称"));
break;
}
let trait_full_path = match extract_trait_path(trait_path) {
Ok(path) => path,
Err(e) => { result.extend(e); break; }
};
let trait_last_ident = match extract_last_ident(trait_path) {
Ok(ident) => ident,
Err(e) => { result.extend(e); break; }
};
if !cursor.is_punct(':') {
result.extend(generate_compile_error(
"batch_trait! 中期望 ':' 分隔 trait 名称和 impl-specs",
));
break;
}
cursor.bump();
let impl_code = parse_batch_trait_entry(
&mut cursor, Op::Semi, &trait_full_path,
trait_last_ident, is_unsafe, None,
);
result.extend(impl_code);
}
result.into()
}
fn parse_batch_trait_entry(
cursor: &mut Cursor,
top_level: Op,
trait_full_path: &TokenStream2,
trait_last_ident: &Ident,
is_unsafe_trait: bool,
start_trait: Option<ItemTrait>,
) -> TokenStream2 {
let mut tys = vec![];
while let Some(ty) = parse_item(cursor, top_level, Some(trait_last_ident)) {
let mut queue = vec![ty];
while let Some(item) = queue.pop() {
match item.expand() {
Ok(expanded) => {
for e in expanded.into_iter().rev() {
queue.push(e);
}
}
Err(leaf) => tys.push(leaf),
}
}
}
let mut impls = start_trait.map_or(quote![], |t| quote![#t]);
for t in tys {
impls.extend(generate_impl(t, trait_full_path, is_unsafe_trait));
}
impls
}
fn extract_trait_path(trait_path: &[TokenTree]) -> Result<TokenStream2, TokenStream2> {
let path: TokenStream2 = trait_path.iter().cloned().collect();
if path.is_empty() {
Err(generate_compile_error("batch_trait! 中期望 trait 名称"))
} else {
Ok(path)
}
}
fn extract_last_ident(trait_path: &[TokenTree]) -> Result<&Ident, TokenStream2> {
trait_path
.iter()
.filter_map(|tt| if let TokenTree::Ident(id) = tt { Some(id) } else { None })
.next_back()
.ok_or_else(|| generate_compile_error("batch_trait! 中期望标识符作为 trait 名称"))
}
fn generate_compile_error(msg: &str) -> TokenStream2 {
quote! { compile_error!(#msg); }
}