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