use std::collections::HashSet;
use proc_macro2::TokenStream;
use quote::quote;
use syn::ItemTrait;
use syn::visit::{self, Visit};
#[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>,
pub(crate) extra_predicates: Vec<(TokenStream, Vec<String>)>,
}
pub(crate) fn generic_param_names(generics: &syn::Generics) -> Vec<TokenStream> {
generics
.params
.iter()
.map(|p| match p {
syn::GenericParam::Lifetime(ld) => quote!(#ld),
syn::GenericParam::Type(tp) => {
let id = &tp.ident;
quote!(#id)
}
syn::GenericParam::Const(cp) => {
let id = &cp.ident;
quote!(#id)
}
})
.collect()
}
pub(crate) fn extract_trait_bounds(trait_item: &ItemTrait) -> TraitBounds {
let type_const_names = 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::<Vec<String>>();
let lt_names = trait_item
.generics
.params
.iter()
.filter_map(|p| match p {
syn::GenericParam::Lifetime(ld) => {
Some(format!("'{}", ld.lifetime.ident))
}
_ => None,
})
.collect::<Vec<String>>();
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 =
collect_bound_refs(&tp.bounds, &type_const_names, <_names);
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![],
}),
}
}
let mut extra_predicates = vec![];
if let Some(wc) = &trait_item.generics.where_clause {
for pred in &wc.predicates {
let tokens = quote!(#pred);
if let syn::WherePredicate::Type(pt) = pred
&& let Some(name) = single_ident_param(&pt.bounded_ty)
&& let Some(pos) = params.iter().position(|p| p.name == name)
{
let b = &pt.bounds;
let extra = quote!(#b);
let extra_refs =
collect_bound_refs(&pt.bounds, &type_const_names, <_names);
let param = &mut params[pos];
param.bound = Some(match ¶m.bound {
Some(inline) => quote!(#inline + #extra),
None => extra,
});
param.refs.extend(extra_refs);
continue;
}
let refs = collect_predicate_refs(pred, &type_const_names, <_names);
extra_predicates.push((tokens, refs));
}
}
TraitBounds { params, extra_predicates }
}
fn single_ident_param(ty: &syn::Type) -> Option<String> {
let syn::Type::Path(tp) = ty else { return None };
if tp.qself.is_some() {
return None;
}
let seg = tp.path.segments.first()?;
if tp.path.segments.len() == 1
&& matches!(&seg.arguments, syn::PathArguments::None)
{
Some(seg.ident.to_string())
} else {
None
}
}
fn collect_bound_refs(
bounds: &syn::punctuated::Punctuated<syn::TypeParamBound, syn::Token![+]>,
type_const_names: &[String], lt_names: &[String],
) -> Vec<String> {
let mut c = Collector::new(type_const_names, lt_names);
for b in bounds {
c.visit_type_param_bound(b);
}
c.refs
}
fn collect_predicate_refs(
pred: &syn::WherePredicate, type_const_names: &[String], lt_names: &[String],
) -> Vec<String> {
let mut c = Collector::new(type_const_names, lt_names);
c.visit_where_predicate(pred);
c.refs
}
struct Collector<'a> {
type_const_names: &'a [String],
lt_names: &'a [String],
hrtb: Vec<HashSet<String>>,
refs: Vec<String>,
}
impl<'a> Collector<'a> {
fn new(type_const_names: &'a [String], lt_names: &'a [String]) -> Self {
Collector { type_const_names, lt_names, hrtb: vec![], refs: vec![] }
}
fn in_hrtb(&self, name: &str) -> bool {
self.hrtb.iter().any(|s| s.contains(name))
}
fn push_hrtb(&mut self, lifetimes: Option<&syn::BoundLifetimes>) -> bool {
if let Some(bl) = lifetimes {
let set = bl
.lifetimes
.iter()
.filter_map(|p| match p {
syn::GenericParam::Lifetime(ld) => {
Some(format!("'{}", ld.lifetime.ident))
}
_ => None,
})
.collect();
self.hrtb.push(set);
true
} else {
false
}
}
}
impl<'ast> Visit<'ast> for Collector<'_> {
fn visit_expr(&mut self, node: &'ast syn::Expr) {
if let syn::Expr::Path(ep) = node
&& ep.qself.is_none()
&& let Some(seg) = ep.path.segments.first()
&& ep.path.segments.len() == 1
&& matches!(&seg.arguments, syn::PathArguments::None)
&& self.type_const_names.contains(&seg.ident.to_string())
{
self.refs.push(seg.ident.to_string());
}
visit::visit_expr(self, node);
}
fn visit_lifetime(&mut self, node: &'ast syn::Lifetime) {
let name = format!("'{}", node.ident);
if !self.in_hrtb(&name) && self.lt_names.contains(&name) {
self.refs.push(name);
}
}
fn visit_trait_bound(&mut self, node: &'ast syn::TraitBound) {
let pushed = self.push_hrtb(node.lifetimes.as_ref());
visit::visit_trait_bound(self, node);
if pushed {
self.hrtb.pop();
}
}
fn visit_type_fn_ptr(&mut self, node: &'ast syn::TypeFnPtr) {
let pushed = self.push_hrtb(node.lifetimes.as_ref());
visit::visit_type_fn_ptr(self, node);
if pushed {
self.hrtb.pop();
}
}
fn visit_type_path(&mut self, node: &'ast syn::TypePath) {
if node.qself.is_none()
&& let Some(seg) = node.path.segments.first()
&& node.path.segments.len() == 1
&& matches!(&seg.arguments, syn::PathArguments::None)
&& self.type_const_names.contains(&seg.ident.to_string())
{
self.refs.push(seg.ident.to_string());
}
visit::visit_type_path(self, node);
}
}