use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc;
use std::task::{Context, Poll, Waker};
use std::thread::ThreadId;
use std::time::{Duration, Instant};
use gpu_handle_types::{BackendKind, Error, SliceFn, SliceOutcome, WaiterThread};
const HANG_BUDGET: Duration = Duration::from_secs(30);
fn poll_once<F: Future>(fut: F) -> Option<F::Output> {
let fut = std::pin::pin!(fut);
match fut.poll(&mut Context::from_waker(Waker::noop())) {
Poll::Ready(v) => Some(v),
Poll::Pending => None,
}
}
fn parking_slice() -> (SliceFn, mpsc::Receiver<()>, mpsc::Sender<()>) {
let (entered_tx, entered_rx) = mpsc::channel::<()>();
let (release_tx, release_rx) = mpsc::channel::<()>();
let slice_fn: SliceFn = Box::new(move |_slice| {
let _ = entered_tx.send(());
let _ = release_rx.recv();
SliceOutcome::TimedOut
});
(slice_fn, entered_rx, release_tx)
}
fn park_once_then_free_run() -> (SliceFn, mpsc::Receiver<()>, mpsc::Sender<()>) {
let (entered_tx, entered_rx) = mpsc::channel::<()>();
let (release_tx, release_rx) = mpsc::channel::<()>();
let mut park = Some((entered_tx, release_rx));
let slice_fn: SliceFn = Box::new(move |_slice| {
match park.take() {
Some((entered_tx, release_rx)) => {
let _ = entered_tx.send(());
let _ = release_rx.recv();
}
None => std::thread::yield_now(),
}
SliceOutcome::TimedOut
});
(slice_fn, entered_rx, release_tx)
}
#[test]
fn signaled_outcome_resolves_future() {
let thread = WaiterThread::new("test-signal");
let probes = Arc::new(AtomicUsize::new(0));
let probes_c = probes.clone();
let slice_fn: SliceFn = Box::new(move |_slice| {
let n = probes_c.fetch_add(1, Ordering::SeqCst);
if n >= 2 { SliceOutcome::Signaled } else { SliceOutcome::TimedOut }
});
let fut = thread.enqueue(slice_fn, None);
let result = pollster::block_on(fut);
assert!(result.is_ok(), "expected Ok, got {result:?}");
assert!(probes.load(Ordering::SeqCst) >= 3);
}
#[test]
fn deadline_expiry_surfaces_timeout() {
let thread = WaiterThread::new("test-timeout");
let slice_fn: SliceFn = Box::new(|_| SliceOutcome::TimedOut);
let deadline = Instant::now().checked_add(Duration::from_millis(50));
let fut = thread.enqueue(slice_fn, deadline);
let result = pollster::block_on(fut);
assert!(matches!(result, Err(Error::Timeout)), "expected Err(Timeout), got {result:?}");
}
#[test]
fn failed_outcome_surfaces_verbatim() {
let thread = WaiterThread::new("test-fail");
let slice_fn: SliceFn = Box::new(|_| SliceOutcome::Failed(Error::DeviceLost { backend: BackendKind::Vulkan }));
let fut = thread.enqueue(slice_fn, None);
let result = pollster::block_on(fut);
match result {
Err(Error::DeviceLost { backend: BackendKind::Vulkan }) => {}
other => panic!("expected Err(DeviceLost {{ Vulkan }}), got {other:?}"),
}
}
#[test]
fn cancellation_observed_within_one_slice() {
let thread = WaiterThread::new("test-cancel");
let (slice_fn, entered_rx, release_tx) = parking_slice();
let fut = thread.enqueue(slice_fn, None);
entered_rx.recv_timeout(HANG_BUDGET).expect("waiter thread never entered its first slice");
drop(fut); release_tx.send(()).expect("waiter thread must still be parked inside slice 1");
assert!(
entered_rx.recv_timeout(HANG_BUDGET).is_err(),
"a slice was issued after the future was dropped; cancellation is not observed at the slice boundary",
);
let done = Arc::new(AtomicUsize::new(0));
let done_c = done.clone();
let slice_fn: SliceFn = Box::new(move |_| {
done_c.fetch_add(1, Ordering::SeqCst);
SliceOutcome::Signaled
});
let result = pollster::block_on(thread.enqueue(slice_fn, None));
assert!(result.is_ok(), "follow-up enqueue must succeed; cancellation broke the loop?");
assert_eq!(done.load(Ordering::SeqCst), 1);
}
#[test]
fn drop_joins_cleanly_when_idle() {
let thread = WaiterThread::new("test-shutdown-idle");
drop(thread); }
#[test]
fn drop_joins_cleanly_with_in_flight_request() {
let thread = WaiterThread::new("test-shutdown-inflight");
let (slice_fn, entered_rx, release_tx) = park_once_then_free_run();
let fut = thread.enqueue(slice_fn, None);
entered_rx.recv_timeout(HANG_BUDGET).expect("waiter thread never entered its first slice");
let (joined_tx, joined_rx) = mpsc::channel::<()>();
let dropper = std::thread::spawn(move || {
drop(thread); let _ = joined_tx.send(());
});
assert!(
joined_rx.recv_timeout(Duration::from_millis(200)).is_err(),
"WaiterThread::drop returned while the waiter thread was still inside a slice — it must join, not detach",
);
release_tx.send(()).expect("waiter thread must still be parked inside slice 1");
joined_rx
.recv_timeout(HANG_BUDGET)
.expect("WaiterThread::drop deadlocked — the slice loop never observed the shutdown flag");
dropper.join().expect("dropper thread panicked");
match poll_once(fut) {
Some(Err(Error::Cancelled)) => {}
other => panic!("an in-flight future must resolve to Err(Cancelled) on shutdown, got {other:?}"),
}
}
#[test]
fn drop_resolves_a_queued_request_instead_of_hanging() {
let thread = WaiterThread::new("test-shutdown-queued");
let (slice_fn, entered_rx, release_tx) = park_once_then_free_run();
let busy = thread.enqueue(slice_fn, None);
entered_rx.recv_timeout(HANG_BUDGET).expect("waiter thread never entered its first slice");
let queued = thread.enqueue(Box::new(|_| SliceOutcome::Signaled), None);
let (joined_tx, joined_rx) = mpsc::channel::<()>();
let dropper = std::thread::spawn(move || {
drop(thread);
let _ = joined_tx.send(());
});
release_tx.send(()).expect("waiter thread must still be parked inside slice 1");
joined_rx.recv_timeout(HANG_BUDGET).expect("WaiterThread::drop deadlocked");
dropper.join().expect("dropper thread panicked");
match poll_once(queued) {
Some(Err(Error::Cancelled)) => {}
other => panic!("a queued request's future must resolve to Err(Cancelled) on shutdown, got {other:?}"),
}
match poll_once(busy) {
Some(Err(Error::Cancelled)) => {}
other => panic!("the in-flight request's future must resolve to Err(Cancelled) on shutdown, got {other:?}"),
}
}
#[test]
fn requests_share_a_single_thread() {
let thread = WaiterThread::new("test-reuse");
let observed: Arc<std::sync::Mutex<Vec<ThreadId>>> = Arc::new(std::sync::Mutex::new(Vec::new()));
for _ in 0..5 {
let observed_c = observed.clone();
let slice_fn: SliceFn = Box::new(move |_| {
observed_c.lock().unwrap().push(std::thread::current().id());
SliceOutcome::Signaled
});
pollster::block_on(thread.enqueue(slice_fn, None)).expect("Signaled outcome should resolve to Ok");
}
let ids = observed.lock().unwrap();
assert_eq!(ids.len(), 5);
let first = ids[0];
for (i, id) in ids.iter().enumerate().skip(1) {
assert_eq!(*id, first, "slice {i} ran on a different ThreadId; thread respawned?");
}
}
#[test]
fn distinct_instances_run_on_distinct_threads() {
let a = WaiterThread::new("test-distinct-a");
let b = WaiterThread::new("test-distinct-b");
let (tid_a, tid_b) = {
let cap_a: Arc<std::sync::Mutex<Option<ThreadId>>> = Arc::new(std::sync::Mutex::new(None));
let cap_b = cap_a.clone();
let cap_a_inner = cap_a.clone();
let slice_fn_a: SliceFn = Box::new(move |_| {
*cap_a_inner.lock().unwrap() = Some(std::thread::current().id());
SliceOutcome::Signaled
});
pollster::block_on(a.enqueue(slice_fn_a, None)).unwrap();
let tid_a = cap_a.lock().unwrap().take().unwrap();
let cap_b_inner = cap_b.clone();
let slice_fn_b: SliceFn = Box::new(move |_| {
*cap_b_inner.lock().unwrap() = Some(std::thread::current().id());
SliceOutcome::Signaled
});
pollster::block_on(b.enqueue(slice_fn_b, None)).unwrap();
let tid_b = cap_b.lock().unwrap().take().unwrap();
(tid_a, tid_b)
};
assert_ne!(tid_a, tid_b, "two WaiterThread instances must spawn distinct OS threads");
}