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 RegDescriptor {
pub addr: usize,
pub len: usize,
pub dev_id: u64,
pub metadata: Vec<u8>,
}
#[derive(Debug, Serialize, Deserialize)]
struct RegDescData {
mem_type: MemType,
descriptors: Vec<RegDescriptor>,
}
impl BackendSyncable for RegDescData {
type Backend = NonNull<bindings::nixl_capi_reg_dlist_s>;
type Error = NixlError;
fn sync_to_backend(&self, backend: &Self::Backend) -> Result<(), Self::Error> {
let status = unsafe { nixl_capi_reg_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_reg_dlist_add_desc(
backend.as_ptr(),
desc.addr as uintptr_t,
desc.len,
desc.dev_id,
desc.metadata.as_ptr() as *const std::ffi::c_void,
desc.metadata.len(),
)
};
match status {
NIXL_CAPI_SUCCESS => {}
NIXL_CAPI_ERROR_INVALID_PARAM => return Err(NixlError::InvalidParam),
_ => return Err(NixlError::BackendError),
}
}
Ok(())
}
}
pub struct RegDescList<'a> {
sync_mgr: SyncManager<RegDescData>,
_phantom: PhantomData<&'a dyn NixlDescriptor>,
mem_type: MemType,
}
impl<'a> RegDescList<'a> {
pub fn new(mem_type: MemType) -> Result<Self, NixlError> {
let mut dlist = ptr::null_mut();
let status = unsafe {
nixl_capi_create_reg_dlist(mem_type as nixl_capi_mem_type_t, &mut dlist)
};
match status {
NIXL_CAPI_SUCCESS => {
if dlist.is_null() {
tracing::error!("Failed to create registration descriptor list");
return Err(NixlError::RegDescListCreationFailed);
}
let backend = NonNull::new(dlist).ok_or(NixlError::RegDescListCreationFailed)?;
let data = RegDescData {
mem_type,
descriptors: Vec::new(),
};
let sync_mgr = SyncManager::new(data, backend);
Ok(Self {
sync_mgr,
_phantom: PhantomData,
mem_type,
})
}
_ => Err(NixlError::RegDescListCreationFailed),
}
}
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.add_desc_with_meta(addr, len, dev_id, &[])
}
pub fn add_desc_with_meta(
&mut self,
addr: usize,
len: usize,
dev_id: u64,
metadata: &[u8],
) {
self.sync_mgr.data_mut().descriptors.push(RegDescriptor {
addr,
len,
dev_id,
metadata: metadata.to_vec(),
});
}
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 print(&self) -> Result<(), NixlError> {
let backend = self.sync_mgr.backend()?;
let status = unsafe { nixl_capi_reg_dlist_print(backend.as_ptr()) };
match status {
NIXL_CAPI_SUCCESS => Ok(()),
NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam),
_ => Err(NixlError::BackendError),
}
}
pub fn clear(&mut self) {
self.sync_mgr.data_mut().descriptors.clear();
}
pub fn resize(&mut self, new_size: usize) {
self.sync_mgr.data_mut().descriptors.resize(new_size, RegDescriptor {
addr: 0,
len: 0,
dev_id: 0,
metadata: Vec::new(),
});
}
pub fn get(&self, index: usize) -> Result<&RegDescriptor, NixlError> {
self.sync_mgr.data().descriptors
.get(index)
.ok_or(NixlError::InvalidParam)
}
pub fn get_mut(&mut self, index: usize) -> Result<&mut RegDescriptor, NixlError> {
self.sync_mgr.data_mut().descriptors
.get_mut(index)
.ok_or(NixlError::InvalidParam)
}
pub fn add_storage_desc(&mut self, desc: &'a dyn NixlDescriptor) -> Result<(), NixlError> {
let desc_mem_type = desc.mem_type();
let list_mem_type = if self.len()? > 0 {
self.get_type()?
} 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_reg_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: RegDescData = bincode::deserialize(bytes)
.map_err(|_| NixlError::RegDescListCreationFailed)?;
let mut list = RegDescList::new(data.mem_type)?;
for desc in data.descriptors {
list.add_desc_with_meta(desc.addr, desc.len, desc.dev_id, &desc.metadata);
}
list.sync_mgr.backend()?;
Ok(list)
}
}
impl std::fmt::Debug for RegDescList<'_> {
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("RegDescList")
.field("mem_type", &mem_type)
.field("len", &len)
.field("desc_count", &desc_count)
.finish()
}
}
impl PartialEq for RegDescList<'_> {
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 RegDescList<'_> {
type Output = RegDescriptor;
fn index(&self, index: usize) -> &Self::Output {
&self.sync_mgr.data().descriptors[index]
}
}
impl IndexMut<usize> for RegDescList<'_> {
fn index_mut(&mut self, index: usize) -> &mut Self::Output {
&mut self.sync_mgr.data_mut().descriptors[index]
}
}
impl Drop for RegDescList<'_> {
fn drop(&mut self) {
tracing::trace!("Dropping registration descriptor list");
if let Ok(backend) = self.sync_mgr.backend() {
unsafe {
nixl_capi_destroy_reg_dlist(backend.as_ptr());
}
}
tracing::trace!("Registration descriptor list dropped");
}
}