use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{GenericArgument, GenericParam, ItemTrait, Path, PathArguments, parse_quote};
use super::input::{Input, SealedType};
use crate::util::{argument, fresh, 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));
let arguments = pinned(entry);
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) -> Vec<(String, TokenStream)> {
entry
.binder
.iter()
.flat_map(|binder| binder.params.iter())
.map(|param| (name_of(param), quote!(#param)))
.collect()
}
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 checks: Vec<TokenStream> = types
.iter()
.enumerate()
.map(|(index, entry)| {
let params = params(entry);
let arguments = match &entry.instantiation {
Some(path) => instantiation(path),
None => Vec::new(),
};
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;
}
let probe = fresh(item.generics.params.iter().map(name_of), "S");
Some(quote! {
#[allow(dead_code)]
const _: () = {
fn assert<#trait_params #probe: #ident #trait_arguments + ?Sized>(_: &#probe) {}
#(#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) -> Vec<TokenStream> {
match &entry.instantiation {
Some(path) => instantiation(path),
None => Vec::new(),
}
}
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()
}