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}