use anyhow::Result;
use nixl_sys::Agent;
use std::collections::{HashMap, HashSet};
use crate::nixl::NixlBackendConfig;
#[derive(Clone, Debug)]
pub struct NixlAgent {
agent: Agent,
available_backends: HashSet<String>,
}
impl NixlAgent {
pub fn new(name: &str) -> Result<Self> {
let agent = Agent::new(name)?;
Ok(Self {
agent,
available_backends: HashSet::new(),
})
}
pub fn from_nixl_backend_config(name: &str, config: NixlBackendConfig) -> Result<Self> {
let mut agent = Self::new(name)?;
for (backend, params) in config.iter() {
agent.add_backend_with_params(backend, params)?;
}
Ok(agent)
}
pub fn add_backend(&mut self, backend: &str) -> Result<()> {
self.add_backend_with_params(backend, &HashMap::new())
}
pub fn add_backend_with_params(
&mut self,
backend: &str,
custom_params: &HashMap<String, String>,
) -> Result<()> {
let backend_upper = backend.to_uppercase();
if self.available_backends.contains(&backend_upper) {
return Ok(());
}
if !custom_params.is_empty() {
anyhow::bail!(
"Custom NIXL backend parameters for {} are not yet supported. \
This feature requires nixl_sys 0.9+. Params provided: {:?}",
backend_upper,
custom_params.keys().collect::<Vec<_>>()
);
}
let (_, params) = match self.agent.get_plugin_params(&backend_upper) {
Ok(result) => result,
Err(_) => anyhow::bail!("No {} plugin found", backend_upper),
};
match self.agent.create_backend(&backend_upper, ¶ms) {
Ok(_) => {
self.available_backends.insert(backend_upper);
Ok(())
}
Err(e) => anyhow::bail!("Failed to create nixl backend: {}", e),
}
}
pub fn with_backends(name: &str, backends: &[&str]) -> Result<Self> {
let mut agent = Self::new(name)?;
let mut failed_backends = Vec::new();
for backend in backends {
let backend_upper = backend.to_uppercase();
match agent.add_backend(&backend_upper) {
Ok(_) => {
tracing::debug!("Initialized NIXL backend: {}", backend_upper);
}
Err(e) => {
tracing::error!("Failed to initialize {} backend: {}", backend_upper, e);
failed_backends.push((backend_upper, e.to_string()));
}
}
}
if !failed_backends.is_empty() {
let error_details: Vec<String> = failed_backends
.iter()
.map(|(name, reason)| format!("{}: {}", name, reason))
.collect();
anyhow::bail!(
"Failed to initialize required backends: [{}]",
error_details.join(", ")
);
}
Ok(agent)
}
pub fn raw_agent(&self) -> &Agent {
&self.agent
}
pub fn into_raw_agent(self) -> Agent {
self.agent
}
pub fn has_backend(&self, backend: &str) -> bool {
self.available_backends.contains(&backend.to_uppercase())
}
pub fn backends(&self) -> &HashSet<String> {
&self.available_backends
}
pub fn require_backend(&self, backend: &str) -> Result<()> {
let backend_upper = backend.to_uppercase();
if self.has_backend(&backend_upper) {
Ok(())
} else {
anyhow::bail!(
"Operation requires {} backend, but it was not initialized. Available backends: {:?}",
backend_upper,
self.available_backends
)
}
}
}
impl std::ops::Deref for NixlAgent {
type Target = Agent;
fn deref(&self) -> &Self::Target {
&self.agent
}
}
#[cfg(all(test, feature = "testing-nixl"))]
mod tests {
use super::*;
#[test]
fn test_agent_backend_tracking() {
let agent = NixlAgent::with_backends("test", &["UCX"]).expect("Need UCX for test");
assert!(agent.has_backend("UCX"));
assert!(agent.has_backend("ucx")); }
#[test]
fn test_require_backend() {
let agent = NixlAgent::with_backends("test", &["UCX"]).expect("Need UCX for test");
assert!(agent.require_backend("UCX").is_ok());
assert!(agent.require_backend("GDS_MT").is_err());
}
#[test]
fn test_require_backends_strict() {
let agent =
NixlAgent::with_backends("test_strict", &["UCX"]).expect("Failed to require backends");
assert!(agent.has_backend("UCX"));
let result = NixlAgent::with_backends("test_strict_fail", &["UCX", "DUDE"]);
assert!(result.is_err());
}
#[test]
fn test_add_backend_with_empty_params() {
let mut agent = NixlAgent::new("test_empty_params").expect("Failed to create agent");
let result = agent.add_backend_with_params("UCX", &HashMap::new());
assert!(result.is_ok());
assert!(agent.has_backend("UCX"));
}
#[test]
fn test_add_backend_with_custom_params_fails() {
let mut agent = NixlAgent::new("test_custom_params").expect("Failed to create agent");
let mut params = HashMap::new();
params.insert("some_key".to_string(), "some_value".to_string());
let result = agent.add_backend_with_params("UCX", ¶ms);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(err_msg.contains("not yet supported"));
assert!(err_msg.contains("nixl_sys 0.9"));
assert!(err_msg.contains("some_key"));
}
#[test]
fn test_from_nixl_backend_config_with_custom_params_fails() {
let mut params = HashMap::new();
params.insert("threads".to_string(), "4".to_string());
let config = NixlBackendConfig::default().with_backend_params("UCX", params);
let result = NixlAgent::from_nixl_backend_config("test_config_params", config);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(err_msg.contains("not yet supported"));
assert!(err_msg.contains("threads"));
}
#[test]
fn test_from_nixl_backend_config_with_empty_params() {
let config = NixlBackendConfig::default().with_backend("UCX");
let result = NixlAgent::from_nixl_backend_config("test_config_empty", config);
assert!(result.is_ok());
let agent = result.unwrap();
assert!(agent.has_backend("UCX"));
}
}