use core::future::Future;
use core::pin::Pin;
use std::any::Any;
use std::ffi::c_void;
use std::sync::Arc;
use std::time::Duration;
#[cfg(feature = "wgpu")]
use crate::SliceOutcome;
use crate::{BackendKind, Error, GlBackend};
#[cfg(not(target_family = "wasm"))]
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[cfg(target_family = "wasm")]
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + 'a>>;
pub const DEFAULT_WAIT_TIMEOUT: Duration = Duration::from_secs(10);
const SATURATED_WAIT_CAP: Duration = Duration::from_secs(60 * 60 * 24 * 365);
pub fn wait_deadline(timeout: Duration) -> Option<std::time::Instant> {
if timeout == Duration::MAX {
return None;
}
let now = std::time::Instant::now();
Some(now.checked_add(timeout).unwrap_or_else(|| now.checked_add(SATURATED_WAIT_CAP).unwrap_or(now)))
}
pub trait SyncWaiter: crate::MaybeSendSync + 'static {
fn wait(&self, timeout: Duration) -> Result<(), Error>;
fn wait_async<'a>(&'a self, timeout: Duration) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(async move {
const SPIN_ITERATIONS: usize = 64;
let deadline = wait_deadline(timeout);
for _ in 0..SPIN_ITERATIONS {
match self.is_signaled() {
Ok(true) => return Ok(()),
Ok(false) => {}
Err(e) => return Err(e),
}
if let Some(d) = deadline
&& std::time::Instant::now() >= d
{
return Err(Error::Timeout);
}
yield_once().await;
}
loop {
match self.is_signaled() {
Ok(true) => return Ok(()),
Ok(false) => {}
Err(e) => return Err(e),
}
if let Some(d) = deadline
&& std::time::Instant::now() >= d
{
return Err(Error::Timeout);
}
yield_once().await;
}
})
}
fn is_signaled(&self) -> Result<bool, Error> {
match self.wait(Duration::ZERO) {
Ok(()) => Ok(true),
Err(Error::Timeout) => Ok(false),
Err(other) => Err(other),
}
}
fn backend(&self) -> BackendKind;
fn as_any(&self) -> &dyn Any;
fn as_cuda_event_waiter(&self) -> Option<&dyn CudaEventWaiter> {
None
}
fn rebind_to_value(&self, value: u64) -> Option<Arc<dyn SyncWaiter>> {
let _ = value;
None
}
}
pub trait CudaEventWaiter: SyncWaiter {
fn wait_on_foreign_stream(&self, foreign_stream: *mut c_void) -> Result<(), Error>;
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum VulkanSemaphoreKind {
Timeline,
Binary,
}
#[cfg(feature = "wgpu")]
#[derive(Debug, Default)]
pub struct DeferredWgpuSlot {
cell: std::sync::OnceLock<wgpu::SubmissionIndex>,
}
#[cfg(feature = "wgpu")]
impl DeferredWgpuSlot {
pub fn new() -> Self {
Self::default()
}
pub fn set(&self, idx: wgpu::SubmissionIndex) -> Result<(), Error> {
self.cell.set(idx).map_err(|_| Error::InvalidArgument("DeferredWgpuSlot already committed".into()))
}
pub fn get(&self) -> Option<&wgpu::SubmissionIndex> {
self.cell.get()
}
}
#[cfg(feature = "wgpu")]
static_assertions::assert_impl_all!(wgpu::SubmissionIndex: Clone, Send, Sync);
#[derive(Clone)]
#[non_exhaustive]
pub enum SyncPoint {
Vulkan { semaphore: u64, kind: VulkanSemaphoreKind, value: u64, device: *mut c_void, waiter: Arc<dyn SyncWaiter> },
D3D12 { fence: *mut c_void, value: u64, waiter: Arc<dyn SyncWaiter> },
D3D11 { keyed_mutex: *mut c_void, key: u64, waiter: Arc<dyn SyncWaiter> },
Cuda { event: *mut c_void, value: Option<u64>, waiter: Arc<dyn SyncWaiter> },
CudaEvent { event: *mut c_void, context: *mut c_void, device: i32, waiter: Arc<dyn SyncWaiter> },
OpenCl { event: *mut c_void, value: Option<u64>, waiter: Arc<dyn SyncWaiter> },
Metal { event: *mut c_void, value: u64, waiter: Arc<dyn SyncWaiter> },
OpenGL { semaphore: u32, value: Option<u64>, context: *mut c_void, waiter: Arc<dyn SyncWaiter> },
OpenGLSync { glsync: *mut c_void, backend: GlBackend, waiter: Arc<dyn SyncWaiter> },
#[cfg(feature = "wgpu")]
Wgpu { submission_index: wgpu::SubmissionIndex, waiter: Arc<dyn SyncWaiter> },
#[cfg(feature = "wgpu")]
DeferredWgpu { slot: Arc<DeferredWgpuSlot>, device: Arc<wgpu::Device>, waiter: Arc<dyn SyncWaiter> },
Cpu,
Noop,
}
unsafe impl Send for SyncPoint {}
unsafe impl Sync for SyncPoint {}
impl SyncPoint {
pub async fn wait(&self) -> Result<(), Error> {
self.wait_with_timeout_async(DEFAULT_WAIT_TIMEOUT).await
}
pub async fn wait_with_timeout_async(&self, timeout: Duration) -> Result<(), Error> {
if let Some(err) = self.binary_vk_cpu_wait_error() {
return Err(err);
}
match self.waiter() {
Some(w) => w.wait_async(timeout).await,
None => Ok(()),
}
}
pub fn wait_blocking(&self) -> Result<(), Error> {
self.wait_with_timeout(DEFAULT_WAIT_TIMEOUT)
}
pub fn wait_with_timeout(&self, timeout: Duration) -> Result<(), Error> {
if let Some(err) = self.binary_vk_cpu_wait_error() {
return Err(err);
}
match self.waiter() {
Some(w) => w.wait(timeout),
None => Ok(()),
}
}
fn binary_vk_cpu_wait_error(&self) -> Option<Error> {
let _ = matches!(self, Self::Vulkan { kind: VulkanSemaphoreKind::Binary, .. });
None
}
pub fn chain(self, then: SyncPoint) -> SyncPoint {
if matches!(self, Self::Cpu | Self::Noop) {
return then;
}
if matches!(then, Self::Cpu | Self::Noop) {
return self;
}
let backend = self.backend();
let waiter: Arc<dyn SyncWaiter> = Arc::new(ChainWaiter { first: self, then, backend });
Self::Cpu .with_chain_waiter(waiter)
}
fn with_chain_waiter(self, waiter: Arc<dyn SyncWaiter>) -> SyncPoint {
let _ = self;
Self::Cuda { event: core::ptr::null_mut(), value: None, waiter }
}
}
#[derive(Debug)]
pub struct RawFenceCarrierWaiter {
backend: BackendKind,
}
impl RawFenceCarrierWaiter {
fn new(backend: BackendKind) -> Self {
Self { backend }
}
fn not_imported() -> Error {
Error::NotSupported(std::borrow::Cow::Borrowed(
"SyncPoint reconstructed from a raw foreign fence handle is a transport \
carrier — import it onto a device before waiting",
))
}
}
impl SyncWaiter for RawFenceCarrierWaiter {
fn wait(&self, _timeout: Duration) -> Result<(), Error> {
Err(Self::not_imported())
}
fn is_signaled(&self) -> Result<bool, Error> {
Err(Self::not_imported())
}
fn backend(&self) -> BackendKind {
self.backend
}
fn as_any(&self) -> &dyn Any {
self
}
}
impl SyncPoint {
pub fn from_raw_d3d12_fence(fence: *mut c_void, value: u64) -> SyncPoint {
let waiter: Arc<dyn SyncWaiter> = Arc::new(RawFenceCarrierWaiter::new(BackendKind::D3D12));
SyncPoint::D3D12 { fence, value, waiter }
}
pub fn from_raw_vulkan_semaphore(
semaphore: u64,
kind: VulkanSemaphoreKind,
value: u64,
device: *mut c_void,
) -> SyncPoint {
let waiter: Arc<dyn SyncWaiter> = Arc::new(RawFenceCarrierWaiter::new(BackendKind::Vulkan));
SyncPoint::Vulkan { semaphore, kind, value, device, waiter }
}
pub fn from_raw_metal_event(event: *mut c_void, value: u64) -> SyncPoint {
let waiter: Arc<dyn SyncWaiter> = Arc::new(RawFenceCarrierWaiter::new(BackendKind::Metal));
SyncPoint::Metal { event, value, waiter }
}
}
#[cfg(feature = "wgpu")]
fn deferred_wgpu_waiter_thread() -> &'static crate::WaiterThread {
static THREAD: std::sync::OnceLock<crate::WaiterThread> = std::sync::OnceLock::new();
THREAD.get_or_init(|| crate::WaiterThread::new("wgpu-deferred"))
}
#[cfg(feature = "wgpu")]
pub(crate) struct DeferredWgpuWaiter {
slot: Arc<DeferredWgpuSlot>,
device: Arc<wgpu::Device>,
}
#[cfg(feature = "wgpu")]
impl DeferredWgpuWaiter {
pub(crate) fn new(slot: Arc<DeferredWgpuSlot>, device: Arc<wgpu::Device>) -> Self {
Self { slot, device }
}
fn poll_slice(&self, timeout: Option<Duration>) -> SliceOutcome {
let submission_index = self.slot.get().cloned();
match self.device.poll(wgpu::PollType::Wait { submission_index, timeout }) {
Ok(_) => SliceOutcome::Signaled,
Err(wgpu::PollError::Timeout) => SliceOutcome::TimedOut,
Err(e) => {
log::warn!("SyncPoint::DeferredWgpu: device.poll returned {e:?}");
SliceOutcome::Failed(Error::NotSupported("SyncPoint::DeferredWgpu: device.poll failed".into()))
}
}
}
}
#[cfg(feature = "wgpu")]
impl SyncWaiter for DeferredWgpuWaiter {
fn wait(&self, timeout: Duration) -> Result<(), Error> {
let timeout_arg = if timeout == Duration::MAX { None } else { Some(timeout) };
match self.poll_slice(timeout_arg) {
SliceOutcome::Signaled => Ok(()),
SliceOutcome::TimedOut => Err(Error::Timeout),
SliceOutcome::Failed(e) => Err(e),
}
}
fn wait_async<'a>(&'a self, timeout: Duration) -> BoxFuture<'a, Result<(), Error>> {
let slot = self.slot.clone();
let device = self.device.clone();
let make_slice = move || -> crate::SliceFn {
Box::new(move |slice: Duration| -> SliceOutcome {
let submission_index = slot.get().cloned();
match device.poll(wgpu::PollType::Wait { submission_index, timeout: Some(slice) }) {
Ok(_) => SliceOutcome::Signaled,
Err(wgpu::PollError::Timeout) => SliceOutcome::TimedOut,
Err(e) => {
log::warn!("SyncPoint::DeferredWgpu::wait_async: device.poll returned {e:?}");
SliceOutcome::Failed(Error::NotSupported(
"SyncPoint::DeferredWgpu::wait_async: device.poll failed".into(),
))
}
}
})
};
Box::pin(crate::run_hybrid_wait(move || self.is_signaled(), deferred_wgpu_waiter_thread(), timeout, make_slice))
}
fn is_signaled(&self) -> Result<bool, Error> {
match self.poll_slice(Some(Duration::ZERO)) {
SliceOutcome::Signaled => Ok(true),
SliceOutcome::TimedOut => Ok(false),
SliceOutcome::Failed(e) => Err(e),
}
}
fn backend(&self) -> BackendKind {
BackendKind::Wgpu
}
fn as_any(&self) -> &dyn Any {
self
}
}
#[cfg(feature = "wgpu")]
pub fn make_deferred_wgpu_sync_point(slot: Arc<DeferredWgpuSlot>, device: Arc<wgpu::Device>) -> SyncPoint {
let waiter: Arc<dyn SyncWaiter> = Arc::new(DeferredWgpuWaiter::new(slot.clone(), device.clone()));
SyncPoint::DeferredWgpu { slot, device, waiter }
}
struct ChainWaiter {
first: SyncPoint,
then: SyncPoint,
backend: BackendKind,
}
impl SyncWaiter for ChainWaiter {
fn wait(&self, timeout: Duration) -> Result<(), Error> {
let start = std::time::Instant::now();
self.first.wait_with_timeout(timeout)?;
let elapsed = start.elapsed();
let remaining = timeout.saturating_sub(elapsed);
self.then.wait_with_timeout(remaining)
}
fn wait_async<'a>(&'a self, timeout: Duration) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(async move {
let start = std::time::Instant::now();
self.first.wait_with_timeout_async(timeout).await?;
let elapsed = start.elapsed();
let remaining = timeout.saturating_sub(elapsed);
self.then.wait_with_timeout_async(remaining).await
})
}
fn is_signaled(&self) -> Result<bool, Error> {
Ok(self.first.is_signaled()? && self.then.is_signaled()?)
}
fn backend(&self) -> BackendKind {
self.backend
}
fn as_any(&self) -> &dyn Any {
self
}
}
impl SyncPoint {
pub fn is_signaled(&self) -> Result<bool, Error> {
match self.waiter() {
Some(w) => w.is_signaled(),
None => Ok(true),
}
}
pub fn backend(&self) -> BackendKind {
match self.waiter() {
Some(w) => w.backend(),
None => BackendKind::Cpu,
}
}
pub fn waiter(&self) -> Option<&Arc<dyn SyncWaiter>> {
match self {
Self::Vulkan { waiter, .. }
| Self::D3D12 { waiter, .. }
| Self::D3D11 { waiter, .. }
| Self::Cuda { waiter, .. }
| Self::CudaEvent { waiter, .. }
| Self::OpenCl { waiter, .. }
| Self::Metal { waiter, .. }
| Self::OpenGL { waiter, .. }
| Self::OpenGLSync { waiter, .. } => Some(waiter),
#[cfg(feature = "wgpu")]
Self::Wgpu { waiter, .. } => Some(waiter),
#[cfg(feature = "wgpu")]
Self::DeferredWgpu { waiter, .. } => Some(waiter),
Self::Cpu | Self::Noop => None,
}
}
}
impl core::fmt::Debug for SyncPoint {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Vulkan { semaphore, kind, value, .. } => f
.debug_struct("Vulkan")
.field("semaphore", semaphore)
.field("kind", kind)
.field("value", value)
.finish(),
Self::D3D12 { value, .. } => f.debug_struct("D3D12").field("value", value).finish(),
Self::D3D11 { key, .. } => f.debug_struct("D3D11").field("key", key).finish(),
Self::Cuda { value, .. } => f.debug_struct("Cuda").field("value", value).finish(),
Self::CudaEvent { event, device, .. } => {
f.debug_struct("CudaEvent").field("event", &(*event as usize)).field("device", device).finish()
}
Self::OpenCl { event, value, .. } => {
f.debug_struct("OpenCl").field("event", &(*event as usize)).field("value", value).finish()
}
Self::Metal { value, .. } => f.debug_struct("Metal").field("value", value).finish(),
Self::OpenGL { semaphore, value, .. } => {
f.debug_struct("OpenGL").field("semaphore", semaphore).field("value", value).finish()
}
Self::OpenGLSync { glsync, backend, .. } => {
f.debug_struct("OpenGLSync").field("glsync", &(*glsync as usize)).field("backend", backend).finish()
}
#[cfg(feature = "wgpu")]
Self::Wgpu { .. } => f.write_str("SyncPoint::Wgpu"),
#[cfg(feature = "wgpu")]
Self::DeferredWgpu { slot, .. } => {
f.debug_struct("DeferredWgpu").field("committed", &slot.get().is_some()).finish()
}
Self::Cpu => f.write_str("SyncPoint::Cpu"),
Self::Noop => f.write_str("SyncPoint::Noop"),
}
}
}
pub async fn yield_once() {
let mut yielded = false;
core::future::poll_fn(|cx| {
if yielded {
core::task::Poll::Ready(())
} else {
yielded = true;
cx.waker().wake_by_ref();
core::task::Poll::Pending
}
})
.await
}
pub fn duration_to_ns(t: Duration) -> u64 {
u64::try_from(t.as_nanos()).unwrap_or(u64::MAX)
}
pub fn duration_to_ms_u32(t: Duration) -> u32 {
u32::try_from(t.as_millis()).unwrap_or(u32::MAX)
}