Skip to main content

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 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                    // single unnamed field: either an ActorPacket (request/response) or a plain payload (fire-and-forget)
28                    Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
29                        let arg_ty = &fields.unnamed[0].ty;
30
31                        // detect ActorPacket<T, R>
32                        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                            // fire-and-forget: generate a synchronous method that sends the variant without awaiting
90                            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                        // currently ignore named-field variants
100                    },
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            // detect ActorPacket<T, R>
122            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                // fire-and-forget: generate a synchronous method that sends the struct without awaiting
180                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}