use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
use syn::{
Error, Expr, Pat, Token, braced,
parse::{Parse, ParseStream},
parse_macro_input,
};
struct SelectBranch {
pat: Pat,
future: Expr,
handler: Expr,
}
struct SelectInput {
cx: Expr,
biased: bool,
branches: Vec<SelectBranch>,
els: Option<Expr>,
}
impl Parse for SelectInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
if input.is_empty() || input.peek(syn::token::Brace) {
return Err(Error::new(input.span(), "select! requires cx argument"));
}
let cx: Expr = input.parse()?;
let _comma: Token![,] = input.parse().map_err(|_| {
Error::new(
input.span(),
"expected comma after cx: select!(cx, { ... })",
)
})?;
let mut biased = false;
if input.peek(syn::Ident) {
let ident: syn::Ident = input.fork().parse()?;
if ident == "biased" {
let _: syn::Ident = input.parse()?;
let _comma: Token![,] = input.parse().map_err(|_| {
Error::new(
input.span(),
"expected comma after biased: select!(cx, biased, { ... })",
)
})?;
biased = true;
}
}
let content;
let _brace = braced!(content in input);
let mut branches = Vec::new();
let mut els = None;
while !content.is_empty() {
if content.peek(Token![else]) {
let _else: Token![else] = content.parse()?;
let _arrow: Token![=>] = content.parse().map_err(|_| {
Error::new(
content.span(),
"expected `=>` after else: `else => fallback`",
)
})?;
els = Some(content.parse()?);
if content.peek(Token![,]) {
let _comma: Token![,] = content.parse()?;
}
if !content.is_empty() {
return Err(Error::new(
content.span(),
"select! else arm must be the last arm",
));
}
break;
}
let pat = Pat::parse_single(&content)?;
let _eq: Token![=] = content.parse().map_err(|_| {
Error::new(
content.span(),
"expected `=` in select branch: `binding = future => handler`",
)
})?;
let future: Expr = content.parse()?;
let _arrow: Token![=>] = content.parse().map_err(|_| {
Error::new(
content.span(),
"expected `=>` in select branch: `binding = future => handler`",
)
})?;
let handler: Expr = content.parse()?;
branches.push(SelectBranch {
pat,
future,
handler,
});
if content.peek(Token![,]) {
let _comma: Token![,] = content.parse()?;
}
}
if branches.is_empty() {
return Err(Error::new(
input.span(),
"select! requires at least one branch",
));
}
if !input.is_empty() {
return Err(Error::new(
input.span(),
"unexpected tokens after select! branches",
));
}
Ok(Self {
cx,
biased,
branches,
els,
})
}
}
pub fn select_impl(input: TokenStream) -> TokenStream {
let parsed = parse_macro_input!(input as SelectInput);
TokenStream::from(generate_select(&parsed))
}
fn branch_future(branch: &SelectBranch) -> TokenStream2 {
let pat = &branch.pat;
let fut = &branch.future;
let handler = &branch.handler;
quote! {
async move {
let #pat = (#fut).await;
#handler
}
}
}
fn generate_select(input: &SelectInput) -> TokenStream2 {
let SelectInput {
cx, branches, els, ..
} = input;
let _ = input.biased;
let branch_futs: Vec<TokenStream2> = branches.iter().map(branch_future).collect();
if els.is_none() {
let boxed: Vec<TokenStream2> = branch_futs
.iter()
.map(|fut| quote! { ::std::boxed::Box::pin(#fut) })
.collect();
quote! {
{
(#cx).race_drained(::std::vec![#(#boxed),*]).await
}
}
} else {
let else_handler = els.as_ref().expect("checked select! else arm exists");
let fut_idents: Vec<_> = (0..branches.len())
.map(|i| syn::Ident::new(&format!("__select_fut_{i}"), proc_macro2::Span::call_site()))
.collect();
let bindings: Vec<TokenStream2> = fut_idents
.iter()
.zip(branch_futs.iter())
.map(|(id, fut)| quote! { let mut #id = ::core::pin::pin!(#fut); })
.collect();
let polls: Vec<TokenStream2> = fut_idents
.iter()
.map(|id| {
quote! {
match ::core::future::Future::poll(
::core::pin::Pin::as_mut(&mut #id),
__select_cx,
) {
::core::task::Poll::Ready(__select_value) => {
return ::core::task::Poll::Ready(__select_value);
}
::core::task::Poll::Pending => {}
}
}
})
.collect();
quote! {
{
let _ = &(#cx);
#(#bindings)*
let mut __select_else = ::core::option::Option::Some(move || #else_handler);
::core::future::poll_fn(|__select_cx| {
#(#polls)*
let __select_default = ::core::option::Option::take(&mut __select_else)
.expect("select! else arm polled after completion");
::core::task::Poll::Ready(__select_default())
})
.await
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn parse_err(input: TokenStream2) -> String {
match syn::parse2::<SelectInput>(input) {
Ok(_) => panic!("expected select! parse error, but parsing succeeded"),
Err(err) => err.to_string(),
}
}
#[test]
fn parse_basic_two_branch() {
let input = quote! { cx, { a = fut_a() => a, b = fut_b() => b } };
let parsed: SelectInput = syn::parse2(input).unwrap();
assert!(!parsed.biased);
assert_eq!(parsed.branches.len(), 2);
assert!(parsed.els.is_none());
}
#[test]
fn parse_single_branch_is_allowed() {
let input = quote! { cx, { a = fut_a() => a } };
let parsed: SelectInput = syn::parse2(input).unwrap();
assert_eq!(parsed.branches.len(), 1);
}
#[test]
fn parse_biased_flag() {
let input = quote! { cx, biased, { a = fut_a() => a, b = fut_b() => b } };
let parsed: SelectInput = syn::parse2(input).unwrap();
assert!(parsed.biased);
assert_eq!(parsed.branches.len(), 2);
}
#[test]
fn parse_else_arm() {
let input = quote! { cx, { a = fut_a() => a, else => fallback() } };
let parsed: SelectInput = syn::parse2(input).unwrap();
assert_eq!(parsed.branches.len(), 1);
assert!(parsed.els.is_some());
}
#[test]
fn parse_patterns_in_binding() {
let input = quote! { cx, { (x, y) = fut_a() => x + y, _ = fut_b() => 0 } };
let parsed: SelectInput = syn::parse2(input).unwrap();
assert_eq!(parsed.branches.len(), 2);
}
#[test]
fn missing_cx_is_rejected() {
let err = parse_err(quote! { { a = fut_a() => a } });
assert!(err.contains("requires cx argument"), "got: {err}");
}
#[test]
fn empty_branches_are_rejected() {
let err = parse_err(quote! { cx, { } });
assert!(err.contains("at least one branch"), "got: {err}");
}
#[test]
fn two_else_arms_are_rejected() {
let err = parse_err(quote! { cx, { a = f() => a, else => 1, else => 2 } });
assert!(err.contains("else arm must be the last arm"), "got: {err}");
}
#[test]
fn else_must_be_last() {
let err = parse_err(quote! { cx, { a = f() => a, else => 1, b = g() => b } });
assert!(err.contains("else arm must be the last arm"), "got: {err}");
}
#[test]
fn blocking_form_routes_through_race_drained() {
let parsed: SelectInput =
syn::parse2(quote! { cx, { a = fut_a() => a, b = fut_b() => b } }).unwrap();
let tokens = generate_select(&parsed).to_string();
assert!(
tokens.contains("race_drained"),
"blocking select! must use the drained engine, got: {tokens}"
);
assert!(
!tokens.replace("race_drained", "").contains("race"),
"no drop-only race* call may survive, got: {tokens}"
);
assert!(
!tokens.contains("poll_fn"),
"blocking select! must not poll inline, got: {tokens}"
);
}
#[test]
fn blocking_form_awaits_future_and_runs_handler() {
let parsed: SelectInput =
syn::parse2(quote! { cx, { x = fetch() => process(x) } }).unwrap();
let tokens = generate_select(&parsed).to_string();
assert!(tokens.contains("let x"), "binding must be bound: {tokens}");
assert!(tokens.contains("fetch"), "future must be awaited: {tokens}");
assert!(
tokens.contains(". await"),
"branch future must be awaited: {tokens}"
);
assert!(
tokens.contains("process"),
"handler must run after await: {tokens}"
);
}
#[test]
fn else_form_polls_inline_and_skips_race_drained() {
let parsed: SelectInput =
syn::parse2(quote! { cx, { a = fut_a() => a, else => fallback() } }).unwrap();
let tokens = generate_select(&parsed).to_string();
assert!(
tokens.contains("poll_fn"),
"else select! must poll inline via poll_fn, got: {tokens}"
);
assert!(
!tokens.contains("race_drained"),
"else select! is non-blocking and must not spawn/drain, got: {tokens}"
);
assert!(
tokens.contains("fallback"),
"else handler must appear in the expansion, got: {tokens}"
);
assert!(
tokens.contains("Pin :: as_mut") || tokens.contains("Pin::as_mut"),
"else select! must re-poll pinned branches via Pin::as_mut, got: {tokens}"
);
}
#[test]
fn biased_blocking_form_still_drains() {
let parsed: SelectInput =
syn::parse2(quote! { cx, biased, { a = fut_a() => a, b = fut_b() => b } }).unwrap();
assert!(parsed.biased);
let tokens = generate_select(&parsed).to_string();
assert!(
tokens.contains("race_drained"),
"biased select! must still use the drained engine, got: {tokens}"
);
}
}