mod fresh;
mod impl_parts;
mod postprocess;
mod repeat;
mod shape;
mod sync_trait;
mod top_level;
mod where_at;
pub(crate) use fresh::*;
pub(crate) use impl_parts::*;
pub(crate) use postprocess::*;
pub(crate) use repeat::*;
pub(crate) use shape::*;
pub(crate) use sync_trait::*;
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_error_str;
use proc_macro2::{Ident, TokenStream, TokenTree};
use quote::{ToTokens, quote};
use std::collections::{HashMap, 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() == '!'
);
return compile_error_str(
if is_top_marked {
"batch-impl: a top-level `{! ...}` block needs an attached type \
(the spec body is prepended to the macro input)"
} else {
"batch-impl: a bare `{...}` block without an attached type \
generates no impl (attach it to a type, e.g. `T { ... }`, or \
use the top-level `{! ...}` macro form)"
},
code.0
.clone()
.into_iter()
.next()
.map_or_else(proc_macro2::Span::call_site, |t| t.span()),
);
}
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);
merge_dup_params(&mut parts);
let impl_name_streams =
parts.impl_generics.iter().map(|(n, _)| bare_param_name(n)).collect::<Vec<TokenStream>>();
let impl_names = impl_name_streams.iter().map(|n| n.to_string()).collect::<HashSet<String>>();
let trait_args =
parts.trait_generic_names.iter().map(|n| n.to_string()).collect::<Vec<String>>();
let mut errs = inherit_trait_bounds(&mut parts, trait_bounds, &trait_args, &impl_names);
if let Some(trait_ident) = trait_last_ident(trait_name) {
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);
match sync_trait_application(t, &trait_args) {
Ok(s) => {
if is_switch {
body_sync = true;
} else {
matched.push(s);
}
}
Err(e) => return e,
}
}
parts.impl_templates = matched;
let mut synced = Vec::with_capacity(parts.where_clauses.len());
for w in &parts.where_clauses {
match sync_trait_application(w.clone(), &trait_args) {
Ok(s) => synced.push(s),
Err(e) => return e,
}
}
parts.where_clauses = synced;
for (_, bound) in &mut parts.impl_generics {
if let Some(b) = bound {
*b = match sync_bound_ty(b, &trait_args) {
Ok(t) => t,
Err(e) => return e,
};
}
}
if body_sync && let Some(b) = &mut parts.body {
*b = match sync_trait_application(b.clone(), &trait_args) {
Ok(s) => s,
Err(e) => return e,
};
}
}
let where_resolved = match resolve_where_predicates(&parts.where_clauses, &impl_name_streams) {
Ok(ws) => ws,
Err(es) => {
errs.extend(es);
vec![]
}
};
errs.extend(validate_at_refs(
&parts.target_type,
&parts.trait_generic_names,
&impl_name_streams,
));
if !errs.is_empty() {
return errs.into_iter().collect();
}
let (shape_entries, var_segs) = if parts.impl_templates.is_empty() {
(Vec::new(), Vec::new())
} else {
match collect_shape_mapping(&parts) {
Ok((m, s)) => (m.entries().to_vec(), s),
Err(e) => return compile_error_str(&e.message(), proc_macro2::Span::call_site()),
}
};
if !shape_entries.is_empty() {
parts.where_clauses =
parts.where_clauses.iter().map(|p| apply_mapping(p.clone(), &shape_entries)).collect();
if let Some(b) = &mut parts.body {
match expand_repeat_blocks(b.clone(), &var_segs) {
Ok(expanded) => *b = apply_mapping(expanded, &shape_entries),
Err(e) => return e,
}
}
}
render_impl(parts, where_resolved, trait_name, is_unsafe_trait, &shape_entries)
}
fn collect_shape_mapping(parts: &ImplParts) -> Result<(Mapping, Vec<VarSeg>), ShapeError> {
let target_tokens = parts.target_type.to_token_stream();
let target: syn::Type = syn::parse2(target_tokens).map_err(|_| {
ShapeError::ShapeMismatch(
"the target type is not a standard Rust type (DSL leftovers cannot be destructured by an `impl{...}` template)"
.into(),
)
})?;
let mut merged = Mapping::default();
let mut segs = vec![];
for t in &parts.impl_templates {
let template: syn::Type = syn::parse2(t.clone()).map_err(|_| {
ShapeError::ShapeMismatch(
"the `impl{...}` template is not a standard Rust type (DSL operators are not allowed inside)"
.into(),
)
})?;
let (m, s) = match_shape(&template, &target)?;
merged.merge(m)?;
segs.extend(s);
}
Ok((merged, segs))
}
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(),
}
}
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);
}
fn render_impl(
parts: ImplParts, where_resolved: Vec<TokenStream>, trait_name: &TokenStream,
is_unsafe_trait: bool, shape_entries: &[(String, TokenStream)],
) -> TokenStream {
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 = if shape_entries.is_empty() {
parts.target_type.to_token_stream()
} else {
apply_mapping(parts.target_type.to_token_stream(), shape_entries)
};
let mut body_tokens: Vec<TokenStream> =
parts.associated_types.iter().map(|(name, value)| quote!(type #name = #value;)).collect();
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)
}