use super::*;
use super::sync_manager::{BackendSyncable, SyncManager};
use std::ops::{Index, IndexMut};
use serde::{Serialize, Deserialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct XferDescriptor {
pub addr: usize,
pub len: usize,
pub dev_id: u64,
}
#[derive(Debug, Serialize, Deserialize)]
struct XferDescData {
mem_type: MemType,
descriptors: Vec<XferDescriptor>,
}
impl BackendSyncable for XferDescData {
type Backend = NonNull<bindings::nixl_capi_xfer_dlist_s>;
type Error = NixlError;
fn sync_to_backend(&self, backend: &Self::Backend) -> Result<(), Self::Error> {
let status = unsafe { nixl_capi_xfer_dlist_clear(backend.as_ptr()) };
match status {
NIXL_CAPI_SUCCESS => {}
NIXL_CAPI_ERROR_INVALID_PARAM => return Err(NixlError::InvalidParam),
_ => return Err(NixlError::BackendError),
}
for desc in &self.descriptors {
let status = unsafe {
nixl_capi_xfer_dlist_add_desc(backend.as_ptr(), desc.addr as uintptr_t, desc.len, desc.dev_id)
};
match status {
NIXL_CAPI_SUCCESS => {}
NIXL_CAPI_ERROR_INVALID_PARAM => return Err(NixlError::InvalidParam),
_ => return Err(NixlError::BackendError),
}
}
Ok(())
}
}
pub struct XferDescList<'a> {
sync_mgr: SyncManager<XferDescData>,
_phantom: PhantomData<&'a dyn NixlDescriptor>,
mem_type: MemType,
}
impl<'a> XferDescList<'a> {
pub fn new(mem_type: MemType) -> Result<Self, NixlError> {
let mut dlist = ptr::null_mut();
let status = unsafe {
nixl_capi_create_xfer_dlist(mem_type as nixl_capi_mem_type_t, &mut dlist)
};
match status {
NIXL_CAPI_SUCCESS => {
let backend = unsafe { NonNull::new_unchecked(dlist) };
let data = XferDescData {
mem_type,
descriptors: Vec::new(),
};
let sync_mgr = SyncManager::new(data, backend);
Ok(Self {
sync_mgr,
_phantom: PhantomData,
mem_type,
})
}
NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam),
_ => Err(NixlError::FailedToCreateXferDlistHandle),
}
}
pub fn as_ptr(&self) -> *mut bindings::nixl_capi_xfer_dlist_s {
self.sync_mgr.backend().map(|b| b.as_ptr()).unwrap_or(ptr::null_mut())
}
pub fn get_type(&self) -> Result<MemType, NixlError> { Ok(self.mem_type) }
pub fn add_desc(&mut self, addr: usize, len: usize, dev_id: u64) {
self.sync_mgr.data_mut().descriptors.push(XferDescriptor { addr, len, dev_id });
}
pub fn is_empty(&self) -> Result<bool, NixlError> {
Ok(self.len()? == 0)
}
pub fn desc_count(&self) -> Result<usize, NixlError> { Ok(self.sync_mgr.data().descriptors.len()) }
pub fn len(&self) -> Result<usize, NixlError> { Ok(self.sync_mgr.data().descriptors.len()) }
pub fn trim(&mut self) {
self.sync_mgr.data_mut().descriptors.shrink_to_fit();
}
pub fn rem_desc(&mut self, index: i32) -> Result<(), NixlError> {
if index < 0 { return Err(NixlError::InvalidParam); }
let idx = index as usize;
let data = self.sync_mgr.data_mut();
if idx >= data.descriptors.len() {
return Err(NixlError::InvalidParam);
}
data.descriptors.remove(idx);
Ok(())
}
pub fn clear(&mut self) {
self.sync_mgr.data_mut().descriptors.clear();
}
pub fn print(&self) -> Result<(), NixlError> {
let backend = self.sync_mgr.backend()?;
let status = unsafe { nixl_capi_xfer_dlist_print(backend.as_ptr()) };
match status {
NIXL_CAPI_SUCCESS => Ok(()),
NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam),
_ => Err(NixlError::BackendError),
}
}
pub fn resize(&mut self, new_size: usize) {
self.sync_mgr.data_mut().descriptors.resize(new_size, XferDescriptor {
addr: 0,
len: 0,
dev_id: 0,
});
}
pub fn get(&self, index: usize) -> Result<&XferDescriptor, NixlError> {
self.sync_mgr.data().descriptors
.get(index)
.ok_or(NixlError::InvalidParam)
}
pub fn get_mut(&mut self, index: usize) -> Result<&mut XferDescriptor, NixlError> {
self.sync_mgr.data_mut().descriptors
.get_mut(index)
.ok_or(NixlError::InvalidParam)
}
pub fn add_storage_desc<D: NixlDescriptor + 'a>(
&mut self,
desc: &'a D,
) -> Result<(), NixlError> {
let desc_mem_type = desc.mem_type();
let list_mem_type = if self.len().unwrap_or(0) > 0 { self.get_type().unwrap() } else { desc_mem_type };
if desc_mem_type != list_mem_type && list_mem_type != MemType::Unknown {
return Err(NixlError::InvalidParam);
}
let addr = unsafe { desc.as_ptr() } as usize;
let len = desc.size();
let dev_id = desc.device_id();
self.add_desc(addr, len, dev_id);
Ok(())
}
pub(crate) fn handle(&self) -> *mut bindings::nixl_capi_xfer_dlist_s {
self.sync_mgr.backend().map(|b| b.as_ptr()).unwrap_or(ptr::null_mut())
}
pub fn serialize(&self) -> Result<Vec<u8>, NixlError> {
bincode::serialize(self.sync_mgr.data()).map_err(|_| NixlError::BackendError)
}
pub fn deserialize(bytes: &[u8]) -> Result<Self, NixlError> {
let data: XferDescData = bincode::deserialize(bytes)
.map_err(|_| NixlError::FailedToCreateXferDlistHandle)?;
let mut list = XferDescList::new(data.mem_type)?;
for desc in data.descriptors {
list.add_desc(desc.addr, desc.len, desc.dev_id);
}
list.sync_mgr.backend()?;
Ok(list)
}
}
impl std::fmt::Debug for XferDescList<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mem_type = self.get_type().unwrap_or(MemType::Unknown);
let len = self.len().unwrap_or(0);
let desc_count = self.desc_count().unwrap_or(0);
f.debug_struct("XferDescList")
.field("mem_type", &mem_type)
.field("len", &len)
.field("desc_count", &desc_count)
.finish()
}
}
impl PartialEq for XferDescList<'_> {
fn eq(&self, other: &Self) -> bool {
if self.mem_type != other.mem_type {
return false;
}
self.sync_mgr.data().descriptors == other.sync_mgr.data().descriptors
}
}
impl Index<usize> for XferDescList<'_> {
type Output = XferDescriptor;
fn index(&self, index: usize) -> &Self::Output {
&self.sync_mgr.data().descriptors[index]
}
}
impl IndexMut<usize> for XferDescList<'_> {
fn index_mut(&mut self, index: usize) -> &mut Self::Output {
&mut self.sync_mgr.data_mut().descriptors[index]
}
}
impl Drop for XferDescList<'_> {
fn drop(&mut self) {
if let Ok(backend) = self.sync_mgr.backend() {
unsafe {
nixl_capi_destroy_xfer_dlist(backend.as_ptr());
}
}
}
}