mod fresh;
mod impl_parts;
mod postprocess;
mod top_level;
mod where_at;
pub(crate) use fresh::*;
pub(crate) use impl_parts::*;
pub(crate) use postprocess::*;
pub(crate) use top_level::*;
pub(crate) use where_at::*;
use crate::TraitBounds;
use crate::ast::types_render::render_param;
use crate::ast::*;
use crate::util::{compile_err, compile_error_str};
use proc_macro2::{Ident, TokenStream, TokenTree};
use quote::quote;
use std::collections::HashSet;
pub(crate) fn generate_impl(
ty: Ty, trait_name: &TokenStream, is_unsafe_trait: bool,
trait_bounds: &TraitBounds, trait_param_names: &[Ident],
) -> TokenStream {
if let Ty { kind: TyKind::WithCode(TyWithCode(None, code)), .. } = &ty {
let is_top_marked = matches!(
code.0.clone().into_iter().next(),
Some(TokenTree::Punct(p)) if p.as_char() == '!'
);
if is_top_marked {
return compile_error_str(
"batch-impl: a top-level `{! ...}` block needs an attached type \
(the spec body is prepended to the macro input)",
code.0
.clone()
.into_iter()
.next()
.map_or_else(proc_macro2::Span::call_site, |t| t.span()),
);
}
return code.0.clone();
}
if let Some(result) = top_level_macro(&ty) {
return match result {
Ok((spec, mac)) => {
if spec.is_empty() {
compile_error_str(
"batch-impl: a top-level `{! ...}` block needs an attached type \
(the spec body is prepended to the macro input)",
proc_macro2::Span::call_site(),
)
} else if mac.is_empty() {
compile_error_str(
"batch-impl: a `{! ...}` top-level block must contain a macro \
call (e.g. `{! my_macro!{...}}`)",
proc_macro2::Span::call_site(),
)
} else {
sweep_fresh_names(rewrite_macro_input(mac, spec))
}
}
Err(e) => e,
};
}
if let Ty { kind: TyKind::Error(e), .. } = ty {
return e.0;
}
let mut parts = extract_impl_parts(ty);
substitute_trait_generics(&mut parts, trait_param_names);
parts.target_type = expand_splat_elems(parts.target_type);
let mut nested_params = vec![];
parts.target_type = hoist_type_params(parts.target_type, &mut nested_params);
parts.impl_generics.extend(nested_params);
let mut errs = vec![];
let trait_args = parts
.trait_generic_names
.iter()
.map(|n| n.to_string())
.collect::<Vec<String>>();
let impl_name_streams = parts
.impl_generics
.iter()
.map(|(n, _)| {
let s = n.to_string();
let bare = s.strip_prefix("const ").unwrap_or(&s);
bare.parse().unwrap()
})
.collect::<Vec<TokenStream>>();
let impl_names =
impl_name_streams.iter().map(|n| n.to_string()).collect::<HashSet<String>>();
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());
}
let mut where_resolved = vec![];
for pred in &parts.where_clauses {
let head = pred.clone().into_iter().collect::<Vec<_>>();
if matches!(head.as_slice(),
[TokenTree::Punct(p), TokenTree::Group(g), ..]
if p.as_char() == '*'
&& matches!(
g.delimiter(),
proc_macro2::Delimiter::Parenthesis
| proc_macro2::Delimiter::Bracket
)
) {
errs.push(compile_err!(
"batch-impl: a bare splat cannot be a where-predicate subject \
(`*(A,B): Trait`); wrap it in a tuple (`(*(A,B)): Trait`) or \
write separate predicates"
));
continue;
}
match resolve_where_at(pred, &impl_name_streams) {
Ok(p) => where_resolved.push(p),
Err(e) => errs.push(e),
}
}
if !errs.is_empty() {
return errs.into_iter().collect();
}
let parts = parts;
let is_unsafe = is_unsafe_trait || parts.is_unsafe_impl;
let unsafe_kw = if is_unsafe { quote!(unsafe) } else { quote!() };
let impl_gen = if parts.impl_generics.is_empty() {
quote!()
} else {
let params = parts
.impl_generics
.iter()
.map(|(name, bound)| render_param(name, bound.as_ref()))
.collect::<Vec<_>>();
quote!(<#(#params),*>)
};
let trait_gen = if parts.trait_generic_names.is_empty() {
quote!()
} else {
let names = &parts.trait_generic_names;
quote!(<#(#names),*>)
};
let target = &parts.target_type;
let mut body_tokens = vec![];
for (name, value) in &parts.associated_types {
body_tokens.push(quote!(type #name = #value;));
}
if let Some(body) = &parts.body {
body_tokens.push(body.clone());
}
let attrs = parts.attrs;
let where_clause = if where_resolved.is_empty() {
quote!()
} else {
let preds = &where_resolved;
quote!(where #(#preds),*)
};
let rendered = quote! {
#(#attrs)*
#unsafe_kw impl #impl_gen #trait_name #trait_gen for #target #where_clause {
#(#body_tokens)*
}
};
sweep_fresh_names(rendered)
}