xasync_macro/
lib.rs

1extern crate proc_macro;
2use proc_macro::TokenStream;
3use quote::quote;
4use syn::{Data, DeriveInput, Fields, parse_macro_input};
5
6#[proc_macro_derive(ActorMessage)]
7pub fn actor_message(input: TokenStream) -> TokenStream {
8    let input = parse_macro_input!(input as DeriveInput);
9    let enum_name = &input.ident;
10    let enum_vis = &input.vis;
11    let caller_name = syn::Ident::new(&format!("{}Caller", enum_name), enum_name.span());
12    let replyer_name = syn::Ident::new(&format!("{}Replyer", enum_name), enum_name.span());
13
14    let variants = match &input.data {
15        Data::Enum(data) => &data.variants,
16        _ => {
17            return syn::Error::new_spanned(&input, "ActorMessage can only be derived for enums")
18                .to_compile_error()
19                .into();
20        },
21    };
22
23    let mut caller_methods = Vec::new();
24
25    for variant in variants {
26        let variant_name = &variant.ident;
27        let snake_case_name = to_snake_case(&variant_name.to_string());
28        let method_name = syn::Ident::new(&snake_case_name, variant_name.span());
29        let timeout_method_name = syn::Ident::new(&format!("{}_timeout", snake_case_name), variant_name.span());
30        let method_name_str = snake_case_name.clone();
31
32        match &variant.fields {
33            Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
34                if let syn::Type::Path(tp) = &fields.unnamed[0].ty {
35                    if let Some(segment) = tp.path.segments.last() {
36                        if segment.ident == "ActorPacket" {
37                            if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
38                                if args.args.len() == 2 {
39                                    let args_type = &args.args[0];
40                                    let ret_type = &args.args[1];
41
42                                    caller_methods.push(quote! {
43                                        pub async fn #method_name(&self, args: #args_type) -> anyhow::Result<#ret_type> {
44                                            let (packet, rx) = ActorPacket::new(args);
45                                            self.0.send(#enum_name::#variant_name(packet))
46                                                .map_err(|e| anyhow::anyhow!("Fail to `{}`, error: {}", #method_name_str, e))?;
47                                            rx.wait().await
48                                                .map_err(|e| anyhow::anyhow!("Fail to `{}`, error: {}", #method_name_str, e))
49                                        }
50                                    });
51
52                                    caller_methods.push(quote! {
53                                        pub async fn #timeout_method_name(&self, args: #args_type, duration: ::std::time::Duration) -> anyhow::Result<#ret_type> {
54                                            let (packet, rx) = ActorPacket::new(args);
55                                            self.0.send(#enum_name::#variant_name(packet))
56                                                .map_err(|e| anyhow::anyhow!("Fail to `{}`, error: {}", #method_name_str, e))?;
57
58                                            rx.wait_timeout(duration).await
59                                                .map_err(|e| anyhow::anyhow!("Fail to `{}`, error: {}", #method_name_str, e))
60                                        }
61                                    });
62                                }
63                            }
64                        }
65                    }
66                }
67            },
68            Fields::Named(fields) => {
69                let mut other_fields = Vec::new();
70
71                for field in &fields.named {
72                    other_fields.push(field);
73                }
74            },
75            _ => {},
76        }
77    }
78
79    let expanded = quote! {
80        #[derive(Clone, Debug)]
81        #enum_vis struct #caller_name(::tokio::sync::mpsc::UnboundedSender<#enum_name>);
82        impl #caller_name {
83            #(#caller_methods)*
84        }
85
86        #enum_vis struct #replyer_name(::tokio::sync::mpsc::UnboundedReceiver<#enum_name>);
87        impl #replyer_name {
88            pub async fn recv(&mut self) -> Option<#enum_name> {
89                self.0.recv().await
90            }
91        }
92
93        impl std::ops::Deref for #replyer_name {
94            type Target = ::tokio::sync::mpsc::UnboundedReceiver<#enum_name>;
95
96            fn deref(&self) -> &Self::Target {
97                &self.0
98            }
99        }
100
101        impl std::ops::DerefMut for #replyer_name {
102            fn deref_mut(&mut self) -> &mut Self::Target {
103                &mut self.0
104            }
105        }
106
107        impl #enum_name {
108            pub fn actor() -> (#caller_name, #replyer_name) {
109                let (tx, rx) = ::tokio::sync::mpsc::unbounded_channel();
110                (#caller_name(tx), #replyer_name(rx))
111            }
112        }
113    };
114
115    TokenStream::from(expanded)
116}
117
118fn to_snake_case(s: &str) -> String {
119    let mut result = String::new();
120    for (i, c) in s.chars().enumerate() {
121        if i == 0 {
122            result.push(c.to_lowercase().next().unwrap());
123        } else if c.is_uppercase() {
124            result.push('_');
125            result.push(c.to_lowercase().next().unwrap());
126        } else {
127            result.push(c);
128        }
129    }
130    result
131}