use proc_macro2::{Delimiter, Group, Ident, TokenStream, TokenTree};
use quote::{quote, ToTokens};
use syn::ItemTrait;
use crate::parse::Cursor;
pub(crate) fn expand_tokens(cursor: &mut Cursor, trait_def: &ItemTrait) -> Result<Vec<TokenTree>, TokenStream> {
let mut result = vec![];
while !cursor.at_end() {
if cursor.is_punct('#')
&& let Some(TokenTree::Ident(name))=cursor.peek_at(1) {
let expanded = expand_directive(name, cursor, trait_def)?;
result.extend(expanded);
continue;
}
if let TokenTree::Group(g) = cursor.peek().unwrap() &&
g.delimiter()==Delimiter::Bracket{
let inner =
expand_tokens(&mut Cursor::new(&g.stream().into_iter().collect::<Vec<_>>()), trait_def)?;
let new_group = Group::new(g.delimiter(), inner.into_iter().collect());
result.push(new_group.into());
cursor.bump();
} else {
result.push(cursor.peek().unwrap().clone());
cursor.bump();
}
}
Ok(result)
}
fn expand_directive(
name:&Ident,
cursor: &mut Cursor,
trait_def: &ItemTrait,
) -> Result<Vec<TokenTree>, TokenStream> {
if let Some(TokenTree::Group(args))=cursor.peek_at(2){
if args.delimiter()==Delimiter::Brace{
cursor.bump(); cursor.bump(); cursor.bump(); expand_single_method(name, args, trait_def)
}else if let Some(TokenTree::Group(body))=cursor.peek_at(3)&&
body.delimiter()==Delimiter::Brace{
cursor.bump(); cursor.bump(); cursor.bump(); cursor.bump(); match name.to_string().as_str(){
"fill"=> expand_fill(args, body, trait_def),
"delegate"=> expand_delegate(args, body, trait_def),
_=> Ok(quote!{
#[#name[#args #body]]#trait_def
}.into_iter().collect())
}
}else{
Err(compile_error(&format!(
"`#{}` 后期望 `(args)` + `{{body}}` 或直接 `{{body}}`",
name
)))
}
}else{
Err(compile_error(&format!(
"`#{}` 后期望括号参数 `(args)` 或代码块 `{{body}}`",
name
)))
}
}
fn expand_single_method(
method_name:&Ident,
body:&Group,
trait_def: &ItemTrait,
) -> Result<Vec<TokenTree>, TokenStream> {
let sig = get_trait_method_sig(trait_def, method_name)?;
Ok(vec![TokenTree::Group(Group::new(
Delimiter::Brace,
build_fn_from_sig(&sig, &body.stream()),
))])
}
fn expand_fill(
args_group:&Group,
body:&Group,
trait_def: &ItemTrait,
) -> Result<Vec<TokenTree>, TokenStream> {
let method_names = parse_method_names_from_tokens(&args_group.stream().into_iter().collect::<Vec<_>>(), trait_def)?;
let mut methods = TokenStream::new();
for name in &method_names {
let sig=get_trait_method_sig(trait_def, name)?;
methods.extend(build_fn_from_sig(&sig, &body.stream()));
}
Ok(vec![TokenTree::Group(Group::new(
Delimiter::Brace,
methods,
))])
}
fn expand_delegate(
args_group:&Group,
target:&Group,
trait_def: &ItemTrait,
) -> Result<Vec<TokenTree>, TokenStream> {
let method_names = parse_method_names_from_tokens(&args_group.stream().into_iter().collect::<Vec<_>>(), trait_def)?;
let mut methods = TokenStream::new();
for name in &method_names {
let sig=get_trait_method_sig(trait_def, name)?;
let call_args = collect_call_args(&sig);
let target=target.stream();
let body = quote! { (#target) . #name ( #(#call_args),* ) };
methods.extend(build_fn_from_sig(&sig, &body));
}
Ok(vec![TokenTree::Group(Group::new(
Delimiter::Brace,
methods,
))])
}
fn parse_method_names_from_tokens(
tokens: &[TokenTree],
trait_def: &ItemTrait,
) -> Result<Vec<Ident>, TokenStream> {
if tokens.is_empty() {
return Err(compile_error("batch-impl: 指令的参数列表不能为空"));
}
if tokens.len() == 2
&& let (TokenTree::Punct(p), TokenTree::Ident(id)) = (&tokens[0], &tokens[1])
&& p.as_char() == '#' && id == "all"
{
return Ok(get_all_trait_methods(trait_def));
}
tokens
.iter()
.map(|t| {
if let TokenTree::Ident(id) = t {
Ok(Ident::new(&id.to_string(), id.span()))
} else if let TokenTree::Punct(p) = t &&
p.as_char()==','{
Err(None)
}else {
Err(Some(compile_error(&format!(
"batch-impl: 指令参数中期望标识符或逗号,得到 `{}`",
t
))))
}
})
.filter_map(|r| match r {
Ok(v) => Some(Ok(v)),
Err(None) => None,
Err(Some(e)) => Some(Err(e)),
})
.collect::<Result<Vec<_>, _>>()
}
fn get_all_trait_methods(trait_def: &ItemTrait) -> Vec<Ident> {
trait_def
.items
.iter()
.filter_map(|item| {
if let syn::TraitItem::Fn(f) = item {
Some(f.sig.ident.clone())
} else {
None
}
})
.collect()
}
fn get_trait_method_sig(trait_def: &ItemTrait, name: &Ident) -> Result<syn::Signature, TokenStream> {
for item in &trait_def.items{
if let syn::TraitItem::Fn(f) = item && f.sig.ident == *name {
return Ok(f.sig.clone());
}
}
Err(compile_error(&format!(
"batch-impl: trait `{}` 中没有找到方法 `{}`",
trait_def.ident, name
)))
}
fn build_fn_from_sig(sig: &syn::Signature, body: &TokenStream) -> TokenStream {
let sig_tokens = sig.to_token_stream();
quote! { #sig_tokens { #body } }
}
fn collect_call_args(sig: &syn::Signature) -> Vec<Ident> {
sig.inputs
.iter()
.filter_map(|arg| {
if let syn::FnArg::Typed(pat_type) = arg
&& let syn::Pat::Ident(pat_ident) = &*pat_type.pat
{
return Some(pat_ident.ident.clone());
}
None
})
.collect()
}
fn compile_error(msg: &str) -> TokenStream {
quote! { compile_error!(#msg); }
}