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 message_name = &input.ident;
10 let message_vis = &input.vis;
11 let caller_name = syn::Ident::new(&format!("{}Caller", message_name), message_name.span());
12 let replyer_name = syn::Ident::new(&format!("{}Replyer", message_name), message_name.span());
13
14 let mut caller_methods = Vec::new();
15
16 match &input.data {
17 Data::Enum(data) => {
18 for variant in &data.variants {
19 let variant_name = &variant.ident;
20 let snake_case_name = to_snake_case(&variant_name.to_string());
21 let method_name = syn::Ident::new(&snake_case_name, variant_name.span());
22 let timeout_method_name = syn::Ident::new(&format!("{}_timeout", snake_case_name), variant_name.span());
23 let method_name_str = snake_case_name.clone();
24
25 match &variant.fields {
26 Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
28 let arg_ty = &fields.unnamed[0].ty;
29
30 let mut is_actor_packet = false;
32 if let syn::Type::Path(tp) = arg_ty {
33 if let Some(segment) = tp.path.segments.last() {
34 if segment.ident == "ActorPacket" {
35 is_actor_packet = true;
36 }
37 }
38 }
39
40 if is_actor_packet {
41 if let syn::Type::Path(tp) = arg_ty {
42 if let Some(segment) = tp.path.segments.last() {
43 if segment.ident == "ActorPacket" {
44 if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
45 if args.args.len() == 2 {
46 let args_type = &args.args[0];
47 let ret_type = &args.args[1];
48
49 caller_methods.push(quote! {
50 pub async fn #method_name(&self, args: #args_type) -> anyhow::Result<#ret_type> {
51 let (packet, rx) = ActorPacket::new(args);
52
53 self.0.send(#message_name::#variant_name(packet))
54 .map_err(|e| {
55 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
56 anyhow::anyhow!(e).context(context)
57 })?;
58 rx.wait().await
59 .map_err(|e| {
60 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
61 anyhow::anyhow!(e).context(context)
62 })
63 }
64 });
65
66 caller_methods.push(quote! {
67 pub async fn #timeout_method_name(&self, args: #args_type, duration: ::std::time::Duration) -> anyhow::Result<#ret_type> {
68 let (packet, rx) = ActorPacket::new(args);
69 self.0.send(#message_name::#variant_name(packet))
70 .map_err(|e| {
71 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
72 anyhow::anyhow!(e).context(context)
73 })?;
74
75 rx.wait_timeout(duration).await
76 .map_err(|e| {
77 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
78 anyhow::anyhow!(e).context(context)
79 })
80 }
81 });
82 }
83 }
84 }
85 }
86 }
87 } else {
88 let args_type = arg_ty;
90 caller_methods.push(quote! {
91 pub fn #method_name(&self, args: #args_type) {
92 let _ = self.0.send(#message_name::#variant_name(args));
93 }
94 });
95 }
96 },
97 Fields::Named(_fields) => {
98 },
100 _ => {},
101 }
102 }
103 },
104 Data::Struct(data) => {
105 let fields = match &data.fields {
106 Fields::Unnamed(fields) if fields.unnamed.len() == 1 => fields,
107 _ => {
108 return syn::Error::new_spanned(&input, "ActorMessage tuple struct must have exactly one field")
109 .to_compile_error()
110 .into();
111 },
112 };
113
114 let arg_ty = &fields.unnamed[0].ty;
115 let snake_case_name = to_snake_case(&message_name.to_string());
116 let method_name = syn::Ident::new(&snake_case_name, message_name.span());
117 let timeout_method_name = syn::Ident::new(&format!("{}_timeout", snake_case_name), message_name.span());
118 let method_name_str = snake_case_name.clone();
119
120 let mut is_actor_packet = false;
122 if let syn::Type::Path(tp) = arg_ty {
123 if let Some(segment) = tp.path.segments.last() {
124 if segment.ident == "ActorPacket" {
125 is_actor_packet = true;
126 }
127 }
128 }
129
130 if is_actor_packet {
131 if let syn::Type::Path(tp) = arg_ty {
132 if let Some(segment) = tp.path.segments.last() {
133 if segment.ident == "ActorPacket" {
134 if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
135 if args.args.len() == 2 {
136 let args_type = &args.args[0];
137 let ret_type = &args.args[1];
138
139 caller_methods.push(quote! {
140 pub async fn #method_name(&self, args: #args_type) -> anyhow::Result<#ret_type> {
141 let (packet, rx) = ActorPacket::new(args);
142
143 self.0.send(#message_name(packet))
144 .map_err(|e| {
145 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
146 anyhow::anyhow!(e).context(context)
147 })?;
148 rx.wait().await
149 .map_err(|e| {
150 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
151 anyhow::anyhow!(e).context(context)
152 })
153 }
154 });
155
156 caller_methods.push(quote! {
157 pub async fn #timeout_method_name(&self, args: #args_type, duration: ::std::time::Duration) -> anyhow::Result<#ret_type> {
158 let (packet, rx) = ActorPacket::new(args);
159 self.0.send(#message_name(packet))
160 .map_err(|e| {
161 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
162 anyhow::anyhow!(e).context(context)
163 })?;
164
165 rx.wait_timeout(duration).await
166 .map_err(|e| {
167 let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
168 anyhow::anyhow!(e).context(context)
169 })
170 }
171 });
172 }
173 }
174 }
175 }
176 }
177 } else {
178 let args_type = arg_ty;
180 caller_methods.push(quote! {
181 pub fn #method_name(&self, args: #args_type) {
182 let _ = self.0.send(#message_name(args));
183 }
184 });
185 }
186 },
187 _ => {
188 return syn::Error::new_spanned(&input, "ActorMessage can only be derived for enums or tuple structs")
189 .to_compile_error()
190 .into();
191 },
192 }
193
194 let expanded = quote! {
195 #[derive(Clone, Debug)]
196 #message_vis struct #caller_name(::tokio::sync::mpsc::UnboundedSender<#message_name>);
197 impl #caller_name {
198 #(#caller_methods)*
199 }
200
201 #message_vis struct #replyer_name(::tokio::sync::mpsc::UnboundedReceiver<#message_name>);
202 impl #replyer_name {
203 pub async fn recv(&mut self) -> Option<#message_name> {
204 self.0.recv().await
205 }
206 }
207
208 impl std::ops::Deref for #replyer_name {
209 type Target = ::tokio::sync::mpsc::UnboundedReceiver<#message_name>;
210
211 fn deref(&self) -> &Self::Target {
212 &self.0
213 }
214 }
215
216 impl std::ops::DerefMut for #replyer_name {
217 fn deref_mut(&mut self) -> &mut Self::Target {
218 &mut self.0
219 }
220 }
221
222 impl #message_name {
223 pub fn actor() -> (#caller_name, #replyer_name) {
224 let (tx, rx) = ::tokio::sync::mpsc::unbounded_channel();
225 (#caller_name(tx), #replyer_name(rx))
226 }
227 }
228 };
229
230 TokenStream::from(expanded)
231}
232
233fn to_snake_case(s: &str) -> String {
234 let mut result = String::new();
235 for (i, c) in s.chars().enumerate() {
236 if i == 0 {
237 result.push(c.to_lowercase().next().unwrap());
238 } else if c.is_uppercase() {
239 result.push('_');
240 result.push(c.to_lowercase().next().unwrap());
241 } else {
242 result.push(c);
243 }
244 }
245 result
246}