mod agent;
mod config;
use super::{MemoryDescriptor, StorageKind};
use std::any::Any;
use std::fmt;
use std::sync::Arc;
pub use agent::NixlAgent;
pub use config::NixlBackendConfig;
pub use nixl_sys::{
Agent, MemType, NotificationMap, OptArgs, RegistrationHandle, XferDescList, XferOp,
XferRequest, is_stub,
};
pub use serde::{Deserialize, Serialize};
pub trait NixlCompatible {
fn nixl_params(&self) -> (*const u8, usize, MemType, u64);
}
pub trait NixlMemory: MemoryDescriptor + NixlCompatible {}
impl<T: MemoryDescriptor + NixlCompatible + ?Sized> NixlMemory for T {}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NixlDescriptor {
pub addr: u64,
pub size: usize,
pub mem_type: MemType,
pub device_id: u64,
}
impl nixl_sys::MemoryRegion for NixlDescriptor {
unsafe fn as_ptr(&self) -> *const u8 {
self.addr as *const u8
}
fn size(&self) -> usize {
self.size
}
}
impl nixl_sys::NixlDescriptor for NixlDescriptor {
fn mem_type(&self) -> MemType {
self.mem_type
}
fn device_id(&self) -> u64 {
self.device_id
}
}
pub trait RegisteredView {
fn agent_name(&self) -> &str;
fn descriptor(&self) -> NixlDescriptor;
}
pub struct NixlRegistered<S: NixlCompatible> {
storage: S,
handle: Option<RegistrationHandle>,
agent_name: String,
}
impl<S: NixlCompatible> Drop for NixlRegistered<S> {
fn drop(&mut self) {
drop(self.handle.take());
}
}
impl<S: NixlCompatible + fmt::Debug> fmt::Debug for NixlRegistered<S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NixlRegistered")
.field("storage", &self.storage)
.field("agent_name", &self.agent_name)
.field("handle", &self.handle.is_some())
.finish()
}
}
impl<S: MemoryDescriptor + NixlCompatible + 'static> MemoryDescriptor for NixlRegistered<S> {
fn addr(&self) -> usize {
self.storage.addr()
}
fn size(&self) -> usize {
self.storage.size()
}
fn storage_kind(&self) -> StorageKind {
self.storage.storage_kind()
}
fn as_any(&self) -> &dyn Any {
self
}
fn nixl_descriptor(&self) -> Option<NixlDescriptor> {
Some(self.descriptor())
}
}
impl<S: MemoryDescriptor + NixlCompatible> RegisteredView for NixlRegistered<S> {
fn agent_name(&self) -> &str {
&self.agent_name
}
fn descriptor(&self) -> NixlDescriptor {
let (ptr, size, mem_type, device_id) = self.storage.nixl_params();
NixlDescriptor {
addr: ptr as u64,
size,
mem_type,
device_id,
}
}
}
impl<S: MemoryDescriptor + NixlCompatible> NixlRegistered<S> {
pub fn storage(&self) -> &S {
&self.storage
}
pub fn storage_mut(&mut self) -> &mut S {
&mut self.storage
}
pub fn is_registered(&self) -> bool {
self.handle.is_some()
}
pub fn into_storage(mut self) -> S {
drop(self.handle.take());
let mut this = std::mem::ManuallyDrop::new(self);
unsafe {
let storage = std::ptr::read(&this.storage);
std::ptr::drop_in_place(&mut this.agent_name);
storage
}
}
}
pub fn register_with_nixl<S>(
storage: S,
agent: &Agent,
opt: Option<&OptArgs>,
) -> std::result::Result<NixlRegistered<S>, S>
where
S: MemoryDescriptor + NixlCompatible,
{
if storage.nixl_descriptor().is_some() {
return Ok(NixlRegistered {
storage,
handle: None,
agent_name: agent.name().to_string(),
});
}
let (ptr, size, mem_type, device_id) = storage.nixl_params();
let descriptor = NixlDescriptor {
addr: ptr as u64,
size,
mem_type,
device_id,
};
match agent.register_memory(&descriptor, opt) {
Ok(handle) => Ok(NixlRegistered {
storage,
handle: Some(handle),
agent_name: agent.name().to_string(),
}),
Err(_) => Err(storage),
}
}
impl NixlCompatible for Arc<dyn NixlMemory + Send + Sync> {
fn nixl_params(&self) -> (*const u8, usize, MemType, u64) {
(**self).nixl_params()
}
}
impl MemoryDescriptor for Arc<dyn NixlMemory + Send + Sync> {
fn addr(&self) -> usize {
(**self).addr()
}
fn size(&self) -> usize {
(**self).size()
}
fn storage_kind(&self) -> StorageKind {
(**self).storage_kind()
}
fn as_any(&self) -> &dyn Any {
(**self).as_any()
}
fn nixl_descriptor(&self) -> Option<NixlDescriptor> {
(**self).nixl_descriptor()
}
}
pub trait NixlRegisterExt: MemoryDescriptor + NixlCompatible + Sized {
fn register(
self,
agent: &NixlAgent,
opt: Option<&OptArgs>,
) -> std::result::Result<NixlRegistered<Self>, Self> {
register_with_nixl(self, agent, opt)
}
}
impl<T: MemoryDescriptor + NixlCompatible + Sized> NixlRegisterExt for T {}