pseudo_backtrace_derive/
lib.rs

1use proc_macro::TokenStream;
2use proc_macro2::{Span, TokenStream as TokenStream2};
3use quote::{format_ident, quote};
4use std::collections::{BTreeMap, BTreeSet};
5use syn::visit::Visit;
6use syn::{
7    Data, DataEnum, DataStruct, DeriveInput, Field, Fields, Generics, Ident, Member,
8    parse_macro_input, spanned::Spanned,
9};
10
11#[proc_macro_derive(StackError, attributes(source, stack_error, location))]
12pub fn derive_stack_error(input: TokenStream) -> TokenStream {
13    let input = parse_macro_input!(input as DeriveInput);
14    match expand(input) {
15        Ok(tokens) => tokens.into(),
16        Err(error) => error.into_compile_error().into(),
17    }
18}
19
20fn expand(input: DeriveInput) -> syn::Result<TokenStream2> {
21    let ident = input.ident;
22    let generics = input.generics;
23
24    match input.data {
25        Data::Struct(data) => expand_struct(ident, generics, data),
26        Data::Enum(data) => expand_enum(ident, generics, data),
27        Data::Union(_) => Err(syn::Error::new(
28            Span::call_site(),
29            "StackError cannot be derived for unions",
30        )),
31    }
32}
33
34fn expand_struct(ident: Ident, generics: Generics, data: DataStruct) -> syn::Result<TokenStream2> {
35    let style = match &data.fields {
36        Fields::Named(_) => FieldsStyle::Named,
37        Fields::Unnamed(_) => FieldsStyle::Unnamed,
38        Fields::Unit => {
39            return Err(syn::Error::new(
40                ident.span(),
41                "unit structs do not support #[derive(StackError)]",
42            ));
43        }
44    };
45
46    let fields = collect_fields(&data.fields)?;
47    let location_index = resolve_location(&fields, style.allows_names(), ident.span())?;
48    let source = resolve_source(&fields, style.allows_names())?;
49
50    let mut generics = generics;
51    let mut bounds = BoundsTracker::new(&generics);
52    if let Some(info) = &source {
53        bounds.collect(&fields[info.index].ty, info.is_terminal);
54    }
55    bounds.apply(&mut generics);
56
57    let location_member = &fields[location_index].member;
58    let next_body = match &source {
59        Some(info) => build_next_struct(&fields[info.index].member, info.is_terminal),
60        None => quote! { ::core::option::Option::None },
61    };
62
63    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
64
65    Ok(quote! {
66        impl #impl_generics ::pseudo_backtrace::StackError for #ident #ty_generics #where_clause {
67            fn location(&self) -> &'static ::core::panic::Location<'static> {
68                self.#location_member
69            }
70
71            fn next<'a>(&'a self) -> ::core::option::Option<::pseudo_backtrace::ErrorDetail<'a>> {
72                #next_body
73            }
74        }
75    })
76}
77
78fn expand_enum(ident: Ident, generics: Generics, data: DataEnum) -> syn::Result<TokenStream2> {
79    let mut variant_infos = Vec::with_capacity(data.variants.len());
80    let mut errors: Option<syn::Error> = None;
81
82    for variant in data.variants {
83        let style = match &variant.fields {
84            Fields::Named(_) => FieldsStyle::Named,
85            Fields::Unnamed(_) => FieldsStyle::Unnamed,
86            Fields::Unit => {
87                errors = combine_error(
88                    errors,
89                    syn::Error::new(
90                        variant.ident.span(),
91                        "unit variants do not support #[derive(StackError)]",
92                    ),
93                );
94                continue;
95            }
96        };
97
98        let fields = match collect_fields(&variant.fields) {
99            Ok(fields) => fields,
100            Err(err) => {
101                errors = combine_error(errors, err);
102                continue;
103            }
104        };
105
106        let location_index =
107            match resolve_location(&fields, style.allows_names(), variant.ident.span()) {
108                Ok(index) => index,
109                Err(err) => {
110                    errors = combine_error(errors, err);
111                    continue;
112                }
113            };
114
115        let source = match resolve_source(&fields, style.allows_names()) {
116            Ok(source) => source,
117            Err(err) => {
118                errors = combine_error(errors, err);
119                continue;
120            }
121        };
122
123        let source_binding = source
124            .as_ref()
125            .map(|_| format_ident!("__stack_error_source"));
126
127        variant_infos.push(VariantInfo {
128            ident: variant.ident,
129            style,
130            fields,
131            location_index,
132            source,
133            location_binding: format_ident!("__stack_error_location"),
134            source_binding,
135        });
136    }
137
138    if let Some(err) = errors {
139        return Err(err);
140    }
141
142    let mut generics = generics;
143    let mut bounds = BoundsTracker::new(&generics);
144    for variant in &variant_infos {
145        if let Some(source) = &variant.source {
146            bounds.collect(&variant.fields[source.index].ty, source.is_terminal);
147        }
148    }
149    bounds.apply(&mut generics);
150
151    let location_arms = variant_infos.iter().map(|variant| {
152        let variant_ident = &variant.ident;
153        let pattern = variant.location_pattern();
154        let value = &variant.location_binding;
155        quote! {
156            Self::#variant_ident #pattern => #value
157        }
158    });
159
160    let next_arms = variant_infos.iter().map(|variant| {
161        let variant_ident = &variant.ident;
162        let pattern = variant.source_pattern();
163        let body = variant.next_body();
164        quote! {
165            Self::#variant_ident #pattern => #body
166        }
167    });
168
169    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
170
171    Ok(quote! {
172        impl #impl_generics ::pseudo_backtrace::StackError for #ident #ty_generics #where_clause {
173            fn location(&self) -> &'static ::core::panic::Location<'static> {
174                match self {
175                    #(#location_arms,)*
176                }
177            }
178
179            fn next<'a>(&'a self) -> ::core::option::Option<::pseudo_backtrace::ErrorDetail<'a>> {
180                match self {
181                    #(#next_arms,)*
182                }
183            }
184        }
185    })
186}
187
188#[derive(Clone)]
189struct FieldInfo {
190    member: Member,
191    ident: Option<Ident>,
192    ty: syn::Type,
193    attrs: FieldAttrs,
194    span: Span,
195}
196
197#[derive(Clone, Copy)]
198enum FieldsStyle {
199    Named,
200    Unnamed,
201}
202
203impl FieldsStyle {
204    fn allows_names(self) -> bool {
205        matches!(self, FieldsStyle::Named)
206    }
207}
208
209#[derive(Default, Clone)]
210struct FieldAttrs {
211    is_source: bool,
212    is_location: bool,
213    is_terminal: bool,
214}
215
216struct SourceInfo {
217    index: usize,
218    is_terminal: bool,
219}
220
221struct VariantInfo {
222    ident: Ident,
223    style: FieldsStyle,
224    fields: Vec<FieldInfo>,
225    location_index: usize,
226    source: Option<SourceInfo>,
227    location_binding: Ident,
228    source_binding: Option<Ident>,
229}
230
231fn collect_fields(fields: &Fields) -> syn::Result<Vec<FieldInfo>> {
232    let mut out = Vec::new();
233
234    match fields {
235        Fields::Named(named) => {
236            for field in named.named.iter() {
237                out.push(build_field_info(field, out.len(), true)?);
238            }
239        }
240        Fields::Unnamed(unnamed) => {
241            for (idx, field) in unnamed.unnamed.iter().enumerate() {
242                out.push(build_field_info(field, idx, false)?);
243            }
244        }
245        Fields::Unit => {}
246    }
247
248    Ok(out)
249}
250
251fn build_field_info(field: &Field, index: usize, named: bool) -> syn::Result<FieldInfo> {
252    let attrs = parse_field_attrs(field)?;
253    let member = if named {
254        Member::Named(field.ident.clone().expect("named field missing ident"))
255    } else {
256        Member::Unnamed(syn::Index::from(index))
257    };
258
259    Ok(FieldInfo {
260        member,
261        ident: field.ident.clone(),
262        ty: field.ty.clone(),
263        attrs,
264        span: field.span(),
265    })
266}
267
268fn parse_field_attrs(field: &Field) -> syn::Result<FieldAttrs> {
269    let mut attrs = FieldAttrs::default();
270
271    for attr in &field.attrs {
272        if attr.path().is_ident("source") {
273            if attrs.is_source {
274                return Err(syn::Error::new_spanned(
275                    attr,
276                    "duplicate #[source] attribute",
277                ));
278            }
279            attrs.is_source = true;
280            continue;
281        }
282
283        if attr.path().is_ident("location") {
284            if attrs.is_location {
285                return Err(syn::Error::new_spanned(
286                    attr,
287                    "duplicate #[location] attribute",
288                ));
289            }
290            attrs.is_location = true;
291            continue;
292        }
293
294        if attr.path().is_ident("stack_error") {
295            match attr.parse_args_with(|input: syn::parse::ParseStream| {
296                let ident: Ident = input.parse()?;
297                if ident == "end" {
298                    Ok(())
299                } else {
300                    Err(syn::Error::new(ident.span(), "expected `end`"))
301                }
302            }) {
303                Ok(()) => {
304                    if attrs.is_terminal {
305                        return Err(syn::Error::new_spanned(
306                            attr,
307                            "duplicate #[stack_error(end)] attribute",
308                        ));
309                    }
310                    attrs.is_terminal = true;
311                }
312                Err(err) => {
313                    return Err(syn::Error::new_spanned(
314                        attr,
315                        format!("invalid #[stack_error] attribute: {}", err),
316                    ));
317                }
318            }
319
320            continue;
321        }
322    }
323
324    Ok(attrs)
325}
326
327fn resolve_location(
328    fields: &[FieldInfo],
329    allow_name: bool,
330    missing_span: Span,
331) -> syn::Result<usize> {
332    let mut index = None;
333
334    for (idx, field) in fields.iter().enumerate() {
335        if field.attrs.is_location {
336            if index.is_some() {
337                return Err(syn::Error::new(
338                    field.span,
339                    "multiple fields marked with #[location]",
340                ));
341            }
342            index = Some(idx);
343        }
344    }
345
346    if let Some(idx) = index {
347        return Ok(idx);
348    }
349
350    if allow_name
351        && let Some((idx, _)) = fields
352            .iter()
353            .enumerate()
354            .find(|(_, field)| matches!(&field.ident, Some(ident) if ident == "location"))
355    {
356        return Ok(idx);
357    }
358
359    Err(syn::Error::new(
360        missing_span,
361        "missing #[location] attribute or field named `location`",
362    ))
363}
364
365fn resolve_source(fields: &[FieldInfo], allow_name: bool) -> syn::Result<Option<SourceInfo>> {
366    let mut source_candidates: Vec<usize> = Vec::new();
367    let mut terminal_candidates: Vec<usize> = Vec::new();
368
369    for (idx, field) in fields.iter().enumerate() {
370        if field.attrs.is_source {
371            source_candidates.push(idx);
372        }
373        if field.attrs.is_terminal {
374            terminal_candidates.push(idx);
375        }
376    }
377
378    if source_candidates.len() > 1 {
379        let span = fields[source_candidates[1]].span;
380        return Err(syn::Error::new(
381            span,
382            "multiple fields marked with #[source]",
383        ));
384    }
385
386    if source_candidates.len() == 1 {
387        let idx = source_candidates[0];
388        let is_terminal = fields[idx].attrs.is_terminal;
389        return Ok(Some(SourceInfo {
390            index: idx,
391            is_terminal,
392        }));
393    }
394
395    if terminal_candidates.len() > 1 {
396        let span = fields[terminal_candidates[1]].span;
397        return Err(syn::Error::new(
398            span,
399            "multiple fields marked with #[stack_error(end)]",
400        ));
401    }
402
403    if let Some(idx) = terminal_candidates.first().copied() {
404        return Ok(Some(SourceInfo {
405            index: idx,
406            is_terminal: true,
407        }));
408    }
409
410    if allow_name
411        && let Some((idx, _)) = fields
412            .iter()
413            .enumerate()
414            .find(|(_, field)| matches!(&field.ident, Some(ident) if ident == "source"))
415    {
416        return Ok(Some(SourceInfo {
417            index: idx,
418            is_terminal: false,
419        }));
420    }
421
422    Ok(None)
423}
424
425fn build_next_struct(member: &Member, is_terminal: bool) -> TokenStream2 {
426    if is_terminal {
427        quote! {
428            ::core::option::Option::Some(::pseudo_backtrace::ErrorDetail::End(
429                &self.#member as &'a dyn ::core::error::Error,
430            ))
431        }
432    } else {
433        quote! {
434            ::core::option::Option::Some(::pseudo_backtrace::ErrorDetail::Stacked(
435                &self.#member as &'a dyn ::pseudo_backtrace::StackError,
436            ))
437        }
438    }
439}
440
441impl VariantInfo {
442    fn location_pattern(&self) -> TokenStream2 {
443        match self.style {
444            FieldsStyle::Named => {
445                let field_ident = self.fields[self.location_index]
446                    .ident
447                    .as_ref()
448                    .expect("named field missing ident")
449                    .clone();
450                let binding = &self.location_binding;
451                quote! { { #field_ident: #binding, .. } }
452            }
453            FieldsStyle::Unnamed => {
454                let binding = &self.location_binding;
455                let patterns = self.fields.iter().enumerate().map(|(idx, _)| {
456                    if idx == self.location_index {
457                        quote! { #binding }
458                    } else {
459                        quote! { _ }
460                    }
461                });
462                quote! { ( #(#patterns),* ) }
463            }
464        }
465    }
466
467    fn source_pattern(&self) -> TokenStream2 {
468        match &self.source {
469            Some(source) => match self.style {
470                FieldsStyle::Named => {
471                    let field_ident = self.fields[source.index]
472                        .ident
473                        .as_ref()
474                        .expect("named field missing ident")
475                        .clone();
476                    let binding = self
477                        .source_binding
478                        .as_ref()
479                        .expect("source binding missing");
480                    quote! { { #field_ident: #binding, .. } }
481                }
482                FieldsStyle::Unnamed => {
483                    let binding = self
484                        .source_binding
485                        .as_ref()
486                        .expect("source binding missing");
487                    let patterns = self.fields.iter().enumerate().map(|(idx, _)| {
488                        if idx == source.index {
489                            quote! { #binding }
490                        } else {
491                            quote! { _ }
492                        }
493                    });
494                    quote! { ( #(#patterns),* ) }
495                }
496            },
497            None => match self.style {
498                FieldsStyle::Named => quote! { { .. } },
499                FieldsStyle::Unnamed => {
500                    let patterns = self.fields.iter().map(|_| quote! { _ });
501                    quote! { ( #(#patterns),* ) }
502                }
503            },
504        }
505    }
506
507    fn next_body(&self) -> TokenStream2 {
508        match &self.source {
509            Some(source) => {
510                let binding = self
511                    .source_binding
512                    .as_ref()
513                    .expect("source binding missing");
514                if source.is_terminal {
515                    quote! {
516                        ::core::option::Option::Some(::pseudo_backtrace::ErrorDetail::End(
517                            #binding as &'a dyn ::core::error::Error,
518                        ))
519                    }
520                } else {
521                    quote! {
522                        ::core::option::Option::Some(::pseudo_backtrace::ErrorDetail::Stacked(
523                            #binding as &'a dyn ::pseudo_backtrace::StackError,
524                        ))
525                    }
526                }
527            }
528            None => quote! { ::core::option::Option::None },
529        }
530    }
531}
532
533struct BoundsTracker {
534    params: BTreeMap<String, Ident>,
535    needs_error: BTreeSet<String>,
536    needs_stack: BTreeSet<String>,
537}
538
539impl BoundsTracker {
540    fn new(generics: &Generics) -> Self {
541        let params = generics
542            .type_params()
543            .map(|param| (param.ident.to_string(), param.ident.clone()))
544            .collect();
545
546        BoundsTracker {
547            params,
548            needs_error: BTreeSet::new(),
549            needs_stack: BTreeSet::new(),
550        }
551    }
552
553    fn collect(&mut self, ty: &syn::Type, is_terminal: bool) {
554        let mut visitor = TypeParamCollector {
555            params: &self.params,
556            found: BTreeSet::new(),
557        };
558        visitor.visit_type(ty);
559
560        for name in visitor.found {
561            self.needs_error.insert(name.clone());
562            if !is_terminal {
563                self.needs_stack.insert(name);
564            }
565        }
566    }
567
568    fn apply(&self, generics: &mut Generics) {
569        for param in generics.type_params_mut() {
570            let name = param.ident.to_string();
571            if self.needs_stack.contains(&name) {
572                param
573                    .bounds
574                    .push(syn::parse_quote!(::pseudo_backtrace::StackError));
575            }
576            if self.needs_error.contains(&name) {
577                param.bounds.push(syn::parse_quote!(::core::error::Error));
578            }
579        }
580    }
581}
582
583struct TypeParamCollector<'a> {
584    params: &'a BTreeMap<String, Ident>,
585    found: BTreeSet<String>,
586}
587
588impl<'a, 'ast> Visit<'ast> for TypeParamCollector<'a> {
589    fn visit_type_path(&mut self, type_path: &'ast syn::TypePath) {
590        if type_path.qself.is_none()
591            && let Some(segment) = type_path.path.segments.first()
592        {
593            let ident = &segment.ident;
594            let name = ident.to_string();
595            if self.params.contains_key(&name) {
596                self.found.insert(name);
597            }
598        }
599
600        syn::visit::visit_type_path(self, type_path);
601    }
602}
603
604fn combine_error(acc: Option<syn::Error>, next: syn::Error) -> Option<syn::Error> {
605    match acc {
606        Some(mut err) => {
607            err.combine(next);
608            Some(err)
609        }
610        None => Some(next),
611    }
612}