Skip to main content

enum_status_code_macros/
lib.rs

1extern crate proc_macro;
2use itertools::Itertools;
3use manyhow::manyhow;
4use proc_macro2::TokenStream;
5use quote::quote;
6use syn::{Data::Enum, DeriveInput, Error, Expr, parse2, spanned::Spanned};
7
8#[manyhow(proc_macro_derive(StatusCode))]
9pub fn derive_status_code(input: TokenStream) -> syn::Result<TokenStream> {
10    let input: DeriveInput = parse2(input)?;
11    let enum_name = &input.ident;
12    let enum_span = input.span();
13    let default = input
14        .attrs
15        .iter()
16        .find_map(|attr| {
17            if attr.path().is_ident("derive") {
18                match attr.parse_args::<Expr>() {
19                    Ok(attr) => match attr {
20                        Expr::Path(path) => path.path.is_ident("StatusCode").then_some(Ok(None)),
21                        Expr::Call(call) => {
22                            if let Expr::Path(path) = *call.func {
23                                path.path.is_ident("StatusCode").then_some(Ok(call
24                                    .args
25                                    .first()
26                                    .cloned()
27                                    .map(Ok)))
28                            } else {
29                                None
30                            }
31                        }
32                        _ => Some(Err(Error::new(attr.span(), "invalid attribute arguments"))),
33                    },
34                    Err(error) => Some(Err(error)),
35                }
36            } else {
37                None
38            }
39        })
40        .ok_or(Error::new(enum_span, "missing derive macro"))??;
41
42    if let Enum(data) = input.data {
43        data.variants
44            .iter()
45            .map(|variant| {
46                if let Some(value) = variant
47                    .attrs
48                    .iter()
49                    .find_map(|attr| {
50                        attr.path()
51                            .is_ident("status_code")
52                            .then_some(attr.parse_args::<Expr>())
53                    })
54                    .or(default.clone())
55                {
56                    let value = value?;
57                    Ok(quote! {
58                        #enum_name::#variant.ident => #value
59                    })
60                } else {
61                    Err(Error::new(variant.span(), "variant is missing status_code"))
62                }
63            })
64            .process_results(|iter| {
65                iter.tree_reduce(|a, b| {
66                    quote! {
67                        #a
68                        #b
69                    }
70                })
71            })?
72            .map(|contents| {
73                quote! {
74                    impl enum_status_code::StatusCode for #enum_name {
75                        fn status_code(&self) -> enum_status_code::http::StatusCode {
76                            match self {
77                                #contents
78                            }
79                        }
80                    }
81                }
82            })
83            .ok_or(Error::new(enum_span, "expcted enum body"))
84    } else {
85        Err(Error::new(enum_span, "expected enum"))
86    }
87}