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);
syn::custom_keyword!(executor);
}
#[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,
executor: Option<Ident>,
}
struct ExecutorItem {
cfg: Vec<Attribute>,
ident: Ident,
}
struct PoolItem {
cfg: Vec<Attribute>,
ident: Ident,
modes: Vec<Ident>,
deps: Vec<Dep>,
spawn: Expr,
policy: Expr,
policy_ty: Option<Type>,
executor: Option<Ident>,
min: LitInt,
max: LitInt,
}
#[allow(clippy::large_enum_variant)]
enum Item {
Node(NodeItem),
Pool(PoolItem),
Executor(ExecutorItem),
}
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 if input.peek(kw::executor) {
input.parse::<kw::executor>()?;
let ident: Ident = input.parse()?;
input.parse::<Token![;]>()?;
items.push(Item::Executor(ExecutorItem { cfg, ident }));
} else {
return Err(input.error(
"expected `node`, `pool`, or `executor` (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;
let mut executor = None;
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 if input.peek(kw::executor) {
input.parse::<kw::executor>()?;
input.parse::<Token![:]>()?;
executor = Some(input.parse::<Ident>()?);
} else {
return Err(input.error("expected `spawn:`, `executor:`, or `disabled`"));
}
}
input.parse::<Token![;]>()?;
Ok(NodeItem {
cfg,
ident,
mode,
deps,
spawn,
disabled,
executor,
})
}
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![,]>()?;
let executor = if input.peek(kw::executor) {
input.parse::<kw::executor>()?;
input.parse::<Token![:]>()?;
let ex: Ident = input.parse()?;
input.parse::<Token![,]>()?;
Some(ex)
} else {
None
};
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,
executor,
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>,
executor: &Option<Ident>,
spawn_fn: &TokenStream2,
) -> SynResult<TokenStream2> {
Ok(match (spawn, executor) {
(None, None) => quote!(::core::option::Option::None),
(None, Some(ex)) => {
return Err(syn::Error::new_spanned(
ex,
"`executor:` requires a `spawn:` (a parked node is spawned by the \
application, which picks its own spawner)",
));
}
(Some(e @ (Expr::Path(_) | Expr::Call(_))), executor) => {
let call = inject_node_call(e, "e!(&#ident))?;
match executor {
None => {
let stmts = spawn_stmts(&call, "e!(&#ident), "e!(s));
quote!(::core::option::Option::Some(
(|s| {
#stmts
::core::result::Result::Ok(())
}) as #spawn_fn
))
}
Some(ex) => {
let stmts = spawn_stmts(&call, "e!(&#ident), "e!(__sp));
quote!(::core::option::Option::Some(
(|_s| {
let __sp = #ex
.get()
.ok_or(::embassy_executor::SpawnError::Busy)?;
#stmts
::core::result::Result::Ok(())
}) as #spawn_fn
))
}
}
}
(Some(_), Some(ex)) => {
return Err(syn::Error::new_spanned(
ex,
"`executor:` cannot be combined with a verbatim spawn closure (the \
closure owns the spawn; use the named SpawnerSlot inside it instead)",
));
}
(Some(e), None) => quote!(::core::option::Option::Some((#e) as #spawn_fn)),
})
}
fn spawn_stmts(call: &TokenStream2, node_ref: &TokenStream2, sp: &TokenStream2) -> TokenStream2 {
if cfg!(feature = "trace") {
quote! {
let __token = #call?;
(#node_ref).adopt(&__token);
#sp.spawn(__token);
}
} else {
quote!(#sp.spawn(#call?);)
}
}
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, &n.executor, spawn_fn)?;
let with_exec = match &n.executor {
Some(ex) => quote!( .with_executor(&#ex) ),
None => quote!(),
};
let def = quote! {
#(#cfg)*
pub static #ident: #cr::TaskNode =
#cr::TaskNode::new(#name, #cr::Mode::#mode, #spawn, #disabled) #with_exec;
};
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 (param, prelude, sp_tokens) = match &p.executor {
None => (quote!(s), quote!(), quote!(s)),
Some(ex) => (
quote!(_s),
quote! {
let __sp = #ex
.get()
.ok_or(::embassy_executor::SpawnError::Busy)?;
},
quote!(__sp),
),
};
let pool_spawn_stmts = spawn_stmts(&call, "e!(&#ident[I]), &sp_tokens);
let wrapper = format_ident!("spawn_{}", lname);
let mut defs: Vec<TokenStream2> = Vec::new();
defs.push(quote! {
#(#cfg)*
fn #wrapper<const I: usize>(
#param: ::embassy_executor::Spawner,
) -> ::core::result::Result<(), ::embassy_executor::SpawnError> {
#prelude
#pool_spawn_stmts
::core::result::Result::Ok(())
}
});
let member_spawn: Vec<TokenStream2> = (0..k).map(|j| quote!(#wrapper::<#j>)).collect();
let member_with_exec = match &p.executor {
Some(ex) => quote!( .with_executor(&#ex) ),
None => quote!(),
};
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,
) #member_with_exec
}
});
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)
}
};
defs.push(quote! {
#(#cfg)*
pub static #pool_static: #cr::ElasticPool<#policy_ty> = #cr::ElasticPool {
nodes: &[ #(#member_refs),* ],
min: #min_v,
max: #max_v,
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();
let mut seen: Vec<(u8, String)> = 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 or pool",
d.ident
),
));
}
};
let cfg = &d.cfg;
let cfg_key = quote!( #(#cfg)* ).to_string();
if seen.iter().any(|(i, k)| *i == idx && *k == cfg_key) {
return Err(syn::Error::new_spanned(
&d.ident,
format!("duplicate dependency `{}`", d.ident),
));
}
seen.push((idx, cfg_key));
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();
let executor_names: Vec<String> = graph
.items
.iter()
.filter_map(|i| match i {
Item::Executor(x) => Some(x.ident.to_string()),
_ => None,
})
.collect();
for item in &graph.items {
match item {
Item::Node(n) => {
if let Some(ex) = &n.executor
&& !executor_names.contains(&ex.to_string())
{
return Err(syn::Error::new_spanned(
ex,
format!(
"unknown executor `{ex}`; declare it in the graph with \
`executor {ex};` (declared: [{}])",
executor_names.join(", ")
),
));
}
if names.insert(n.ident.to_string(), slots.len()).is_some() {
return Err(syn::Error::new_spanned(
&n.ident,
format!("duplicate node/pool name `{}`", n.ident),
));
}
let (def, slot) = emit_node(n, &cr, &spawn_fn)?;
defs.push(def);
slots.push(slot);
}
Item::Executor(x) => {
let (cfg, ident) = (&x.cfg, &x.ident);
defs.push(quote! {
#(#cfg)*
pub static #ident: #cr::SpawnerSlot = #cr::SpawnerSlot::new();
});
}
Item::Pool(p) => {
if cfg!(feature = "pool") {
if let Some(ex) = &p.executor
&& !executor_names.contains(&ex.to_string())
{
return Err(syn::Error::new_spanned(
ex,
format!(
"unknown executor `{ex}`; declare it in the graph with \
`executor {ex};` (declared: [{}])",
executor_names.join(", ")
),
));
}
let (pool_defs, pool_entry, pool_slots) = emit_pool(p, &cr, &spawn_fn)?;
if names.insert(p.ident.to_string(), slots.len()).is_some() {
return Err(syn::Error::new_spanned(
&p.ident,
format!("duplicate node/pool name `{}`", p.ident),
));
}
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!()
};
let trace_hooks = if cfg!(feature = "trace-hooks") {
quote! {
#[unsafe(no_mangle)]
fn _embassy_trace_poll_start(executor_id: u32) {
#cr::trace::on_poll_start(executor_id);
}
#[unsafe(no_mangle)]
fn _embassy_trace_task_new(_executor_id: u32, _task_id: u32) {}
#[unsafe(no_mangle)]
fn _embassy_trace_task_end(executor_id: u32, task_id: u32) {
#cr::trace::on_task_end(executor_id, task_id);
}
#[unsafe(no_mangle)]
fn _embassy_trace_task_exec_begin(executor_id: u32, task_id: u32) {
#cr::trace::on_task_exec_begin(executor_id, task_id);
}
#[unsafe(no_mangle)]
fn _embassy_trace_task_exec_end(executor_id: u32, task_id: u32) {
#cr::trace::on_task_exec_end(executor_id, task_id);
}
#[unsafe(no_mangle)]
fn _embassy_trace_task_ready_begin(_executor_id: u32, _task_id: u32) {}
#[unsafe(no_mangle)]
fn _embassy_trace_executor_idle(executor_id: u32) {
#cr::trace::on_executor_idle(executor_id);
}
}
} 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
};
#trace_hooks
})
}
#[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()
}