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 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 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 flawless::logger::init_logger().unwrap();
60 flawless::panic::override_panic(#export_name);
62 flawless::init::init(#function_name)
64 }
65
66 #[export_name = #run_export]
67 fn #run_name() {
68 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 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}