use std::any::Any;
use std::future::Future;
use std::panic::{AssertUnwindSafe, resume_unwind};
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex, MutexGuard};
use std::task::{Context, Poll, Wake, Waker};
use std::thread::{self, Thread};
struct ThreadWaker {
thread: Thread,
notified: AtomicBool,
}
impl Wake for ThreadWaker {
fn wake(self: Arc<Self>) {
self.wake_by_ref();
}
fn wake_by_ref(self: &Arc<Self>) {
self.notified.store(true, Ordering::Release);
self.thread.unpark();
}
}
pub fn block_on<F: Future>(future: F) -> F::Output {
let mut future = std::pin::pin!(future);
let signal = Arc::new(ThreadWaker {
thread: thread::current(),
notified: AtomicBool::new(false),
});
let waker = Waker::from(Arc::clone(&signal));
let mut cx = Context::from_waker(&waker);
loop {
match future.as_mut().poll(&mut cx) {
Poll::Ready(val) => return val,
Poll::Pending => {
while !signal.notified.swap(false, Ordering::Acquire) {
thread::park();
}
}
}
}
}
pub async fn join_all<F: Future>(futures: impl IntoIterator<Item = F>) -> Vec<F::Output> {
join_all_boxed(futures.into_iter().map(Box::pin)).await
}
pub async fn join_all_boxed<F: Future + ?Sized>(
futures: impl IntoIterator<Item = Pin<Box<F>>>,
) -> Vec<F::Output> {
let mut pending: Vec<Option<Pin<Box<F>>>> = futures.into_iter().map(Some).collect();
let count = pending.len();
let mut done: Vec<Option<F::Output>> = (0..count).map(|_| None).collect();
std::future::poll_fn(move |cx| {
let mut remaining = 0usize;
for (slot, output) in pending.iter_mut().zip(done.iter_mut()) {
let Some(future) = slot.as_mut() else {
continue;
};
match future.as_mut().poll(cx) {
Poll::Ready(value) => {
*output = Some(value);
*slot = None;
}
Poll::Pending => remaining = remaining.saturating_add(1),
}
}
if remaining == 0 {
Poll::Ready(done.drain(..).flatten().collect())
} else {
Poll::Pending
}
})
.await
}
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
enum Job<T> {
Running,
Done(T),
Panicked(Box<dyn Any + Send>),
Taken,
}
struct Shared<T> {
job: Job<T>,
waker: Option<Waker>,
}
pub struct JoinHandle<T> {
shared: Arc<Mutex<Shared<T>>>,
}
impl<T> Future for JoinHandle<T> {
type Output = T;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<T> {
let mut state = lock(&self.shared);
match std::mem::replace(&mut state.job, Job::Taken) {
Job::Done(value) => Poll::Ready(value),
Job::Panicked(payload) => {
drop(state);
resume_unwind(payload)
}
Job::Taken => resume_unwind(Box::new(
"spawn_blocking JoinHandle polled after completion",
)),
Job::Running => {
state.job = Job::Running;
let replace = match state.waker.as_ref() {
Some(existing) => !existing.will_wake(cx.waker()),
None => true,
};
if replace {
state.waker = Some(cx.waker().clone());
}
Poll::Pending
}
}
}
}
pub fn spawn_blocking<F, T>(job: F) -> JoinHandle<T>
where
F: FnOnce() -> T + Send + 'static,
T: Send + 'static,
{
let shared = Arc::new(Mutex::new(Shared {
job: Job::Running,
waker: None,
}));
let worker = Arc::clone(&shared);
let spawned = thread::Builder::new()
.name("lgwks-blocking".into())
.spawn(move || {
let outcome = std::panic::catch_unwind(AssertUnwindSafe(job));
let mut state = lock(&worker);
state.job = match outcome {
Ok(value) => Job::Done(value),
Err(payload) => Job::Panicked(payload),
};
if let Some(waker) = state.waker.take() {
drop(state);
waker.wake();
}
});
if let Err(error) = spawned {
let mut state = lock(&shared);
state.job = Job::Panicked(Box::new(error));
}
JoinHandle { shared }
}
#[cfg(test)]
#[expect(
clippy::disallowed_methods,
reason = "tests of a thread-parking executor must spawn a waker thread and sleep to let the driver park"
)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
use std::time::Duration;
#[test]
fn immediate_future_returns_value() {
let result = block_on(async { 42 });
assert_eq!(result, 42);
}
#[test]
fn yields_and_resumes() {
async fn step() -> String {
let first_part = async { "hello" }.await;
let second_part = async { "world" }.await;
format!("{first_part} {second_part}")
}
assert_eq!(block_on(step()), "hello world");
}
struct DeferredValue {
value: Arc<Mutex<Option<i32>>>,
waker_slot: Arc<Mutex<Option<Waker>>>,
}
impl Future for DeferredValue {
type Output = i32;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
if let Some(value) = lock(&self.value).take() {
return Poll::Ready(value);
}
*lock(&self.waker_slot) = Some(cx.waker().clone());
Poll::Pending
}
}
#[test]
fn threaded_waker_unparks() {
let value: Arc<Mutex<Option<i32>>> = Arc::new(Mutex::new(None));
let waker_slot: Arc<Mutex<Option<Waker>>> = Arc::new(Mutex::new(None));
let sender_value = Arc::clone(&value);
let sender_slot = Arc::clone(&waker_slot);
thread::scope(|scope| {
scope.spawn(move || {
thread::sleep(Duration::from_millis(5));
*lock(&sender_value) = Some(100);
if let Some(waker) = lock(&sender_slot).take() {
waker.wake();
}
});
let result = block_on(DeferredValue { value, waker_slot });
assert_eq!(result, 100);
});
}
#[test]
fn join_all_empty_resolves_immediately() {
let output: Vec<u8> = block_on(join_all(std::iter::empty::<std::future::Ready<u8>>()));
assert!(output.is_empty());
}
#[test]
fn join_all_preserves_input_order_across_completion_order() {
let output = block_on(join_all(vec![
spawn_blocking(|| {
thread::sleep(Duration::from_millis(30));
1u32
}),
spawn_blocking(|| 2u32),
spawn_blocking(|| {
thread::sleep(Duration::from_millis(10));
3u32
}),
]));
assert_eq!(output, vec![1, 2, 3]);
}
#[test]
fn join_all_runs_children_concurrently() {
let barrier = Arc::new(std::sync::Barrier::new(2));
let handles: Vec<_> = (0..2)
.map(|_| {
let barrier = Arc::clone(&barrier);
spawn_blocking(move || {
barrier.wait();
7u32
})
})
.collect();
let output = block_on(join_all(handles));
assert_eq!(output, vec![7, 7]);
}
struct CountPolls {
polls: Arc<AtomicUsize>,
}
impl Future for CountPolls {
type Output = usize;
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<usize> {
let n = self
.polls
.fetch_add(1, AtomicOrdering::SeqCst)
.saturating_add(1);
Poll::Ready(n)
}
}
struct PendingThenReady {
polls: Arc<AtomicUsize>,
waker_slot: Arc<Mutex<Option<Waker>>>,
}
impl Future for PendingThenReady {
type Output = usize;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<usize> {
let n = self
.polls
.fetch_add(1, AtomicOrdering::SeqCst)
.saturating_add(1);
if n > 1 {
return Poll::Ready(n);
}
*lock(&self.waker_slot) = Some(cx.waker().clone());
let slot = Arc::clone(&self.waker_slot);
thread::spawn(move || {
thread::sleep(Duration::from_millis(5));
if let Some(waker) = lock(&slot).take() {
waker.wake();
}
});
Poll::Pending
}
}
#[test]
fn join_all_never_repolls_a_completed_future() {
let fast_polls = Arc::new(AtomicUsize::new(0));
let slow_polls = Arc::new(AtomicUsize::new(0));
let fast: Pin<Box<dyn Future<Output = usize>>> = Box::pin(CountPolls {
polls: Arc::clone(&fast_polls),
});
let slow: Pin<Box<dyn Future<Output = usize>>> = Box::pin(PendingThenReady {
polls: Arc::clone(&slow_polls),
waker_slot: Arc::new(Mutex::new(None)),
});
let output = block_on(join_all(vec![fast, slow]));
assert_eq!(output, vec![1, 2]);
assert_eq!(fast_polls.load(AtomicOrdering::SeqCst), 1);
assert_eq!(slow_polls.load(AtomicOrdering::SeqCst), 2);
}
#[test]
fn join_all_boxed_matches_join_all_and_accepts_boxes() {
let plain = block_on(join_all(vec![std::future::ready(1), std::future::ready(2)]));
let boxed: Vec<Pin<Box<dyn Future<Output = i32>>>> = vec![
Box::pin(std::future::ready(1)),
Box::pin(std::future::ready(2)),
];
assert_eq!(plain, vec![1, 2]);
assert_eq!(block_on(join_all_boxed(boxed)), vec![1, 2]);
let empty: Vec<Pin<Box<dyn Future<Output = i32>>>> = Vec::new();
assert_eq!(block_on(join_all_boxed(empty)), Vec::<i32>::new());
}
#[test]
fn spawn_blocking_returns_value() {
assert_eq!(block_on(spawn_blocking(|| 99u32)), 99);
}
#[test]
fn spawn_blocking_result_ready_before_first_poll() {
let handle = spawn_blocking(|| 1234u32);
thread::sleep(Duration::from_millis(30));
assert_eq!(block_on(handle), 1234);
}
#[test]
#[should_panic(expected = "worker exploded")]
fn spawn_blocking_panic_resumes_on_joiner() {
let _: u32 = block_on(spawn_blocking(|| -> u32 {
assert_eq!(1, 2, "worker exploded");
0u32
}));
}
}