1use proc_macro::{Span, TokenStream};
2use quote::quote;
3use syn::{parse::Parse, parse_macro_input, Error, Ident, ItemFn, LitInt, Token};
4
5const DEFAULT_NUM_WORKERS: usize = 1;
6
7struct WorkersAttr {
8 workers: usize,
9}
10
11impl Parse for WorkersAttr {
12 fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
16 let name: Ident = input.parse()?;
17 input.parse::<Token![=]>()?;
18 let value: LitInt = input.parse()?;
19
20 if name == "workers" {
21 let workers = value.base10_parse()?;
22 Ok(Self { workers })
23 } else {
24 Err(input.error("Expected `workers` argument"))
25 }
26 }
27}
28
29#[proc_macro_attribute]
30pub fn hooch_main(attr: TokenStream, item: TokenStream) -> TokenStream {
31 let mut input = parse_macro_input!(item as ItemFn);
32 let is_async = input.sig.asyncness.is_some();
33
34 if input.sig.ident != "main" {
35 let error = Error::new_spanned(input, "hooch_main can only be used on 'main'");
36 return error.to_compile_error().into();
37 }
38
39 let workers: usize = if attr.is_empty() {
40 DEFAULT_NUM_WORKERS
41 } else {
42 let WorkersAttr { workers } = parse_macro_input!(attr as WorkersAttr);
43 workers
44 };
45
46 input.sig.ident = syn::Ident::new("main_hooch", input.sig.ident.span());
47 let main_hooch_fn = syn::Ident::new("main_hooch", Span::call_site().into());
48
49 if !is_async {
50 let error = Error::new_spanned(input.sig.fn_token, "main must be async");
51 return error.to_compile_error().into();
52 }
53
54 let output = quote! {
55 use hooch::runtime::RuntimeBuilder;
56 #input
57 fn main() {
58 let handle = RuntimeBuilder::new().num_workers(#workers).build();
59
60 handle.run_blocking(async {
61 #main_hooch_fn().await
62 });
63 }
64
65 };
66 output.into()
67}