use anyhow::{Result, bail};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use dynamo_config::parse_bool;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct NixlBackendConfig {
#[serde(default)]
backends: HashMap<String, HashMap<String, String>>,
}
impl NixlBackendConfig {
pub fn new(backends: HashMap<String, HashMap<String, String>>) -> Self {
Self { backends }
}
pub fn from_env() -> Result<Self> {
let mut backends = HashMap::new();
for (key, value) in std::env::vars() {
if let Some(remainder) = key.strip_prefix("DYN_KVBM_NIXL_BACKEND_") {
if remainder.contains('_') {
bail!(
"Custom NIXL backend parameters are not yet supported. \
Found: {}. Please use only DYN_KVBM_NIXL_BACKEND_<backend>=true \
to enable backends with default parameters.",
key
);
}
let backend_name = remainder.to_uppercase();
match parse_bool(&value) {
Ok(true) => {
backends.insert(backend_name, HashMap::new());
}
Ok(false) => {
continue;
}
Err(e) => bail!("Invalid value for {}: {}", key, e),
}
}
}
Ok(Self { backends })
}
pub fn with_backend(mut self, backend: impl Into<String>) -> Self {
self.backends
.insert(backend.into().to_uppercase(), HashMap::new());
self
}
pub fn with_backend_params(
mut self,
backend: impl Into<String>,
params: HashMap<String, String>,
) -> Self {
self.backends.insert(backend.into().to_uppercase(), params);
self
}
pub fn backends(&self) -> Vec<String> {
self.backends.keys().cloned().collect()
}
pub fn backend_params(&self, backend: &str) -> Option<&HashMap<String, String>> {
self.backends.get(&backend.to_uppercase())
}
pub fn has_backend(&self, backend: &str) -> bool {
self.backends.contains_key(&backend.to_uppercase())
}
pub fn merge(mut self, other: NixlBackendConfig) -> Self {
self.backends.extend(other.backends);
self
}
pub fn iter(&self) -> impl Iterator<Item = (&String, &HashMap<String, String>)> {
self.backends.iter()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_new_config_is_empty() {
let config = NixlBackendConfig::default();
assert_eq!(config.backends().len(), 0);
}
#[test]
fn test_default_is_empty() {
let config = NixlBackendConfig::default();
assert!(config.backends().is_empty()); }
#[test]
fn test_with_backend() {
let config = NixlBackendConfig::default()
.with_backend("ucx")
.with_backend("gds_mt");
assert!(config.has_backend("ucx"));
assert!(config.has_backend("UCX"));
assert!(config.has_backend("gds_mt"));
assert!(config.has_backend("GDS_MT"));
assert!(!config.has_backend("other"));
}
#[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 = NixlBackendConfig::default()
.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_merge_configs() {
let config1 = NixlBackendConfig::default().with_backend("ucx");
let config2 = NixlBackendConfig::default().with_backend("gds");
let merged = config1.merge(config2);
assert!(merged.has_backend("ucx"));
assert!(merged.has_backend("gds"));
}
#[test]
fn test_backend_name_case_insensitive() {
let config = NixlBackendConfig::default()
.with_backend("ucx")
.with_backend("Gds_mt")
.with_backend("OTHER");
assert!(config.has_backend("UCX"));
assert!(config.has_backend("ucx"));
assert!(config.has_backend("GDS_MT"));
assert!(config.has_backend("gds_mt"));
assert!(config.has_backend("OTHER"));
assert!(config.has_backend("other"));
}
#[test]
fn test_iter() {
let mut params = HashMap::new();
params.insert("key".to_string(), "value".to_string());
let config = NixlBackendConfig::default()
.with_backend("UCX")
.with_backend_params("GDS", params);
let items: Vec<_> = config.iter().collect();
assert_eq!(items.len(), 2);
}
}