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 enum_name = &input.ident;
10    let caller_name = syn::Ident::new(&format!("{}Caller", enum_name), enum_name.span());
11    let replyer_name = syn::Ident::new(&format!("{}Replyer", enum_name), enum_name.span());
12
13    let variants = match &input.data {
14        Data::Enum(data) => &data.variants,
15        _ => {
16            return syn::Error::new_spanned(&input, "ActorMessage can only be derived for enums")
17                .to_compile_error()
18                .into();
19        },
20    };
21
22    let mut caller_methods = Vec::new();
23
24    for variant in variants {
25        let variant_name = &variant.ident;
26        let snake_case_name = to_snake_case(&variant_name.to_string());
27        let method_name = syn::Ident::new(&snake_case_name, variant_name.span());
28
29        match &variant.fields {
30            Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
31                if let syn::Type::Path(tp) = &fields.unnamed[0].ty {
32                    if let Some(segment) = tp.path.segments.last() {
33                        if segment.ident == "ActorPacket" {
34                            if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
35                                if args.args.len() == 2 {
36                                    let args_type = &args.args[0];
37                                    let ret_type = &args.args[1];
38
39                                    caller_methods.push(quote! {
40                                        pub async fn #method_name(&self, args: #args_type) -> anyhow::Result<#ret_type> {
41                                            let (packet, rx) = cyberex_xasync::call::ActorPacket::new(args);
42                                            self.0.send(#enum_name::#variant_name(packet))?;
43                                            rx.wait().await
44                                        }
45                                    });
46                                }
47                            }
48                        }
49                    }
50                }
51            },
52            Fields::Named(fields) => {
53                let mut sender_field = None;
54                let mut other_fields = Vec::new();
55
56                for field in &fields.named {
57                    if let syn::Type::Path(tp) = &field.ty {
58                        if let Some(segment) = tp.path.segments.last() {
59                            if segment.ident == "Sender" {
60                                sender_field = Some(field);
61                            } else {
62                                other_fields.push(field);
63                            }
64                        }
65                    } else {
66                        other_fields.push(field);
67                    }
68                }
69
70                if let Some(sender_field) = sender_field {
71                    let sender_field_name = &sender_field.ident;
72                    let arg_types: Vec<_> = other_fields.iter().map(|f| &f.ty).collect();
73                    let arg_names: Vec<_> = other_fields.iter().map(|f| &f.ident).collect();
74
75                    if arg_types.len() == 1 {
76                        let arg_type = arg_types[0];
77                        let arg_name = arg_names[0];
78
79                        caller_methods.push(quote! {
80                            pub async fn #method_name(&self, #arg_name: #arg_type) -> anyhow::Result<String> {
81                                let (tx, rx) = ::tokio::sync::oneshot::channel();
82                                self.0.send(#enum_name::#variant_name { #arg_name, #sender_field_name: tx })?;
83                                rx.await?
84                            }
85                        });
86                    }
87                }
88            },
89            _ => {},
90        }
91    }
92
93    let expanded = quote! {
94        pub struct #caller_name(::tokio::sync::mpsc::UnboundedSender<#enum_name>);
95        impl #caller_name {
96            #(#caller_methods)*
97        }
98
99        pub struct #replyer_name(::tokio::sync::mpsc::UnboundedReceiver<#enum_name>);
100        impl #replyer_name {
101            pub async fn recv(&mut self) -> Option<#enum_name> {
102                self.0.recv().await
103            }
104        }
105
106        impl #enum_name {
107            pub fn actor() -> (#caller_name, #replyer_name) {
108                let (tx, rx) = ::tokio::sync::mpsc::unbounded_channel();
109                (#caller_name(tx), #replyer_name(rx))
110            }
111        }
112    };
113
114    TokenStream::from(expanded)
115}
116
117fn to_snake_case(s: &str) -> String {
118    let mut result = String::new();
119    for (i, c) in s.chars().enumerate() {
120        if i == 0 {
121            result.push(c.to_lowercase().next().unwrap());
122        } else if c.is_uppercase() {
123            result.push('_');
124            result.push(c.to_lowercase().next().unwrap());
125        } else {
126            result.push(c);
127        }
128    }
129    result
130}