use crate::simt::device_operation::{DeviceOperation, ExecutionContext};
use crate::simt::error::DeviceError;
use crate::simt::reclaim;
use futures::task::AtomicWaker;
use std::future::Future;
use std::io::{self, Write};
use std::mem;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll};
#[derive(Debug, Default, Eq, PartialEq, Copy, Clone)]
pub enum DeviceFutureState {
Failed,
#[default]
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::Relaxed);
self.waker.wake();
}
}
#[derive(Debug)]
pub struct DeviceFuture<T: Send + 'static, DO: DeviceOperation<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 + 'static, DO: DeviceOperation<Output = T>> DeviceFuture<T, DO> {
pub fn new() -> Self {
Self::default()
}
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_else(|| {
DeviceError::Internal("Cannot execute future without an execution context.".to_string())
})?;
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_else(|| {
DeviceError::Internal("Cannot execute future without an execution context.".to_string())
})?;
let operation = self
.device_operation
.take()
.ok_or_else(|| DeviceError::Internal("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 reclaim_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()));
if let Some(stream) = &stream {
if let Ok(event) = stream.record_event(None) {
let result = self.result.take().expect("Checked above.");
reclaim::park(event, Box::new(result));
return;
}
}
self.cleanup_executing_result_with(move || {
let stream = stream.ok_or_else(|| {
DeviceError::Internal(
"Cannot clean up an in-flight future without an execution context.".to_string(),
)
})?;
stream.synchronize().map_err(DeviceError::Driver)
});
}
fn cleanup_executing_result_with<F>(&mut self, synchronize: F)
where
F: FnOnce() -> Result<(), DeviceError>,
{
if !self.has_undelivered_submission() {
return;
}
let Some(result) = self.result.take() else {
return;
};
if let Err(error) = synchronize() {
let mut stderr = io::stderr().lock();
let _ = writeln!(
stderr,
"cuda-async: leaking in-flight future result after cleanup failure: {}",
error
);
mem::forget(result);
return;
}
drop(result);
}
}
impl<T: Send + 'static, DO: DeviceOperation<Output = T>> Default for DeviceFuture<T, DO> {
fn default() -> Self {
Self {
device_operation: Default::default(),
execution_context: Default::default(),
result: Default::default(),
error: Default::default(),
state: Default::default(),
callback_state: Default::default(),
}
}
}
impl<T: Send + 'static, DO: DeviceOperation<Output = T>> Unpin for DeviceFuture<T, DO> {}
impl<T: Send + 'static, DO: DeviceOperation<Output = T>> Drop for DeviceFuture<T, DO> {
fn drop(&mut self) {
reclaim::sweep();
self.reclaim_in_flight_result();
}
}
impl<T: Send + 'static, DO: DeviceOperation<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> {
reclaim::sweep();
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()
.map(Arc::clone)
.expect("Impossible.");
match self.state {
DeviceFutureState::Idle => {
waker_state.waker.register(cx.waker());
if let Err(e) = self.execute() {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Err(e));
}
if let Err(e) = unsafe { self.register_callback(Arc::clone(&waker_state)) } {
self.state = DeviceFutureState::Complete;
self.reclaim_in_flight_result();
return Poll::Ready(Err(e));
}
self.state = DeviceFutureState::Executing;
Poll::Pending
}
DeviceFutureState::Executing => {
if waker_state.complete.load(Ordering::Relaxed) {
self.state = DeviceFutureState::Complete;
return Poll::Ready(Ok(self.result.take().expect("Expected result.")));
}
waker_state.waker.register(cx.waker());
if waker_state.complete.load(Ordering::Relaxed) {
self.state = DeviceFutureState::Complete;
Poll::Ready(Ok(self.result.take().expect("Expected result.")))
} else {
Poll::Pending
}
}
DeviceFutureState::Complete => panic!("Poll called after completion."),
DeviceFutureState::Failed => unreachable!(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::simt::device_operation::Value;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
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");
}
}
#[test]
fn cleanup_executing_result_synchronizes_before_drop() {
let events = Arc::new(Mutex::new(Vec::new()));
let tracker = DropTracker {
events: Arc::clone(&events),
};
let mut future: DeviceFuture<DropTracker, Value<DropTracker>> = DeviceFuture {
device_operation: None,
execution_context: None,
result: Some(tracker),
error: None,
state: DeviceFutureState::Executing,
callback_state: None,
};
future.cleanup_executing_result_with(|| {
events.lock().unwrap().push("sync");
Ok(())
});
assert_eq!(events.lock().unwrap().as_slice(), ["sync", "drop"]);
assert!(future.result.is_none());
}
#[test]
fn cleanup_executing_result_leaks_when_synchronize_fails() {
let drops = Arc::new(AtomicUsize::new(0));
struct CountDrop(Arc<AtomicUsize>);
impl Drop for CountDrop {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
let mut future: DeviceFuture<CountDrop, Value<CountDrop>> = DeviceFuture {
device_operation: None,
execution_context: None,
result: Some(CountDrop(Arc::clone(&drops))),
error: None,
state: DeviceFutureState::Executing,
callback_state: None,
};
future.cleanup_executing_result_with(|| Err(DeviceError::Internal("boom".to_string())));
assert_eq!(drops.load(Ordering::Relaxed), 0);
assert!(future.result.is_none());
}
#[test]
fn cleanup_covers_complete_future_left_by_callback_registration_failure() {
let events = Arc::new(Mutex::new(Vec::new()));
let tracker = DropTracker {
events: Arc::clone(&events),
};
let mut future: DeviceFuture<DropTracker, Value<DropTracker>> = DeviceFuture {
device_operation: None,
execution_context: None,
result: Some(tracker),
error: None,
state: DeviceFutureState::Complete,
callback_state: None,
};
assert!(future.has_undelivered_submission());
future.cleanup_executing_result_with(|| {
events.lock().unwrap().push("sync");
Ok(())
});
assert_eq!(events.lock().unwrap().as_slice(), ["sync", "drop"]);
assert!(future.result.is_none());
assert!(!future.has_undelivered_submission());
}
#[test]
fn cleanup_is_noop_after_result_delivery() {
let mut future: DeviceFuture<u32, Value<u32>> = DeviceFuture {
device_operation: None,
execution_context: None,
result: None,
error: None,
state: DeviceFutureState::Complete,
callback_state: None,
};
assert!(!future.has_undelivered_submission());
future.cleanup_executing_result_with(|| {
panic!("delivered futures must not synchronize during cleanup")
});
future.reclaim_in_flight_result();
}
#[test]
fn cleanup_is_noop_for_idle_future() {
let drops = Arc::new(AtomicUsize::new(0));
struct CountDrop(Arc<AtomicUsize>);
impl Drop for CountDrop {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
let mut future: DeviceFuture<CountDrop, Value<CountDrop>> = DeviceFuture {
device_operation: None,
execution_context: None,
result: Some(CountDrop(Arc::clone(&drops))),
error: None,
state: DeviceFutureState::Idle,
callback_state: None,
};
future.cleanup_executing_result_with(|| {
panic!("idle futures should not synchronize during cleanup")
});
assert_eq!(drops.load(Ordering::Relaxed), 0);
assert!(future.result.is_some());
drop(future);
assert_eq!(drops.load(Ordering::Relaxed), 1);
}
}