use std::collections::HashMap;
use std::os::raw::c_uint;
use std::sync::Mutex;
use nvml_wrapper::Nvml;
use nvml_wrapper::enum_wrappers::nv_link::IntDeviceType;
use nvml_wrapper::error::{NvmlError, nvml_try};
use crate::device::types::{GpmMetrics, NvLinkRemoteDevice, NvLinkRemoteType};
pub const NVML_NVLINK_MAX_LINKS: u32 = 18;
#[derive(Debug, Clone, Default)]
pub struct HardwareDetails {
pub numa_node_id: Option<i32>,
pub gsp_firmware_mode: Option<u8>,
pub gsp_firmware_version: Option<String>,
}
fn is_permanent_unavailable(err: &NvmlError) -> bool {
matches!(err, NvmlError::NotSupported | NvmlError::FunctionNotFound)
}
type FieldCache<T> = Mutex<HashMap<u32, Option<T>>>;
pub struct HardwareDetailCache {
numa: FieldCache<i32>,
gsp_mode: FieldCache<u8>,
gsp_version: FieldCache<String>,
}
impl Default for HardwareDetailCache {
fn default() -> Self {
Self::new()
}
}
impl HardwareDetailCache {
pub fn new() -> Self {
Self {
numa: Mutex::new(HashMap::new()),
gsp_mode: Mutex::new(HashMap::new()),
gsp_version: Mutex::new(HashMap::new()),
}
}
pub fn get_or_fetch(&self, device: &nvml_wrapper::Device, index: u32) -> HardwareDetails {
HardwareDetails {
numa_node_id: self
.get_or_fetch_field(&self.numa, index, || numa_node_id_result(device)),
gsp_firmware_mode: self
.get_or_fetch_field(&self.gsp_mode, index, || gsp_firmware_mode_result(device)),
gsp_firmware_version: self.get_or_fetch_field(&self.gsp_version, index, || {
gsp_firmware_version_result(device)
}),
}
}
fn get_or_fetch_field<T, F>(&self, cache: &FieldCache<T>, index: u32, fetch: F) -> Option<T>
where
T: Clone,
F: FnOnce() -> Result<T, NvmlError>,
{
if let Ok(map) = cache.lock()
&& let Some(cached) = map.get(&index)
{
return cached.clone();
}
match fetch() {
Ok(value) => {
if let Ok(mut map) = cache.lock() {
map.insert(index, Some(value.clone()));
}
Some(value)
}
Err(ref e) if is_permanent_unavailable(e) => {
if let Ok(mut map) = cache.lock() {
map.insert(index, None);
}
None
}
Err(_transient) => {
None
}
}
}
}
fn numa_node_id_result(device: &nvml_wrapper::Device) -> Result<i32, NvmlError> {
let raw = device.numa_node_id()?;
if raw == u32::MAX {
return Err(NvmlError::NotSupported);
}
i32::try_from(raw).map_err(|_| NvmlError::NotSupported)
}
fn gsp_firmware_mode_result(device: &nvml_wrapper::Device) -> Result<u8, NvmlError> {
let mode = device.gsp_firmware_mode()?;
if mode.default {
Ok(2)
} else if mode.enabled {
Ok(1)
} else {
Ok(0)
}
}
fn gsp_firmware_version_result(device: &nvml_wrapper::Device) -> Result<String, NvmlError> {
let raw = device.gsp_firmware_version()?;
let trimmed = raw.trim_end_matches('\0').trim().to_string();
if trimmed.is_empty() {
return Err(NvmlError::NotSupported);
}
Ok(trimmed)
}
pub fn collect_nvlink_remote_devices(
nvml: &Nvml,
device: &nvml_wrapper::Device,
) -> Vec<NvLinkRemoteDevice> {
let mut out = Vec::new();
for link in 0..NVML_NVLINK_MAX_LINKS {
let link_wrapper = device.link_wrapper_for(link);
match link_wrapper.is_active() {
Ok(true) => {}
Ok(false) => continue,
Err(NvmlError::InvalidArg | NvmlError::NotSupported) => break,
Err(_transient) => continue,
}
let remote_type = match nvlink_remote_device_type_ffi(nvml, device, link) {
Some(t) => t,
None => NvLinkRemoteType::Unknown,
};
out.push(NvLinkRemoteDevice {
link_index: link,
remote_type,
bandwidth_mb_s: None,
});
}
out
}
fn nvlink_remote_device_type_ffi(
nvml: &Nvml,
device: &nvml_wrapper::Device,
link: u32,
) -> Option<NvLinkRemoteType> {
let sym = nvml
.lib()
.nvmlDeviceGetNvLinkRemoteDeviceType
.as_ref()
.ok()?;
unsafe {
let mut value: c_uint = 0;
let rc = sym(device.handle(), link, &mut value);
nvml_try(rc).ok()?;
Some(map_remote_device_type(value))
}
}
fn map_remote_device_type(value: c_uint) -> NvLinkRemoteType {
match value {
0 => NvLinkRemoteType::Gpu,
1 => NvLinkRemoteType::IbmNpu,
2 => NvLinkRemoteType::Switch,
_ => NvLinkRemoteType::Unknown,
}
}
#[allow(dead_code)]
pub(crate) fn nvlink_remote_type_from_wrapper(value: IntDeviceType) -> NvLinkRemoteType {
match value {
IntDeviceType::Gpu => NvLinkRemoteType::Gpu,
IntDeviceType::Ibmnpu => NvLinkRemoteType::IbmNpu,
IntDeviceType::Switch => NvLinkRemoteType::Switch,
IntDeviceType::Unknown => NvLinkRemoteType::Unknown,
}
}
pub fn gpm_is_supported(device: &nvml_wrapper::Device) -> bool {
device.gpm_support().unwrap_or(false)
}
pub fn collect_gpm_metrics(device: &nvml_wrapper::Device) -> Option<GpmMetrics> {
if !gpm_is_supported(device) {
return None;
}
Some(GpmMetrics::default())
}
#[allow(dead_code)]
fn err_unsupported() -> NvmlError {
NvmlError::NotSupported
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn field_cache_permanent_error_caches_none() {
let cache: FieldCache<i32> = Mutex::new(HashMap::new());
let hw = HardwareDetailCache::new();
let mut call_count = 0u32;
let result = hw.get_or_fetch_field(&cache, 0, || {
call_count += 1;
Err(NvmlError::NotSupported)
});
assert!(result.is_none());
assert_eq!(call_count, 1);
let result2 = hw.get_or_fetch_field(&cache, 0, || {
call_count += 1;
Ok(42) });
assert!(result2.is_none(), "cached None should persist");
assert_eq!(call_count, 1, "fetcher must not be called on cache hit");
}
#[test]
fn field_cache_transient_error_is_not_cached() {
let cache: FieldCache<i32> = Mutex::new(HashMap::new());
let hw = HardwareDetailCache::new();
let mut call_count = 0u32;
let result = hw.get_or_fetch_field(&cache, 0, || {
call_count += 1;
Err(NvmlError::Unknown)
});
assert!(result.is_none());
assert_eq!(call_count, 1);
let result2 = hw.get_or_fetch_field(&cache, 0, || {
call_count += 1;
Ok(7)
});
assert_eq!(result2, Some(7));
assert_eq!(
call_count, 2,
"fetcher must be called again after transient miss"
);
}
#[test]
fn field_cache_success_is_cached() {
let cache: FieldCache<i32> = Mutex::new(HashMap::new());
let hw = HardwareDetailCache::new();
let mut call_count = 0u32;
let _ = hw.get_or_fetch_field(&cache, 0, || {
call_count += 1;
Ok(42)
});
let result2 = hw.get_or_fetch_field(&cache, 0, || {
call_count += 1;
Ok(99)
});
assert_eq!(result2, Some(42));
assert_eq!(call_count, 1, "fetcher must not be called on cache hit");
}
#[test]
fn field_caches_are_independent() {
let hw = HardwareDetailCache::new();
let _ = hw.get_or_fetch_field(&hw.gsp_mode, 0, || Ok(2u8));
let numa_result = hw.get_or_fetch_field(&hw.numa, 0, || Err(NvmlError::Unknown));
assert!(numa_result.is_none());
let mut mode_calls = 0u32;
let mode_result = hw.get_or_fetch_field(&hw.gsp_mode, 0, || {
mode_calls += 1;
Ok(99u8) });
assert_eq!(mode_result, Some(2u8));
assert_eq!(
mode_calls, 0,
"gsp_mode fetcher must not be called — already cached"
);
}
#[test]
fn permanent_unavailable_classification() {
assert!(is_permanent_unavailable(&NvmlError::NotSupported));
assert!(is_permanent_unavailable(&NvmlError::FunctionNotFound));
assert!(!is_permanent_unavailable(&NvmlError::Unknown));
assert!(!is_permanent_unavailable(&NvmlError::GpuLost));
assert!(!is_permanent_unavailable(&NvmlError::DriverNotLoaded));
}
#[test]
fn remote_device_type_mapping_is_stable() {
assert_eq!(map_remote_device_type(0), NvLinkRemoteType::Gpu);
assert_eq!(map_remote_device_type(1), NvLinkRemoteType::IbmNpu);
assert_eq!(map_remote_device_type(2), NvLinkRemoteType::Switch);
assert_eq!(map_remote_device_type(255), NvLinkRemoteType::Unknown);
}
#[test]
fn remote_device_type_unknown_future_values_degrade_to_unknown() {
assert_eq!(map_remote_device_type(17), NvLinkRemoteType::Unknown);
assert_eq!(map_remote_device_type(u32::MAX), NvLinkRemoteType::Unknown);
}
#[test]
fn nvlink_remote_type_label_round_trip() {
for v in [
NvLinkRemoteType::Gpu,
NvLinkRemoteType::IbmNpu,
NvLinkRemoteType::Switch,
NvLinkRemoteType::Unknown,
] {
assert_eq!(NvLinkRemoteType::from_label(v.as_label()), v);
}
}
#[test]
fn nvlink_remote_type_from_label_unknown_inputs_degrade() {
assert_eq!(NvLinkRemoteType::from_label(""), NvLinkRemoteType::Unknown);
assert_eq!(
NvLinkRemoteType::from_label("garbage"),
NvLinkRemoteType::Unknown
);
}
#[test]
fn nvlink_max_links_matches_nvml_header() {
assert_eq!(NVML_NVLINK_MAX_LINKS, 18);
}
#[test]
fn wrapper_enum_mapping_covers_every_variant() {
assert_eq!(
nvlink_remote_type_from_wrapper(IntDeviceType::Gpu),
NvLinkRemoteType::Gpu
);
assert_eq!(
nvlink_remote_type_from_wrapper(IntDeviceType::Ibmnpu),
NvLinkRemoteType::IbmNpu
);
assert_eq!(
nvlink_remote_type_from_wrapper(IntDeviceType::Switch),
NvLinkRemoteType::Switch
);
assert_eq!(
nvlink_remote_type_from_wrapper(IntDeviceType::Unknown),
NvLinkRemoteType::Unknown
);
}
}