use std::collections::{HashMap, HashSet};
use proc_macro2::{TokenStream, TokenTree};
use quote::quote;
use crate::TraitBounds;
use crate::ast::{Ty, TyPrimitive};
use crate::codegen::extract::ImplParts;
use crate::util::compile_err;
pub(crate) fn inherit_trait_bounds(
parts: &mut ImplParts, trait_bounds: &TraitBounds, trait_args: &[String],
impl_names: &HashSet<String>,
) -> Vec<TokenStream> {
let mut errs = vec![];
for (name, bound) in &mut parts.impl_generics {
if bound.is_some() {
continue;
}
let key = name.to_string();
let Some(pos) = trait_args.iter().position(|a| a == &key) else {
continue;
};
let Some(tp) = trait_bounds.params.get(pos) else {
continue;
};
let Some(b) = &tp.bound else {
continue;
};
if tp.name != key {
errs.push(compile_err!(
"batch-impl: trait argument `{}` maps to parameter `{}` (bound `{}`); automatic \
inheritance requires the same name; rename to `{}` or write the bound manually",
key,
tp.name,
b,
tp.name
));
continue;
}
if let Some(r) = tp.refs.iter().find(|r| !impl_names.contains(*r)) {
errs.push(compile_err!(
"batch-impl: inherited bound `{}` references parameter `{}`, but the impl declares \
no such name; declare `{}` or write the bound manually",
b,
r,
r
));
continue;
}
*bound = Some(TyPrimitive(b.clone()).to_ty());
}
for (pred, refs) in &trait_bounds.extra_predicates {
if let Some(r) = refs.iter().find(|r| !impl_names.contains(*r)) {
errs.push(compile_err!(
"batch-impl: inherited where predicate `{}` references parameter `{}`, \
but the impl declares no such name; declare `{}` or hand-write the where clause",
pred,
r,
r
));
continue;
}
parts.where_clauses.push(pred.clone());
}
errs
}
pub(crate) fn bare_param_name(name: &TokenStream) -> TokenStream {
let mut tokens = name.clone().into_iter();
match (tokens.next(), tokens.next()) {
(Some(TokenTree::Ident(id)), None) => quote!(#id),
(Some(TokenTree::Ident(kw)), Some(TokenTree::Ident(id)))
if kw == "const" && tokens.next().is_none() =>
{
quote!(#id)
}
_ => name.clone(),
}
}
pub(crate) fn merge_dup_params(parts: &mut ImplParts) {
let mut counts: HashMap<String, usize> = HashMap::new();
for (name, _) in &parts.impl_generics {
*counts.entry(bare_param_name(name).to_string()).or_insert(0) += 1;
}
let mut merged: Vec<(TokenStream, Option<Ty>)> = Vec::new();
let mut extra_where: Vec<TokenStream> = Vec::new();
let mut seen: HashSet<String> = HashSet::new();
for (name, bound) in std::mem::take(&mut parts.impl_generics) {
let name_str = name.to_string();
let is_const = name_str.starts_with("const");
let key = bare_param_name(&name).to_string();
if counts.get(&key).copied().unwrap_or(0) > 1 {
if is_const {
if !seen.insert(key) {
continue; }
merged.push((name, bound));
} else {
if seen.insert(key.clone()) {
merged.push((name.clone(), None));
}
if let Some(b) = bound {
extra_where.push(quote!(#name: #b));
}
}
} else {
merged.push((name, bound));
}
}
parts.impl_generics = merged;
parts.where_clauses.extend(extra_where);
}