use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::{format_ident, quote};
use std::collections::HashMap;
use syn::parse::{Parse, ParseStream};
use syn::punctuated::Punctuated;
use syn::{
Attribute, Expr, Ident, LitInt, Meta, Path, Result as SynResult, Token, Type, bracketed,
};
mod kw {
syn::custom_keyword!(node);
syn::custom_keyword!(pool);
syn::custom_keyword!(deps);
syn::custom_keyword!(spawn);
syn::custom_keyword!(policy);
syn::custom_keyword!(min);
syn::custom_keyword!(max);
syn::custom_keyword!(disabled);
}
#[derive(Clone)]
struct Dep {
cfg: Vec<Attribute>,
ident: Ident,
}
fn parse_dep_list(input: ParseStream) -> SynResult<Vec<Dep>> {
let content;
bracketed!(content in input);
let mut deps = Vec::new();
while !content.is_empty() {
let cfg = content.call(Attribute::parse_outer)?;
let ident: Ident = content.parse()?;
deps.push(Dep { cfg, ident });
if content.peek(Token![,]) {
content.parse::<Token![,]>()?;
}
}
Ok(deps)
}
fn parse_mode_list(input: ParseStream) -> SynResult<Vec<Ident>> {
let content;
bracketed!(content in input);
let punct = Punctuated::<Ident, Token![,]>::parse_terminated(&content)?;
Ok(punct.into_iter().collect())
}
struct NodeItem {
cfg: Vec<Attribute>,
ident: Ident,
mode: Ident,
deps: Vec<Dep>,
spawn: Option<Expr>,
disabled: bool,
}
struct PoolItem {
cfg: Vec<Attribute>,
ident: Ident,
modes: Vec<Ident>,
deps: Vec<Dep>,
spawn: Expr,
policy: Expr,
policy_ty: Option<Type>,
min: LitInt,
max: LitInt,
}
#[allow(clippy::large_enum_variant)]
enum Item {
Node(NodeItem),
Pool(PoolItem),
}
struct GraphSpec {
items: Vec<Item>,
}
impl Parse for GraphSpec {
fn parse(input: ParseStream) -> SynResult<Self> {
let mut items = Vec::new();
while !input.is_empty() {
let cfg = input.call(Attribute::parse_outer)?;
if input.peek(kw::node) {
items.push(Item::Node(parse_node(input, cfg)?));
} else if input.peek(kw::pool) {
items.push(Item::Pool(parse_pool(input, cfg)?));
} else {
return Err(
input.error("expected `node` or `pool` (optionally `#[cfg(...)]`-prefixed)")
);
}
}
Ok(GraphSpec { items })
}
}
fn parse_node(input: ParseStream, cfg: Vec<Attribute>) -> SynResult<NodeItem> {
input.parse::<kw::node>()?;
let ident: Ident = input.parse()?;
input.parse::<Token![=]>()?;
let mode: Ident = input.parse()?;
input.parse::<Token![,]>()?;
input.parse::<kw::deps>()?;
input.parse::<Token![:]>()?;
let deps = parse_dep_list(input)?;
let mut spawn = None;
let mut disabled = false;
while input.peek(Token![,]) {
input.parse::<Token![,]>()?;
if input.peek(kw::spawn) {
input.parse::<kw::spawn>()?;
input.parse::<Token![:]>()?;
spawn = Some(input.parse::<Expr>()?);
} else if input.peek(kw::disabled) {
input.parse::<kw::disabled>()?;
disabled = true;
} else {
return Err(input.error("expected `spawn:` or `disabled`"));
}
}
input.parse::<Token![;]>()?;
Ok(NodeItem {
cfg,
ident,
mode,
deps,
spawn,
disabled,
})
}
fn parse_pool(input: ParseStream, cfg: Vec<Attribute>) -> SynResult<PoolItem> {
input.parse::<kw::pool>()?;
let ident: Ident = input.parse()?;
input.parse::<Token![=]>()?;
let modes = parse_mode_list(input)?;
input.parse::<Token![,]>()?;
input.parse::<kw::deps>()?;
input.parse::<Token![:]>()?;
let deps = parse_dep_list(input)?;
input.parse::<Token![,]>()?;
input.parse::<kw::spawn>()?;
input.parse::<Token![:]>()?;
let spawn: Expr = input.parse()?;
input.parse::<Token![,]>()?;
input.parse::<kw::policy>()?;
input.parse::<Token![:]>()?;
let policy_ty = {
let fork = input.fork();
if fork.parse::<Type>().is_ok() && fork.peek(Token![=]) {
let ty: Type = input.parse()?;
input.parse::<Token![=]>()?;
Some(ty)
} else {
None
}
};
let policy: Expr = input.parse()?;
input.parse::<Token![,]>()?;
input.parse::<kw::min>()?;
input.parse::<Token![:]>()?;
let min: LitInt = input.parse()?;
input.parse::<Token![,]>()?;
input.parse::<kw::max>()?;
input.parse::<Token![:]>()?;
let max: LitInt = input.parse()?;
input.parse::<Token![;]>()?;
Ok(PoolItem {
cfg,
ident,
modes,
deps,
spawn,
policy,
policy_ty,
min,
max,
})
}
fn name_string(ident: &Ident) -> String {
ident.to_string().to_lowercase().replace('_', "-")
}
fn inject_node_call(task: &Expr, node_ref: &TokenStream2) -> SynResult<TokenStream2> {
match task {
Expr::Path(_) => Ok(quote!(#task(#node_ref))),
Expr::Call(c) => {
let f = &c.func;
let args = c.args.iter();
Ok(quote!(#f(#node_ref #(, #args)*)))
}
other => Err(syn::Error::new_spanned(
other,
"expected a task-fn path or a partial call like `f(extra_args)`",
)),
}
}
fn cfg_predicate(attrs: &[Attribute]) -> Option<TokenStream2> {
let preds: Vec<TokenStream2> = attrs
.iter()
.filter_map(|a| match &a.meta {
Meta::List(ml) if ml.path.is_ident("cfg") => Some(ml.tokens.clone()),
_ => None,
})
.collect();
match preds.len() {
0 => None,
1 => Some(preds[0].clone()),
_ => Some(quote!(all(#(#preds),*))),
}
}
fn policy_type(expr: &Expr) -> SynResult<Path> {
if let Expr::Call(call) = expr
&& let Expr::Path(p) = &*call.func
{
let n = p.path.segments.len();
if n >= 2 {
let segs: Punctuated<_, Token![::]> =
p.path.segments.iter().take(n - 1).cloned().collect();
return Ok(Path {
leading_colon: p.path.leading_colon,
segments: segs,
});
}
}
Err(syn::Error::new_spanned(
expr,
"pool `policy:` must be a `Type::new(..)` constructor (e.g. `DeferredShrink::new(..)`), \
or give the type explicitly: `policy: <Type> = <expr>`",
))
}
struct Slot {
cfg_pred: Option<TokenStream2>,
reference: TokenStream2,
deps: Vec<Dep>,
}
fn node_spawn(
ident: &Ident,
spawn: &Option<Expr>,
spawn_fn: &TokenStream2,
) -> SynResult<TokenStream2> {
Ok(match spawn {
None => quote!(::core::option::Option::None),
Some(e @ (Expr::Path(_) | Expr::Call(_))) => {
let call = inject_node_call(e, "e!(&#ident))?;
quote!(::core::option::Option::Some(
(|s| {
s.spawn(#call?);
::core::result::Result::Ok(())
}) as #spawn_fn
))
}
Some(e) => quote!(::core::option::Option::Some((#e) as #spawn_fn)),
})
}
fn emit_node(
n: &NodeItem,
cr: &TokenStream2,
spawn_fn: &TokenStream2,
) -> SynResult<(TokenStream2, Slot)> {
let ident = &n.ident;
let cfg = &n.cfg;
let mode = &n.mode;
let name = name_string(&n.ident);
let disabled = n.disabled;
let spawn = node_spawn(ident, &n.spawn, spawn_fn)?;
let def = quote! {
#(#cfg)*
pub static #ident: #cr::TaskNode =
#cr::TaskNode::new(#name, #cr::Mode::#mode, #spawn, #disabled);
};
let slot = Slot {
cfg_pred: cfg_predicate(cfg),
reference: quote!(&#ident),
deps: n.deps.clone(),
};
Ok((def, slot))
}
fn emit_pool(
p: &PoolItem,
cr: &TokenStream2,
spawn_fn: &TokenStream2,
) -> SynResult<(Vec<TokenStream2>, TokenStream2, Vec<Slot>)> {
let ident = &p.ident;
let cfg = &p.cfg;
let lname = name_string(&p.ident);
let pool_static = format_ident!("{}_POOL", ident);
let k = p.modes.len();
let min_v: u8 = p.min.base10_parse()?;
let max_v: u8 = p.max.base10_parse()?;
if min_v > max_v {
return Err(syn::Error::new_spanned(
&p.min,
format!("pool `min:` ({min_v}) must not exceed `max:` ({max_v})"),
));
}
if usize::from(max_v) > k {
return Err(syn::Error::new_spanned(
&p.max,
format!("pool `max:` ({max_v}) exceeds the declared member count ({k})"),
));
}
let call = inject_node_call(&p.spawn, "e!(&#ident[I]))?;
let wrapper = format_ident!("spawn_{}", lname);
let mut defs: Vec<TokenStream2> = Vec::new();
defs.push(quote! {
#(#cfg)*
fn #wrapper<const I: usize>(
s: ::embassy_executor::Spawner,
) -> ::core::result::Result<(), ::embassy_executor::SpawnError> {
s.spawn(#call?);
::core::result::Result::Ok(())
}
});
let member_spawn: Vec<TokenStream2> = (0..k).map(|j| quote!(#wrapper::<#j>)).collect();
let members = p
.modes
.iter()
.zip(&member_spawn)
.enumerate()
.map(|(j, (mode, sp))| {
let nm = format!("{lname}{j}");
quote! {
#cr::TaskNode::new(
#nm, #cr::Mode::#mode,
::core::option::Option::Some((#sp) as #spawn_fn), false,
)
}
});
defs.push(quote! {
#(#cfg)*
pub static #ident: [#cr::TaskNode; #k] = [ #(#members),* ];
});
let member_refs = (0..k).map(|j| quote!(&#ident[#j]));
let policy = &p.policy;
let policy_ty = match &p.policy_ty {
Some(ty) => quote!(#ty),
None => {
let path = policy_type(policy)?;
quote!(#path)
}
};
let (min, max) = (&p.min, &p.max);
defs.push(quote! {
#(#cfg)*
pub static #pool_static: #cr::ElasticPool<#policy_ty> = #cr::ElasticPool {
nodes: &[ #(#member_refs),* ],
min: #min,
max: #max,
policy: #policy,
};
});
let pool_entry = quote!( #(#cfg)* &#pool_static );
let pred = cfg_predicate(cfg);
let slots = (0..k)
.map(|j| Slot {
cfg_pred: pred.clone(),
reference: quote!(&#ident[#j]),
deps: p.deps.clone(),
})
.collect();
Ok((defs, pool_entry, slots))
}
fn slot_tables(
slots: &[Slot],
names: &HashMap<String, usize>,
) -> SynResult<(Vec<TokenStream2>, Vec<TokenStream2>)> {
let mut all_entries: Vec<TokenStream2> = Vec::new();
let mut deps_entries: Vec<TokenStream2> = Vec::new();
for slot in slots {
let reference = &slot.reference;
all_entries.push(match &slot.cfg_pred {
None => quote!(::core::option::Option::Some(#reference)),
Some(pred) => quote!({
#[cfg(#pred)]
{ ::core::option::Option::Some(#reference) }
#[cfg(not(#pred))]
{ ::core::option::Option::None }
}),
});
let mut dep_toks: Vec<TokenStream2> = Vec::new();
for d in &slot.deps {
let idx = match names.get(&d.ident.to_string()) {
Some(&i) => i as u8,
None => {
return Err(syn::Error::new_spanned(
&d.ident,
format!("unknown dependency `{}` — not a declared node", d.ident),
));
}
};
let cfg = &d.cfg;
dep_toks.push(quote!( #(#cfg)* #idx ));
}
deps_entries.push(quote!( &[ #(#dep_toks),* ] ));
}
Ok((all_entries, deps_entries))
}
fn expand(graph: GraphSpec) -> SynResult<TokenStream2> {
let cr = quote!(::embassy_supervisor);
let spawn_fn = quote!(
fn(
::embassy_executor::Spawner,
) -> ::core::result::Result<(), ::embassy_executor::SpawnError>
);
let mut defs: Vec<TokenStream2> = Vec::new();
let mut pool_entries: Vec<TokenStream2> = Vec::new();
let mut slots: Vec<Slot> = Vec::new();
let mut names: HashMap<String, usize> = HashMap::new();
for item in &graph.items {
match item {
Item::Node(n) => {
names.insert(n.ident.to_string(), slots.len());
let (def, slot) = emit_node(n, &cr, &spawn_fn)?;
defs.push(def);
slots.push(slot);
}
Item::Pool(p) => {
if cfg!(feature = "pool") {
let (pool_defs, pool_entry, pool_slots) = emit_pool(p, &cr, &spawn_fn)?;
defs.extend(pool_defs);
pool_entries.push(pool_entry);
slots.extend(pool_slots);
} else {
return Err(syn::Error::new_spanned(
&p.ident,
"a `pool` requires enabling embassy-supervisor's `pool` feature",
));
}
}
}
}
let m = slots.len();
if m > 256 {
return Err(syn::Error::new(
proc_macro2::Span::call_site(),
format!(
"supervisor_graph!: {m} node slots declared, but at most 256 are supported \
(including pool members) — graph indices are `u8`"
),
));
}
let (all_entries, deps_entries) = slot_tables(&slots, &names)?;
let pools_field = if cfg!(feature = "pool") {
quote!( pools: &[ #(#pool_entries),* ], )
} else {
quote!()
};
Ok(quote! {
#(#defs)*
static NODES: [::core::option::Option<&'static #cr::TaskNode>; #m] = [ #(#all_entries),* ];
const DEPS: [&'static [u8]; #m] = [ #(#deps_entries),* ];
pub static GRAPH: #cr::Graph<#m> = #cr::Graph {
nodes: &NODES,
deps: &DEPS,
order: #cr::topo_sort_const(&DEPS),
#pools_field
};
})
}
#[proc_macro]
pub fn supervisor_graph(input: TokenStream) -> TokenStream {
let graph = syn::parse_macro_input!(input as GraphSpec);
expand(graph)
.unwrap_or_else(syn::Error::into_compile_error)
.into()
}