use crate::queue::Queue;
use allocator_api2::alloc::{AllocError, Allocator};
use sycl_rs_sys::usm::ffi;
use std::{alloc::Layout, marker::PhantomData, ptr::NonNull};
type CxxResult<T> = cxx::core::result::Result<T, cxx::Exception>;
pub struct UsmAllocator<T: UsmAllocatorKind> {
queue: Queue,
_kind: PhantomData<T>,
}
pub unsafe trait UsmAlloc: Allocator {}
unsafe impl<T: UsmAllocatorKind> UsmAlloc for UsmAllocator<T> {}
pub trait UsmAllocatorKind {
unsafe fn alloc(alignment: usize, num_bytes: usize, queue: &Queue) -> CxxResult<*mut u8>;
}
pub unsafe trait HostAccessible {}
impl<T: UsmAllocatorKind> From<&Queue> for UsmAllocator<T> {
fn from(queue: &Queue) -> Self {
Self {
queue: queue.clone(),
_kind: PhantomData,
}
}
}
unsafe impl<T: UsmAllocatorKind> Allocator for UsmAllocator<T> {
fn allocate(&self, layout: Layout) -> Result<NonNull<[u8]>, AllocError> {
let ptr = unsafe { T::alloc(layout.align(), layout.size(), &self.queue) }
.map_err(|_e| AllocError)?;
let ptr = NonNull::new(ptr).ok_or(AllocError)?;
let slice = NonNull::slice_from_raw_parts(ptr, layout.size());
Ok(slice)
}
unsafe fn deallocate(&self, ptr: NonNull<u8>, _layout: Layout) {
unsafe {
ffi::free(ptr.as_ptr(), &self.queue.0);
}
}
}
#[allow(dead_code)]
pub struct DeviceAllocator;
impl UsmAllocatorKind for DeviceAllocator {
unsafe fn alloc(alignment: usize, num_bytes: usize, queue: &Queue) -> CxxResult<*mut u8> {
unsafe { ffi::aligned_alloc_device(alignment, num_bytes, &queue.0) }
}
}
pub struct HostAllocator;
impl UsmAllocatorKind for HostAllocator {
unsafe fn alloc(alignment: usize, num_bytes: usize, queue: &Queue) -> CxxResult<*mut u8> {
unsafe { ffi::aligned_alloc_host(alignment, num_bytes, &queue.0) }
}
}
unsafe impl HostAccessible for UsmAllocator<HostAllocator> {}
pub struct SharedAllocator;
impl UsmAllocatorKind for SharedAllocator {
unsafe fn alloc(alignment: usize, num_bytes: usize, queue: &Queue) -> CxxResult<*mut u8> {
unsafe { ffi::aligned_alloc_shared(alignment, num_bytes, &queue.0) }
}
}
unsafe impl HostAccessible for UsmAllocator<SharedAllocator> {}