use dynamo_memory::nixl::NixlBackendConfig;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use validator::Validate;
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct NixlConfig {
#[serde(default = "default_backends")]
pub backends: HashMap<String, HashMap<String, String>>,
}
fn default_backends() -> HashMap<String, HashMap<String, String>> {
let mut backends = HashMap::new();
backends.insert("UCX".to_string(), HashMap::new());
backends.insert("POSIX".to_string(), HashMap::new());
backends
}
impl Default for NixlConfig {
fn default() -> Self {
Self {
backends: default_backends(),
}
}
}
impl NixlConfig {
pub fn new(backends: HashMap<String, HashMap<String, String>>) -> Self {
Self { backends }
}
pub fn empty() -> Self {
Self {
backends: HashMap::new(),
}
}
pub fn from_nixl_backend_config(config: NixlBackendConfig) -> Self {
let backends: HashMap<String, HashMap<String, String>> = config
.iter()
.map(|(backend, params)| (backend.to_string(), params.clone()))
.collect();
Self { backends }
}
pub fn with_backend(mut self, name: impl Into<String>) -> Self {
self.backends
.insert(name.into().to_uppercase(), HashMap::new());
self
}
pub fn with_backend_params(
mut self,
name: impl Into<String>,
params: HashMap<String, String>,
) -> Self {
self.backends.insert(name.into().to_uppercase(), params);
self
}
pub fn enabled_backends(&self) -> Vec<&String> {
self.backends.keys().collect()
}
pub fn has_backend(&self, backend: &str) -> bool {
self.backends.contains_key(&backend.to_uppercase())
}
pub fn backend_params(&self, backend: &str) -> Option<&HashMap<String, String>> {
self.backends.get(&backend.to_uppercase())
}
pub fn iter(&self) -> impl Iterator<Item = (&String, &HashMap<String, String>)> {
self.backends.iter()
}
}
impl From<NixlConfig> for NixlBackendConfig {
fn from(config: NixlConfig) -> Self {
NixlBackendConfig::new(config.backends)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config() {
let config = NixlConfig::default();
assert!(config.has_backend("UCX"));
assert!(!config.has_backend("GDS"));
}
#[test]
fn test_new_default() {
let config = NixlConfig::default();
assert!(config.has_backend("UCX"));
assert!(config.has_backend("POSIX"));
assert!(!config.enabled_backends().is_empty());
}
#[test]
fn test_with_backend() {
let config = NixlConfig::empty().with_backend("ucx").with_backend("gds");
assert!(config.has_backend("UCX"));
assert!(config.has_backend("GDS"));
assert!(!config.has_backend("POSIX"));
assert!(config.backends.contains_key("UCX"));
assert!(config.backends.contains_key("GDS"));
}
#[test]
fn test_with_backend_params() {
let mut params = HashMap::new();
params.insert("threads".to_string(), "4".to_string());
params.insert("buffer_size".to_string(), "1048576".to_string());
let config = NixlConfig::empty()
.with_backend("UCX")
.with_backend_params("GDS", params);
let ucx_params = config.backend_params("UCX").unwrap();
assert!(ucx_params.is_empty());
let gds_params = config.backend_params("GDS").unwrap();
assert_eq!(gds_params.get("threads"), Some(&"4".to_string()));
assert_eq!(gds_params.get("buffer_size"), Some(&"1048576".to_string()));
}
#[test]
fn test_lookup_normalizes_to_uppercase() {
let config = NixlConfig::empty().with_backend("ucx");
assert!(config.has_backend("ucx"));
assert!(config.has_backend("UCX"));
assert!(config.has_backend("Ucx"));
assert!(config.backend_params("ucx").is_some());
assert!(config.backend_params("UCX").is_some());
}
#[test]
fn test_enabled_backends() {
let config = NixlConfig::empty().with_backend("ucx").with_backend("gds");
let backends = config.enabled_backends();
assert_eq!(backends.len(), 2);
assert!(backends.contains(&&"UCX".to_string()));
assert!(backends.contains(&&"GDS".to_string()));
}
#[test]
fn test_iter() {
let mut params = HashMap::new();
params.insert("key".to_string(), "value".to_string());
let config = NixlConfig::empty()
.with_backend("UCX")
.with_backend_params("GDS", params);
let items: Vec<_> = config.iter().collect();
assert_eq!(items.len(), 2);
}
#[test]
fn test_serde_roundtrip() {
let mut params = HashMap::new();
params.insert("threads".to_string(), "4".to_string());
let config = NixlConfig::empty()
.with_backend("UCX")
.with_backend_params("GDS", params);
let json = serde_json::to_string(&config).unwrap();
let parsed: NixlConfig = serde_json::from_str(&json).unwrap();
assert!(parsed.has_backend("UCX"));
assert!(parsed.has_backend("GDS"));
assert_eq!(
parsed.backend_params("GDS").unwrap().get("threads"),
Some(&"4".to_string())
);
}
}