Skip to main content

pn_dcp_macro/
lib.rs

1use proc_macro::TokenStream;
2use proc_macro2::{Ident, Literal};
3use quote::{quote, ToTokens, TokenStreamExt};
4use syn::__private::TokenStream2;
5use syn::parse::Parser;
6use syn::{ItemStruct, Lit, Meta, NestedMeta, Type};
7
8fn impl_derefmut(attr_ty: AttrTy, name: Ident, item: ItemStruct) -> TokenStream {
9    let gen = quote! {
10        #item
11        impl DerefMut for #name {
12            fn deref_mut(&mut self) -> &mut Self::Target {
13                &mut self.#attr_ty
14            }
15        }
16    };
17    gen.into()
18}
19
20fn impl_deref(attr_ty: AttrTy, ty: Type, name: Ident, item: ItemStruct) -> TokenStream {
21    let gen = quote! {
22        #item
23        impl Deref for #name {
24            type Target = #ty;
25
26            fn deref(&self) -> &Self::Target {
27                &self.#attr_ty
28            }
29        }
30        impl DerefMut for #name {
31            fn deref_mut(&mut self) -> &mut Self::Target {
32                &mut self.#attr_ty
33            }
34        }
35    };
36    gen.into()
37}
38
39type AttrAlisa = syn::punctuated::Punctuated<syn::NestedMeta, syn::Token![,]>;
40#[proc_macro_attribute]
41pub fn derefmut(attr: TokenStream, item: TokenStream) -> TokenStream {
42    let attr_ty = resolve_attr(attr.into());
43    let item: ItemStruct = syn::parse(item).unwrap();
44    let ty = get_field_ty(&item, attr_ty.clone());
45    match &attr_ty {
46        AttrTy::Ident(_) => {
47            impl_derefmut(attr_ty, item.ident.clone(), item)
48        }
49        AttrTy::Index(_) => {
50            impl_deref(attr_ty, ty, item.ident.clone(), item)
51        }
52    }
53
54}
55
56fn get_field_ty(a: &ItemStruct, right: AttrTy) -> Type {
57    match right {
58        AttrTy::Index(index_right) => {
59            for (index, field) in a.fields.iter().enumerate() {
60                if index_right == index {
61                    return field.ty.clone();
62                }
63            }
64            panic!("字段索引越界!")
65        }
66        AttrTy::Ident(ref ident) => {
67            for field in a.fields.iter() {
68                if let Some(c) = field.ident.as_ref() {
69                    if c == ident {
70                        return field.ty.clone();
71                    }
72                }
73            }
74            panic!("无法找到字段!")
75        }
76    }
77}
78#[derive(Clone)]
79enum AttrTy {
80    Index(usize),
81    Ident(Ident),
82}
83
84impl ToTokens for AttrTy {
85    fn to_tokens(&self, tokens: &mut proc_macro2::TokenStream) {
86        match self {
87            AttrTy::Ident(ident) => ident.to_tokens(tokens),
88            AttrTy::Index(index) => {
89                let lit = Literal::usize_unsuffixed(*index);
90                tokens.append(lit)
91            }
92        }
93    }
94}
95
96fn resolve_attr(attr: TokenStream2) -> AttrTy {
97    let attr_vals = AttrAlisa::parse_terminated.parse2(attr).unwrap();
98    // println!("{:?}", attr_vals);
99    for attr_val in attr_vals.iter() {
100        match attr_val {
101            NestedMeta::Meta(Meta::Path(meta)) => {
102                return AttrTy::Ident(meta.get_ident().unwrap().clone());
103            }
104            NestedMeta::Lit(Lit::Int(lit)) => {
105                return AttrTy::Index(lit.base10_parse::<usize>().unwrap());
106            }
107            _ => {
108                panic!("无法解析属性值:只能为字段索引或者字段名称")
109            }
110        }
111    }
112    panic!("必须设置属性值(字段索引或者字段名称)")
113}