use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::{parse_macro_input, ItemFn, LitInt};
#[proc_macro_attribute]
pub fn main(attr: TokenStream, item: TokenStream) -> TokenStream {
if !attr.is_empty() {
return syn::Error::new_spanned(
proc_macro2::TokenStream::from(attr),
"#[rivet::main] takes no arguments",
)
.to_compile_error()
.into();
}
let input = parse_macro_input!(item as ItemFn);
if !input.sig.inputs.is_empty() {
return syn::Error::new_spanned(
&input.sig.inputs,
"#[rivet::main] functions take no parameters",
)
.to_compile_error()
.into();
}
let fn_attrs = &input.attrs;
let fn_block = &input.block;
let expanded = quote! {
#[no_mangle]
#(#fn_attrs)*
extern "C" fn rivet_main() -> ! {
::rivet::init();
#fn_block
}
};
TokenStream::from(expanded)
}
#[proc_macro_attribute]
pub fn task(attr: TokenStream, item: TokenStream) -> TokenStream {
let input = parse_macro_input!(item as ItemFn);
let fn_name = &input.sig.ident;
let fn_visibility = &input.vis;
let fn_sig = &input.sig;
let fn_attrs = &input.attrs;
let task_body = &input.block;
if input.sig.asyncness.is_none() {
return syn::Error::new_spanned(&input.sig, "#[rivet::task] requires an `async fn`")
.to_compile_error()
.into();
}
if !input.sig.inputs.is_empty() {
return syn::Error::new_spanned(
&input.sig.inputs,
"#[rivet::task] functions currently take no parameters; \
use a `static` for shared state (peripherals, queues, etc.)",
)
.to_compile_error()
.into();
}
let mut priority: u8 = 0;
let mut stack_size: usize = 512;
let mut saw_priority = false;
let parser = syn::meta::parser(|meta| {
if meta.path.is_ident("priority") {
let value = meta.value()?;
let lit: LitInt = value.parse()?;
priority = lit.base10_parse::<u8>()?;
saw_priority = true;
} else if meta.path.is_ident("stack") {
let value = meta.value()?;
let lit: LitInt = value.parse()?;
stack_size = lit.base10_parse::<usize>()?;
} else {
return Err(
meta.error("unsupported #[rivet::task] attribute; expected `priority` or `stack`")
);
}
Ok(())
});
parse_macro_input!(attr with parser);
if !saw_priority {
return syn::Error::new_spanned(
fn_name,
"#[rivet::task] requires `priority = N`, e.g. #[rivet::task(priority = 1)]",
)
.to_compile_error()
.into();
}
let poll_fn_name = format_ident!("__rivet_poll_{}", fn_name);
let completed_fn_name = format_ident!("__rivet_completed_{}", fn_name);
let cell_name = format_ident!("__RIVET_CELL_{}", fn_name);
let reg_name = format_ident!("__RIVET_REG_{}", fn_name);
let expanded = quote! {
#(#fn_attrs)*
#fn_visibility #fn_sig #task_body
#[allow(non_upper_case_globals)]
static #cell_name: ::rivet::task::TaskCell<#stack_size> = ::rivet::task::TaskCell::new();
#[allow(non_snake_case)]
unsafe fn #poll_fn_name(
_user_data: *mut (),
waker: &::core::task::Waker,
) -> ::core::task::Poll<()> {
#cell_name.poll(#fn_name, waker)
}
#[allow(non_snake_case)]
unsafe fn #completed_fn_name(_user_data: *mut ()) -> bool {
#cell_name.is_completed()
}
#[link_section = ".rivet_tasks"]
#[used]
#[allow(non_upper_case_globals)]
static #reg_name: ::rivet::task::TaskReg = ::rivet::task::TaskReg {
priority: #priority,
index_in_priority: 0,
_reserved: [0; 2],
poll_fn: #poll_fn_name as unsafe fn(*mut (), &::core::task::Waker) -> ::core::task::Poll<()>,
completed_fn: #completed_fn_name as unsafe fn(*mut ()) -> bool,
user_data: ::core::ptr::null_mut(),
};
};
TokenStream::from(expanded)
}