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}