use proc_macro2::{Group, Ident, TokenStream, TokenTree};
use quote::quote;
use crate::ast::{Ty, TyGeneric, TyKind, TyPrimitive, TyTrait, TyTypeParam};
use crate::codegen::extract::ImplParts;
use crate::util::is_punct_at;
pub(crate) fn sync_impl_parts(
parts: &mut ImplParts, trait_name: &TokenStream,
) -> Result<(), TokenStream> {
let Some(trait_ident) = trait_last_ident(trait_name) else {
return Ok(());
};
let trait_args = parts.trait_generic_names.clone();
let mut body_sync = false;
let mut matched = Vec::new();
for t in std::mem::take(&mut parts.impl_templates) {
let is_switch = is_switch_template(&t.clone().into_iter().collect::<Vec<_>>(), &trait_ident);
let s = sync_trait_application(t, &trait_args)?;
if is_switch {
body_sync = true;
} else {
matched.push(s);
}
}
parts.impl_templates = matched;
let mut synced = Vec::with_capacity(parts.where_clauses.len());
for w in &parts.where_clauses {
synced.push(sync_trait_application(w.clone(), &trait_args)?);
}
parts.where_clauses = synced;
for (_, bound) in &mut parts.impl_generics {
if let Some(b) = bound {
*b = sync_bound_ty(b, &trait_args)?;
}
}
if body_sync && let Some(b) = &mut parts.body {
*b = sync_trait_application(b.clone(), &trait_args)?;
}
Ok(())
}
pub(crate) fn trait_last_ident(trait_name: &TokenStream) -> Option<Ident> {
let mut last = None;
for t in trait_name.clone() {
if let TokenTree::Ident(id) = t {
last = Some(id);
}
}
last
}
pub(crate) fn sync_trait_application(
tokens: TokenStream, args: &[TokenStream],
) -> Result<TokenStream, TokenStream> {
let v = tokens.into_iter().collect::<Vec<_>>();
sync_at(&v, args, 0).map(|o| o.into_iter().collect())
}
fn empty_angle_at(tokens: &[TokenTree], i: usize) -> bool {
matches!(tokens.get(i), Some(TokenTree::Group(g))
if g.delimiter() == delimiter![<>] && g.stream().is_empty())
|| (is_punct_at(tokens, i, '<') && is_punct_at(tokens, i + 1, '>'))
}
fn sync_at(
tokens: &[TokenTree], args: &[TokenStream], depth: usize,
) -> Result<Vec<TokenTree>, TokenStream> {
if depth > crate::util::MAX_NEST_DEPTH {
return Err(crate::util::depth_err(tokens, ""));
}
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
let is_ident_angle =
matches!(&tokens[i], TokenTree::Ident(_)) && empty_angle_at(tokens, i + 1);
if is_ident_angle {
let Some(TokenTree::Ident(id)) = tokens.get(i) else {
return Err(crate::util::depth_err(tokens, ""));
};
let mut ts = quote!(#id);
if !args.is_empty() {
ts.extend(quote!(<#(#args),*>));
}
out.extend(ts);
i += if matches!(tokens[i + 1], TokenTree::Group(_)) { 2 } else { 3 };
continue;
}
if let TokenTree::Group(g) = &tokens[i] {
if depth + 1 > crate::util::MAX_NEST_DEPTH {
return Err(crate::util::depth_err(&tokens[i..i + 1], ""));
}
let inner = g.stream().into_iter().collect::<Vec<_>>();
let synced = sync_at(&inner, args, depth + 1)?;
let mut ng = Group::new(g.delimiter(), synced.into_iter().collect());
ng.set_span(g.span());
out.push(TokenTree::Group(ng));
i += 1;
continue;
}
out.push(tokens[i].clone());
i += 1;
}
Ok(out)
}
pub(crate) fn is_switch_template(tokens: &[TokenTree], trait_ident: &Ident) -> bool {
let Some(idx) = tokens.iter().rposition(|t| matches!(t, TokenTree::Ident(_))) else {
return false;
};
let Some(TokenTree::Ident(id)) = tokens.get(idx) else {
return false;
};
if id != trait_ident {
return false;
}
match &tokens[idx + 1..] {
[TokenTree::Punct(lt), TokenTree::Punct(gt)] => lt.as_char() == '<' && gt.as_char() == '>',
[TokenTree::Group(g)] => g.delimiter() == delimiter![<>] && g.stream().is_empty(),
_ => false,
}
}
pub(crate) fn sync_bound_ty(ty: &Ty, args: &[TokenStream]) -> Result<Ty, TokenStream> {
match &ty.kind {
TyKind::Generic(g) if g.1.params.is_empty() && g.1.bindings.is_empty() => {
Ok(TyGeneric(g.0.clone(), filled_params(args)).to_ty().with_span(ty.span))
}
TyKind::Trait(t) if t.1.params.is_empty() && t.1.bindings.is_empty() => {
Ok(TyTrait(t.0.clone(), filled_params(args)).to_ty().with_span(ty.span))
}
_ => Ok(ty.clone()),
}
}
fn filled_params(args: &[TokenStream]) -> TyTypeParam {
TyTypeParam {
params: args.iter().map(|a| (Box::new(TyPrimitive(a.clone()).to_ty()), None)).collect(),
bindings: vec![],
}
}
#[cfg(test)]
mod tests {
use super::*;
use quote::ToTokens;
fn args(list: &[&str]) -> Vec<TokenStream> {
list.iter().map(|a| a.parse::<TokenStream>().unwrap()).collect()
}
#[test]
fn where_predicate_fills_args() {
let ts = "@0.. : Semiring < >".parse::<TokenStream>().unwrap();
let out = sync_trait_application(ts, &args(&["Additive", "Multiplicative"])).unwrap();
assert_eq!(out.to_string(), "@ 0 .. : Semiring < Additive , Multiplicative >");
}
#[test]
fn bare_trait_without_args_drops_brackets() {
let ts = "@0.. : Sized < >".parse::<TokenStream>().unwrap();
let out = sync_trait_application(ts, &[]).unwrap();
assert_eq!(out.to_string(), "@ 0 .. : Sized");
}
#[test]
fn other_ident_fills() {
let ts = "@0.. : Other < >".parse::<TokenStream>().unwrap();
let out = sync_trait_application(ts, &args(&["Additive"])).unwrap();
assert_eq!(out.to_string(), "@ 0 .. : Other < Additive >");
}
#[test]
fn flat_template_shape() {
let ts = "impl { Semiring < > }".parse::<TokenStream>().unwrap();
let out = sync_trait_application(ts, &args(&["Additive", "Multiplicative"])).unwrap();
assert_eq!(out.to_string(), "impl { Semiring < Additive , Multiplicative > }");
}
#[test]
fn switch_template_flat() {
let ts = "Tr < >".parse::<TokenStream>().unwrap();
let v = ts.into_iter().collect::<Vec<_>>();
assert!(is_switch_template(&v, &Ident::new("Tr", proc_macro2::Span::call_site())));
}
#[test]
fn switch_template_group() {
let ts = "Tr < >".parse::<TokenStream>().unwrap();
let v = ts.into_iter().collect::<Vec<_>>();
assert!(is_switch_template(&v, &Ident::new("Tr", proc_macro2::Span::call_site())));
}
#[test]
fn switch_template_path_qualified() {
let ts = "mod :: Tr < >".parse::<TokenStream>().unwrap();
let v = ts.into_iter().collect::<Vec<_>>();
assert!(is_switch_template(&v, &Ident::new("Tr", proc_macro2::Span::call_site())));
let ts = "crate :: ext :: Tr < >".parse::<TokenStream>().unwrap();
let v = ts.into_iter().collect::<Vec<_>>();
assert!(is_switch_template(&v, &Ident::new("Tr", proc_macro2::Span::call_site())));
}
#[test]
fn switch_template_not_recognized() {
let ts = "Tr < Additive >".parse::<TokenStream>().unwrap();
let v = ts.into_iter().collect::<Vec<_>>();
assert!(!is_switch_template(&v, &Ident::new("Tr", proc_macro2::Span::call_site())));
let ts = "Other < >".parse::<TokenStream>().unwrap();
let v = ts.into_iter().collect::<Vec<_>>();
assert!(!is_switch_template(&v, &Ident::new("Tr", proc_macro2::Span::call_site())));
let ts = "Tr".parse::<TokenStream>().unwrap();
let v = ts.into_iter().collect::<Vec<_>>();
assert!(!is_switch_template(&v, &Ident::new("Tr", proc_macro2::Span::call_site())));
}
#[test]
fn other_trait_untouched() {
let ts = "@0.. : Module < (), () >".parse::<TokenStream>().unwrap();
let out = sync_trait_application(ts, &args(&["Additive"])).unwrap();
assert_eq!(out.to_string(), "@ 0 .. : Module < () , () >");
}
#[test]
fn bound_ty_fills_args() {
let base = TyPrimitive(quote!(BoundSync)).to_ty();
let empty = TyTypeParam { params: vec![], bindings: vec![] };
let bound = TyGeneric(Box::new(base), empty).to_ty();
let out = sync_bound_ty(&bound, &args(&["Additive", "Multiplicative"])).unwrap();
assert_eq!(out.to_token_stream().to_string(), "BoundSync < Additive , Multiplicative >");
}
#[test]
fn bound_trait_ty_fills_args() {
let tp = TyTypeParam { params: vec![], bindings: vec![] };
let bound = TyTrait(quote!(BoundSync), tp).to_ty();
let out = sync_bound_ty(&bound, &args(&["Additive", "Multiplicative"])).unwrap();
assert_eq!(out.to_token_stream().to_string(), "BoundSync < Additive , Multiplicative >");
}
#[test]
fn bound_ty_wrong_name_untouched() {
let base = TyPrimitive(quote!(Module)).to_ty();
let params = vec![(Box::new(TyPrimitive(quote!(A)).to_ty()), None)];
let tp = TyTypeParam { params, bindings: vec![] };
let bound = TyGeneric(Box::new(base), tp).to_ty();
let out = sync_bound_ty(&bound, &args(&["Additive"])).unwrap();
assert_eq!(out.to_token_stream().to_string(), "Module < A >");
}
#[test]
fn bound_other_name_fills() {
let tp = TyTypeParam { params: vec![], bindings: vec![] };
let bound = TyTrait(quote!(Module), tp).to_ty();
let out = sync_bound_ty(&bound, &args(&["Additive", "Multiplicative"])).unwrap();
assert_eq!(out.to_token_stream().to_string(), "Module < Additive , Multiplicative >");
}
}