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 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}