Skip to main content

simple_enum/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{parse_macro_input, Data, DeriveInput, Expr, ExprLit, Lit, Meta, MetaNameValue};
4
5#[proc_macro_derive(SimpleEnum, attributes(simple_enum))]
6pub fn derive_simple_enum(input: TokenStream) -> TokenStream {
7    let input = parse_macro_input!(input as DeriveInput);
8    let enum_name = input.ident;
9
10    let mut to_str_match_arms = Vec::new();
11    let mut get_code_match_arms = Vec::new();
12    let mut from_code_match_arms = Vec::new();
13    let mut desc_match_arms = Vec::new();
14    let mut code_type = None; // 用来推断code的类型
15
16    if let Data::Enum(data_enum) = input.data {
17        for variant in data_enum.variants {
18            let var_ident = variant.ident;
19
20            let mut code_expr = None;
21            let mut desc_val = None;
22
23            // 读取 #[simple_enum(code=X, desc="Y")] - code可以是任意类型
24            for attr in variant.attrs {
25                if attr.path().is_ident("simple_enum") {
26                    if let Meta::List(_) = &attr.meta {
27                        // 使用 parse_args 来解析属性参数
28                        let parsed_args: Result<
29                            syn::punctuated::Punctuated<Meta, syn::Token![,]>,
30                            _,
31                        > = attr.parse_args_with(syn::punctuated::Punctuated::parse_terminated);
32
33                        if let Ok(args) = parsed_args {
34                            for meta in args {
35                                if let Meta::NameValue(MetaNameValue { path, value, .. }) = meta {
36                                    if path.is_ident("code") {
37                                        // 保存整个表达式,并推断类型
38                                        code_expr = Some(value.clone());
39
40                                        // 推断code的类型
41                                        if code_type.is_none() {
42                                            match &value {
43                                                Expr::Lit(ExprLit {
44                                                    lit: Lit::Str(_), ..
45                                                }) => {
46                                                    code_type = Some(quote! { &str });
47                                                }
48                                                Expr::Lit(ExprLit {
49                                                    lit: Lit::Int(_), ..
50                                                }) => {
51                                                    code_type = Some(quote! { i32 });
52                                                }
53                                                _ => {
54                                                    code_type = Some(quote! { i32 });
55                                                    // 默认类型
56                                                }
57                                            }
58                                        }
59                                    } else if path.is_ident("desc") {
60                                        if let Expr::Lit(ExprLit {
61                                            lit: Lit::Str(lit_str),
62                                            ..
63                                        }) = value
64                                        {
65                                            desc_val = Some(lit_str.value());
66                                        }
67                                    }
68                                }
69                            }
70                        }
71                    }
72                }
73            }
74
75            let code_expr = code_expr.expect("simple_enum must have code");
76            let desc = desc_val.expect("simple_enum must have desc");
77            let variant_name = var_ident.to_string();
78
79            to_str_match_arms.push(quote! {
80                #enum_name::#var_ident => #variant_name,
81            });
82            get_code_match_arms.push(quote! {
83                #enum_name::#var_ident => #code_expr,
84            });
85            from_code_match_arms.push(quote! {
86                #code_expr => Some(#enum_name::#var_ident),
87            });
88            desc_match_arms.push(quote! {
89                #enum_name::#var_ident => #desc,
90            });
91        }
92    } else {
93        panic!("SimpleEnum only works on enums");
94    }
95
96    let code_type = code_type.expect("Could not determine code type");
97
98    let expanded = quote! {
99        impl #enum_name {
100            pub fn to_str(&self) -> &'static str {
101                match self {
102                    #(#to_str_match_arms)*
103                }
104            }
105
106            pub fn get_code(&self) -> #code_type {
107                match self {
108                    #(#get_code_match_arms)*
109                }
110            }
111
112            pub fn from_code(c: #code_type) -> Option<Self> {
113                match c {
114                    #(#from_code_match_arms)*
115                    _ => None,
116                }
117            }
118
119            pub fn desc(&self) -> &'static str {
120                match self {
121                    #(#desc_match_arms)*
122                }
123            }
124        }
125    };
126
127    TokenStream::from(expanded)
128}