use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use nvml_wrapper::Nvml;
use nvml_wrapper::enum_wrappers::device::{Clock, TemperatureSensor};
use nvml_wrapper::enums::device::UsedGpuMemory;
use crate::gpu::{
GpuBackend, GpuDeviceSnapshot, GpuProcessKind, GpuProcessSnapshot, GpuVendor, GpusSnapshot,
};
use crate::gpu_engine::{GpuEngine, GpuError};
pub struct NvmlEngine {
nvml: Arc<Nvml>,
driver_version: Option<String>,
}
impl NvmlEngine {
pub fn connect() -> Result<Self, GpuError> {
let nvml = Nvml::init().map_err(|e| GpuError::DriverUnavailable {
vendor: "NVIDIA",
reason: e.to_string(),
})?;
let count = nvml
.device_count()
.map_err(|e| GpuError::DriverUnavailable {
vendor: "NVIDIA",
reason: format!("device enumeration failed: {e}"),
})?;
if count == 0 {
return Err(GpuError::DriverUnavailable {
vendor: "NVIDIA",
reason: "NVML loaded but reports no devices".to_string(),
});
}
let driver_version = nvml.sys_driver_version().ok();
Ok(Self {
nvml: Arc::new(nvml),
driver_version,
})
}
}
#[async_trait]
impl GpuEngine for NvmlEngine {
async fn snapshot(&self) -> Result<GpusSnapshot, GpuError> {
let nvml = Arc::clone(&self.nvml);
let driver_version = self.driver_version.clone();
tokio::task::spawn_blocking(move || collect(&nvml, driver_version))
.await
.map_err(|e| {
GpuError::Query(format!(
"NVML collection task panicked or was cancelled: {e}"
))
})?
}
fn backend(&self) -> GpuBackend {
GpuBackend::Nvml
}
}
fn collect(nvml: &Nvml, driver_version: Option<String>) -> Result<GpusSnapshot, GpuError> {
let count = nvml
.device_count()
.map_err(|e| GpuError::Query(format!("nvmlDeviceGetCount failed: {e}")))?;
let mut devices = Vec::with_capacity(count as usize);
let mut processes = Vec::new();
for index in 0..count {
let device = match nvml.device_by_index(index) {
Ok(d) => d,
Err(err) => {
tracing::debug!(
target: "muxtop::gpu",
index,
error = %err,
"skipping unreadable NVML device"
);
continue;
}
};
let utilization = device.utilization_rates().ok();
let memory = device.memory_info().ok();
let power_limit_mw = device
.enforced_power_limit()
.or_else(|_| device.power_management_limit())
.ok();
devices.push(GpuDeviceSnapshot {
index: devices.len() as u32,
vendor: GpuVendor::Nvidia,
backend: GpuBackend::Nvml,
name: device.name().unwrap_or_default(),
bus_id: device.pci_info().map(|p| p.bus_id).unwrap_or_default(),
driver_version: driver_version.clone(),
utilization_pct: utilization.as_ref().map(|u| u.gpu as f32),
mem_utilization_pct: utilization.as_ref().map(|u| u.memory as f32),
mem_used_bytes: memory.as_ref().map(|m| m.used),
mem_total_bytes: memory.as_ref().map(|m| m.total),
temperature_c: device
.temperature(TemperatureSensor::Gpu)
.ok()
.map(|t| t as f32),
power_watts: device.power_usage().ok().map(|mw| mw as f32 / 1000.0),
power_limit_watts: power_limit_mw.map(|mw| mw as f32 / 1000.0),
graphics_clock_mhz: device.clock_info(Clock::Graphics).ok(),
memory_clock_mhz: device.clock_info(Clock::Memory).ok(),
fan_pct: device.fan_speed(0).ok().map(|f| f as f32),
encoder_pct: device
.encoder_utilization()
.ok()
.map(|u| u.utilization as f32),
decoder_pct: device
.decoder_utilization()
.ok()
.map(|u| u.utilization as f32),
supports_process_stats: true,
});
let device_index = (devices.len() - 1) as u32;
collect_processes(&device, device_index, &mut processes);
}
if devices.is_empty() {
return Ok(GpusSnapshot::unavailable_with(
"NVML reported devices but none could be read",
));
}
Ok(GpusSnapshot {
backends: vec![GpuBackend::Nvml],
available: true,
devices,
processes,
detail: String::new(),
})
}
fn collect_processes(
device: &nvml_wrapper::Device<'_>,
device_index: u32,
out: &mut Vec<GpuProcessSnapshot>,
) {
let mut merged: HashMap<u32, (GpuProcessKind, Option<u64>)> = HashMap::new();
let mut order: Vec<u32> = Vec::new();
fn absorb(
procs: Vec<nvml_wrapper::struct_wrappers::device::ProcessInfo>,
kind: GpuProcessKind,
merged: &mut HashMap<u32, (GpuProcessKind, Option<u64>)>,
order: &mut Vec<u32>,
) {
for proc_info in procs {
let mem = match proc_info.used_gpu_memory {
UsedGpuMemory::Unavailable => None,
UsedGpuMemory::Used(bytes) => Some(bytes),
};
merged
.entry(proc_info.pid)
.and_modify(|(k, m)| {
*k = k.merge(kind);
if m.is_none() {
*m = mem;
}
})
.or_insert_with(|| {
order.push(proc_info.pid);
(kind, mem)
});
}
}
if let Ok(procs) = device.running_compute_processes() {
absorb(procs, GpuProcessKind::Compute, &mut merged, &mut order);
}
if let Ok(procs) = device.running_graphics_processes() {
absorb(procs, GpuProcessKind::Graphics, &mut merged, &mut order);
}
for pid in order {
let (kind, mem_bytes) = merged[&pid];
out.push(GpuProcessSnapshot {
pid,
device_index,
name: String::new(),
kind,
mem_bytes,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn nvml_handle_is_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Nvml>();
assert_send_sync::<NvmlEngine>();
}
#[test]
fn connect_without_driver_reports_driver_unavailable() {
match NvmlEngine::connect() {
Ok(engine) => {
assert_eq!(engine.backend(), GpuBackend::Nvml);
}
Err(err) => {
assert!(
matches!(
err,
GpuError::DriverUnavailable {
vendor: "NVIDIA",
..
}
),
"expected DriverUnavailable, got {err:?}"
);
assert!(!err.to_string().is_empty());
}
}
}
#[tokio::test]
async fn snapshot_invariants_hold_on_real_hardware() {
let Ok(engine) = NvmlEngine::connect() else {
return;
};
let snap = engine.snapshot().await.expect("snapshot must not error");
assert!(snap.available);
assert!(!snap.devices.is_empty());
assert_eq!(snap.backends, vec![GpuBackend::Nvml]);
for (i, device) in snap.devices.iter().enumerate() {
assert_eq!(device.index, i as u32);
assert_eq!(device.vendor, GpuVendor::Nvidia);
assert!(device.supports_process_stats);
if let Some(pct) = device.utilization_pct {
assert!(
(0.0..=100.0).contains(&pct),
"utilisation out of range: {pct}"
);
}
if let (Some(used), Some(total)) = (device.mem_used_bytes, device.mem_total_bytes) {
assert!(used <= total, "used {used} exceeds total {total}");
}
}
for p in &snap.processes {
assert!(
snap.devices.iter().any(|d| d.index == p.device_index),
"process {} points at missing device {}",
p.pid,
p.device_index
);
}
assert!(
snap.processes.iter().all(|p| p.name.is_empty()),
"the NVML backend must not resolve process names itself"
);
}
#[test]
fn process_kind_merge_collapses_dual_use() {
let merged = GpuProcessKind::Compute.merge(GpuProcessKind::Graphics);
assert_eq!(merged, GpuProcessKind::Both);
}
}