use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::parse::{Parse, ParseStream};
use syn::punctuated::Punctuated;
use syn::{Error, Expr, LitInt, Result, Token, Type, parse_macro_input};
struct Spec {
count: LitInt,
ty: Type,
build: Expr,
}
impl Parse for Spec {
fn parse(input: ParseStream) -> Result<Self> {
let content;
syn::parenthesized!(content in input);
let count: Expr = content.parse()?;
let Expr::Lit(syn::ExprLit {
lit: syn::Lit::Int(count),
..
}) = count
else {
return Err(Error::new_spanned(
count,
"worker count must be a usize literal (ranges are computed at expansion)",
));
};
content.parse::<Token![,]>()?;
let ty: Type = content.parse()?;
content.parse::<Token![,]>()?;
let build: Expr = content.parse()?;
Ok(Spec { count, ty, build })
}
}
struct Specs(Punctuated<Spec, Token![,]>);
impl Parse for Specs {
fn parse(input: ParseStream) -> Result<Self> {
Ok(Specs(Punctuated::parse_terminated(input)?))
}
}
#[proc_macro]
pub fn workers(input: TokenStream) -> TokenStream {
let Specs(specs) = parse_macro_input!(input as Specs);
let specs: Vec<Spec> = specs.into_iter().collect();
if specs.is_empty() {
return Error::new(
proc_macro2::Span::call_site(),
"workers! needs at least one (count, Type, build_fn) spec",
)
.to_compile_error()
.into();
}
let first_ty = &specs[0].ty;
let variants: Vec<_> = specs
.iter()
.enumerate()
.map(|(i, _)| format_ident!("V{i}"))
.collect();
let variant_defs = specs.iter().zip(&variants).map(|(s, v)| {
let ty = &s.ty;
quote! { #v(#ty) }
});
let mut start = 0_usize;
let mut build_arms = Vec::new();
for (s, v) in specs.iter().zip(&variants) {
let n: usize = match s.count.base10_parse() {
Ok(n) => n,
Err(e) => return e.to_compile_error().into(),
};
let end = start + n;
let build = &s.build;
build_arms.push(quote! { #start..#end => Crew::#v((#build)(i)) });
start = end;
}
let total = start;
let step_arms = variants
.iter()
.map(|v| quote! { Crew::#v(b) => b.step(ev).await });
let init_arms = variants
.iter()
.map(|v| quote! { Crew::#v(b) => b.init().await });
let out = quote! {
{
enum Crew {
#(#variant_defs),*
}
impl ::behavior::Behavior for Crew {
type Addr = <#first_ty as ::behavior::Behavior>::Addr;
type Msg = <#first_ty as ::behavior::Behavior>::Msg;
type Event = <#first_ty as ::behavior::Behavior>::Event;
type Sends = <#first_ty as ::behavior::Behavior>::Sends;
type Ph = ::behavior::Never;
type Error = <#first_ty as ::behavior::Behavior>::Error;
type Birth = <#first_ty as ::behavior::Behavior>::Birth;
async fn init(&mut self) -> ::core::result::Result<
::behavior::Actions<Self::Addr, Self::Ph, Self::Sends, Self::Birth>,
Self::Error,
> {
match self {
#(#init_arms),*
}
}
async fn step(
&mut self,
ev: Self::Event,
) -> ::core::result::Result<
::behavior::Actions<Self::Addr, Self::Ph, Self::Sends, Self::Birth>,
Self::Error,
> {
match self {
#(#step_arms),*
}
}
}
fn crew_build(i: usize) -> Crew {
match i {
#(#build_arms,)*
_ => unreachable!("workers!: fleet index out of range — driver/behavior desync"),
}
}
(#total, crew_build as fn(usize) -> Crew)
}
};
out.into()
}