use cuda_core::{CudaEvent, DriverError};
use std::io::{self, Write};
use std::mem;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Mutex, MutexGuard};
pub(crate) trait ReclaimGate: Send {
fn passed(&self) -> Result<bool, DriverError>;
fn wait(&self) -> Result<(), DriverError>;
}
impl ReclaimGate for CudaEvent {
fn passed(&self) -> Result<bool, DriverError> {
self.query()
}
fn wait(&self) -> Result<(), DriverError> {
self.synchronize()
}
}
struct LimboEntry {
gate: Box<dyn ReclaimGate>,
payload: Box<dyn Send>,
}
static PENDING: AtomicUsize = AtomicUsize::new(0);
static LIMBO: Mutex<Vec<LimboEntry>> = Mutex::new(Vec::new());
fn limbo() -> MutexGuard<'static, Vec<LimboEntry>> {
LIMBO
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub(crate) fn park(gate: impl ReclaimGate + 'static, payload: Box<dyn Send>) {
let mut entries = limbo();
entries.push(LimboEntry {
gate: Box::new(gate),
payload,
});
PENDING.store(entries.len(), Ordering::Release);
}
pub fn pending() -> usize {
PENDING.load(Ordering::Acquire)
}
pub fn sweep() -> usize {
if PENDING.load(Ordering::Acquire) == 0 {
return 0;
}
let completed: Vec<LimboEntry> = {
let mut entries = limbo();
let mut kept = Vec::with_capacity(entries.len());
let mut done = Vec::new();
for entry in entries.drain(..) {
match entry.gate.passed() {
Ok(true) => done.push(entry),
Ok(false) | Err(_) => kept.push(entry),
}
}
*entries = kept;
PENDING.store(entries.len(), Ordering::Release);
done
};
let reclaimed = completed.len();
drop(completed);
reclaimed
}
pub fn drain() -> usize {
let entries: Vec<LimboEntry> = {
let mut entries = limbo();
let drained = mem::take(&mut *entries);
PENDING.store(0, Ordering::Release);
drained
};
let mut reclaimed = 0;
for entry in entries {
let waited = match entry.gate.passed() {
Ok(true) => Ok(()),
_ => entry.gate.wait(),
};
match waited {
Ok(()) => {
drop(entry.payload);
reclaimed += 1;
}
Err(error) => {
let mut stderr = io::stderr().lock();
let _ = writeln!(
stderr,
"cuda-async: leaking a cancelled in-flight result; the driver \
could not prove the GPU work finished: {error}"
);
mem::forget(entry.payload);
}
}
}
reclaimed
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicBool;
use std::sync::Arc;
static TEST_LOCK: Mutex<()> = Mutex::new(());
fn test_lock() -> MutexGuard<'static, ()> {
TEST_LOCK
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
struct MockGate {
passed: Arc<AtomicBool>,
wait_fails: bool,
}
impl ReclaimGate for MockGate {
fn passed(&self) -> Result<bool, DriverError> {
Ok(self.passed.load(Ordering::Relaxed))
}
fn wait(&self) -> Result<(), DriverError> {
if self.wait_fails {
Err(DriverError(
cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE,
))
} else {
self.passed.store(true, Ordering::Relaxed);
Ok(())
}
}
}
struct CountDrop(Arc<AtomicUsize>);
impl Drop for CountDrop {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
#[test]
fn sweep_keeps_payload_parked_until_gate_passes() {
let _guard = test_lock();
let drops = Arc::new(AtomicUsize::new(0));
let gate_passed = Arc::new(AtomicBool::new(false));
park(
MockGate {
passed: Arc::clone(&gate_passed),
wait_fails: false,
},
Box::new(CountDrop(Arc::clone(&drops))),
);
sweep();
assert_eq!(
drops.load(Ordering::Relaxed),
0,
"payload must stay parked while the GPU work is in flight"
);
gate_passed.store(true, Ordering::Relaxed);
sweep();
assert_eq!(
drops.load(Ordering::Relaxed),
1,
"payload must drop once the gate reports completion"
);
}
#[test]
fn drain_waits_for_gate_then_drops_payload() {
let _guard = test_lock();
let drops = Arc::new(AtomicUsize::new(0));
let gate_passed = Arc::new(AtomicBool::new(false));
park(
MockGate {
passed: Arc::clone(&gate_passed),
wait_fails: false,
},
Box::new(CountDrop(Arc::clone(&drops))),
);
let reclaimed = drain();
assert_eq!(reclaimed, 1);
assert_eq!(drops.load(Ordering::Relaxed), 1);
assert!(
gate_passed.load(Ordering::Relaxed),
"drain must have waited on the gate before dropping"
);
}
#[test]
fn drain_leaks_payload_when_wait_fails() {
let _guard = test_lock();
let drops = Arc::new(AtomicUsize::new(0));
park(
MockGate {
passed: Arc::new(AtomicBool::new(false)),
wait_fails: true,
},
Box::new(CountDrop(Arc::clone(&drops))),
);
let reclaimed = drain();
assert_eq!(reclaimed, 0);
assert_eq!(
drops.load(Ordering::Relaxed),
0,
"an unprovable gate must leak the payload, never drop it early"
);
}
}