use std::collections::HashSet;
use syn::{
DeriveInput, GenericParam, Ident, TypePath,
visit::{self, Visit},
};
use super::field::CtxField;
pub(crate) fn validate_generics<'a>(
input: &DeriveInput,
context_fields: impl IntoIterator<Item = &'a CtxField<'a>>,
) -> syn::Result<()> {
for param in &input.generics.params {
match param {
GenericParam::Type(_) => {}
GenericParam::Lifetime(lifetime) => {
return Err(syn::Error::new_spanned(
lifetime,
"`#[derive(FromContext)]` does not support lifetime parameters; the context must be `'static`",
));
}
GenericParam::Const(const_param) => {
return Err(syn::Error::new_spanned(
const_param,
"`#[derive(FromContext)]` does not support const generic parameters",
));
}
}
}
let type_params: HashSet<Ident> = input
.generics
.type_params()
.map(|param| param.ident.clone())
.collect();
if type_params.is_empty() {
return Ok(());
}
let mut used_in_fields = HashSet::new();
let mut used_in_borrow_targets = HashSet::new();
for cf in context_fields {
let mut collector = UsedTypeParams {
type_params: &type_params,
used: &mut used_in_fields,
};
collector.visit_type(cf.ty());
if let Some(target) = &cf.borrow_target {
let mut collector = UsedTypeParams {
type_params: &type_params,
used: &mut used_in_borrow_targets,
};
collector.visit_type(target);
}
}
for param in input.generics.type_params() {
let ident = ¶m.ident;
if !used_in_fields.contains(ident) && !used_in_borrow_targets.contains(ident) {
return Err(syn::Error::new_spanned(
ident,
"type parameter does not appear in any context field",
));
}
if !used_in_fields.contains(ident) {
return Err(syn::Error::new_spanned(
ident,
"type parameter appears in a `#[context(borrow = ...)]` target, but no fields",
));
}
}
Ok(())
}
struct UsedTypeParams<'a> {
type_params: &'a HashSet<Ident>,
used: &'a mut HashSet<Ident>,
}
impl<'ast> Visit<'ast> for UsedTypeParams<'_> {
fn visit_type_path(&mut self, ty: &'ast TypePath) {
if ty.qself.is_none()
&& ty.path.segments.len() == 1
&& let Some(segment) = ty.path.segments.first()
&& segment.arguments.is_none()
&& self.type_params.contains(&segment.ident)
{
self.used.insert(segment.ident.clone());
}
visit::visit_type_path(self, ty);
}
}