dynamo_memory/nixl/
agent.rs1use anyhow::Result;
11use nixl_sys::{Agent, is_stub};
12use std::collections::{HashMap, HashSet};
13
14use crate::nixl::NixlBackendConfig;
15
16#[derive(Clone, Debug)]
29pub struct NixlAgent {
30 agent: Agent,
31 available_backends: HashSet<String>,
32}
33
34impl NixlAgent {
35 pub fn new(name: &str) -> Result<Self> {
37 if is_stub() {
38 return Err(anyhow::anyhow!("NIXL is not supported in stub mode"));
39 }
40 let agent = Agent::new(name)?;
41
42 Ok(Self {
43 agent,
44 available_backends: HashSet::new(),
45 })
46 }
47
48 pub fn from_nixl_backend_config(name: &str, config: NixlBackendConfig) -> Result<Self> {
54 let mut agent = Self::new(name)?;
55 for (backend, params) in config.iter() {
56 agent.add_backend_with_params(backend, params)?;
57 }
58 Ok(agent)
59 }
60
61 pub fn add_backend(&mut self, backend: &str) -> Result<()> {
63 self.add_backend_with_params(backend, &HashMap::new())
64 }
65
66 pub fn add_backend_with_params(
74 &mut self,
75 backend: &str,
76 custom_params: &HashMap<String, String>,
77 ) -> Result<()> {
78 let backend_upper = backend.to_uppercase();
79 if self.available_backends.contains(&backend_upper) {
80 return Ok(());
81 }
82
83 if !custom_params.is_empty() {
85 anyhow::bail!(
86 "Custom NIXL backend parameters for {} are not yet supported. \
87 This feature requires nixl_sys 0.9+. Params provided: {:?}",
88 backend_upper,
89 custom_params.keys().collect::<Vec<_>>()
90 );
91 }
92
93 let (_, params) = match self.agent.get_plugin_params(&backend_upper) {
95 Ok(result) => result,
96 Err(_) => anyhow::bail!("No {} plugin found", backend_upper),
97 };
98
99 match self.agent.create_backend(&backend_upper, ¶ms) {
100 Ok(_) => {
101 self.available_backends.insert(backend_upper);
102 Ok(())
103 }
104 Err(e) => anyhow::bail!("Failed to create nixl backend: {}", e),
105 }
106 }
107
108 pub fn with_backends(name: &str, backends: &[&str]) -> Result<Self> {
126 let mut agent = Self::new(name)?;
127 let mut failed_backends = Vec::new();
128
129 for backend in backends {
130 let backend_upper = backend.to_uppercase();
131 match agent.add_backend(&backend_upper) {
132 Ok(_) => {
133 tracing::debug!("Initialized NIXL backend: {}", backend_upper);
134 }
135 Err(e) => {
136 tracing::error!("Failed to initialize {} backend: {}", backend_upper, e);
137 failed_backends.push((backend_upper, e.to_string()));
138 }
139 }
140 }
141
142 if !failed_backends.is_empty() {
143 let error_details: Vec<String> = failed_backends
144 .iter()
145 .map(|(name, reason)| format!("{}: {}", name, reason))
146 .collect();
147
148 anyhow::bail!(
149 "Failed to initialize required backends: [{}]",
150 error_details.join(", ")
151 );
152 }
153
154 Ok(agent)
155 }
156
157 pub fn raw_agent(&self) -> &Agent {
159 &self.agent
160 }
161
162 pub fn into_raw_agent(self) -> Agent {
167 self.agent
168 }
169
170 pub fn has_backend(&self, backend: &str) -> bool {
172 self.available_backends.contains(&backend.to_uppercase())
173 }
174
175 pub fn backends(&self) -> &HashSet<String> {
177 &self.available_backends
178 }
179
180 pub fn require_backend(&self, backend: &str) -> Result<()> {
188 let backend_upper = backend.to_uppercase();
189 if self.has_backend(&backend_upper) {
190 Ok(())
191 } else {
192 anyhow::bail!(
193 "Operation requires {} backend, but it was not initialized. Available backends: {:?}",
194 backend_upper,
195 self.available_backends
196 )
197 }
198 }
199}
200
201impl std::ops::Deref for NixlAgent {
203 type Target = Agent;
204
205 fn deref(&self) -> &Self::Target {
206 &self.agent
207 }
208}
209
210#[cfg(all(test, feature = "testing-nixl"))]
211mod tests {
212 use super::*;
213
214 #[test]
215 fn test_agent_backend_tracking() {
216 let agent = NixlAgent::with_backends("test", &["UCX"]).expect("Need UCX for test");
218
219 assert!(agent.has_backend("UCX"));
221 assert!(agent.has_backend("ucx")); }
223
224 #[test]
225 fn test_require_backend() {
226 let agent = NixlAgent::with_backends("test", &["UCX"]).expect("Need UCX for test");
227
228 assert!(agent.require_backend("UCX").is_ok());
230
231 assert!(agent.require_backend("GDS_MT").is_err());
233 }
234
235 #[test]
236 fn test_require_backends_strict() {
237 let agent =
239 NixlAgent::with_backends("test_strict", &["UCX"]).expect("Failed to require backends");
240 assert!(agent.has_backend("UCX"));
241
242 let result = NixlAgent::with_backends("test_strict_fail", &["UCX", "DUDE"]);
244 assert!(result.is_err());
245 }
246
247 #[test]
248 fn test_add_backend_with_empty_params() {
249 let mut agent = NixlAgent::new("test_empty_params").expect("Failed to create agent");
250
251 let result = agent.add_backend_with_params("UCX", &HashMap::new());
253 assert!(result.is_ok());
254 assert!(agent.has_backend("UCX"));
255 }
256
257 #[test]
258 fn test_add_backend_with_custom_params_fails() {
259 let mut agent = NixlAgent::new("test_custom_params").expect("Failed to create agent");
260
261 let mut params = HashMap::new();
263 params.insert("some_key".to_string(), "some_value".to_string());
264
265 let result = agent.add_backend_with_params("UCX", ¶ms);
266 assert!(result.is_err());
267
268 let err_msg = result.unwrap_err().to_string();
269 assert!(err_msg.contains("not yet supported"));
270 assert!(err_msg.contains("nixl_sys 0.9"));
271 assert!(err_msg.contains("some_key"));
272 }
273
274 #[test]
275 fn test_from_nixl_backend_config_with_custom_params_fails() {
276 let mut params = HashMap::new();
278 params.insert("threads".to_string(), "4".to_string());
279
280 let config = NixlBackendConfig::default().with_backend_params("UCX", params);
281
282 let result = NixlAgent::from_nixl_backend_config("test_config_params", config);
283 assert!(result.is_err());
284
285 let err_msg = result.unwrap_err().to_string();
286 assert!(err_msg.contains("not yet supported"));
287 assert!(err_msg.contains("threads"));
288 }
289
290 #[test]
291 fn test_from_nixl_backend_config_with_empty_params() {
292 let config = NixlBackendConfig::default().with_backend("UCX");
294
295 let result = NixlAgent::from_nixl_backend_config("test_config_empty", config);
296 assert!(result.is_ok());
297
298 let agent = result.unwrap();
299 assert!(agent.has_backend("UCX"));
300 }
301}