Skip to main content

this_error_from_box/
lib.rs

1//! Procedural macro for automatic From implementation for #[from] Box<T>
2
3use proc_macro::TokenStream;
4use quote::quote;
5use syn::{
6    Data, DeriveInput, Fields,
7    parse::{Parse, ParseStream},
8    parse_macro_input,
9};
10
11struct WrapperArg {
12    wrapper: Option<syn::Path>,
13}
14
15impl Parse for WrapperArg {
16    fn parse(input: ParseStream) -> syn::Result<Self> {
17        if input.is_empty() {
18            Ok(WrapperArg { wrapper: None })
19        } else {
20            let wrapper: syn::Path = input.parse()?;
21            Ok(WrapperArg {
22                wrapper: Some(wrapper),
23            })
24        }
25    }
26}
27
28#[proc_macro_attribute]
29pub fn this_error_from_box(attr: TokenStream, item: TokenStream) -> TokenStream {
30    let WrapperArg { wrapper } = parse_macro_input!(attr as WrapperArg);
31    let wrapper_ident = wrapper.unwrap_or_else(|| syn::parse_str("Box").unwrap());
32    let input = parse_macro_input!(item as DeriveInput);
33    let enum_name = &input.ident;
34    let mut from_impls = Vec::new();
35
36    if let Data::Enum(data_enum) = &input.data {
37        for variant in &data_enum.variants {
38            let Fields::Unnamed(fields) = &variant.fields else {
39                continue;
40            };
41            if fields.unnamed.len() != 1 {
42                continue;
43            }
44            let field = &fields.unnamed[0];
45            let has_from = field.attrs.iter().any(|attr| attr.path().is_ident("from"));
46            if !has_from {
47                continue;
48            }
49            let syn::Type::Path(type_path) = &field.ty else {
50                continue;
51            };
52
53            let Some(last_segment) = type_path.path.segments.last() else {
54                continue;
55            };
56
57            if type_path.path.leading_colon.is_some() ^ wrapper_ident.leading_colon.is_some() {
58                continue;
59            }
60
61            if type_path.path.segments.len() != wrapper_ident.segments.len() {
62                continue;
63            }
64
65            let paths_equal = type_path
66                .path
67                .segments
68                .iter()
69                .zip(wrapper_ident.segments.iter())
70                .all(|(a, b)| a.ident == b.ident);
71
72            if !paths_equal {
73                continue;
74            }
75
76            let syn::PathArguments::AngleBracketed(args) = &last_segment.arguments else {
77                continue;
78            };
79            if args.args.len() != 1 {
80                continue;
81            }
82            let syn::GenericArgument::Type(inner_ty) = &args.args[0] else {
83                continue;
84            };
85            let variant_ident = &variant.ident;
86            from_impls.push(quote! {
87                impl ::std::convert::From<#inner_ty> for #enum_name {
88                    fn from(e: #inner_ty) -> Self {
89                        #enum_name::#variant_ident(#wrapper_ident::from(e))
90                    }
91                }
92            });
93        }
94    }
95    let expanded = quote! {
96        #input
97        #(#from_impls)*
98    };
99    TokenStream::from(expanded)
100}