use core::sync::atomic::{AtomicPtr, AtomicUsize, Ordering};
const MAX_TRACKED_CUDA_ALLOCATIONS: usize = 256;
pub struct CudaAllocationRegistry {
slots: [AtomicPtr<u8>; MAX_TRACKED_CUDA_ALLOCATIONS],
count: AtomicUsize,
}
impl CudaAllocationRegistry {
const fn new() -> Self {
Self {
slots: [const { AtomicPtr::new(core::ptr::null_mut()) }; MAX_TRACKED_CUDA_ALLOCATIONS],
count: AtomicUsize::new(0),
}
}
}
pub(super) static CUDA_ALLOCATIONS: CudaAllocationRegistry = CudaAllocationRegistry::new();
pub(super) static CUDA_DEVICE_ALLOCATIONS: CudaAllocationRegistry = CudaAllocationRegistry::new();
pub(super) static CUDA_HOST_PINNED_ALLOCATIONS: CudaAllocationRegistry =
CudaAllocationRegistry::new();
pub(super) fn register_cuda_ptr_in(registry: &CudaAllocationRegistry, ptr: *mut u8) -> bool {
let start_idx = (ptr as usize >> 12) % MAX_TRACKED_CUDA_ALLOCATIONS;
for i in 0..MAX_TRACKED_CUDA_ALLOCATIONS {
let idx = (start_idx + i) % MAX_TRACKED_CUDA_ALLOCATIONS;
let slot = ®istry.slots[idx];
if slot.load(Ordering::Relaxed).is_null()
&& slot
.compare_exchange(
core::ptr::null_mut(),
ptr,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
registry.count.fetch_add(1, Ordering::Release);
return true;
}
}
false
}
pub(super) fn unregister_cuda_ptr_in(registry: &CudaAllocationRegistry, ptr: *mut u8) -> bool {
let start_idx = (ptr as usize >> 12) % MAX_TRACKED_CUDA_ALLOCATIONS;
for i in 0..MAX_TRACKED_CUDA_ALLOCATIONS {
let idx = (start_idx + i) % MAX_TRACKED_CUDA_ALLOCATIONS;
let slot = ®istry.slots[idx];
if slot.load(Ordering::Relaxed) == ptr
&& slot
.compare_exchange(
ptr,
core::ptr::null_mut(),
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
registry.count.fetch_sub(1, Ordering::Release);
return true;
}
}
false
}
#[cfg(test)]
mod tests {
use super::*;
fn test_registry() -> CudaAllocationRegistry {
CudaAllocationRegistry::new()
}
fn is_cuda_ptr_in(registry: &CudaAllocationRegistry, ptr: *mut u8) -> bool {
let start_idx = (ptr as usize >> 12) % MAX_TRACKED_CUDA_ALLOCATIONS;
for i in 0..MAX_TRACKED_CUDA_ALLOCATIONS {
let idx = (start_idx + i) % MAX_TRACKED_CUDA_ALLOCATIONS;
if registry.slots[idx].load(Ordering::Relaxed) == ptr {
return true;
}
}
false
}
#[repr(align(4096))]
struct Page([u8; 4096]);
#[test]
fn cuda_registry_is_bounded_and_reusable() {
let registry = test_registry();
let mut bytes = [0_u8; MAX_TRACKED_CUDA_ALLOCATIONS + 1];
for byte in bytes.iter_mut().take(MAX_TRACKED_CUDA_ALLOCATIONS) {
assert!(register_cuda_ptr_in(®istry, byte as *mut u8));
}
assert!(!register_cuda_ptr_in(
®istry,
&mut bytes[MAX_TRACKED_CUDA_ALLOCATIONS] as *mut u8
));
assert!(unregister_cuda_ptr_in(®istry, &mut bytes[7] as *mut u8));
assert!(register_cuda_ptr_in(
®istry,
&mut bytes[MAX_TRACKED_CUDA_ALLOCATIONS] as *mut u8
));
}
#[test]
fn cuda_registry_rejects_unknown_pointers() {
let registry = test_registry();
let mut byte = 0_u8;
assert!(!unregister_cuda_ptr_in(®istry, &mut byte as *mut u8));
}
#[test]
fn cuda_registry_hashing_and_fallback_forwarding() {
let registry = test_registry();
let mut byte1 = 0_u8;
let mut byte2 = 0_u8;
let ptr1 = &mut byte1 as *mut u8;
let ptr2 = &mut byte2 as *mut u8;
assert!(register_cuda_ptr_in(®istry, ptr1));
assert!(is_cuda_ptr_in(®istry, ptr1));
assert!(!is_cuda_ptr_in(®istry, ptr2));
assert!(unregister_cuda_ptr_in(®istry, ptr1));
assert!(!is_cuda_ptr_in(®istry, ptr1));
}
#[test]
fn cuda_registry_unregister_survives_register_race_and_reclaims_all_slots() {
let registry = test_registry();
let mut page = Page([0; 4096]);
let base = page.0.as_mut_ptr();
let start_idx = (base as usize >> 12) % MAX_TRACKED_CUDA_ALLOCATIONS;
let mut ptrs = [core::ptr::null_mut::<u8>(); MAX_TRACKED_CUDA_ALLOCATIONS];
for (k, slot_ptr) in ptrs.iter_mut().enumerate() {
*slot_ptr = unsafe { base.add(k) };
assert!(register_cuda_ptr_in(®istry, *slot_ptr));
}
assert_eq!(
registry.count.load(Ordering::Acquire),
MAX_TRACKED_CUDA_ALLOCATIONS
);
for &ptr in ptrs.iter().take(MAX_TRACKED_CUDA_ALLOCATIONS - 1) {
assert!(unregister_cuda_ptr_in(®istry, ptr));
}
assert_eq!(registry.count.load(Ordering::Acquire), 1);
let target = ptrs[MAX_TRACKED_CUDA_ALLOCATIONS - 1];
let racer = unsafe { base.add(4095) };
registry.slots[start_idx].store(racer, Ordering::Release);
assert!(unregister_cuda_ptr_in(®istry, target));
assert!(!is_cuda_ptr_in(®istry, target));
registry.count.fetch_add(1, Ordering::Release);
assert!(unregister_cuda_ptr_in(®istry, racer));
assert_eq!(registry.count.load(Ordering::Acquire), 0);
for k in 0..MAX_TRACKED_CUDA_ALLOCATIONS {
assert!(register_cuda_ptr_in(®istry, unsafe { base.add(k) }));
}
assert_eq!(
registry.count.load(Ordering::Acquire),
MAX_TRACKED_CUDA_ALLOCATIONS
);
}
}