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