use crate::device_operation::{DeviceOp, ExecutionContext};
use crate::error::DeviceError;
use cuda_core::{DriverError, Stream};
use futures::task::AtomicWaker;
use std::future::Future;
use std::io::{self, Write};
use std::mem::{self, MaybeUninit};
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll};
#[derive(Debug, Eq, PartialEq, Copy, Clone)]
pub enum DeviceFutureState {
Failed,
Idle,
Executing,
Complete,
}
#[derive(Debug, Default)]
pub struct StreamCallbackState {
pub(crate) waker: AtomicWaker,
pub(crate) complete: AtomicBool,
}
impl StreamCallbackState {
pub fn new() -> Self {
Self::default()
}
pub fn signal(&self) {
self.complete.store(true, Ordering::Release);
self.waker.wake();
}
pub fn wake(&self) {
self.waker.wake();
}
}
pub(crate) enum StreamHealth {
Capturing,
Busy,
Idle,
Faulted(DriverError),
}
pub(crate) fn probe_stream(stream: &Stream) -> StreamHealth {
let mut status = MaybeUninit::uninit();
let code =
unsafe { cuda_bindings::cuStreamIsCapturing(stream.cu_stream(), status.as_mut_ptr()) };
if code != cuda_bindings::cudaError_enum_CUDA_SUCCESS {
return StreamHealth::Faulted(DriverError(code));
}
let status = unsafe { status.assume_init() };
if status != cuda_bindings::CUstreamCaptureStatus_enum_CU_STREAM_CAPTURE_STATUS_NONE {
return StreamHealth::Capturing;
}
match unsafe { stream.query() } {
Ok(true) => StreamHealth::Idle,
Ok(false) => StreamHealth::Busy,
Err(e) => StreamHealth::Faulted(e),
}
}
#[derive(Debug)]
pub struct DeviceFuture<T: Send, DO: DeviceOp<Output = T>> {
pub(crate) device_operation: Option<DO>,
pub(crate) execution_context: Option<ExecutionContext>,
pub(crate) result: Option<T>,
pub(crate) error: Option<DeviceError>,
pub(crate) state: DeviceFutureState,
pub(crate) callback_state: Option<Arc<StreamCallbackState>>,
}
impl<T: Send, DO: DeviceOp<Output = T>> DeviceFuture<T, DO> {
pub fn new() -> Self {
Self::default()
}
pub fn scheduled(op: DO, ctx: ExecutionContext) -> Self {
Self {
device_operation: Some(op),
execution_context: Some(ctx),
result: None,
error: None,
state: DeviceFutureState::Idle,
callback_state: None,
}
}
pub fn failed(error: DeviceError) -> Self {
Self {
execution_context: None,
device_operation: None,
state: DeviceFutureState::Failed,
callback_state: None,
result: None,
error: Some(error),
}
}
unsafe fn register_callback(
&self,
waker_state: Arc<StreamCallbackState>,
) -> Result<(), DeviceError> {
let ctx = self
.execution_context
.as_ref()
.ok_or(DeviceError::Internal(
"Cannot execute future without setting stream on which to execute.".to_string(),
))?;
fn host_task_sync_mode() -> Option<::core::ffi::c_uint> {
static MODE: std::sync::OnceLock<Option<::core::ffi::c_uint>> =
std::sync::OnceLock::new();
*MODE.get_or_init(
|| match std::env::var("CUDA_ASYNC_HOST_SYNC").ok()?.as_str() {
"spin" | "spinwait" => Some(cuda_bindings::CU_HOST_TASK_SPINWAIT),
"block" | "blocking" => Some(cuda_bindings::CU_HOST_TASK_BLOCKING),
_ => None,
},
)
}
if let Some(mode) = host_task_sync_mode() {
ctx.get_cuda_stream()
.launch_host_function_with_sync_mode(move || waker_state.signal(), mode)?;
return Ok(());
}
#[cfg(not(loom))]
{
if crate::reactor::register(ctx.get_cuda_stream(), waker_state.clone()).is_ok() {
return Ok(());
}
}
ctx.get_cuda_stream()
.launch_host_function(move || waker_state.signal())?;
Ok(())
}
fn execute(&mut self) -> Result<(), DeviceError> {
let ctx = self
.execution_context
.as_ref()
.ok_or(DeviceError::Internal(
"Cannot execute future without setting stream on which to execute.".to_string(),
))?;
let operation = self.device_operation.take().ok_or(DeviceError::Internal(
"Unable to execute future: No operation has been set.".to_string(),
))?;
let out = unsafe { operation.execute(ctx) }?;
self.result = Some(out);
Ok(())
}
fn has_undelivered_submission(&self) -> bool {
matches!(
self.state,
DeviceFutureState::Executing | DeviceFutureState::Complete
) && self.result.is_some()
}
fn release_in_flight_result(&mut self) {
if !self.has_undelivered_submission() {
return;
}
let stream = self
.execution_context
.as_ref()
.map(|ctx| Arc::clone(ctx.get_cuda_stream()));
self.release_in_flight_result_with(move || {
let stream = stream.ok_or_else(|| {
DeviceError::Internal(
"Cannot release an in-flight future without an execution context.".to_string(),
)
})?;
stream.device().bind_to_thread()?;
match probe_stream(&stream) {
StreamHealth::Idle => Ok(()),
StreamHealth::Busy => unsafe { stream.synchronize() }.map_err(DeviceError::Driver),
StreamHealth::Capturing => Err(DeviceError::Internal(
"the future's stream is recording a graph; it cannot be synchronized".into(),
)),
StreamHealth::Faulted(e) => Err(DeviceError::Driver(e)),
}
});
}
fn release_in_flight_result_with<F>(&mut self, wait: F)
where
F: FnOnce() -> Result<(), DeviceError>,
{
if !self.has_undelivered_submission() {
return;
}
let Some(result) = self.result.take() else {
return;
};
if let Err(error) = wait() {
let mut stderr = io::stderr().lock();
let _ = writeln!(
stderr,
"cuda-async: leaking the result of a dropped in-flight future; the driver \
could not prove its GPU work finished: {error}"
);
mem::forget(result);
return;
}
drop(result);
}
}
impl<T: Send, DO: DeviceOp<Output = T>> Drop for DeviceFuture<T, DO> {
fn drop(&mut self) {
self.release_in_flight_result();
}
}
impl<T: Send, DO: DeviceOp<Output = T>> Default for DeviceFuture<T, DO> {
fn default() -> Self {
Self {
device_operation: None,
execution_context: None,
result: None,
error: None,
state: DeviceFutureState::Idle,
callback_state: None,
}
}
}
impl<T: Send, DO: DeviceOp<Output = T>> Unpin for DeviceFuture<T, DO> {}
impl<T: Send, DO: DeviceOp<Output = T>> Future for DeviceFuture<T, DO> {
type Output = Result<T, DeviceError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
if self.state == DeviceFutureState::Failed {
self.state = DeviceFutureState::Complete;
let error = self
.error
.take()
.expect("Failed state must carry an error.");
return Poll::Ready(Err(error));
}
if self.callback_state.is_none() {
self.callback_state = Some(Arc::new(StreamCallbackState::new()));
}
let waker_state = self.callback_state.as_ref().cloned().expect("Impossible.");
match self.state {
DeviceFutureState::Idle => {
let _execution_lock = match crate::device_operation::acquire_execution_lock() {
Ok(guard) => guard,
Err(e) => {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Err(e));
}
};
waker_state.waker.register(cx.waker());
if let Err(e) = self.execute() {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Err(e));
}
fn inline_spin_budget_us() -> u64 {
static BUDGET: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
*BUDGET.get_or_init(|| {
std::env::var("CUDA_ASYNC_SPIN_BUDGET_US")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(20)
})
}
let spin_outcome: Result<bool, DriverError> = 'spin: {
if inline_spin_budget_us() == 0 {
break 'spin Ok(false);
}
let Some(ctx) = self.execution_context.as_ref() else {
break 'spin Ok(false);
};
let deadline = std::time::Instant::now()
+ std::time::Duration::from_micros(inline_spin_budget_us());
loop {
match unsafe { ctx.get_cuda_stream().query() } {
Ok(true) => break 'spin Ok(true),
Ok(false) => {}
Err(e) => break 'spin Err(e),
}
if std::time::Instant::now() >= deadline {
break 'spin Ok(false);
}
std::hint::spin_loop();
}
};
match spin_outcome {
Ok(true) => {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Ok(self
.result
.take()
.expect("Expected future result to be Some.")));
}
Ok(false) => {}
Err(e) => {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Err(DeviceError::Driver(e)));
}
}
if let Err(e) = unsafe { self.register_callback(waker_state.clone()) } {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Err(e));
}
self.state = DeviceFutureState::Executing;
Poll::Pending
}
DeviceFutureState::Executing => {
if waker_state.complete.load(Ordering::Acquire) {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Ok(self
.result
.take()
.expect("Expected future result to be Some.")));
}
let health = match self.execution_context.as_ref() {
Some(ctx) => {
let _ = ctx.device().bind_to_thread();
probe_stream(ctx.get_cuda_stream())
}
None => StreamHealth::Busy,
};
match health {
StreamHealth::Faulted(e) => {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Err(DeviceError::Driver(e)));
}
StreamHealth::Idle => {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Ok(self
.result
.take()
.expect("Expected future result to be Some.")));
}
StreamHealth::Busy | StreamHealth::Capturing => {}
}
waker_state.waker.register(cx.waker());
if waker_state.complete.load(Ordering::Acquire) {
self.state = DeviceFutureState::Complete;
Poll::Ready(Ok(self
.result
.take()
.expect("Expected future result to be Some.")))
} else {
Poll::Pending
}
}
DeviceFutureState::Complete => {
panic!("Poll called after completion.");
}
DeviceFutureState::Failed => {
unreachable!();
}
}
}
}
#[cfg(test)]
mod release_tests {
use super::*;
use crate::device_operation::Value;
use std::sync::atomic::AtomicUsize;
use std::sync::Mutex;
#[derive(Clone)]
struct DropTracker {
events: Arc<Mutex<Vec<&'static str>>>,
}
impl Drop for DropTracker {
fn drop(&mut self) {
self.events.lock().unwrap().push("drop");
}
}
struct CountDrop(Arc<AtomicUsize>);
impl Drop for CountDrop {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
fn future_in_state<T: Send>(
state: DeviceFutureState,
result: Option<T>,
) -> DeviceFuture<T, Value<T>> {
DeviceFuture {
device_operation: None,
execution_context: None,
result,
error: None,
state,
callback_state: None,
}
}
#[test]
fn release_waits_before_dropping_the_result() {
let events = Arc::new(Mutex::new(Vec::new()));
let tracker = DropTracker {
events: Arc::clone(&events),
};
let mut future = future_in_state(DeviceFutureState::Executing, Some(tracker));
future.release_in_flight_result_with(|| {
events.lock().unwrap().push("wait");
Ok(())
});
assert_eq!(events.lock().unwrap().as_slice(), ["wait", "drop"]);
assert!(future.result.is_none());
assert!(!future.has_undelivered_submission());
}
#[test]
fn release_leaks_when_the_wait_fails() {
let drops = Arc::new(AtomicUsize::new(0));
let mut future = future_in_state(
DeviceFutureState::Executing,
Some(CountDrop(Arc::clone(&drops))),
);
future.release_in_flight_result_with(|| Err(DeviceError::Internal("boom".to_string())));
assert_eq!(drops.load(Ordering::Relaxed), 0);
assert!(future.result.is_none());
}
#[test]
fn release_covers_complete_future_left_by_registration_failure() {
let events = Arc::new(Mutex::new(Vec::new()));
let tracker = DropTracker {
events: Arc::clone(&events),
};
let mut future = future_in_state(DeviceFutureState::Complete, Some(tracker));
assert!(future.has_undelivered_submission());
future.release_in_flight_result_with(|| {
events.lock().unwrap().push("wait");
Ok(())
});
assert_eq!(events.lock().unwrap().as_slice(), ["wait", "drop"]);
assert!(!future.has_undelivered_submission());
}
#[test]
fn release_is_noop_after_result_delivery() {
let mut future: DeviceFuture<u32, Value<u32>> =
future_in_state(DeviceFutureState::Complete, None);
assert!(!future.has_undelivered_submission());
future.release_in_flight_result_with(|| {
panic!("delivered futures must not wait during release")
});
future.release_in_flight_result();
}
#[test]
fn release_is_noop_for_idle_future() {
let drops = Arc::new(AtomicUsize::new(0));
let mut future =
future_in_state(DeviceFutureState::Idle, Some(CountDrop(Arc::clone(&drops))));
future
.release_in_flight_result_with(|| panic!("idle futures must not wait during release"));
assert_eq!(drops.load(Ordering::Relaxed), 0);
assert!(future.result.is_some());
drop(future);
assert_eq!(drops.load(Ordering::Relaxed), 1);
}
#[test]
fn dropping_failed_future_is_a_noop() {
let future: DeviceFuture<u32, Value<u32>> =
DeviceFuture::failed(DeviceError::Internal("never scheduled".into()));
drop(future);
}
}
#[cfg(test)]
mod callback_state_tests {
use super::StreamCallbackState;
use futures::task::{waker, ArcWake};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
struct CountingWaker(AtomicUsize);
impl ArcWake for CountingWaker {
fn wake_by_ref(arc_self: &Arc<Self>) {
arc_self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[test]
fn signal_without_registered_waker_is_a_noop() {
let state = StreamCallbackState::new();
state.signal(); assert!(state.complete.load(Ordering::Relaxed));
state.signal(); assert!(state.complete.load(Ordering::Relaxed));
}
#[test]
fn signal_wakes_registered_waker_and_sets_complete() {
let state = StreamCallbackState::new();
let counter = Arc::new(CountingWaker(AtomicUsize::new(0)));
state.waker.register(&waker(counter.clone()));
assert_eq!(counter.0.load(Ordering::SeqCst), 0);
state.signal();
assert_eq!(counter.0.load(Ordering::SeqCst), 1);
assert!(state.complete.load(Ordering::Relaxed));
}
#[test]
fn signal_wakes_only_the_latest_registered_waker() {
let state = StreamCallbackState::new();
let first = Arc::new(CountingWaker(AtomicUsize::new(0)));
let second = Arc::new(CountingWaker(AtomicUsize::new(0)));
state.waker.register(&waker(first.clone()));
state.waker.register(&waker(second.clone()));
state.signal();
assert_eq!(first.0.load(Ordering::SeqCst), 0, "stale waker was woken");
assert_eq!(second.0.load(Ordering::SeqCst), 1);
}
}