use crate::device_future::{probe_stream, StreamCallbackState, StreamHealth};
use crate::error::DeviceError;
use crate::slot_table::{FlagArray, SlotTable};
use cuda_core::Stream;
use std::mem::MaybeUninit;
use std::sync::atomic::AtomicU32;
use std::sync::{Arc, OnceLock};
use std::thread;
use std::time::{Duration, Instant};
const NUM_SLOTS: usize = 1024;
const SPIN_PASSES: u32 = 10_000;
const STALE_PROBE_INTERVAL: Duration = Duration::from_millis(2);
struct CudaFlags {
host: *mut u32,
}
unsafe impl Send for CudaFlags {}
unsafe impl Sync for CudaFlags {}
impl FlagArray for CudaFlags {
fn flag(&self, slot: usize) -> &AtomicU32 {
unsafe { AtomicU32::from_ptr(self.host.add(slot)) }
}
}
struct Registration {
waker_state: Arc<StreamCallbackState>,
stream: Arc<Stream>,
}
struct Reactor {
table: SlotTable<Registration, CudaFlags>,
dptr: cuda_bindings::CUdeviceptr,
scanner: thread::Thread,
}
fn internal(msg: String) -> DeviceError {
DeviceError::Internal(msg)
}
fn reactor() -> Result<&'static Reactor, DeviceError> {
static REACTOR: OnceLock<Result<Reactor, String>> = OnceLock::new();
let result = REACTOR.get_or_init(|| unsafe {
let mut host = MaybeUninit::uninit();
let flags =
cuda_bindings::CU_MEMHOSTALLOC_PORTABLE | cuda_bindings::CU_MEMHOSTALLOC_DEVICEMAP;
let code = cuda_bindings::cuMemHostAlloc(
host.as_mut_ptr(),
NUM_SLOTS * std::mem::size_of::<u32>(),
flags,
);
if code != cuda_bindings::cudaError_enum_CUDA_SUCCESS {
return Err(format!("cuMemHostAlloc failed: {code}"));
}
let host = host.assume_init() as *mut u32;
std::ptr::write_bytes(host, 0, NUM_SLOTS);
let mut dptr = MaybeUninit::uninit();
let code = cuda_bindings::cuMemHostGetDevicePointer_v2(dptr.as_mut_ptr(), host as _, 0);
if code != cuda_bindings::cudaError_enum_CUDA_SUCCESS {
return Err(format!("cuMemHostGetDevicePointer failed: {code}"));
}
let dptr = dptr.assume_init();
let handle = thread::Builder::new()
.name("cuda-async-reactor".into())
.spawn(scan_loop)
.map_err(|e| format!("failed to spawn reactor thread: {e}"))?;
Ok(Reactor {
table: SlotTable::new(NUM_SLOTS, CudaFlags { host }),
dptr,
scanner: handle.thread().clone(),
})
});
result.as_ref().map_err(|e| internal(e.clone()))
}
fn scan_loop() {
let reactor = loop {
if let Ok(r) = reactor() {
break r;
}
thread::yield_now();
};
let mut idle_passes: u32 = 0;
let mut woken: Vec<Registration> = Vec::new();
let mut faulted: Vec<Registration> = Vec::new();
let mut last_probe = Instant::now();
loop {
reactor.table.scan_once(&mut woken);
if !woken.is_empty() {
for reg in woken.drain(..) {
reg.waker_state.signal();
}
idle_passes = 0;
continue;
}
if reactor.table.is_idle() {
thread::park();
idle_passes = 0;
last_probe = Instant::now();
continue;
}
idle_passes += 1;
if idle_passes < SPIN_PASSES {
std::hint::spin_loop();
continue;
}
if last_probe.elapsed() >= STALE_PROBE_INTERVAL {
last_probe = Instant::now();
probe_stale_slots(&reactor.table, &mut woken, &mut faulted);
for reg in woken.drain(..) {
reg.waker_state.signal();
}
for reg in faulted.drain(..) {
reg.waker_state.wake();
}
}
thread::yield_now();
}
}
fn probe_stale_slots(
table: &SlotTable<Registration, CudaFlags>,
woken: &mut Vec<Registration>,
faulted: &mut Vec<Registration>,
) {
let mut memo: Vec<(cuda_bindings::CUstream, bool)> = Vec::new();
table.scan_probing(woken, faulted, &mut |reg: &Registration| {
let handle = reg.stream.cu_stream();
if let Some((_, dead)) = memo.iter().find(|(h, _)| *h == handle) {
return *dead;
}
let dead = stream_is_faulted(®.stream);
memo.push((handle, dead));
dead
});
}
fn stream_is_faulted(stream: &Stream) -> bool {
if stream.device().bind_to_thread().is_err() {
return false;
}
matches!(probe_stream(stream), StreamHealth::Faulted(_))
}
pub(crate) unsafe fn register(
stream: &Arc<Stream>,
waker_state: Arc<StreamCallbackState>,
) -> Result<(), DeviceError> {
let reactor = reactor()?;
let slot = reactor
.table
.claim()
.ok_or_else(|| internal("reactor slot pool exhausted".into()))?;
reactor.table.reset_flag(slot);
let addr = reactor.dptr + (slot * std::mem::size_of::<u32>()) as u64;
let code = cuda_bindings::cuStreamWriteValue32_v2(stream.cu_stream(), addr, 1, 0);
if code != cuda_bindings::cudaError_enum_CUDA_SUCCESS {
reactor.table.release(slot);
return Err(internal(format!("cuStreamWriteValue32 failed: {code}")));
}
let registration = Registration {
waker_state,
stream: Arc::clone(stream),
};
if reactor.table.publish(slot, registration) {
reactor.scanner.unpark();
}
Ok(())
}