use std::collections::BTreeSet;
use proc_macro2::Span;
use proc_macro2::TokenStream;
use proc_macro2::TokenTree;
use quote::ToTokens;
use quote::format_ident;
use quote::quote;
use syn::Data;
use syn::DeriveInput;
use syn::Field;
use syn::GenericArgument;
use syn::GenericParam;
use syn::Generics;
use syn::Ident;
use syn::Lifetime;
use syn::Meta;
use syn::Path;
use syn::PathArguments;
use syn::Token;
use syn::Type;
use syn::WhereClause;
use syn::WherePredicate;
use syn::parse_quote;
use syn::punctuated::Punctuated;
use syn::token::Comma;
use crate::model::ContainerData;
use crate::model::FieldMode;
use crate::model::FieldsData;
pub(crate) fn add_redact_bounds(generics: &mut Generics, model: &ContainerData<'_>, runtime: &Path) {
for_each_field(model, &mut |field, mode, _serialize_with| match mode {
FieldMode::Unmarked => {
add_trait_bound(generics, field, quote!(::core::fmt::Debug));
}
FieldMode::KeyedBy(_) => {
add_trait_bound(generics, field, quote!(::core::fmt::Debug));
add_trait_bound(generics, field, quote!(#runtime::domain::RedactLevelValue));
}
FieldMode::DisplayLevel(_) => add_trait_bound(generics, field, quote!(::core::fmt::Display)),
FieldMode::Level(_) => {
add_trait_bound(generics, field, quote!(#runtime::domain::RedactLevelValue));
}
FieldMode::Nested => {
add_trait_bound(generics, field, quote!(#runtime::Redact));
}
FieldMode::Map => add_trait_bound(generics, field, quote!(#runtime::domain::RedactMapValue)),
FieldMode::MapLevels { .. } => add_trait_bound(generics, field, quote!(#runtime::domain::RedactMapKeyValue)),
FieldMode::Json => add_trait_bound(generics, field, quote!(#runtime::domain::RedactJsonValue)),
FieldMode::Skip => {}
});
}
fn for_each_field(model: &ContainerData<'_>, callback: &mut impl FnMut(&Field, &FieldMode, Option<&Path>)) {
match model {
ContainerData::Struct(fields) => for_each_fields(fields, callback),
ContainerData::Enum(variants) => {
for variant in variants {
for_each_fields(variant.fields(), callback);
}
}
}
}
fn for_each_fields(fields: &FieldsData<'_>, callback: &mut impl FnMut(&Field, &FieldMode, Option<&Path>)) {
match fields {
FieldsData::Named(fields) => {
for field in fields {
callback(
field.field(),
field.attributes().mode(),
field.serde_attributes().serialize_with(),
);
}
}
FieldsData::Unnamed(fields) => {
for field in fields {
callback(
field.field(),
field.attributes().mode(),
field.serde_attributes().serialize_with(),
);
}
}
FieldsData::Unit => {}
}
}
fn add_trait_bound(generics: &mut Generics, field: &Field, trait_path: TokenStream) {
if !uses_type_parameter(generics, &field.ty) {
return;
}
let field_type = &field.ty;
let predicate: WherePredicate = parse_quote!(#field_type: #trait_path);
let candidate = predicate.to_token_stream().to_string();
let where_clause = generics.make_where_clause();
if where_clause
.predicates
.iter()
.any(|item| item.to_token_stream().to_string() == candidate)
{
return;
}
where_clause.predicates.push(predicate);
}
#[must_use]
fn uses_type_parameter(generics: &Generics, field_type: &impl ToTokens) -> bool {
let parameters: Vec<String> = generics
.params
.iter()
.filter_map(|parameter| {
let GenericParam::Type(parameter) = parameter else {
return None;
};
Some(parameter.ident.to_string())
})
.collect();
token_stream_uses_parameter(field_type.to_token_stream(), ¶meters)
}
#[must_use]
pub(crate) fn generics_for_field(generics: &Generics, field_type: &Type) -> Generics {
let parameter_names = generic_parameter_names(generics);
let mut used = BTreeSet::new();
collect_parameter_names(field_type.to_token_stream(), ¶meter_names, &mut used);
loop {
let mut changed = false;
for parameter in &generics.params {
if used.contains(&generic_parameter_name(parameter)) {
for name in parameter_names_in(parameter, ¶meter_names) {
changed |= used.insert(name);
}
}
}
if let Some(where_clause) = &generics.where_clause {
for predicate in &where_clause.predicates {
let names = parameter_names_in(predicate, ¶meter_names);
if names.iter().any(|name| used.contains(name)) {
changed |= names.iter().any(|name| used.insert(name.clone()));
}
}
}
if !changed {
break;
}
}
let mut filtered = generics.clone();
filtered.params = generics
.params
.iter()
.filter(|parameter| used.contains(&generic_parameter_name(parameter)))
.cloned()
.collect();
filtered.where_clause = filtered.where_clause.and_then(|where_clause| {
let predicates: Punctuated<WherePredicate, Comma> = where_clause
.predicates
.into_iter()
.filter(|predicate| {
let names = parameter_names_in(predicate, ¶meter_names);
names.iter().any(|name| used.contains(name))
})
.collect();
if predicates.is_empty() {
None
} else {
Some(WhereClause {
where_token: where_clause.where_token,
predicates,
})
}
});
filtered
}
#[must_use]
pub(crate) fn fresh_identifier(generics: &Generics, base: &str) -> Ident {
let used = generic_parameter_names(generics);
if !used.contains(base) {
return format_ident!("{base}");
}
(0..)
.map(|index| format_ident!("{base}_{index}"))
.find(|candidate| !used.contains(&candidate.to_string()))
.expect("an unused generated identifier should always exist")
}
#[must_use]
pub(crate) fn fresh_lifetime(generics: &Generics) -> Lifetime {
let used = generic_parameter_names(generics);
let base = "__qubit_redact_lifetime";
let name = if !used.contains(base) {
base.to_owned()
} else {
(0..)
.map(|index| format!("{base}_{index}"))
.find(|candidate| !used.contains(candidate))
.expect("an unused generated lifetime should always exist")
};
Lifetime::new(&format!("'{name}"), Span::call_site())
}
#[must_use]
fn generic_parameter_names(generics: &Generics) -> BTreeSet<String> {
generics.params.iter().map(generic_parameter_name).collect()
}
#[must_use]
fn generic_parameter_name(parameter: &GenericParam) -> String {
match parameter {
GenericParam::Type(parameter) => parameter.ident.to_string(),
GenericParam::Const(parameter) => parameter.ident.to_string(),
GenericParam::Lifetime(parameter) => parameter.lifetime.ident.to_string(),
}
}
#[must_use]
fn parameter_names_in(tokens: &impl ToTokens, candidates: &BTreeSet<String>) -> BTreeSet<String> {
let mut names = BTreeSet::new();
collect_parameter_names(tokens.to_token_stream(), candidates, &mut names);
names
}
fn collect_parameter_names(tokens: TokenStream, candidates: &BTreeSet<String>, names: &mut BTreeSet<String>) {
for token in tokens {
match token {
TokenTree::Ident(identifier) => {
let name = identifier.to_string();
if candidates.contains(&name) {
names.insert(name);
}
}
TokenTree::Group(group) => {
collect_parameter_names(group.stream(), candidates, names);
}
TokenTree::Punct(_) | TokenTree::Literal(_) => {}
}
}
}
#[must_use]
fn token_stream_uses_parameter(tokens: TokenStream, parameters: &[String]) -> bool {
tokens.into_iter().any(|token| match token {
TokenTree::Ident(identifier) => {
let name = identifier.to_string();
parameters.iter().any(|parameter| parameter == &name)
}
TokenTree::Group(group) => token_stream_uses_parameter(group.stream(), parameters),
TokenTree::Punct(_) | TokenTree::Literal(_) => false,
})
}
pub(crate) fn non_recursive_predicates(input: &DeriveInput, predicate: WherePredicate) -> Vec<WherePredicate> {
let WherePredicate::Type(bound) = &predicate else {
return vec![predicate];
};
let name = &input.ident;
let (_, arguments, _) = input.generics.split_for_impl();
let target = quote!(#name #arguments).to_string();
let Some(types) = recursive_types(&bound.bounded_ty, &target) else {
return vec![predicate];
};
types
.into_iter()
.map(|ty| {
let mut replacement = bound.clone();
replacement.bounded_ty = ty.clone();
WherePredicate::Type(replacement)
})
.collect()
}
fn recursive_types<'a>(ty: &'a Type, target: &str) -> Option<Vec<&'a Type>> {
let rendered = ty.to_token_stream().to_string();
if rendered == target || rendered == "Self" {
return Some(Vec::new());
}
let children: Vec<&Type> = match ty {
Type::Path(path) if path.qself.is_none() => {
let segment = path.path.segments.last()?;
let root = &path.path.segments.first()?.ident;
if path.path.segments.len() > 1 && root != "std" && root != "alloc" && root != "core" {
return None;
}
if !matches!(
segment.ident.to_string().as_str(),
"Option"
| "Box"
| "Rc"
| "Arc"
| "Vec"
| "VecDeque"
| "LinkedList"
| "BTreeMap"
| "BTreeSet"
| "HashMap"
| "HashSet"
) {
return None;
}
let PathArguments::AngleBracketed(arguments) = &segment.arguments else {
return None;
};
arguments
.args
.iter()
.filter_map(|argument| match argument {
GenericArgument::Type(ty) => Some(ty),
_ => None,
})
.collect()
}
Type::Tuple(tuple) => tuple.elems.iter().collect(),
Type::Array(array) => vec![&array.elem],
Type::Paren(paren) => vec![&paren.elem],
_ => return None,
};
let mut found = false;
let bounds = children
.into_iter()
.flat_map(|child| {
if let Some(bounds) = recursive_types(child, target) {
found = true;
bounds
} else {
vec![child]
}
})
.collect();
found.then_some(bounds)
}
pub(crate) fn normalize_recursive_fields(input: &mut DeriveInput) {
let name = &input.ident;
let (_, arguments, _) = input.generics.split_for_impl();
let target = quote!(#name #arguments).to_string();
let fields: Vec<_> = match &mut input.data {
Data::Struct(data) => data.fields.iter_mut().collect(),
Data::Enum(data) => data
.variants
.iter_mut()
.flat_map(|variant| variant.fields.iter_mut())
.collect(),
Data::Union(_) => return,
};
for field in fields {
let custom_serializer = field
.attrs
.iter()
.filter(|attribute| attribute.path().is_ident("serde"))
.any(|attribute| {
attribute
.parse_args_with(Punctuated::<Meta, Token![,]>::parse_terminated)
.is_ok_and(|items| {
items
.iter()
.any(|item| item.path().is_ident("with") || item.path().is_ident("serialize_with"))
})
});
if !custom_serializer
&& !field.attrs.iter().any(|attribute| attribute.path().is_ident("redact"))
&& recursive_types(&field.ty, &target).is_some()
{
field.attrs.push(parse_quote!(#[redact(nested)]));
}
}
}