flawless-macros 1.0.0-beta.2

Macros for flawless.
Documentation
use proc_macro::TokenStream;

use quote::quote;
use syn::{
    parse, punctuated::Punctuated, spanned::Spanned, token::Comma, Error, Ident, Item, ItemFn, LitStr, Type,
};

#[proc_macro_attribute]
pub fn workflow(args: TokenStream, item: TokenStream) -> TokenStream {
    let export_name: LitStr = match parse(args) {
        Ok(it) => it,
        Err(e) => return token_stream_with_error(item, e),
    };

    let function: ItemFn = match parse(item.clone()) {
        Ok(function) => function,
        Err(e) => return token_stream_with_error(item, e),
    };

    // Extract just the workflow function argument types.
    // E.g `fn test(i: Input<T>, r: Response)` -> `Input<T>, Response`.
    let function_arg_types: Punctuated<Box<Type>, Comma> = function
        .sig
        .inputs
        .iter()
        .map(|arg| match arg {
            syn::FnArg::Receiver(receiver) => receiver.ty.clone(),
            syn::FnArg::Typed(pat_type) => pat_type.ty.clone(),
        })
        .collect();

    // The workflow generic type is the `T` in `Workflow<T>`. It's the same as the function arguments, with an
    // additional trailing comma indicating a tuple.
    let mut workflow_t = function_arg_types.clone();
    if !workflow_t.empty_or_trailing() {
        workflow_t.push_punct(Comma::default());
    }

    let function_name = function.sig.ident.clone();
    let function_vis = function.vis.clone();
    let type_alias_name = Ident::new(&format!("_type_alias_{}", function_name), function_name.span());

    let input_export = format!("#flawless#workflow#init#{}", export_name.value());
    let input_export = LitStr::new(&input_export, export_name.span());
    let input_name = format!("{}_flawless_workflow_init", function_name.to_string());
    let input_name = Ident::new(&input_name, function_name.span());

    let run_export = format!("#flawless#workflow#run#{}", export_name.value());
    let run_export = LitStr::new(&run_export, export_name.span());
    let run_name = format!("{}_flawless_workflow_run", function_name.to_string());
    let run_name = Ident::new(&run_name, function_name.span());

    let expanded = quote! {
        #function

        #[export_name = #input_export]
        fn #input_name() -> i32 {
            // Set correct logger for flawless.
            flawless::logger::init_logger().unwrap();
            // Override panic hook.
            flawless::panic::override_panic(#export_name);
            // Parese input.
            flawless::init::init(#function_name)
        }

        #[export_name = #run_export]
        fn #run_name() {
            // Run workflow.
            flawless::init::run(#function_name);
        }

        #[allow(non_camel_case_types)]
        #function_vis type #function_name = #type_alias_name;

        #[doc(hidden)]
        #[allow(non_camel_case_types)]
        #function_vis struct #type_alias_name;

        impl flawless::workflow::WorkflowTypeAlias for #function_name {
            const NAME: &'static str =  #export_name;
            type InputArg =
                <fn(#function_arg_types) as flawless::workflow::Workflow<(#workflow_t)>>::InputArg;
        }
    };

    TokenStream::from(expanded)
}

#[proc_macro_attribute]
pub fn message(args: TokenStream, item: TokenStream) -> TokenStream {
    let msg_type: LitStr = match parse(args) {
        Ok(it) => it,
        Err(e) => return token_stream_with_error(item, e),
    };

    let msg_item: Item = match parse(item.clone()) {
        Ok(item) => item,
        Err(e) => return token_stream_with_error(item, e),
    };

    // Only a struct or enum can be a message.
    let name = match msg_item.clone() {
        Item::Enum(enum_item) => enum_item.ident,
        Item::Struct(struct_item) => struct_item.ident,
        _ => {
            return token_stream_with_error(
                item,
                Error::new(
                    msg_item.span(),
                    "The `#[message(\"<type>\")]` macro can only be used on an enum or struct",
                ),
            )
        }
    };

    let expanded = quote! {
        #[derive(flawless::serde::Serialize, flawless::serde::Deserialize)]
        #msg_item

        impl flawless::message::Message for #name {
            const TYPE: &'static str =  #msg_type;
        }
    };

    TokenStream::from(expanded)
}

fn token_stream_with_error(mut tokens: TokenStream, error: Error) -> TokenStream {
    tokens.extend(TokenStream::from(error.into_compile_error()));
    tokens
}