use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
use syn::{
Expr, Token,
parse::{Parse, ParseStream},
parse_macro_input,
punctuated::Punctuated,
};
struct JoinInput {
cx: Option<Expr>,
futures: Punctuated<Expr, Token![,]>,
}
impl Parse for JoinInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let fork = input.fork();
let cx = if let Ok(cx_expr) = fork.parse::<Expr>() {
if fork.peek(Token![;]) {
let _ = input.parse::<Expr>()?;
let _semi: Token![;] = input.parse()?;
Some(cx_expr)
} else {
None
}
} else {
None
};
let futures = Punctuated::parse_terminated(input)?;
Ok(Self { cx, futures })
}
}
pub fn join_impl(input: TokenStream) -> TokenStream {
let JoinInput { cx, futures } = parse_macro_input!(input as JoinInput);
let expanded = generate_join(cx.as_ref(), &futures);
TokenStream::from(expanded)
}
fn generate_join(cx: Option<&Expr>, futures: &Punctuated<Expr, Token![,]>) -> TokenStream2 {
let future_count = futures.len();
let cx_ack = generate_cx_ack(cx);
if future_count == 0 {
return quote! {
{
#cx_ack
()
}
};
}
if future_count == 1 {
let fut = futures
.first()
.expect("future_count == 1 guarantees first element exists");
return quote! {
{
#cx_ack
(#fut.await,)
}
};
}
let (bindings, decls, polls, idents) = concurrent_branch_tokens(futures);
quote! {
{
#cx_ack
#(#bindings)*
#(#decls)*
::core::future::poll_fn(|__join_cx| {
let mut __join_pending = false;
#(#polls)*
if __join_pending {
::core::task::Poll::Pending
} else {
::core::task::Poll::Ready(())
}
})
.await;
( #(#idents.expect("join! branch completed before the join resolved")),* )
}
}
}
pub fn join_all_impl(input: TokenStream) -> TokenStream {
let JoinInput { cx, futures } = parse_macro_input!(input as JoinInput);
let expanded = generate_join_all(cx.as_ref(), &futures);
TokenStream::from(expanded)
}
fn generate_join_all(cx: Option<&Expr>, futures: &Punctuated<Expr, Token![,]>) -> TokenStream2 {
let future_count = futures.len();
let cx_ack = generate_cx_ack(cx);
if future_count == 0 {
return quote! {
{
#cx_ack
[]
}
};
}
if future_count == 1 {
let fut = futures
.first()
.expect("future_count == 1 guarantees first element exists");
return quote! {
{
#cx_ack
[#fut.await]
}
};
}
let (bindings, decls, polls, idents) = concurrent_branch_tokens(futures);
quote! {
{
#cx_ack
#(#bindings)*
#(#decls)*
::core::future::poll_fn(|__join_cx| {
let mut __join_pending = false;
#(#polls)*
if __join_pending {
::core::task::Poll::Pending
} else {
::core::task::Poll::Ready(())
}
})
.await;
[ #(#idents.expect("join_all! branch completed before the join resolved")),* ]
}
}
}
fn concurrent_branch_tokens(
futures: &Punctuated<Expr, Token![,]>,
) -> (
Vec<TokenStream2>,
Vec<TokenStream2>,
Vec<TokenStream2>,
Vec<syn::Ident>,
) {
let future_count = futures.len();
let fut_idents: Vec<_> = (0..future_count)
.map(|i| syn::Ident::new(&format!("__join_fut_{i}"), proc_macro2::Span::call_site()))
.collect();
let out_idents: Vec<syn::Ident> = (0..future_count)
.map(|i| syn::Ident::new(&format!("__join_out_{i}"), proc_macro2::Span::call_site()))
.collect();
let bindings = futures
.iter()
.zip(fut_idents.iter())
.map(|(future, ident)| quote! { let mut #ident = ::core::pin::pin!(#future); })
.collect();
let decls = out_idents
.iter()
.map(|ident| quote! { let mut #ident = ::core::option::Option::None; })
.collect();
let polls = fut_idents
.iter()
.zip(out_idents.iter())
.map(|(fut_ident, out_ident)| {
quote! {
if ::core::option::Option::is_none(&#out_ident) {
match ::core::future::Future::poll(
::core::pin::Pin::as_mut(&mut #fut_ident),
__join_cx,
) {
::core::task::Poll::Ready(__join_value) => {
#out_ident = ::core::option::Option::Some(__join_value);
}
::core::task::Poll::Pending => {
__join_pending = true;
}
}
}
}
})
.collect();
(bindings, decls, polls, out_idents)
}
fn generate_cx_ack(cx: Option<&Expr>) -> TokenStream2 {
if cx.is_some() {
quote! {
let _ = &#cx;
}
} else {
quote! {}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_single_future() {
let input: proc_macro2::TokenStream = quote! { future_a };
let parsed: JoinInput = syn::parse2(input).unwrap();
assert_eq!(parsed.futures.len(), 1);
}
#[test]
fn test_parse_multiple_futures() {
let input: proc_macro2::TokenStream = quote! { future_a, future_b, future_c };
let parsed: JoinInput = syn::parse2(input).unwrap();
assert_eq!(parsed.futures.len(), 3);
}
#[test]
fn test_parse_trailing_comma() {
let input: proc_macro2::TokenStream = quote! { future_a, future_b, };
let parsed: JoinInput = syn::parse2(input).unwrap();
assert_eq!(parsed.futures.len(), 2);
}
#[test]
fn test_parse_with_cx() {
let input: proc_macro2::TokenStream = quote! { cx; future_a, future_b };
let parsed: JoinInput = syn::parse2(input).unwrap();
assert!(parsed.cx.is_some());
assert_eq!(parsed.futures.len(), 2);
}
#[test]
fn test_join_single_future_keeps_cx_expr() {
let input: JoinInput = syn::parse2(quote! { make_cx(); future_a }).unwrap();
let tokens = generate_join(input.cx.as_ref(), &input.futures).to_string();
assert!(
tokens.contains("make_cx"),
"single-future join must still typecheck the cx expression"
);
}
#[test]
fn test_join_empty_keeps_cx_expr() {
let input: JoinInput = syn::parse2(quote! { make_cx(); }).unwrap();
let tokens = generate_join(input.cx.as_ref(), &input.futures).to_string();
assert!(
tokens.contains("make_cx"),
"empty join must still typecheck the cx expression"
);
}
#[test]
fn test_join_all_single_future_keeps_cx_expr() {
let input: JoinInput = syn::parse2(quote! { make_cx(); future_a }).unwrap();
let tokens = generate_join_all(input.cx.as_ref(), &input.futures).to_string();
assert!(
tokens.contains("make_cx"),
"single-future join_all must still typecheck the cx expression"
);
}
#[test]
fn test_join_all_empty_keeps_cx_expr() {
let input: JoinInput = syn::parse2(quote! { make_cx(); }).unwrap();
let tokens = generate_join_all(input.cx.as_ref(), &input.futures).to_string();
assert!(
tokens.contains("make_cx"),
"empty join_all must still typecheck the cx expression"
);
}
#[test]
fn join_multi_polls_branches_concurrently() {
let input: JoinInput = syn::parse2(quote! { a, b, c }).unwrap();
let tokens = generate_join(input.cx.as_ref(), &input.futures).to_string();
assert!(
tokens.contains("poll_fn"),
"multi-branch join! must drive a concurrent poll_fn, not sequential awaits"
);
assert!(
tokens.contains("Future :: poll") || tokens.contains("Future::poll"),
"multi-branch join! must poll each branch directly"
);
assert!(
!tokens.contains("__join_result_"),
"join! must not fall back to the sequential await chain"
);
}
#[test]
fn join_all_multi_polls_branches_concurrently() {
let input: JoinInput = syn::parse2(quote! { a, b, c }).unwrap();
let tokens = generate_join_all(input.cx.as_ref(), &input.futures).to_string();
assert!(
tokens.contains("poll_fn"),
"multi-branch join_all! must drive a concurrent poll_fn"
);
assert!(
tokens.contains("Pin :: as_mut") || tokens.contains("Pin::as_mut"),
"concurrent join_all! must re-poll pinned branches via Pin::as_mut"
);
}
#[test]
fn join_multi_pins_each_branch_once() {
let input: JoinInput = syn::parse2(quote! { a, b }).unwrap();
let tokens = generate_join(input.cx.as_ref(), &input.futures).to_string();
let pins = tokens.matches("pin !").count() + tokens.matches("pin!").count();
assert!(
pins >= 2,
"each branch must be pinned exactly once: {tokens}"
);
}
}