Skip to main content

flawless_macros/
lib.rs

1use proc_macro::TokenStream;
2
3use quote::quote;
4use syn::{
5    parse, punctuated::Punctuated, spanned::Spanned, token::Comma, Error, Ident, Item, ItemFn, LitStr, Type,
6};
7
8#[proc_macro_attribute]
9pub fn workflow(args: TokenStream, item: TokenStream) -> TokenStream {
10    let export_name: LitStr = match parse(args) {
11        Ok(it) => it,
12        Err(e) => return token_stream_with_error(item, e),
13    };
14
15    let function: ItemFn = match parse(item.clone()) {
16        Ok(function) => function,
17        Err(e) => return token_stream_with_error(item, e),
18    };
19
20    // Extract just the workflow function argument types.
21    // E.g `fn test(i: Input<T>, r: Response)` -> `Input<T>, Response`.
22    let function_arg_types: Punctuated<Box<Type>, Comma> = function
23        .sig
24        .inputs
25        .iter()
26        .map(|arg| match arg {
27            syn::FnArg::Receiver(receiver) => receiver.ty.clone(),
28            syn::FnArg::Typed(pat_type) => pat_type.ty.clone(),
29        })
30        .collect();
31
32    // The workflow generic type is the `T` in `Workflow<T>`. It's the same as the function arguments, with an
33    // additional trailing comma indicating a tuple.
34    let mut workflow_t = function_arg_types.clone();
35    if !workflow_t.empty_or_trailing() {
36        workflow_t.push_punct(Comma::default());
37    }
38
39    let function_name = function.sig.ident.clone();
40    let function_vis = function.vis.clone();
41    let type_alias_name = Ident::new(&format!("_type_alias_{}", function_name), function_name.span());
42
43    let input_export = format!("#flawless#workflow#init#{}", export_name.value());
44    let input_export = LitStr::new(&input_export, export_name.span());
45    let input_name = format!("{}_flawless_workflow_init", function_name);
46    let input_name = Ident::new(&input_name, function_name.span());
47
48    let run_export = format!("#flawless#workflow#run#{}", export_name.value());
49    let run_export = LitStr::new(&run_export, export_name.span());
50    let run_name = format!("{}_flawless_workflow_run", function_name);
51    let run_name = Ident::new(&run_name, function_name.span());
52
53    let expanded = quote! {
54        #function
55
56        #[export_name = #input_export]
57        fn #input_name() -> i32 {
58            // Set correct logger for flawless.
59            flawless::logger::init_logger().unwrap();
60            // Override panic hook.
61            flawless::panic::override_panic(#export_name);
62            // Parese input.
63            flawless::init::init(#function_name)
64        }
65
66        #[export_name = #run_export]
67        fn #run_name() {
68            // Run workflow.
69            flawless::init::run(#function_name);
70        }
71
72        #[allow(non_camel_case_types)]
73        #function_vis type #function_name = #type_alias_name;
74
75        #[doc(hidden)]
76        #[allow(non_camel_case_types)]
77        #function_vis struct #type_alias_name;
78
79        impl flawless::workflow::WorkflowTypeAlias for #function_name {
80            const NAME: &'static str =  #export_name;
81            type InputArg =
82                <fn(#function_arg_types) as flawless::workflow::Workflow<(#workflow_t)>>::InputArg;
83        }
84    };
85
86    TokenStream::from(expanded)
87}
88
89#[proc_macro_attribute]
90pub fn message(args: TokenStream, item: TokenStream) -> TokenStream {
91    let msg_type: LitStr = match parse(args) {
92        Ok(it) => it,
93        Err(e) => return token_stream_with_error(item, e),
94    };
95
96    let msg_item: Item = match parse(item.clone()) {
97        Ok(item) => item,
98        Err(e) => return token_stream_with_error(item, e),
99    };
100
101    // Only a struct or enum can be a message.
102    let name = match msg_item.clone() {
103        Item::Enum(enum_item) => enum_item.ident,
104        Item::Struct(struct_item) => struct_item.ident,
105        _ => {
106            return token_stream_with_error(
107                item,
108                Error::new(
109                    msg_item.span(),
110                    "The `#[message(\"<type>\")]` macro can only be used on an enum or struct",
111                ),
112            )
113        }
114    };
115
116    let expanded = quote! {
117        #[derive(flawless::serde::Serialize, flawless::serde::Deserialize)]
118        #msg_item
119
120        impl flawless::message::Message for #name {
121            const TYPE: &'static str =  #msg_type;
122        }
123    };
124
125    TokenStream::from(expanded)
126}
127
128fn token_stream_with_error(mut tokens: TokenStream, error: Error) -> TokenStream {
129    tokens.extend(TokenStream::from(error.into_compile_error()));
130    tokens
131}