use std::process::Command;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GpuInfo {
pub index: u8,
pub name: String,
pub sm_major: u32,
pub sm_minor: u32,
pub total_memory_mb: u64,
}
impl GpuInfo {
pub fn sm_version(&self) -> String {
format!("sm_{}{}", self.sm_major, self.sm_minor)
}
pub fn vram_bytes(&self) -> u64 {
self.total_memory_mb * 1024 * 1024
}
pub fn short_name(&self) -> String {
self.name.replace("NVIDIA ", "").replace("GeForce ", "")
}
}
pub fn detect_gpus() -> Vec<GpuInfo> {
let all = detect_gpus_raw();
let Ok(visible) = std::env::var("CUDA_VISIBLE_DEVICES") else {
return all;
};
let trimmed = visible.trim();
if trimmed.is_empty() {
return Vec::new();
}
let mut allowed: std::collections::HashSet<u8> = std::collections::HashSet::new();
for entry in trimmed.split(',') {
let entry = entry.trim();
match entry.parse::<u8>() {
Ok(idx) => {
allowed.insert(idx);
}
Err(_) => {
eprintln!(
"flodl sys: CUDA_VISIBLE_DEVICES entry '{entry}' is not a \
numeric index (UUID/MIG forms are not resolved by \
detect_gpus); GPU detection may under-count"
);
}
}
}
all.into_iter()
.filter(|g| allowed.contains(&g.index))
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemInfo {
pub total_bytes: u64,
pub available_bytes: u64,
}
pub fn mem_info() -> Option<MemInfo> {
parse_meminfo(&std::fs::read_to_string("/proc/meminfo").ok()?)
}
fn parse_meminfo(text: &str) -> Option<MemInfo> {
let mut total = None;
let mut available = None;
for line in text.lines() {
if let Some(rest) = line.strip_prefix("MemTotal:") {
total = parse_meminfo_kb(rest);
} else if let Some(rest) = line.strip_prefix("MemAvailable:") {
available = parse_meminfo_kb(rest);
}
if total.is_some() && available.is_some() {
break;
}
}
Some(MemInfo {
total_bytes: total? * 1024,
available_bytes: available? * 1024,
})
}
fn parse_meminfo_kb(rest: &str) -> Option<u64> {
rest.trim().strip_suffix("kB")?.trim().parse().ok()
}
fn detect_gpus_raw() -> Vec<GpuInfo> {
let output = match Command::new("nvidia-smi")
.args([
"--query-gpu=index,compute_cap,memory.total,name",
"--format=csv,noheader,nounits",
])
.output()
{
Ok(o) if o.status.success() => o,
Ok(o) => {
eprintln!(
"flodl sys: nvidia-smi exited non-zero ({}); reporting 0 GPUs. \
stderr: {}",
o.status,
String::from_utf8_lossy(&o.stderr).trim(),
);
return Vec::new();
}
Err(_) => return Vec::new(), };
let stdout = String::from_utf8_lossy(&output.stdout);
stdout
.lines()
.filter(|line| !line.trim().is_empty())
.filter_map(|line| {
let parsed = parse_gpu_csv_row(line);
if parsed.is_none() {
eprintln!("flodl sys: could not parse an nvidia-smi GPU row, skipping: {line:?}");
}
parsed
})
.collect()
}
fn parse_gpu_csv_row(line: &str) -> Option<GpuInfo> {
let parts: Vec<&str> = line.splitn(4, ", ").collect();
if parts.len() < 4 {
return None;
}
let index: u8 = parts[0].trim().parse().ok()?;
let cap_parts: Vec<&str> = parts[1].trim().split('.').collect();
let sm_major: u32 = cap_parts.first()?.parse().ok()?;
let sm_minor: u32 = cap_parts.get(1)?.parse().ok()?;
let total_memory_mb: u64 = parts[2].trim().parse().ok()?;
let name = parts[3].trim().to_string();
Some(GpuInfo {
index,
name,
sm_major,
sm_minor,
total_memory_mb,
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
static ENV_LOCK: Mutex<()> = Mutex::new(());
struct CudaVisibleGuard {
prev: Option<String>,
}
impl CudaVisibleGuard {
fn set(value: &str) -> Self {
let prev = std::env::var("CUDA_VISIBLE_DEVICES").ok();
unsafe { std::env::set_var("CUDA_VISIBLE_DEVICES", value); }
Self { prev }
}
fn unset() -> Self {
let prev = std::env::var("CUDA_VISIBLE_DEVICES").ok();
unsafe { std::env::remove_var("CUDA_VISIBLE_DEVICES"); }
Self { prev }
}
}
impl Drop for CudaVisibleGuard {
fn drop(&mut self) {
unsafe {
match &self.prev {
Some(v) => std::env::set_var("CUDA_VISIBLE_DEVICES", v),
None => std::env::remove_var("CUDA_VISIBLE_DEVICES"),
}
}
}
}
#[test]
fn detect_gpus_returns_empty_or_valid() {
let _lock = ENV_LOCK.lock().unwrap();
let _g = CudaVisibleGuard::unset();
let gpus = detect_gpus();
for g in &gpus {
assert!(!g.name.is_empty(), "name parsed");
assert!(g.total_memory_mb > 0, "VRAM parsed");
assert!(g.sm_version().starts_with("sm_"), "sm version formatted");
}
}
#[test]
fn gpu_info_sm_version_format() {
let g = GpuInfo {
index: 0,
name: "NVIDIA Test".into(),
sm_major: 12,
sm_minor: 0,
total_memory_mb: 16000,
};
assert_eq!(g.sm_version(), "sm_120");
assert_eq!(g.short_name(), "Test");
assert_eq!(g.vram_bytes(), 16000 * 1024 * 1024);
}
#[test]
fn detect_gpus_empty_cuda_visible_devices_returns_empty() {
let _lock = ENV_LOCK.lock().unwrap();
let _g = CudaVisibleGuard::set("");
assert!(detect_gpus().is_empty());
}
#[test]
fn detect_gpus_cuda_visible_devices_filters_by_index() {
let _lock = ENV_LOCK.lock().unwrap();
let _g_unset = CudaVisibleGuard::unset();
let physical = detect_gpus();
if physical.is_empty() {
return;
}
drop(_g_unset);
let pick = physical[0].index;
let _g_set = CudaVisibleGuard::set(&pick.to_string());
let filtered = detect_gpus();
assert_eq!(filtered.len(), 1, "single-index filter narrows to one");
assert_eq!(filtered[0].index, pick);
}
#[test]
fn mem_info_parses_meminfo_format() {
let text = "MemTotal: 131781120 kB\n\
MemFree: 8123456 kB\n\
MemAvailable: 98765432 kB\n\
Buffers: 123456 kB\n";
let m = parse_meminfo(text).unwrap();
assert_eq!(m.total_bytes, 131_781_120 * 1024);
assert_eq!(m.available_bytes, 98_765_432 * 1024);
assert!(parse_meminfo("MemTotal: 100 kB\nMemFree: 50 kB\n").is_none());
assert!(parse_meminfo("").is_none());
}
#[test]
fn mem_info_reads_live_host() {
if let Some(m) = mem_info() {
assert!(m.total_bytes > 0);
assert!(m.available_bytes <= m.total_bytes);
}
}
#[test]
fn detect_gpus_cuda_visible_devices_excludes_nonexistent_index() {
let _lock = ENV_LOCK.lock().unwrap();
let _g = CudaVisibleGuard::set("99");
assert!(detect_gpus().is_empty());
}
#[test]
fn parse_gpu_csv_row_parses_a_well_formed_row() {
let g = parse_gpu_csv_row("1, 6.1, 6078, NVIDIA GeForce GTX 1060 6GB").unwrap();
assert_eq!(g.index, 1);
assert_eq!(g.name, "NVIDIA GeForce GTX 1060 6GB");
assert_eq!(g.sm_major, 6);
assert_eq!(g.sm_minor, 1);
assert_eq!(g.total_memory_mb, 6078);
}
#[test]
fn parse_gpu_csv_row_keeps_comma_in_name() {
let g = parse_gpu_csv_row("0, 8.0, 81920, NVIDIA A100, 80GB").unwrap();
assert_eq!(g.name, "NVIDIA A100, 80GB");
assert_eq!(g.total_memory_mb, 81920);
}
#[test]
fn parse_gpu_csv_row_rejects_malformed() {
assert!(parse_gpu_csv_row("0, 8.9, three").is_none());
assert!(parse_gpu_csv_row("x, 8.9, 24564, name").is_none());
}
}