use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
#[cfg(not(target_family = "wasm"))]
use std::sync::mpsc;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};
#[cfg(not(target_family = "wasm"))]
use std::thread::JoinHandle;
use std::time::{Duration, Instant};
use crate::Error;
pub const WAITER_SLICE: Duration = Duration::from_millis(10);
#[derive(Debug)]
pub enum SliceOutcome {
Signaled,
TimedOut,
Failed(Error),
}
#[cfg(not(target_family = "wasm"))]
pub type SliceFn = Box<dyn FnMut(Duration) -> SliceOutcome + Send + 'static>;
#[cfg(target_family = "wasm")]
pub type SliceFn = Box<dyn FnMut(Duration) -> SliceOutcome + 'static>;
struct WaitCompletion {
result: Option<Result<(), Error>>,
waker: Option<Waker>,
}
#[cfg(not(target_family = "wasm"))]
struct WaitRequest {
slice_fn: SliceFn,
deadline: Option<Instant>,
cancelled: Arc<AtomicBool>,
completion: Arc<Mutex<WaitCompletion>>,
}
pub struct WaiterThread {
#[cfg(not(target_family = "wasm"))]
sender: Option<mpsc::Sender<WaitRequest>>,
#[cfg(not(target_family = "wasm"))]
join: Option<JoinHandle<()>>,
shutdown: Arc<AtomicBool>,
}
impl std::fmt::Debug for WaiterThread {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
#[cfg(not(target_family = "wasm"))]
let alive = self.sender.is_some();
#[cfg(target_family = "wasm")]
let alive = false;
f.debug_struct("WaiterThread").field("alive", &alive).finish()
}
}
impl WaiterThread {
pub fn new(name: &str) -> Self {
#[cfg(target_family = "wasm")]
{
let _ = name;
Self { shutdown: Arc::new(AtomicBool::new(false)) }
}
#[cfg(not(target_family = "wasm"))]
{
let (tx, rx) = mpsc::channel::<WaitRequest>();
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_thread = shutdown.clone();
let join = std::thread::Builder::new()
.name(format!("gpu-wait-{name}"))
.spawn(move || waiter_loop(&rx, &shutdown_thread))
.expect("gpu-handle-types: failed to spawn waiter thread");
Self { sender: Some(tx), join: Some(join), shutdown }
}
}
#[cfg_attr(target_family = "wasm", allow(unused_variables, clippy::needless_pass_by_value))]
pub fn enqueue(&self, slice_fn: SliceFn, deadline: Option<Instant>) -> BackendWaitFuture {
let cancelled = Arc::new(AtomicBool::new(false));
let completion = Arc::new(Mutex::new(WaitCompletion { result: None, waker: None }));
#[cfg(target_family = "wasm")]
finish_completion(
&completion,
Err(Error::NotSupported(
"blocking waits are unavailable on wasm (no waiter thread); drive the async wait \
path instead"
.into(),
)),
);
#[cfg(not(target_family = "wasm"))]
{
let req = WaitRequest { slice_fn, deadline, cancelled: cancelled.clone(), completion: completion.clone() };
match self.sender.as_ref() {
Some(s) => {
if s.send(req).is_err() {
finish_completion(
&completion,
Err(Error::NotSupported("waiter thread died before request was accepted".into())),
);
}
}
None => {
finish_completion(&completion, Err(Error::NotSupported("waiter thread already shut down".into())));
}
}
}
BackendWaitFuture { completion, cancelled }
}
}
impl Drop for WaiterThread {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::Release);
#[cfg(not(target_family = "wasm"))]
{
self.sender.take();
if let Some(h) = self.join.take() {
let _ = h.join();
}
}
}
}
#[cfg(not(target_family = "wasm"))]
fn waiter_loop(rx: &mpsc::Receiver<WaitRequest>, shutdown: &AtomicBool) {
while let Ok(mut req) = rx.recv() {
if shutdown.load(Ordering::Acquire) {
abandon_on_shutdown(&req);
break;
}
let mut shutting_down = false;
loop {
if shutdown.load(Ordering::Acquire) {
abandon_on_shutdown(&req);
shutting_down = true;
break;
}
if req.cancelled.load(Ordering::Acquire) {
break;
}
let remaining = match req.deadline {
Some(d) => match d.checked_duration_since(Instant::now()) {
Some(r) => r,
None => {
finish_completion(&req.completion, Err(Error::Timeout));
break;
}
},
None => Duration::MAX,
};
let slice = std::cmp::min(WAITER_SLICE, remaining);
match (req.slice_fn)(slice) {
SliceOutcome::Signaled => {
finish_completion(&req.completion, Ok(()));
break;
}
SliceOutcome::TimedOut => continue,
SliceOutcome::Failed(e) => {
finish_completion(&req.completion, Err(e));
break;
}
}
}
if shutting_down {
break;
}
}
while let Ok(req) = rx.try_recv() {
abandon_on_shutdown(&req);
}
}
#[cfg(not(target_family = "wasm"))]
fn abandon_on_shutdown(req: &WaitRequest) {
finish_completion(&req.completion, Err(Error::Cancelled));
}
fn finish_completion(completion: &Arc<Mutex<WaitCompletion>>, result: Result<(), Error>) {
let mut c = completion.lock().expect("WaitCompletion mutex poisoned");
if c.result.is_none() {
c.result = Some(result);
}
let waker = c.waker.take();
drop(c);
if let Some(w) = waker {
w.wake();
}
}
pub struct BackendWaitFuture {
completion: Arc<Mutex<WaitCompletion>>,
cancelled: Arc<AtomicBool>,
}
impl std::fmt::Debug for BackendWaitFuture {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BackendWaitFuture")
.field("cancelled", &self.cancelled.load(Ordering::Relaxed))
.field("resolved", &self.completion.lock().map(|c| c.result.is_some()).unwrap_or(false))
.finish()
}
}
impl Future for BackendWaitFuture {
type Output = Result<(), Error>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut c = self.completion.lock().expect("WaitCompletion mutex poisoned");
if let Some(r) = c.result.take() {
return Poll::Ready(r);
}
c.waker = Some(cx.waker().clone());
Poll::Pending
}
}
impl Drop for BackendWaitFuture {
fn drop(&mut self) {
self.cancelled.store(true, Ordering::Release);
}
}
pub async fn run_hybrid_wait<F, M>(
is_signaled: F,
waiter_thread: &WaiterThread,
timeout: Duration,
make_slice_fn: M,
) -> Result<(), Error>
where
F: Fn() -> Result<bool, Error> + crate::MaybeSend,
M: FnOnce() -> SliceFn + crate::MaybeSend,
{
const SPIN_ITERATIONS: usize = 64;
let deadline = crate::wait_deadline(timeout);
for _ in 0..SPIN_ITERATIONS {
match is_signaled() {
Ok(true) => return Ok(()),
Ok(false) => {}
Err(e) => return Err(e),
}
if let Some(d) = deadline
&& Instant::now() >= d
{
return Err(Error::Timeout);
}
crate::yield_once().await;
}
let slice_fn = make_slice_fn();
waiter_thread.enqueue(slice_fn, deadline).await
}