Skip to main content

hooch_macro/
lib.rs

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    // ParseStream acts as an internal cursor that keeps track of the current position in the
13    // token stream. Each `input.parse::<Type>()?` call parses a part of the token stream and
14    // advances the cursor
15    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}