#![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 generic::matching_angle;
use preprocess_helpers::{build_from_item, get_trait_item, parse_names_from_tokens};
use scan::{Cursor, is_arrow, scan_stop};
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))?;
let is_unsafe = trait_item.unsafety.is_some();
let trait_bounds = extract_trait_bounds(&trait_item);
let expanded = expand_empty_trait_generics(&expanded, &trait_item)?;
cursor = Cursor::new(&expanded);
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,
&trait_bounds,
);
Ok(impls.into())
}
#[derive(Default)]
pub(crate) struct TraitParam {
pub(crate) name: String,
pub(crate) bound: Option<TokenStream>,
pub(crate) refs: Vec<String>,
}
#[derive(Default)]
pub(crate) struct TraitBounds {
pub(crate) params: Vec<TraitParam>,
}
fn extract_trait_bounds(trait_item: &ItemTrait) -> TraitBounds {
let type_const_names: Vec<String> = trait_item
.generics
.params
.iter()
.filter_map(|p| match p {
syn::GenericParam::Type(tp) => Some(tp.ident.to_string()),
syn::GenericParam::Const(cp) => Some(cp.ident.to_string()),
_ => None,
})
.collect();
let lt_names: Vec<String> = trait_item
.generics
.params
.iter()
.filter_map(|p| match p {
syn::GenericParam::Lifetime(ld) => {
Some(format!("'{}", ld.lifetime.ident))
}
_ => None,
})
.collect();
let mut params = vec![];
for p in &trait_item.generics.params {
match p {
syn::GenericParam::Type(tp) => {
let bound = if tp.bounds.is_empty() {
None
} else {
let b = &tp.bounds;
Some(quote!(#b))
};
let refs = bound
.as_ref()
.map(|b| bound_refs(b, &type_const_names, <_names))
.unwrap_or_default();
params.push(TraitParam { name: tp.ident.to_string(), bound, refs });
}
syn::GenericParam::Lifetime(ld) => params.push(TraitParam {
name: format!("'{}", ld.lifetime.ident),
bound: None,
refs: vec![],
}),
syn::GenericParam::Const(cp) => params.push(TraitParam {
name: cp.ident.to_string(),
bound: None,
refs: vec![],
}),
}
}
TraitBounds { params }
}
fn bound_refs(
bound: &TokenStream, type_const_names: &[String], lt_names: &[String],
) -> Vec<String> {
let mut refs = vec![];
let mut iter = bound.clone().into_iter().peekable();
while let Some(tt) = iter.next() {
match tt {
TokenTree::Ident(id) if type_const_names.contains(&id.to_string()) => {
refs.push(id.to_string())
}
TokenTree::Punct(p) if p.as_char() == '\'' => {
if let Some(TokenTree::Ident(id)) = iter.peek() {
let name = format!("'{}", id);
if lt_names.contains(&name) {
refs.push(name);
}
}
}
_ => {}
}
}
refs
}
fn args_all_bindings(args: &[TokenTree]) -> bool {
let mut rest = args;
while let Some(idx) = scan_stop(rest, &[',']) {
if scan_stop(&rest[..idx], &['=']).is_none() {
return false;
}
rest = &rest[idx + 1..];
}
scan_stop(rest, &['=']).is_some()
}
fn expand_empty_trait_generics(
tokens: &[TokenTree], trait_def: &ItemTrait,
) -> Result<Vec<TokenTree>, TokenStream> {
if trait_def.generics.params.is_empty() {
return Ok(tokens.to_vec());
}
let generics = &trait_def.generics;
let impl_gen = quote!(#generics);
let mut arg_names: Vec<TokenStream> = vec![];
for p in &trait_def.generics.params {
match p {
syn::GenericParam::Lifetime(ld) => arg_names.push(quote!(#ld)),
syn::GenericParam::Type(tp) => {
let id = &tp.ident;
arg_names.push(quote!(#id));
}
syn::GenericParam::Const(cp) => {
let id = &cp.ident;
arg_names.push(quote!(#id));
}
}
}
let mut out = vec![];
let mut depth = 0usize;
let mut i = 0;
while i < tokens.len() {
match &tokens[i] {
TokenTree::Punct(p) if p.as_char() == '<' => {
depth += 1;
out.push(tokens[i].clone());
i += 1;
}
TokenTree::Punct(p) if p.as_char() == '>' => {
if !is_arrow(tokens, i) {
depth = depth.saturating_sub(1);
}
out.push(tokens[i].clone());
i += 1;
}
TokenTree::Ident(id)
if depth == 0
&& matches!(tokens.get(i + 1), Some(TokenTree::Punct(p)) if p.as_char() == '<') =>
{
let Some(close) = matching_angle(tokens, i + 1) else {
out.push(tokens[i].clone());
i += 1;
continue;
};
let args = &tokens[i + 2..close];
let bindings_only = !args.is_empty() && args_all_bindings(args);
if args.is_empty() || bindings_only {
let name = quote!(#id);
out.extend(impl_gen.clone());
if args.is_empty() {
out.extend(quote!(#name < #(#arg_names),* >));
} else {
let args_ts: TokenStream = args.iter().cloned().collect();
out.extend(quote!(#name < #(#arg_names),* , #args_ts >));
}
i = close + 1;
} else {
out.push(tokens[i].clone());
i += 1;
}
}
_ => {
out.push(tokens[i].clone());
i += 1;
}
}
}
Ok(out)
}
#[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,
&Default::default(),
);
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()
}