use proc_macro2::{Group, Ident, TokenStream, TokenTree};
use quote::quote;
use syn::ItemTrait;
use syn::parse::Parser;
use super::*;
use crate::util::compile_err;
pub(crate) fn expand_directive(
name: &Ident, tokens: &[TokenTree], i: usize, trait_def: &ItemTrait,
trait_full_path: &TokenStream,
) -> Result<(Vec<TokenTree>, usize), TokenStream> {
if let Some(TokenTree::Group(args)) = tokens.get(i + 2) {
match args.delimiter() {
delimiter![{}] => {
check_builtin_typo(name)?;
expand_single(name, args, trait_def).map(|tt| (vec![tt], 3))
}
_ => {
let Some(TokenTree::Group(body)) = tokens.get(i + 3) else {
return Err(compile_err!(
"`#{}` must be followed by `(args)` or `[args]` + \
`{{body}}` (or directly `{{body}}`)",
name
));
};
if body.delimiter() != delimiter![{}] {
return Err(compile_err!(
"`#{}` must be followed by `(args)` or `[args]` + \
`{{body}}` (or directly `{{body}}`)",
name
));
}
let consumed = 4;
match name.to_string().as_str() {
"fill" => expand_fill(args, body, trait_def)
.map(|tt| (vec![tt], consumed)),
"delegate" => expand_delegate(args, body, trait_def)
.map(|tt| (vec![tt], consumed)),
"blanket" => {
expand_blanket(args, body, trait_def, trait_full_path)
.map(|v| (v, consumed))
}
_ => {
check_builtin_typo(name)?;
let inner = quote! {
#name ! { #args #body #trait_def }
};
Ok((
vec![Group::new(delimiter![{}], quote!(! #inner)).into()],
consumed,
))
}
}
}
}
} else {
Err(compile_err!(
"`#{}` must be followed by `(args)` / `[args]` or a code \
block `{{body}}`",
name
))
}
}
fn expand_single(
method_name: &Ident, body: &Group, trait_def: &ItemTrait,
) -> Result<TokenTree, TokenStream> {
let item = get_trait_item(trait_def, method_name)?;
Ok(Group::new(delimiter![{}], build_from_item(item, &body.stream())).into())
}
fn expand_many(
args_group: &Group, trait_def: &ItemTrait,
build: impl Fn(&Ident, &syn::TraitItem) -> Result<TokenStream, TokenStream>,
) -> Result<TokenTree, TokenStream> {
let method_names = parse_names_from_tokens(
&args_group.stream().into_iter().collect::<Vec<_>>(),
trait_def,
)?;
let mut methods = TokenStream::new();
for name in &method_names {
let item = get_trait_item(trait_def, name)?;
methods.extend(build(name, item)?);
}
Ok(Group::new(delimiter![{}], methods).into())
}
fn expand_fill(
args_group: &Group, body: &Group, trait_def: &ItemTrait,
) -> Result<TokenTree, TokenStream> {
let body_stream = body.stream();
expand_many(args_group, trait_def, |_name, item| {
Ok(build_from_item(item, &body_stream))
})
}
fn expand_delegate(
args_group: &Group, target: &Group, trait_def: &ItemTrait,
) -> Result<TokenTree, TokenStream> {
let target_stream = target.stream();
expand_many(args_group, trait_def, |name, item| {
let syn::TraitItem::Fn(f) = item else {
return Err(compile_err!(
"batch-impl: #delegate only works on methods; `{}` in trait \
`{}` is not a method",
trait_def.ident,
name
));
};
let mut sig = f.sig.clone();
let mut arg_idx = 0usize;
for input in &mut sig.inputs {
if let syn::FnArg::Typed(pat_type) = input
&& !pat_is_forwardable(&pat_type.pat)
{
pat_type.pat = syn::Pat::parse_single
.parse_str(&format!("arg{}", arg_idx))
.expect(
"generated arg names are always valid identifier patterns",
)
.into();
arg_idx += 1;
}
}
let call_args = collect_call_args(&sig).map_err(|pat| {
compile_err!(
"batch-impl: #delegate method `{}::{}` param `{}` cannot be \
forwarded (unsupported parameter pattern); please rename it \
to a plain identifier",
trait_def.ident,
name,
pat
)
})?;
let body = quote! { (#target_stream) . #name ( #(#call_args),* ) };
Ok(build_from_item_sig(item, Some(&sig), &body))
})
}
fn check_builtin_typo(name: &Ident) -> Result<(), TokenStream> {
let name_str = name.to_string();
for builtin in ["fill", "delegate", "blanket"] {
if levenshtein(&name_str, builtin) <= 2 {
return Err(compile_err!(
"batch-impl: unknown directive `#{}` — did you mean `#{}`?",
name,
builtin
));
}
}
Ok(())
}
fn levenshtein(a: &str, b: &str) -> usize {
let a: Vec<char> = a.chars().collect();
let b: Vec<char> = b.chars().collect();
let mut dp: Vec<usize> = (0..=b.len()).collect();
for (i, ca) in a.iter().enumerate() {
let mut prev = dp[0];
dp[0] = i + 1;
for (j, cb) in b.iter().enumerate() {
let cur = dp[j + 1];
dp[j + 1] =
if ca == cb { prev } else { 1 + prev.min(dp[j + 1]).min(dp[j]) };
prev = cur;
}
}
dp[b.len()]
}