use proc_macro2::{Span, TokenStream};
use quote::{format_ident, quote};
use syn::{
GenericArgument, GenericParam, Ident, ItemTrait, Lifetime, Path, PathArguments, parse_quote,
};
use super::input::{Input, SealedType};
use crate::util::{argument, lifetimes, mentions, name_of, render};
pub(crate) fn expand(input: Input) -> TokenStream {
let Input { mut item, types } = input;
let module = format_ident!("__sealed_{}", item.ident);
let marker_params: Vec<TokenStream> = item
.generics
.params
.iter()
.filter_map(marker_param)
.collect();
let marker_args: Vec<TokenStream> = item
.generics
.params
.iter()
.filter(|param| !matches!(param, GenericParam::Lifetime(_)))
.map(argument)
.collect();
let marker_declarations = (!marker_params.is_empty()).then(|| quote!(<#(#marker_params),*>));
let marker_arguments = (!marker_args.is_empty()).then(|| quote!(<#(#marker_args),*>));
item.supertraits
.push(parse_quote!(#module::Sealed #marker_arguments));
let assertion = assertion(&item, &types);
let mut written = Vec::new();
let impls: Vec<TokenStream> = types
.iter()
.filter_map(|entry| {
let generics = declarations(¶ms(entry, &item));
let arguments = pinned(entry, &item);
let arguments = (!arguments.is_empty()).then(|| quote!(<#(#arguments),*>));
let ty = &entry.ty;
let already = render("e!(#generics #arguments #ty));
if written.contains(&already) {
return None;
}
written.push(already);
Some(quote! {
impl #generics Sealed #arguments for #ty {}
impl #generics unforgeable::Unforgeable #arguments for #ty {}
})
})
.collect();
let ident = &item.ident;
let message = format!("`{{Self}}` cannot implement `{ident}`");
let label = format!("not permitted to implement `{ident}` here");
let note = format!(
"only the types listed in `#[sealed(..)]` on `{ident}` may implement it, and only at the \
instantiations listed there"
);
let forged =
format!("`{{Self}}` cannot be given `{ident}`'s seal from outside the macro that wrote it");
let forged_note =
format!("add the type to `#[sealed(..)]` on `{ident}` instead of implementing this");
quote! {
#[doc(hidden)]
#[allow(non_snake_case)]
mod #module {
#[allow(unused_imports)]
use super::*;
#[diagnostic::on_unimplemented(
message = #message,
label = #label,
note = #note,
)]
pub trait Sealed #marker_declarations: unforgeable::Unforgeable #marker_arguments {}
mod unforgeable {
#[diagnostic::on_unimplemented(
message = #forged,
note = #forged_note,
)]
pub trait Unforgeable #marker_declarations {}
}
#(#impls)*
}
#assertion
#item
}
}
fn params(entry: &SealedType, item: &ItemTrait) -> Vec<(String, TokenStream)> {
let ty = &entry.ty;
let named = lifetimes(ty);
let mut params: Vec<(String, TokenStream)> = Vec::new();
let declared = |params: &[(String, TokenStream)], name: &str| {
params.iter().any(|(known, _)| known == name)
};
if let Some(binder) = &entry.binder {
for param in &binder.params {
params.push((name_of(param), quote!(#param)));
}
}
for param in &item.generics.params {
let name = name_of(param);
if declared(¶ms, &name) {
continue;
}
let used = match param {
GenericParam::Lifetime(param) => named.contains(¶m.lifetime.ident),
GenericParam::Type(_) | GenericParam::Const(_) => mentions(ty, &name),
};
if used {
params.push((name, quote!(#param)));
}
}
for lifetime in named {
let name = lifetime.to_string();
if declared(¶ms, &name) {
continue;
}
let lifetime = Lifetime::new(&format!("'{lifetime}"), lifetime.span());
params.push((name, quote!(#lifetime)));
}
params
}
fn declarations(params: &[(String, TokenStream)]) -> Option<TokenStream> {
let declarations = params.iter().map(|(_, tokens)| tokens);
(!params.is_empty()).then(|| quote!(<#(#declarations),*>))
}
fn assertion(item: &ItemTrait, types: &[SealedType]) -> Option<TokenStream> {
let ident = &item.ident;
let trait_params = &item.generics.params;
let names = item.generics.params.iter().map(argument);
let trait_arguments = (!trait_params.is_empty()).then(|| quote!(<#(#names),*>));
let trait_params = (!trait_params.is_empty()).then(|| quote!(#trait_params,));
let supplied: Vec<String> = item
.generics
.params
.iter()
.filter(|param| !matches!(param, GenericParam::Lifetime(_)))
.map(name_of)
.collect();
let checks: Vec<TokenStream> = types
.iter()
.enumerate()
.map(|(index, entry)| {
let params = params(entry, item);
let arguments = match &entry.instantiation {
Some(path) => instantiation(path),
None => supplied
.iter()
.map(|name| {
let name = Ident::new(name, Span::call_site());
quote!(#name)
})
.collect(),
};
let check = format_ident!("check_{index}");
let declarations = declarations(¶ms);
let ty = &entry.ty;
quote! {
fn #check #declarations(value: &#ty) {
assert::<#(#arguments,)* _>(value);
}
}
})
.collect();
if checks.is_empty() {
return None;
}
Some(quote! {
#[allow(dead_code)]
const _: () = {
fn assert<#trait_params S: #ident #trait_arguments + ?Sized>(_: &S) {}
#(#checks)*
};
})
}
fn marker_param(param: &GenericParam) -> Option<TokenStream> {
match param {
GenericParam::Lifetime(_) => None,
GenericParam::Type(param) => {
let ident = ¶m.ident;
Some(quote!(#ident))
}
GenericParam::Const(param) => {
let ident = ¶m.ident;
let ty = ¶m.ty;
Some(quote!(const #ident: #ty))
}
}
}
fn pinned(entry: &SealedType, item: &ItemTrait) -> Vec<TokenStream> {
match &entry.instantiation {
Some(path) => instantiation(path),
None => item
.generics
.params
.iter()
.filter(|param| !matches!(param, GenericParam::Lifetime(_)))
.map(argument)
.collect(),
}
}
fn instantiation(path: &Path) -> Vec<TokenStream> {
let Some(segment) = path.segments.last() else {
return Vec::new();
};
let PathArguments::AngleBracketed(arguments) = &segment.arguments else {
return Vec::new();
};
arguments
.args
.iter()
.filter(|arg| !matches!(arg, GenericArgument::Lifetime(_)))
.map(|arg| quote!(#arg))
.collect()
}