use convert_case::{Case, Casing};
use proc_macro2::TokenStream;
use quote::{ToTokens, format_ident, quote};
use syn::{
Expr, Fields, Ident, LitStr, Token, Type,
parse::{Parse, ParseStream},
};
pub(crate) enum DisplayDefault {
Code,
Name,
}
pub(crate) struct Unit<'a> {
pub pat: TokenStream,
pub fields: &'a Fields,
pub error: Option<ErrorSpec>,
pub name: Ident,
pub source: Option<TokenStream>,
}
#[derive(Clone)]
pub(crate) struct ErrorSpec {
pub lit: LitStr,
pub args: Vec<Expr>,
}
fn validate_field_access_path(expr: &Expr) -> syn::Result<()> {
match expr {
Expr::Path(p) if p.path.get_ident().is_some() => Ok(()),
Expr::Field(f) => validate_field_access_path(&f.base),
_ => Err(syn::Error::new_spanned(
expr,
"#[error(\"..\", ..)] trailing arguments must be field access paths (e.g. \
`field.sub_field`); arbitrary expressions are not supported",
)),
}
}
impl Parse for ErrorSpec {
fn parse(input: ParseStream) -> syn::Result<Self> {
let lit: LitStr = input.parse()?;
let mut args = Vec::new();
while !input.is_empty() {
input.parse::<Token![,]>()?;
if input.is_empty() {
break;
}
let arg = input.parse::<Expr>()?;
validate_field_access_path(&arg)?;
args.push(arg);
}
Ok(ErrorSpec { lit, args })
}
}
fn root_ident(expr: &Expr) -> &Ident {
match expr {
Expr::Path(p) => p.path.get_ident().expect("validated as a bare ident"),
Expr::Field(f) => root_ident(&f.base),
_ => unreachable!("ErrorSpec::args are validated to be field access paths"),
}
}
fn check_trailing_args(args: &[Expr], fields: &Fields) -> syn::Result<()> {
let names = bound_fields(fields);
for arg in args {
let root = root_ident(arg);
if !names.iter().any(|n| n == root) {
return Err(syn::Error::new_spanned(
root,
format!(
"#[error(\"..\", ..)] trailing argument starts from `{root}`, which is not \
a field here"
),
));
}
}
Ok(())
}
fn field_list(fields: &Fields) -> Vec<&syn::Field> {
match fields {
Fields::Named(f) => f.named.iter().collect(),
Fields::Unnamed(f) => f.unnamed.iter().collect(),
Fields::Unit => Vec::new(),
}
}
pub(crate) fn field_binding(fields: &Fields, index: usize) -> Ident {
match field_list(fields)[index].ident.clone() {
Some(name) => name,
None => format_ident!("field_{index}"),
}
}
fn bound_fields(fields: &Fields) -> Vec<Ident> {
(0..field_list(fields).len())
.map(|i| field_binding(fields, i))
.collect()
}
pub(crate) fn bind_pattern(path: TokenStream, fields: &Fields) -> TokenStream {
let names = bound_fields(fields);
match fields {
Fields::Named(_) => quote! { #path { #(#names),* } },
Fields::Unnamed(_) => quote! { #path(#(#names),*) },
Fields::Unit => path,
}
}
pub(crate) fn marked_source(fields: &Fields, span: &Ident) -> syn::Result<Option<TokenStream>> {
let list = field_list(fields);
let mut marked = list.iter().enumerate().filter(|(_, f)| {
f.attrs.iter().any(|a| a.path().is_ident("source"))
|| f.ident.as_ref().is_some_and(|i| i == "source")
});
match (marked.next(), marked.next()) {
(None, _) => Ok(None),
(Some((i, _)), None) => {
let name = field_binding(fields, i);
Ok(Some(quote! { #name }))
}
(Some(_), Some(_)) => Err(syn::Error::new_spanned(
span,
"more than one `#[source]`/`source` field; errlanes picks at most one",
)),
}
}
pub(crate) fn take_error_lit(attrs: &[syn::Attribute]) -> syn::Result<Option<ErrorSpec>> {
let mut found = None;
for attr in attrs {
if attr.path().is_ident("error") {
if found.is_some() {
return Err(syn::Error::new_spanned(
attr,
"at most one #[error(\"..\")] per type or variant",
));
}
found = Some(attr.parse_args::<ErrorSpec>()?);
}
}
Ok(found)
}
pub(crate) fn reject_from_field_attr(fields: &Fields) -> syn::Result<()> {
for field in field_list(fields) {
if field.attrs.iter().any(|a| a.path().is_ident("from")) {
return Err(syn::Error::new_spanned(
field,
"`#[from]` is not supported here; use `from` inside `#[classify(..)]` / \
`#[rejection(..)]` instead, or `error = manual` to let thiserror own this type",
));
}
}
Ok(())
}
type FieldRef = (Ident, Type, bool);
fn rewrite_literal(
lit: &LitStr,
fields: &Fields,
allow_auto_index: bool,
) -> syn::Result<(String, Vec<FieldRef>)> {
let names = bound_fields(fields);
let types: Vec<Type> = field_list(fields).iter().map(|f| f.ty.clone()).collect();
let src = lit.value();
let mut out = String::with_capacity(src.len());
let mut refs = Vec::new();
let mut chars = src.chars().peekable();
while let Some(c) = chars.next() {
match c {
'{' if chars.peek() == Some(&'{') => {
chars.next();
out.push_str("{{");
}
'}' if chars.peek() == Some(&'}') => {
chars.next();
out.push_str("}}");
}
'{' => {
let mut body = String::new();
let mut closed = false;
for c2 in chars.by_ref() {
if c2 == '}' {
closed = true;
break;
}
body.push(c2);
}
if !closed {
return Err(syn::Error::new(
lit.span(),
"unterminated `{` in #[error(\"..\")]",
));
}
let (name_part, spec) = match body.find(':') {
Some(i) => (&body[..i], &body[i..]),
None => (body.as_str(), ""),
};
if name_part.is_empty() {
if allow_auto_index {
out.push('{');
out.push_str(&body);
out.push('}');
continue;
}
return Err(syn::Error::new(
lit.span(),
"#[error(\"..\")] does not support auto-indexed `{}`; name the field, \
e.g. `{0}` or `{field}`",
));
}
let is_debug = spec.contains('?');
if let Ok(index) = name_part.parse::<usize>() {
let Some(name) = names.get(index) else {
return Err(syn::Error::new(
lit.span(),
format!(
"#[error(\"..\")] references field {index}, but this has {} \
field(s)",
names.len()
),
));
};
refs.push((name.clone(), types[index].clone(), is_debug));
out.push('{');
out.push_str(&name.to_string());
out.push_str(spec);
out.push('}');
} else {
let Some(index) = names.iter().position(|n| n == name_part) else {
return Err(syn::Error::new(
lit.span(),
format!(
"#[error(\"..\")] references `{{{name_part}}}`, which is not a \
field here"
),
));
};
refs.push((names[index].clone(), types[index].clone(), is_debug));
out.push('{');
out.push_str(&body);
out.push('}');
}
}
other => out.push(other),
}
}
Ok((out, refs))
}
fn type_mentions_param(ty: &Type, param: &Ident) -> bool {
fn contains(ts: TokenStream, param: &Ident) -> bool {
ts.into_iter().any(|t| match t {
proc_macro2::TokenTree::Ident(i) => i == *param,
proc_macro2::TokenTree::Group(g) => contains(g.stream(), param),
_ => false,
})
}
contains(ty.to_token_stream(), param)
}
pub(crate) fn field_types(fields: &Fields) -> Vec<Type> {
field_list(fields).iter().map(|f| f.ty.clone()).collect()
}
pub(crate) fn merge_where(
where_clause: Option<&syn::WhereClause>,
extra: &[TokenStream],
) -> TokenStream {
if extra.is_empty() {
quote! { #where_clause }
} else if let Some(wc) = where_clause {
quote! { #wc #(, #extra)* }
} else {
quote! { where #(#extra),* }
}
}
pub(crate) fn classify_bounds(generics: &syn::Generics, field_types: &[Type]) -> Vec<TokenStream> {
generics
.type_params()
.map(|p| &p.ident)
.filter(|p| field_types.iter().any(|ty| type_mentions_param(ty, p)))
.map(|p| quote! { #p: std::fmt::Debug + Send + Sync + 'static })
.collect()
}
fn extra_bounds(generics: &syn::Generics, refs: &[FieldRef]) -> Vec<TokenStream> {
let params: Vec<Ident> = generics.type_params().map(|p| p.ident.clone()).collect();
if params.is_empty() {
return Vec::new();
}
let mut bounds: Vec<(String, TokenStream)> = Vec::new();
for (_, ty, is_debug) in refs {
if !params.iter().any(|p| type_mentions_param(ty, p)) {
continue;
}
let bound = if *is_debug {
quote! { #ty: std::fmt::Debug }
} else {
quote! { #ty: std::fmt::Display }
};
let key = bound.to_string();
if !bounds.iter().any(|(k, _)| *k == key) {
bounds.push((key, bound));
}
}
bounds.into_iter().map(|(_, b)| b).collect()
}
pub(crate) fn emit(
target: &Ident,
generics: &syn::Generics,
default: DisplayDefault,
units: &[Unit<'_>],
) -> syn::Result<(TokenStream, Vec<TokenStream>)> {
let mut display_arms = Vec::with_capacity(units.len());
let mut source_arms = Vec::with_capacity(units.len());
let mut all_refs = Vec::new();
for unit in units {
let pat = &unit.pat;
let body = match &unit.error {
Some(spec) => {
check_trailing_args(&spec.args, unit.fields)?;
let (rewritten, refs) =
rewrite_literal(&spec.lit, unit.fields, !spec.args.is_empty())?;
let mut named_args = Vec::new();
for (name, _, _) in &refs {
if !named_args.contains(name) {
named_args.push(name.clone());
}
}
all_refs.extend(refs);
let rewritten = LitStr::new(&rewritten, spec.lit.span());
let args = &spec.args;
quote! { write!(f, #rewritten #(, #args)* #(, #named_args = #named_args)*) }
}
None => match default {
DisplayDefault::Code => {
quote! { f.write_str(<Self as errlanes::Rejection>::code(self).into()) }
}
DisplayDefault::Name => {
let name = unit.name.to_string().to_case(Case::Snake);
quote! { f.write_str(#name) }
}
},
};
display_arms.push(quote! { #pat => #body });
let source_body = match &unit.source {
Some(name) => quote! { Some(#name as &(dyn std::error::Error + 'static)) },
None => quote! { None },
};
source_arms.push(quote! { #pat => #source_body });
}
let extra = extra_bounds(generics, &all_refs);
let all_field_types: Vec<Type> = units.iter().flat_map(|u| field_types(u.fields)).collect();
let mut error_extra = extra.clone();
for param in generics.type_params().map(|p| &p.ident) {
if !all_field_types
.iter()
.any(|ty| type_mentions_param(ty, param))
{
continue;
}
let bound = quote! { #param: std::fmt::Debug };
let key = bound.to_string();
if !error_extra.iter().any(|b| b.to_string() == key) {
error_extra.push(bound);
}
}
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let display_where_tokens = merge_where(where_clause, &extra);
let error_where_tokens = merge_where(where_clause, &error_extra);
let tokens = quote! {
#[automatically_derived]
impl #impl_generics std::fmt::Display for #target #ty_generics #display_where_tokens {
#[allow(unused_variables)]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
#(#display_arms,)*
}
}
}
#[automatically_derived]
impl #impl_generics std::error::Error for #target #ty_generics #error_where_tokens {
#[allow(unused_variables)]
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
#(#source_arms,)*
}
}
}
};
Ok((tokens, extra))
}