#![doc = include_str!(concat!(env!("CARGO_MANIFEST_DIR"), "/README.md"))]
#![forbid(unsafe_code)]
#![deny(missing_docs)]
#![allow(linker_messages)]
#[cfg(test)]
mod fuzz;
use proc_macro2::{TokenStream, TokenTree};
use quote::quote;
use syn::{ItemTrait, parse_macro_input};
mod apply;
mod apply_tuple;
mod batch_trait_entry;
mod codegen;
mod diagnostic;
mod generic;
mod parse;
mod parse_atom;
mod path_prefix;
mod preprocess;
mod preprocess_helpers;
mod scan;
mod types;
mod types_render;
mod where_process;
use batch_trait_entry::parse_batch_trait_entry;
use diagnostic::compile_error_str;
use preprocess_helpers::{build_from_item, get_trait_item, parse_names_from_tokens};
use scan::Cursor;
use types::{Op, reset_fresh_counter};
use where_process::where_process;
#[proc_macro_attribute]
pub fn batch_impl(
attr: proc_macro::TokenStream, item: proc_macro::TokenStream,
) -> proc_macro::TokenStream {
let trait_item = parse_macro_input!(item as ItemTrait);
expand_attr_macro(attr, trait_item, true).unwrap_or_else(Into::into)
}
#[proc_macro_attribute]
pub fn batch_impl_only(
attr: proc_macro::TokenStream, item: proc_macro::TokenStream,
) -> proc_macro::TokenStream {
let trait_item = parse_macro_input!(item as ItemTrait);
expand_attr_macro(attr, trait_item, false).unwrap_or_else(Into::into)
}
fn expand_attr_macro(
attr: proc_macro::TokenStream, trait_item: ItemTrait, include_trait: bool,
) -> Result<proc_macro::TokenStream, TokenStream> {
reset_fresh_counter();
let trait_name = trait_item.ident.clone();
let attr_vec = TokenStream::from(attr).into_iter().collect::<Vec<_>>();
let (trait_full_path, trait_last_ident, rest_tokens) = if !include_trait {
match path_prefix::try_parse_path_prefix(&attr_vec) {
Some((path, last_ident, rest)) => {
match last_ident {
Some(id) if id == trait_name => {
let path_ts = path.into_iter().collect();
(path_ts, trait_name.clone(), rest)
}
Some(id) => {
let msg = format!(
"batch-impl: 路径前缀 `#...{}` \
的末尾标识符与 trait 名 `{}` \
不一致;二者必须相同",
id, trait_name,
);
return Err(compile_error_str(&msg));
}
None => {
let msg = "batch-impl: 路径前缀 `#` 后 \
期望至少一个标识符作为 trait 路径";
return Err(compile_error_str(msg));
}
}
}
None => (quote![#trait_name], trait_name.clone(), attr_vec.clone()),
}
} else {
(quote![#trait_name], trait_name.clone(), attr_vec.clone())
};
let mut cursor = Cursor::new(&rest_tokens);
let expanded = preprocess::expand_tokens(&mut cursor, &trait_item)?;
let expanded = where_process(&mut Cursor::new(&expanded))?;
cursor = Cursor::new(&expanded);
let is_unsafe = trait_item.unsafety.is_some();
let start_trait = if include_trait { trait_item.into() } else { None };
let impls = parse_batch_trait_entry(
&mut cursor,
Op::Comma,
&trait_full_path,
&trait_last_ident,
is_unsafe,
start_trait,
);
Ok(impls.into())
}
#[proc_macro]
pub fn batch_trait(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
expand_batch_trait(input).unwrap_or_else(Into::into)
}
fn expand_batch_trait(
input: proc_macro::TokenStream,
) -> Result<proc_macro::TokenStream, TokenStream> {
reset_fresh_counter();
let tokens = TokenStream::from(input).into_iter().collect::<Vec<_>>();
let tokens = where_process(&mut Cursor::new(&tokens))?;
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 cursor.is_single_colon() {
break;
} else {
cursor.bump();
cursor.bump();
}
}
_ => cursor.bump(),
}
}
let trait_path = cursor.slice_since(path_start);
if trait_path.is_empty() {
result.extend(compile_error_str("batch_trait! 中期望 trait 名称"));
break;
}
let trait_full_path = trait_path.iter().cloned().collect();
let trait_last_ident =
match trait_path
.iter()
.filter_map(|tt| {
if let TokenTree::Ident(id) = tt { id.into() } else { None }
})
.next_back()
{
Some(ident) => ident,
None => {
result.extend(compile_error_str(
"batch_trait! 中期望标识符作为 trait 名称",
));
break;
}
};
if !cursor.is_punct(':') {
result.extend(compile_error_str(
"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);
}
Ok(result.into())
}
#[doc(hidden)]
#[proc_macro]
pub fn batch_preprocess_test(
input: proc_macro::TokenStream,
) -> proc_macro::TokenStream {
let tokens = TokenStream::from(input).into_iter().collect::<Vec<_>>();
let Some(TokenTree::Group(names_group)) = tokens.first() else {
return compile_error_str(
"batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
)
.into();
};
if names_group.delimiter() != proc_macro2::Delimiter::Parenthesis {
return compile_error_str(
"batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
)
.into();
}
let Some(TokenTree::Group(body_group)) = tokens.get(1) else {
return compile_error_str(
"batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
)
.into();
};
if body_group.delimiter() != proc_macro2::Delimiter::Brace {
return compile_error_str(
"batch-impl: batch_preprocess_test 期望 `(方法名列表){body} trait ...`",
)
.into();
}
let trait_ts = tokens[2..].iter().cloned().collect();
let trait_item = match syn::parse2(trait_ts) {
Ok(t) => t,
Err(_) => {
return compile_error_str(
"batch-impl: batch_preprocess_test 无法解析 trait 定义",
)
.into();
}
};
let names = match parse_names_from_tokens(
&names_group.stream().into_iter().collect::<Vec<_>>(),
&trait_item,
) {
Ok(names) => names,
Err(e) => return e.into(),
};
let body = body_group.stream();
let mut methods = TokenStream::new();
for name in &names {
let item = match get_trait_item(&trait_item, name) {
Ok(item) => item,
Err(e) => return e.into(),
};
methods.extend(build_from_item(item, &body));
}
methods.into()
}