use std::any::Any;
use std::ffi::c_void;
use std::ptr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::Duration;
use gpu_handle_types::{BackendKind, CudaEventWaiter, Error, SyncPoint, SyncWaiter};
struct StubCudaEventWaiter {
dropped: Arc<AtomicBool>,
foreign_stream_calls: Arc<AtomicUsize>,
last_foreign_stream: parking_lot::Mutex<*mut c_void>,
}
unsafe impl Send for StubCudaEventWaiter {}
unsafe impl Sync for StubCudaEventWaiter {}
impl Drop for StubCudaEventWaiter {
fn drop(&mut self) {
self.dropped.store(true, Ordering::SeqCst);
}
}
impl SyncWaiter for StubCudaEventWaiter {
fn wait(&self, _timeout: Duration) -> Result<(), Error> {
Ok(())
}
fn is_signaled(&self) -> Result<bool, Error> {
Ok(true)
}
fn backend(&self) -> BackendKind {
BackendKind::Cpu
}
fn as_any(&self) -> &dyn Any {
self
}
}
impl CudaEventWaiter for StubCudaEventWaiter {
fn wait_on_foreign_stream(&self, foreign_stream: *mut c_void) -> Result<(), Error> {
self.foreign_stream_calls.fetch_add(1, Ordering::SeqCst);
*self.last_foreign_stream.lock() = foreign_stream;
Ok(())
}
}
fn make_sync_point(waiter_inner: Arc<StubCudaEventWaiter>) -> (SyncPoint, *mut c_void, *mut c_void) {
let event: *mut c_void = 0x1000 as *mut c_void;
let context: *mut c_void = 0x2000 as *mut c_void;
let waiter: Arc<dyn SyncWaiter> = waiter_inner;
let sp = SyncPoint::CudaEvent { event, context, device: 0, waiter };
(sp, event, context)
}
#[test]
fn cuda_event_variant_fields_round_trip() {
let stub = Arc::new(StubCudaEventWaiter {
dropped: Arc::new(AtomicBool::new(false)),
foreign_stream_calls: Arc::new(AtomicUsize::new(0)),
last_foreign_stream: parking_lot::Mutex::new(ptr::null_mut()),
});
let (sp, expect_event, expect_context) = make_sync_point(stub);
match &sp {
SyncPoint::CudaEvent { event, context, device, waiter: _ } => {
assert_eq!(*event, expect_event);
assert_eq!(*context, expect_context);
assert_eq!(*device, 0);
}
_ => panic!("expected CudaEvent variant"),
}
}
#[test]
fn cuda_event_generic_dispatch_routes_to_inner_waiter() {
let stub = Arc::new(StubCudaEventWaiter {
dropped: Arc::new(AtomicBool::new(false)),
foreign_stream_calls: Arc::new(AtomicUsize::new(0)),
last_foreign_stream: parking_lot::Mutex::new(ptr::null_mut()),
});
let (sp, _, _) = make_sync_point(stub);
assert_eq!(sp.backend(), BackendKind::Cpu);
assert!(sp.is_signaled().expect("stub is_signaled returns Ok"));
sp.wait_blocking().expect("stub wait returns Ok");
}
#[test]
fn cuda_event_drop_discipline_fires_on_last_clone() {
let dropped = Arc::new(AtomicBool::new(false));
let stub = Arc::new(StubCudaEventWaiter {
dropped: dropped.clone(),
foreign_stream_calls: Arc::new(AtomicUsize::new(0)),
last_foreign_stream: parking_lot::Mutex::new(ptr::null_mut()),
});
let (sp, _, _) = make_sync_point(stub);
let clone = sp.clone();
assert!(!dropped.load(Ordering::SeqCst));
drop(sp);
assert!(!dropped.load(Ordering::SeqCst));
drop(clone);
assert!(dropped.load(Ordering::SeqCst), "CudaEventWaiter Drop must fire when the last SyncPoint clone is released");
}
#[test]
fn cuda_event_waiter_trait_object_dispatch_foreign_stream() {
let foreign_calls = Arc::new(AtomicUsize::new(0));
let stub = Arc::new(StubCudaEventWaiter {
dropped: Arc::new(AtomicBool::new(false)),
foreign_stream_calls: foreign_calls.clone(),
last_foreign_stream: parking_lot::Mutex::new(ptr::null_mut()),
});
let (sp, _, _) = make_sync_point(stub);
let waiter = sp.waiter().expect("CudaEvent carries a waiter");
let concrete =
waiter.as_any().downcast_ref::<StubCudaEventWaiter>().expect("downcast to concrete CudaEventWaiter impl");
concrete.wait_on_foreign_stream(0xdead_beef_usize as *mut c_void).expect("stub forwards Ok");
assert_eq!(foreign_calls.load(Ordering::SeqCst), 1);
assert_eq!(*concrete.last_foreign_stream.lock(), 0xdead_beef_usize as *mut c_void);
}