use proc_macro::TokenStream;
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::{format_ident, quote};
use std::collections::{BTreeMap, BTreeSet};
use syn::visit::Visit;
use syn::{
Data, DataEnum, DataStruct, DeriveInput, Field, Fields, Generics, Ident, Member,
parse_macro_input, spanned::Spanned,
};
#[proc_macro_derive(StackError, attributes(source, stack_error, location))]
pub fn derive_stack_error(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
match expand(input) {
Ok(tokens) => tokens.into(),
Err(error) => error.into_compile_error().into(),
}
}
fn expand(input: DeriveInput) -> syn::Result<TokenStream2> {
let ident = input.ident;
let generics = input.generics;
match input.data {
Data::Struct(data) => expand_struct(ident, generics, data),
Data::Enum(data) => expand_enum(ident, generics, data),
Data::Union(_) => Err(syn::Error::new(
Span::call_site(),
"StackError cannot be derived for unions",
)),
}
}
fn expand_struct(ident: Ident, generics: Generics, data: DataStruct) -> syn::Result<TokenStream2> {
let style = match &data.fields {
Fields::Named(_) => FieldsStyle::Named,
Fields::Unnamed(_) => FieldsStyle::Unnamed,
Fields::Unit => {
return Err(syn::Error::new(
ident.span(),
"unit structs do not support #[derive(StackError)]",
));
}
};
let fields = collect_fields(&data.fields)?;
let allow_name = style.allows_names();
let has_explicit_location = fields.iter().any(|f| {
f.attrs.is_location || (allow_name && matches!(&f.ident, Some(id) if id == "location"))
});
let location_index = if has_explicit_location {
resolve_location(&fields, allow_name, ident.span())?
} else {
resolve_location_from_located_source(&fields, allow_name)
.or_else(|| single_field_located(&fields))
.ok_or_else(|| {
syn::Error::new(
ident.span(),
"missing #[location] attribute or field named `location`",
)
})?
};
let source = resolve_source(&fields, style.allows_names())?;
let mut generics = generics;
let mut bounds = BoundsTracker::new(&generics);
if let Some(info) = &source {
bounds.collect(&fields[info.index].ty, info.is_terminal);
}
bounds.apply(&mut generics);
let location_member = &fields[location_index].member;
let location_expr = if is_located_error(&fields[location_index].ty) {
quote! { ::pseudo_backtrace::StackError::location(&self.#location_member) }
} else {
quote! { self.#location_member }
};
let next_body = match &source {
Some(info) => build_next_struct(
&fields[info.index].member,
&fields[info.index].ty,
info.is_terminal,
),
None => quote! { ::core::option::Option::None },
};
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
Ok(quote! {
impl #impl_generics ::pseudo_backtrace::StackError for #ident #ty_generics #where_clause {
fn location(&self) -> &'static ::core::panic::Location<'static> {
#location_expr
}
fn next<'pseudo_backtrace>(&'pseudo_backtrace self) -> ::core::option::Option<::pseudo_backtrace::Chain<'pseudo_backtrace>> {
use ::pseudo_backtrace::private::AsDynStdError as _;
use ::pseudo_backtrace::private::AsDynStackError as _;
#next_body
}
}
})
}
fn expand_enum(ident: Ident, generics: Generics, data: DataEnum) -> syn::Result<TokenStream2> {
let mut variant_infos = Vec::with_capacity(data.variants.len());
let mut errors: Option<syn::Error> = None;
for variant in data.variants {
let style = match &variant.fields {
Fields::Named(_) => FieldsStyle::Named,
Fields::Unnamed(_) => FieldsStyle::Unnamed,
Fields::Unit => {
errors = combine_error(
errors,
syn::Error::new(
variant.ident.span(),
"unit variants do not support #[derive(StackError)]",
),
);
continue;
}
};
let fields = match collect_fields(&variant.fields) {
Ok(fields) => fields,
Err(err) => {
errors = combine_error(errors, err);
continue;
}
};
let allow_name = style.allows_names();
let has_explicit_location = fields.iter().any(|f| {
f.attrs.is_location || (allow_name && matches!(&f.ident, Some(id) if id == "location"))
});
let location_index = if has_explicit_location {
match resolve_location(&fields, allow_name, variant.ident.span()) {
Ok(index) => Some(index),
Err(err) => {
errors = combine_error(errors, err);
None
}
}
} else {
resolve_location_from_located_source(&fields, allow_name)
.or_else(|| single_field_located(&fields))
};
let Some(location_index) = location_index else {
errors = combine_error(
errors,
syn::Error::new(
variant.ident.span(),
"missing #[location] attribute or field named `location`",
),
);
continue;
};
let source = match resolve_source(&fields, style.allows_names()) {
Ok(source) => source,
Err(err) => {
errors = combine_error(errors, err);
continue;
}
};
let source_binding = source
.as_ref()
.map(|_| format_ident!("__stack_error_source"));
variant_infos.push(VariantInfo {
ident: variant.ident,
style,
fields,
location_index,
source,
location_binding: format_ident!("__stack_error_location"),
source_binding,
});
}
if let Some(err) = errors {
return Err(err);
}
let mut generics = generics;
let mut bounds = BoundsTracker::new(&generics);
for variant in &variant_infos {
if let Some(source) = &variant.source {
bounds.collect(&variant.fields[source.index].ty, source.is_terminal);
}
}
bounds.apply(&mut generics);
let location_arms = variant_infos.iter().map(|variant| {
let variant_ident = &variant.ident;
let pattern = variant.location_pattern();
let value = variant.location_value_expr();
quote! {
Self::#variant_ident #pattern => #value
}
});
let next_arms = variant_infos.iter().map(|variant| {
let variant_ident = &variant.ident;
let pattern = variant.source_pattern();
let body = variant.next_body();
quote! {
Self::#variant_ident #pattern => #body
}
});
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
Ok(quote! {
impl #impl_generics ::pseudo_backtrace::StackError for #ident #ty_generics #where_clause {
fn location(&self) -> &'static ::core::panic::Location<'static> {
match self {
#(#location_arms,)*
}
}
fn next<'pseudo_backtrace>(&'pseudo_backtrace self) -> ::core::option::Option<::pseudo_backtrace::Chain<'pseudo_backtrace>> {
use ::pseudo_backtrace::private::AsDynStdError as _;
use ::pseudo_backtrace::private::AsDynStackError as _;
match self {
#(#next_arms,)*
}
}
}
})
}
#[derive(Clone)]
struct FieldInfo {
member: Member,
ident: Option<Ident>,
ty: syn::Type,
attrs: FieldAttrs,
span: Span,
}
#[derive(Clone, Copy)]
enum FieldsStyle {
Named,
Unnamed,
}
impl FieldsStyle {
fn allows_names(self) -> bool {
matches!(self, FieldsStyle::Named)
}
}
#[derive(Default, Clone)]
struct FieldAttrs {
is_source: bool,
is_location: bool,
is_terminal: bool,
}
struct SourceInfo {
index: usize,
is_terminal: bool,
}
struct VariantInfo {
ident: Ident,
style: FieldsStyle,
fields: Vec<FieldInfo>,
location_index: usize,
source: Option<SourceInfo>,
location_binding: Ident,
source_binding: Option<Ident>,
}
fn collect_fields(fields: &Fields) -> syn::Result<Vec<FieldInfo>> {
let mut out = Vec::new();
match fields {
Fields::Named(named) => {
for field in named.named.iter() {
out.push(build_field_info(field, out.len(), true)?);
}
}
Fields::Unnamed(unnamed) => {
for (idx, field) in unnamed.unnamed.iter().enumerate() {
out.push(build_field_info(field, idx, false)?);
}
}
Fields::Unit => {}
}
Ok(out)
}
fn build_field_info(field: &Field, index: usize, named: bool) -> syn::Result<FieldInfo> {
let attrs = parse_field_attrs(field)?;
let member = if named {
Member::Named(field.ident.clone().expect("named field missing ident"))
} else {
Member::Unnamed(syn::Index::from(index))
};
Ok(FieldInfo {
member,
ident: field.ident.clone(),
ty: field.ty.clone(),
attrs,
span: field.span(),
})
}
fn parse_field_attrs(field: &Field) -> syn::Result<FieldAttrs> {
let mut attrs = FieldAttrs::default();
for attr in &field.attrs {
if attr.path().is_ident("source") {
if attrs.is_source {
return Err(syn::Error::new_spanned(
attr,
"duplicate #[source] attribute",
));
}
attrs.is_source = true;
continue;
}
if attr.path().is_ident("location") {
if attrs.is_location {
return Err(syn::Error::new_spanned(
attr,
"duplicate #[location] attribute",
));
}
attrs.is_location = true;
continue;
}
if attr.path().is_ident("stack_error") {
match attr.parse_args_with(|input: syn::parse::ParseStream| {
let ident: Ident = input.parse()?;
if ident == "std" {
Ok((true, ident.span()))
} else if ident == "stacked" {
Ok((false, ident.span()))
} else {
Err(syn::Error::new(ident.span(), "expected `std` or `stacked`"))
}
}) {
Ok((is_std, _span)) => {
if is_std {
if attrs.is_terminal {
return Err(syn::Error::new_spanned(
attr,
"duplicate #[stack_error(std)] attribute",
));
}
attrs.is_terminal = true;
} else {
if attrs.is_source {
return Err(syn::Error::new_spanned(
attr,
"duplicate #[stack_error(stacked)] attribute",
));
}
attrs.is_source = true;
}
}
Err(err) => {
return Err(syn::Error::new_spanned(
attr,
format!("invalid #[stack_error] attribute: {}", err),
));
}
}
continue;
}
}
Ok(attrs)
}
fn resolve_location(
fields: &[FieldInfo],
allow_name: bool,
missing_span: Span,
) -> syn::Result<usize> {
let mut index = None;
for (idx, field) in fields.iter().enumerate() {
if field.attrs.is_location {
if index.is_some() {
return Err(syn::Error::new(
field.span,
"multiple fields marked with #[location]",
));
}
index = Some(idx);
}
}
if let Some(idx) = index {
return Ok(idx);
}
if allow_name
&& let Some((idx, _)) = fields
.iter()
.enumerate()
.find(|(_, field)| matches!(&field.ident, Some(ident) if ident == "location"))
{
return Ok(idx);
}
Err(syn::Error::new(
missing_span,
"missing #[location] attribute or field named `location`",
))
}
fn resolve_source(fields: &[FieldInfo], allow_name: bool) -> syn::Result<Option<SourceInfo>> {
let mut source_candidates: Vec<usize> = Vec::new();
let mut terminal_candidates: Vec<usize> = Vec::new();
for (idx, field) in fields.iter().enumerate() {
if field.attrs.is_source {
source_candidates.push(idx);
}
if field.attrs.is_terminal {
terminal_candidates.push(idx);
}
}
if source_candidates.len() > 1 {
let span = fields[source_candidates[1]].span;
return Err(syn::Error::new(
span,
"multiple fields marked with #[source]",
));
}
if source_candidates.len() == 1 {
let idx = source_candidates[0];
let is_terminal = fields[idx].attrs.is_terminal;
return Ok(Some(SourceInfo {
index: idx,
is_terminal,
}));
}
if terminal_candidates.len() > 1 {
let span = fields[terminal_candidates[1]].span;
return Err(syn::Error::new(
span,
"multiple fields marked with #[stack_error(std)]",
));
}
if let Some(idx) = terminal_candidates.first().copied() {
return Ok(Some(SourceInfo {
index: idx,
is_terminal: true,
}));
}
if allow_name
&& let Some((idx, _)) = fields
.iter()
.enumerate()
.find(|(_, field)| matches!(&field.ident, Some(ident) if ident == "source"))
{
return Ok(Some(SourceInfo {
index: idx,
is_terminal: false,
}));
}
Ok(None)
}
fn build_next_struct(member: &Member, ty: &syn::Type, is_terminal: bool) -> TokenStream2 {
if is_terminal {
if type_parameter_of_option(ty).is_some() {
quote! {
self.#member
.as_ref()
.map(|__s| ::pseudo_backtrace::Chain::Std(__s.as_dyn_std_error()))
}
} else {
quote! {
::core::option::Option::Some(::pseudo_backtrace::Chain::Std(
self.#member.as_dyn_std_error(),
))
}
}
} else if type_parameter_of_option(ty).is_some() {
quote! {
self.#member
.as_ref()
.map(|__s| ::pseudo_backtrace::Chain::Stacked(__s.as_dyn_stack_error()))
}
} else {
quote! {
::core::option::Option::Some(::pseudo_backtrace::Chain::Stacked(
self.#member.as_dyn_stack_error(),
))
}
}
}
impl VariantInfo {
fn location_pattern(&self) -> TokenStream2 {
match self.style {
FieldsStyle::Named => {
let field_ident = self.fields[self.location_index]
.ident
.as_ref()
.expect("named field missing ident")
.clone();
let binding = &self.location_binding;
quote! { { #field_ident: #binding, .. } }
}
FieldsStyle::Unnamed => {
let binding = &self.location_binding;
let patterns = self.fields.iter().enumerate().map(|(idx, _)| {
if idx == self.location_index {
quote! { #binding }
} else {
quote! { _ }
}
});
quote! { ( #(#patterns),* ) }
}
}
}
fn location_value_expr(&self) -> TokenStream2 {
let binding = &self.location_binding;
let ty = &self.fields[self.location_index].ty;
if is_located_error(ty) {
quote! { ::pseudo_backtrace::StackError::location(#binding) }
} else {
quote! { #binding }
}
}
fn source_pattern(&self) -> TokenStream2 {
match &self.source {
Some(source) => match self.style {
FieldsStyle::Named => {
let field_ident = self.fields[source.index]
.ident
.as_ref()
.expect("named field missing ident")
.clone();
let binding = self
.source_binding
.as_ref()
.expect("source binding missing");
quote! { { #field_ident: #binding, .. } }
}
FieldsStyle::Unnamed => {
let binding = self
.source_binding
.as_ref()
.expect("source binding missing");
let patterns = self.fields.iter().enumerate().map(|(idx, _)| {
if idx == source.index {
quote! { #binding }
} else {
quote! { _ }
}
});
quote! { ( #(#patterns),* ) }
}
},
None => match self.style {
FieldsStyle::Named => quote! { { .. } },
FieldsStyle::Unnamed => {
let patterns = self.fields.iter().map(|_| quote! { _ });
quote! { ( #(#patterns),* ) }
}
},
}
}
fn next_body(&self) -> TokenStream2 {
match &self.source {
Some(source) => {
let binding = self
.source_binding
.as_ref()
.expect("source binding missing");
let ty = &self.fields[source.index].ty;
if source.is_terminal {
if type_parameter_of_option(ty).is_some() {
quote! {
#binding
.as_ref()
.map(|__s| ::pseudo_backtrace::Chain::Std(__s.as_dyn_std_error()))
}
} else {
quote! {
::core::option::Option::Some(::pseudo_backtrace::Chain::Std(
#binding.as_dyn_std_error(),
))
}
}
} else if type_parameter_of_option(ty).is_some() {
quote! {
#binding
.as_ref()
.map(|__s| ::pseudo_backtrace::Chain::Stacked(__s.as_dyn_stack_error()))
}
} else {
quote! {
::core::option::Option::Some(::pseudo_backtrace::Chain::Stacked(
#binding.as_dyn_stack_error(),
))
}
}
}
None => quote! { ::core::option::Option::None },
}
}
}
fn type_parameter_of_option(ty: &syn::Type) -> Option<&syn::Type> {
let path = match ty {
syn::Type::Path(ty) => &ty.path,
_ => return None,
};
let last = path.segments.last()?;
if last.ident != "Option" {
return None;
}
let args = match &last.arguments {
syn::PathArguments::AngleBracketed(args) => args,
_ => return None,
};
if args.args.len() != 1 {
return None;
}
match &args.args[0] {
syn::GenericArgument::Type(inner) => Some(inner),
_ => None,
}
}
fn is_located_error(ty: &syn::Type) -> bool {
let ty = match ty {
syn::Type::Reference(r) => &*r.elem,
_ => ty,
};
let syn::Type::Path(type_path) = ty else {
return false;
};
let Some(last) = type_path.path.segments.last() else {
return false;
};
last.ident == "LocatedError"
}
fn resolve_location_from_located_source(fields: &[FieldInfo], allow_name: bool) -> Option<usize> {
let mut candidates: Vec<usize> = Vec::new();
for (idx, field) in fields.iter().enumerate() {
if field.attrs.is_source || field.attrs.is_terminal {
candidates.push(idx);
continue;
}
if allow_name
&& let Some(ident) = &field.ident
&& ident == "source"
{
candidates.push(idx);
}
}
let picked = candidates
.into_iter()
.find(|&idx| is_located_error(&fields[idx].ty));
if picked.is_some() {
return picked;
}
None
}
fn single_field_located(fields: &[FieldInfo]) -> Option<usize> {
if fields.len() == 1 && is_located_error(&fields[0].ty) {
Some(0)
} else {
None
}
}
struct BoundsTracker {
params: BTreeMap<String, Ident>,
needs_error: BTreeSet<String>,
needs_stack: BTreeSet<String>,
}
impl BoundsTracker {
fn new(generics: &Generics) -> Self {
let params = generics
.type_params()
.map(|param| (param.ident.to_string(), param.ident.clone()))
.collect();
BoundsTracker {
params,
needs_error: BTreeSet::new(),
needs_stack: BTreeSet::new(),
}
}
fn collect(&mut self, ty: &syn::Type, is_terminal: bool) {
let mut visitor = TypeParamCollector {
params: &self.params,
found: BTreeSet::new(),
};
visitor.visit_type(ty);
for name in visitor.found {
self.needs_error.insert(name.clone());
if !is_terminal {
self.needs_stack.insert(name);
}
}
}
fn apply(&self, generics: &mut Generics) {
for param in generics.type_params_mut() {
let name = param.ident.to_string();
if self.needs_stack.contains(&name) {
param
.bounds
.push(syn::parse_quote!(::pseudo_backtrace::StackError));
}
if self.needs_error.contains(&name) {
param.bounds.push(syn::parse_quote!(::core::error::Error));
}
}
}
}
struct TypeParamCollector<'a> {
params: &'a BTreeMap<String, Ident>,
found: BTreeSet<String>,
}
impl<'a, 'ast> Visit<'ast> for TypeParamCollector<'a> {
fn visit_type_path(&mut self, type_path: &'ast syn::TypePath) {
if type_path.qself.is_none()
&& let Some(segment) = type_path.path.segments.first()
{
let ident = &segment.ident;
let name = ident.to_string();
if self.params.contains_key(&name) {
self.found.insert(name);
}
}
syn::visit::visit_type_path(self, type_path);
}
}
fn combine_error(acc: Option<syn::Error>, next: syn::Error) -> Option<syn::Error> {
match acc {
Some(mut err) => {
err.combine(next);
Some(err)
}
None => Some(next),
}
}