use serde::{Deserialize, Serialize};
pub const HW_ENV: &str = "MODELSHELF_HW";
pub const FALLBACK_RAM_BYTES: u64 = 8 * GIB;
const GIB: u64 = 1024 * 1024 * 1024;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Gpu {
pub name: String,
pub vram_bytes: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Hardware {
pub ram_bytes: u64,
#[serde(default)]
pub ram_assumed: bool,
#[serde(default)]
pub gpus: Vec<Gpu>,
#[serde(default)]
pub unified_memory: bool,
}
impl Hardware {
pub fn detect() -> Hardware {
if let Ok(json) = std::env::var(HW_ENV) {
if let Ok(hw) = serde_json::from_str::<Hardware>(&json) {
return hw;
}
tracing::warn!("ignoring unparseable {HW_ENV} override");
}
let (ram_bytes, ram_assumed) = match detect_ram() {
Some(bytes) if bytes > 0 => (bytes, false),
_ => (FALLBACK_RAM_BYTES, true),
};
Hardware {
ram_bytes,
ram_assumed,
gpus: detect_nvidia(),
unified_memory: cfg!(all(target_os = "macos", target_arch = "aarch64")),
}
}
}
#[cfg(target_os = "linux")]
fn detect_ram() -> Option<u64> {
parse_meminfo(&std::fs::read_to_string("/proc/meminfo").ok()?)
}
#[cfg(target_os = "macos")]
fn detect_ram() -> Option<u64> {
let out = std::process::Command::new("sysctl")
.args(["-n", "hw.memsize"])
.output()
.ok()?;
if !out.status.success() {
return None;
}
String::from_utf8_lossy(&out.stdout).trim().parse().ok()
}
#[cfg(windows)]
fn detect_ram() -> Option<u64> {
#[repr(C)]
struct MemoryStatusEx {
length: u32,
memory_load: u32,
total_phys: u64,
avail_phys: u64,
total_page_file: u64,
avail_page_file: u64,
total_virtual: u64,
avail_virtual: u64,
avail_extended_virtual: u64,
}
#[link(name = "kernel32")]
extern "system" {
fn GlobalMemoryStatusEx(buffer: *mut MemoryStatusEx) -> i32;
}
let mut status = MemoryStatusEx {
length: std::mem::size_of::<MemoryStatusEx>() as u32,
memory_load: 0,
total_phys: 0,
avail_phys: 0,
total_page_file: 0,
avail_page_file: 0,
total_virtual: 0,
avail_virtual: 0,
avail_extended_virtual: 0,
};
let ok = unsafe { GlobalMemoryStatusEx(&mut status) };
(ok != 0).then_some(status.total_phys)
}
#[cfg(not(any(target_os = "linux", target_os = "macos", windows)))]
fn detect_ram() -> Option<u64> {
None
}
#[cfg_attr(not(target_os = "linux"), allow(dead_code))]
fn parse_meminfo(contents: &str) -> Option<u64> {
let line = contents.lines().find_map(|l| l.strip_prefix("MemTotal:"))?;
let kib: u64 = line.split_whitespace().next()?.parse().ok()?;
Some(kib * 1024)
}
fn detect_nvidia() -> Vec<Gpu> {
let out = match std::process::Command::new("nvidia-smi")
.args([
"--query-gpu=name,memory.total",
"--format=csv,noheader,nounits",
])
.output()
{
Ok(out) if out.status.success() => out,
_ => return Vec::new(),
};
parse_nvidia_smi(&String::from_utf8_lossy(&out.stdout))
}
fn parse_nvidia_smi(output: &str) -> Vec<Gpu> {
let mut gpus = Vec::new();
for line in output.lines() {
let line = line.trim();
if line.is_empty() {
continue;
}
let Some((name, mib)) = line.rsplit_once(',') else {
continue;
};
let Ok(mib) = mib.trim().parse::<u64>() else {
continue;
};
gpus.push(Gpu {
name: name.trim().to_owned(),
vram_bytes: mib * 1024 * 1024,
});
}
gpus
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn meminfo_memtotal_is_parsed_as_kib() {
let sample = "MemTotal: 32657844 kB\nMemFree: 1074956 kB\n";
assert_eq!(parse_meminfo(sample), Some(32_657_844 * 1024));
assert_eq!(parse_meminfo("MemFree: 1 kB\n"), None);
assert_eq!(parse_meminfo(""), None);
assert_eq!(parse_meminfo("MemTotal: garbage kB"), None);
}
#[test]
fn nvidia_smi_single_and_multi_gpu() {
let one = parse_nvidia_smi("NVIDIA GeForce RTX 4090, 24564\n");
assert_eq!(
one,
vec![Gpu {
name: "NVIDIA GeForce RTX 4090".into(),
vram_bytes: 24564 * 1024 * 1024,
}]
);
let two = parse_nvidia_smi("NVIDIA RTX A6000, 49140\nNVIDIA RTX A6000, 49140\n");
assert_eq!(two.len(), 2);
assert_eq!(two[1].vram_bytes, 49140 * 1024 * 1024);
}
#[test]
fn nvidia_smi_skips_na_and_junk() {
assert!(parse_nvidia_smi("").is_empty());
assert!(parse_nvidia_smi("NVIDIA GeForce GT 710, [N/A]\n").is_empty());
assert!(parse_nvidia_smi("no comma here\n").is_empty());
let mixed = parse_nvidia_smi("Broken GPU, [N/A]\nNVIDIA T4, 15360\n");
assert_eq!(mixed.len(), 1);
assert_eq!(mixed[0].name, "NVIDIA T4");
}
#[test]
fn hardware_env_json_shape_roundtrips() {
let hw: Hardware = serde_json::from_str(
r#"{"ram_bytes": 68719476736,
"gpus": [{"name": "NVIDIA GeForce RTX 4090", "vram_bytes": 25757220864}]}"#,
)
.unwrap();
assert_eq!(hw.ram_bytes, 64 * GIB);
assert!(!hw.ram_assumed);
assert!(!hw.unified_memory);
assert_eq!(hw.gpus.len(), 1);
let json = serde_json::to_string(&hw).unwrap();
assert_eq!(serde_json::from_str::<Hardware>(&json).unwrap(), hw);
}
#[test]
fn detect_never_panics_and_reports_positive_ram() {
let hw = Hardware::detect();
assert!(hw.ram_bytes > 0);
}
}