use crate::device_future::StreamCallbackState;
use crate::error::DeviceError;
use crate::slot_table::{FlagArray, SlotTable};
use std::mem::MaybeUninit;
use std::sync::atomic::AtomicU32;
use std::sync::{Arc, OnceLock};
use std::thread;
const NUM_SLOTS: usize = 1024;
const SPIN_PASSES: u32 = 10_000;
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 Reactor {
table: SlotTable<Arc<StreamCallbackState>, 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<Arc<StreamCallbackState>> = Vec::new();
loop {
reactor.table.scan_once(&mut woken);
if !woken.is_empty() {
for state in woken.drain(..) {
state.signal();
}
idle_passes = 0;
continue;
}
if reactor.table.is_idle() {
thread::park();
idle_passes = 0;
continue;
}
idle_passes += 1;
if idle_passes < SPIN_PASSES {
std::hint::spin_loop();
} else {
thread::yield_now();
}
}
}
pub(crate) unsafe fn register(
stream: cuda_bindings::CUstream,
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, addr, 1, 0);
if code != cuda_bindings::cudaError_enum_CUDA_SUCCESS {
reactor.table.release(slot);
return Err(internal(format!("cuStreamWriteValue32 failed: {code}")));
}
if reactor.table.publish(slot, waker_state) {
reactor.scanner.unpark();
}
Ok(())
}