use std::sync::Arc;
use cudarc::driver::sys;
use xlog_cuda::device_runtime::{
AllocTag, AsyncCudaResource, DeviceMemoryResource, StreamId, StreamPool,
};
use xlog_cuda::CudaDevice;
const BYTES: usize = 4096;
const REUSE_ITERATIONS: usize = 32;
fn try_setup() -> Option<(Arc<CudaDevice>, Arc<StreamPool>)> {
let device = Arc::new(CudaDevice::new(0).ok()?);
let pool = Arc::new(StreamPool::with_defaults(Arc::clone(&device)));
Some((device, pool))
}
unsafe fn htod_async(stream: sys::CUstream, dst: u64, src: &[u8]) {
let res = sys::cuMemcpyHtoDAsync_v2(dst, src.as_ptr() as *const _, src.len(), stream);
assert_eq!(
res,
sys::cudaError_enum::CUDA_SUCCESS,
"cuMemcpyHtoDAsync_v2 returned {:?}",
res
);
}
unsafe fn dtoh_sync(dst: &mut [u8], src: u64) {
let res = sys::cuMemcpyDtoH_v2(dst.as_mut_ptr() as *mut _, src, dst.len());
assert_eq!(
res,
sys::cudaError_enum::CUDA_SUCCESS,
"cuMemcpyDtoH_v2 returned {:?}",
res
);
}
#[test]
fn stream_ordered_alloc_write_free_realloc_no_host_sync_between_phases() {
let Some((device, pool)) = try_setup() else {
eprintln!("Skipping stream-ordered allocation lifetime test: CUDA runtime unavailable");
return;
};
let resource = AsyncCudaResource::new(Arc::clone(&device), 0, Arc::clone(&pool));
let stream_id = match pool.acquire() {
Ok(id) => id,
Err(e) => {
eprintln!(
"Skipping stream-ordered allocation lifetime test: StreamPool::acquire failed: {}",
e
);
return;
}
};
assert_ne!(stream_id, StreamId::DEFAULT);
let stream = pool
.resolve(stream_id)
.expect("acquired StreamId must resolve");
let cu_stream = stream.cu_stream();
let block_a = resource
.allocate(BYTES, stream_id, AllocTag("stream-order-A"))
.expect("alloc A");
assert_eq!(block_a.alloc_stream, stream_id);
assert_eq!(block_a.bytes, BYTES);
let bytes_after_a = resource.bytes_outstanding();
assert_eq!(bytes_after_a, BYTES);
let pattern_a = vec![0xAAu8; BYTES];
unsafe {
htod_async(cu_stream, block_a.ptr, &pattern_a);
}
resource.deallocate(block_a).expect("dealloc A");
assert_eq!(
resource.bytes_outstanding(),
BYTES,
"queued cuMemFreeAsync must remain counted as pending until reaped"
);
let block_b = resource
.allocate(BYTES, stream_id, AllocTag("stream-order-B"))
.expect("alloc B");
assert_eq!(block_b.alloc_stream, stream_id);
assert_eq!(block_b.bytes, BYTES);
let pattern_b = vec![0xBBu8; BYTES];
unsafe {
htod_async(cu_stream, block_b.ptr, &pattern_b);
}
stream.synchronize().expect("stream sync");
let mut readback = vec![0u8; BYTES];
unsafe {
dtoh_sync(&mut readback, block_b.ptr);
}
assert_eq!(
readback, pattern_b,
"stream-ordered reuse violated: block B contains stale bytes from A's queued write"
);
resource.deallocate(block_b).expect("dealloc B");
resource.reap_pending().expect("reap pending");
assert_eq!(resource.bytes_outstanding(), 0);
}
#[test]
fn repeated_alloc_free_realloc_on_same_stream_stays_stream_ordered() {
let Some((device, pool)) = try_setup() else {
return;
};
let resource = AsyncCudaResource::new(Arc::clone(&device), 0, Arc::clone(&pool));
let stream_id = match pool.acquire() {
Ok(id) => id,
Err(e) => {
eprintln!(
"Skipping stream-ordered allocation reuse stress: StreamPool::acquire failed: {}",
e
);
return;
}
};
assert_ne!(stream_id, StreamId::DEFAULT);
let stream = pool
.resolve(stream_id)
.expect("acquired StreamId must resolve");
let cu_stream = stream.cu_stream();
let mut current = resource
.allocate(BYTES, stream_id, AllocTag("stream-order-stress-init"))
.expect("initial alloc");
let mut last_pattern = vec![0u8; BYTES];
for iter in 0..REUSE_ITERATIONS {
let stamp: u8 = ((iter as u32) & 0xFF) as u8;
let pattern: Vec<u8> = (0..BYTES)
.map(|i| stamp.wrapping_add((i & 0xFF) as u8))
.collect();
unsafe {
htod_async(cu_stream, current.ptr, &pattern);
}
if iter == REUSE_ITERATIONS - 1 {
last_pattern = pattern;
break;
}
resource.deallocate(current).expect("dealloc mid-loop");
current = resource
.allocate(BYTES, stream_id, AllocTag("stream-order-stress-iter"))
.expect("alloc mid-loop");
}
stream.synchronize().expect("stream sync");
let mut readback = vec![0u8; BYTES];
unsafe {
dtoh_sync(&mut readback, current.ptr);
}
assert_eq!(
readback, last_pattern,
"stream-ordered reuse violated under repeated alloc/free on the same stream"
);
resource.deallocate(current).expect("dealloc final");
resource.reap_pending().expect("reap pending");
assert_eq!(resource.bytes_outstanding(), 0);
}