xasync-macro 0.1.10

Procedural Macros for Cyberex Async
Documentation
extern crate proc_macro;
use proc_macro::TokenStream;
use quote::quote;
use syn::{Data, DeriveInput, Fields, parse_macro_input};

#[proc_macro_derive(ActorMessage)]
pub fn actor_message(input: TokenStream) -> TokenStream {
    let input = parse_macro_input!(input as DeriveInput);
    let message_name = &input.ident;
    let message_vis = &input.vis;
    let caller_name = syn::Ident::new(&format!("{}Caller", message_name), message_name.span());
    let weak_caller_name = syn::Ident::new(&format!("{}WeakCaller", message_name), message_name.span());
    let replyer_name = syn::Ident::new(&format!("{}Replyer", message_name), message_name.span());

    let mut caller_methods = Vec::new();

    match &input.data {
        Data::Enum(data) => {
            for variant in &data.variants {
                let variant_name = &variant.ident;
                let snake_case_name = to_snake_case(&variant_name.to_string());
                let method_name = syn::Ident::new(&snake_case_name, variant_name.span());
                let timeout_method_name = syn::Ident::new(&format!("{}_timeout", snake_case_name), variant_name.span());
                let method_name_str = snake_case_name.clone();

                match &variant.fields {
                    // single unnamed field: either an ActorPacket (request/response) or a plain payload (fire-and-forget)
                    Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
                        let arg_ty = &fields.unnamed[0].ty;

                        // detect ActorPacket<T, R>
                        let mut is_actor_packet = false;
                        if let syn::Type::Path(tp) = arg_ty {
                            if let Some(segment) = tp.path.segments.last() {
                                if segment.ident == "ActorPacket" {
                                    is_actor_packet = true;
                                }
                            }
                        }

                        if is_actor_packet {
                            if let syn::Type::Path(tp) = arg_ty {
                                if let Some(segment) = tp.path.segments.last() {
                                    if segment.ident == "ActorPacket" {
                                        if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
                                            if args.args.len() == 2 {
                                                let args_type = &args.args[0];
                                                let ret_type = &args.args[1];

                                                caller_methods.push(quote! {
                                                    pub async fn #method_name(&self, args: #args_type) -> anyhow::Result<#ret_type> {
                                                        let (packet, rx) = ActorPacket::new(args);

                                                        self.0.send(#message_name::#variant_name(packet))
                                                            .map_err(|e| {
                                                                let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                                anyhow::anyhow!(e).context(context)
                                                            })?;
                                                        rx.wait().await
                                                            .map_err(|e| {
                                                                let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                                anyhow::anyhow!(e).context(context)
                                                            })
                                                    }
                                                });

                                                caller_methods.push(quote! {
                                                    pub async fn #timeout_method_name(&self, args: #args_type, duration: ::std::time::Duration) -> anyhow::Result<#ret_type> {
                                                        let (packet, rx) = ActorPacket::new(args);
                                                        self.0.send(#message_name::#variant_name(packet))
                                                            .map_err(|e| {
                                                                let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                                anyhow::anyhow!(e).context(context)
                                                            })?;

                                                        rx.wait_timeout(duration).await
                                                            .map_err(|e| {
                                                                let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                                anyhow::anyhow!(e).context(context)
                                                            })
                                                    }
                                                });
                                            }
                                        }
                                    }
                                }
                            }
                        } else {
                            // fire-and-forget: generate a synchronous method that sends the variant without awaiting
                            let args_type = arg_ty;
                            caller_methods.push(quote! {
                                pub fn #method_name(&self, args: #args_type) {
                                    let _ = self.0.send(#message_name::#variant_name(args));
                                }
                            });
                        }
                    },
                    Fields::Named(_fields) => {
                        // currently ignore named-field variants
                    },
                    _ => {},
                }
            }
        },
        Data::Struct(data) => {
            let fields = match &data.fields {
                Fields::Unnamed(fields) if fields.unnamed.len() == 1 => fields,
                _ => {
                    return syn::Error::new_spanned(&input, "ActorMessage tuple struct must have exactly one field")
                        .to_compile_error()
                        .into();
                },
            };

            let arg_ty = &fields.unnamed[0].ty;
            let snake_case_name = to_snake_case(&message_name.to_string());
            let method_name = syn::Ident::new(&snake_case_name, message_name.span());
            let timeout_method_name = syn::Ident::new(&format!("{}_timeout", snake_case_name), message_name.span());
            let method_name_str = snake_case_name.clone();

            // detect ActorPacket<T, R>
            let mut is_actor_packet = false;
            if let syn::Type::Path(tp) = arg_ty {
                if let Some(segment) = tp.path.segments.last() {
                    if segment.ident == "ActorPacket" {
                        is_actor_packet = true;
                    }
                }
            }

            if is_actor_packet {
                if let syn::Type::Path(tp) = arg_ty {
                    if let Some(segment) = tp.path.segments.last() {
                        if segment.ident == "ActorPacket" {
                            if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
                                if args.args.len() == 2 {
                                    let args_type = &args.args[0];
                                    let ret_type = &args.args[1];

                                    caller_methods.push(quote! {
                                        pub async fn #method_name(&self, args: #args_type) -> anyhow::Result<#ret_type> {
                                            let (packet, rx) = ActorPacket::new(args);

                                            self.0.send(#message_name(packet))
                                                .map_err(|e| {
                                                    let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                    anyhow::anyhow!(e).context(context)
                                                })?;
                                            rx.wait().await
                                                .map_err(|e| {
                                                    let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                    anyhow::anyhow!(e).context(context)
                                                })
                                        }
                                    });

                                    caller_methods.push(quote! {
                                        pub async fn #timeout_method_name(&self, args: #args_type, duration: ::std::time::Duration) -> anyhow::Result<#ret_type> {
                                            let (packet, rx) = ActorPacket::new(args);
                                            self.0.send(#message_name(packet))
                                                .map_err(|e| {
                                                    let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                    anyhow::anyhow!(e).context(context)
                                                })?;

                                            rx.wait_timeout(duration).await
                                                .map_err(|e| {
                                                    let context = format!("Fail to `{}`, error: {}", #method_name_str, e);
                                                    anyhow::anyhow!(e).context(context)
                                                })
                                        }
                                    });
                                }
                            }
                        }
                    }
                }
            } else {
                // fire-and-forget: generate a synchronous method that sends the struct without awaiting
                let args_type = arg_ty;
                caller_methods.push(quote! {
                    pub fn #method_name(&self, args: #args_type) {
                        let _ = self.0.send(#message_name(args));
                    }
                });
            }
        },
        _ => {
            return syn::Error::new_spanned(&input, "ActorMessage can only be derived for enums or tuple structs")
                .to_compile_error()
                .into();
        },
    }

    let expanded = quote! {
        #[derive(Clone, Debug)]
        #message_vis struct #caller_name(::tokio::sync::mpsc::UnboundedSender<#message_name>);
        impl #caller_name {
            pub fn downgrade(&self) -> #weak_caller_name {
                #weak_caller_name(self.0.downgrade())
            }

            #(#caller_methods)*
        }

        impl ::std::convert::AsRef<::tokio::sync::mpsc::UnboundedSender<#message_name>> for #caller_name {
            fn as_ref(&self) -> &::tokio::sync::mpsc::UnboundedSender<#message_name> {
                &self.0
            }
        }

        #[derive(Clone, Debug)]
        #message_vis struct #weak_caller_name(::tokio::sync::mpsc::WeakUnboundedSender<#message_name>);
        impl #weak_caller_name {
            pub fn upgrade(&self) -> Option<#caller_name> {
                self.0.upgrade().map(#caller_name)
            }
        }

        #message_vis struct #replyer_name(::tokio::sync::mpsc::UnboundedReceiver<#message_name>);
        impl #replyer_name {
            pub async fn recv(&mut self) -> Option<#message_name> {
                self.0.recv().await
            }
        }

        impl std::ops::Deref for #replyer_name {
            type Target = ::tokio::sync::mpsc::UnboundedReceiver<#message_name>;

            fn deref(&self) -> &Self::Target {
                &self.0
            }
        }

        impl std::ops::DerefMut for #replyer_name {
            fn deref_mut(&mut self) -> &mut Self::Target {
                &mut self.0
            }
        }

        impl #message_name {
            pub fn actor() -> (#caller_name, #replyer_name) {
                let (tx, rx) = ::tokio::sync::mpsc::unbounded_channel();
                (#caller_name(tx), #replyer_name(rx))
            }
        }
    };

    TokenStream::from(expanded)
}

fn to_snake_case(s: &str) -> String {
    let mut result = String::new();
    for (i, c) in s.chars().enumerate() {
        if i == 0 {
            result.push(c.to_lowercase().next().unwrap());
        } else if c.is_uppercase() {
            result.push('_');
            result.push(c.to_lowercase().next().unwrap());
        } else {
            result.push(c);
        }
    }
    result
}