use proc_macro2::{Ident, TokenStream, TokenTree};
use quote::quote;
use syn::ItemTrait;
use crate::ast::{Op, reset_fresh_counter};
use crate::batch_trait_entry::parse_batch_trait_entry;
use crate::diagnostic::compile_error_str;
use crate::empty_generics::expand_empty_trait_generics;
use crate::preprocess::{angle_collect, expand_tokens, render_angles, where_process};
use crate::scan::Cursor;
use crate::trait_bounds::TraitBounds;
fn run_pipeline(
tokens: &[TokenTree], top_level: Op, trait_full_path: &TokenStream,
trait_last_ident: &Ident, is_unsafe: bool, start_trait: Option<ItemTrait>,
trait_bounds: &TraitBounds,
) -> Result<TokenStream, TokenStream> {
let mut cursor = Cursor::new(tokens);
let impls = parse_batch_trait_entry(
&mut cursor,
top_level,
trait_full_path,
trait_last_ident,
is_unsafe,
start_trait,
trait_bounds,
);
Ok(render_angles(impls))
}
pub(crate) 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 attr_vec = angle_collect(&attr_vec)?;
let (trait_full_path, trait_last_ident, rest_tokens) = if !include_trait {
match crate::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 expanded = crate::consts::expand_consts(&rest_tokens, None)?;
let expanded =
expand_tokens(&mut Cursor::new(&expanded), &trait_item, &trait_full_path)?;
let expanded = where_process(&mut Cursor::new(&expanded))?;
let is_unsafe = trait_item.unsafety.is_some();
let trait_bounds = crate::trait_bounds::extract_trait_bounds(&trait_item);
let expanded =
expand_empty_trait_generics(&expanded, &trait_item, &trait_bounds)?;
let start_trait = if include_trait { trait_item.into() } else { None };
run_pipeline(
&expanded,
Op::Comma,
&trait_full_path,
&trait_last_ident,
is_unsafe,
start_trait,
&trait_bounds,
)
.map(Into::into)
}
pub(crate) 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 = angle_collect(&tokens)?;
let (tokens, user_consts) = crate::consts::collect_user_consts(&tokens)?;
let tokens = crate::consts::expand_consts(&tokens, Some(&user_consts))?;
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();
while let Some(token) = cursor.peek() {
match token {
TokenTree::Punct(p) if p.as_char() == ':' => {
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() {
return Err(compile_error_str("batch_trait! 中期望 trait 名称"));
}
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 => {
return Err(compile_error_str(
"batch_trait! 中期望标识符作为 trait 名称",
));
}
};
if !cursor.is_punct(':') {
return Err(compile_error_str(
"batch_trait! 中期望 ':' 分隔 trait 名称和 impl-specs",
));
}
cursor.bump();
let spec = cursor.take_segment(&[';']);
result.extend(run_pipeline(
spec,
Op::Comma,
&trait_full_path,
trait_last_ident,
is_unsafe,
None,
&Default::default(),
)?);
}
Ok(result.into())
}