use crate::TraitBounds;
use crate::ast::*;
use crate::diagnostic::compile_error_str;
use proc_macro2::TokenStream;
use quote::{ToTokens, quote};
pub(crate) struct ImplParts {
pub(crate) impl_generics: Vec<(TokenStream, Option<Ty>)>,
pub(crate) trait_generic_names: Vec<TokenStream>,
pub(crate) associated_types: Vec<(TokenStream, TokenStream)>,
pub(crate) target_type: Ty,
pub(crate) body: Option<TokenStream>,
pub(crate) attrs: Vec<TokenStream>,
pub(crate) is_unsafe_impl: bool,
pub(crate) where_clauses: Vec<TokenStream>,
}
impl ImplParts {
fn leaf(target_type: Ty) -> Self {
ImplParts {
impl_generics: vec![],
trait_generic_names: vec![],
associated_types: vec![],
target_type,
body: None,
attrs: vec![],
is_unsafe_impl: false,
where_clauses: vec![],
}
}
}
pub(crate) fn extract_impl_parts(ty: Ty) -> ImplParts {
match ty {
Ty::WithType(wt) => {
let mut parts = extract_impl_parts(*wt.1);
let (impl_generics, associated_types) =
(parts.impl_generics, parts.associated_types);
parts.impl_generics = wt.0.params;
parts.associated_types = wt.0.bindings;
parts.impl_generics.extend(impl_generics);
parts.associated_types.extend(associated_types);
parts
}
Ty::WithTrait(wt) => {
let mut parts = extract_impl_parts(*wt.1);
parts.trait_generic_names.extend(wt.0.1.params.into_iter().map(|p| p.0));
parts.associated_types.extend(wt.0.1.bindings);
parts
}
Ty::WithCode(wc) => match wc.0 {
Some(inner) => {
let mut parts = extract_impl_parts(*inner);
match &mut parts.body {
Some(t) => t.extend(wc.1.0),
None => parts.body = wc.1.0.into(),
}
parts
}
None => ImplParts::leaf(wc.into()),
},
Ty::WithWhere(ww) => match ww.0 {
Some(inner) => {
let mut parts = extract_impl_parts(*inner);
parts.where_clauses.push(ww.1.0);
parts
}
None => ImplParts::leaf(ww.into()),
},
Ty::WithAttr(wa) => match wa.1 {
Some(inner) => {
let mut parts = extract_impl_parts(*inner);
let stream = &wa.0.0;
parts.attrs.push(quote!(#[#stream]));
parts
}
None => ImplParts::leaf(wa.into()),
},
Ty::WithPrefix(wp) => match wp.1 {
Some(inner) => {
let mut parts = extract_impl_parts(*inner);
match wp.0 {
TyPrefix::Unsafe => parts.is_unsafe_impl = true,
_ => {
parts.target_type =
TyWithPrefix(wp.0, parts.target_type.into()).into()
}
}
parts
}
None => ImplParts::leaf(wp.into()),
},
Ty::Error(e) => ImplParts::leaf(e.into()),
o => ImplParts::leaf(o),
}
}
fn hoist_type_params(ty: Ty, out: &mut Vec<(TokenStream, Option<Ty>)>) -> Ty {
match ty {
Ty::WithType(wt) => {
out.extend(wt.0.params);
hoist_type_params(*wt.1, out)
}
Ty::Array(a) => {
TyArray(a.0.into_iter().map(|e| hoist_type_params(e, out)).collect())
.into()
}
Ty::Tuple(t) => {
TyTuple(t.0.into_iter().map(|e| hoist_type_params(e, out)).collect())
.into()
}
Ty::Group(g) => TyGroup(hoist_type_params(*g.0, out).into()).into(),
Ty::PrimitiveArray(pa) => {
TyPrimitiveArray(pa.0.map(|e| hoist_type_params(*e, out).into()), pa.1)
.into()
}
Ty::Generic(g) => {
let base = hoist_type_params(*g.0, out);
TyGeneric(base.into(), g.1).into()
}
Ty::WithPrefix(wp) => {
TyWithPrefix(wp.0, wp.1.map(|e| hoist_type_params(*e, out).into())).into()
}
Ty::WithTrait(wt) => {
TyWithTrait(wt.0, hoist_type_params(*wt.1, out).into()).into()
}
Ty::WithCode(wc) => {
TyWithCode(wc.0.map(|e| hoist_type_params(*e, out).into()), wc.1).into()
}
Ty::WithWhere(ww) => {
TyWithWhere(ww.0.map(|e| hoist_type_params(*e, out).into()), ww.1).into()
}
Ty::WithAttr(wa) => {
TyWithAttr(wa.0, wa.1.map(|e| hoist_type_params(*e, out).into())).into()
}
Ty::Fn(f) => TyFn(
f.0.map(|params| {
params.into_iter().map(|p| hoist_type_params(p, out)).collect()
}),
f.1.map(|r| hoist_type_params(*r, out).into()),
f.2,
)
.into(),
other => other,
}
}
pub(crate) fn generate_impl(
ty: Ty, trait_name: &TokenStream, is_unsafe_trait: bool,
trait_bounds: &TraitBounds,
) -> TokenStream {
if let Ty::Error(e) = ty {
return e.0;
}
if let Ty::WithCode(TyWithCode(None, code)) = &ty {
return code.0.clone();
}
let mut parts = extract_impl_parts(ty);
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<TokenStream> = vec![];
let trait_args: Vec<String> =
parts.trait_generic_names.iter().map(|n| n.to_string()).collect();
let impl_names: std::collections::HashSet<String> = parts
.impl_generics
.iter()
.map(|(n, _)| {
let s = n.to_string();
s.strip_prefix("const ").unwrap_or(&s).to_string()
})
.collect();
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_error_str(&format!(
"batch-impl: trait 实参 `{}` 对应形参 `{}`(bound `{}`),\
自动继承要求同名;请改名为 `{}` 或手写 bound",
key, tp.name, b, tp.name
)));
continue;
}
if let Some(r) = tp.refs.iter().find(|r| !impl_names.contains(*r)) {
errs.push(compile_error_str(&format!(
"batch-impl: 继承的 bound `{}` 引用形参 `{}`,impl 未声明同名参数;\
请声明 `{}` 或手写 bound",
b, r, r
)));
continue;
}
*bound = Some(Ty::Primitive(TyPrimitive(b.clone())));
}
for (pred, refs) in &trait_bounds.extra_predicates {
if let Some(r) = refs.iter().find(|r| !impl_names.contains(*r)) {
errs.push(compile_error_str(&format!(
"batch-impl: 继承的 where 谓词 `{}` 引用形参 `{}`,\
impl 未声明同名参数;请声明 `{}` 或手写 where",
pred, r, r
)));
continue;
}
parts.where_clauses.push(pred.clone());
}
if !errs.is_empty() {
return errs.into_iter().collect();
}
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)| match bound {
Some(b) => {
let b_tokens = b.to_token_stream();
quote!(#name: #b_tokens)
}
None => name.clone(),
})
.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 parts.where_clauses.is_empty() {
quote!()
} else {
let preds = &parts.where_clauses;
quote!(where #(#preds),*)
};
quote! {
#(#attrs)*
#unsafe_kw impl #impl_gen #trait_name #trait_gen for #target #where_clause {
#(#body_tokens)*
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::trait_bounds::extract_trait_bounds;
use syn::parse_quote;
#[test]
fn const_param_where_predicate_no_error() {
let trait_def: syn::ItemTrait = parse_quote!(
trait WhereArr<T, const N: usize>
where
[T; N]: Sized,
{
}
);
let tb = extract_trait_bounds(&trait_def);
let target: Ty = TyTuple(vec![]).into();
let trait_ty = TyTrait(
quote!(WhereArr),
TyTypeParam {
params: vec![(quote!(T), None), (quote!(N), None)],
bindings: vec![],
},
);
let wrapped = TyWithTrait(trait_ty, target.into());
let impl_ty = TyWithType(
TyTypeParam {
params: vec![
(quote!(T), None),
(
quote!(const N),
Some(Ty::Primitive(TyPrimitive(quote!(usize)))),
),
],
bindings: vec![],
},
wrapped.into(),
)
.into();
let out = generate_impl(impl_ty, "e!(WhereArr), false, &tb).to_string();
assert!(
!out.contains("compile_error"),
"展开不应含 compile_error:{out}"
);
assert!(
out.contains("where [T ; N] : Sized"),
"缺少 where 谓词:{out}"
);
assert!(
out.contains("impl < T , const N : usize > WhereArr < T , N >"),
"impl 泛型异常:{out}"
);
}
}