use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicU8, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::{Duration, Instant};
pub const PHASE_LOADING: u8 = 0;
pub const PHASE_IDLE: u8 = 1;
pub const PHASE_BUSY: u8 = 2;
pub const PHASE_DEAD: u8 = 3;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PeerProbeIntegrity {
Ok,
Deferred(u64),
Degraded,
}
impl PeerProbeIntegrity {
pub fn detail(self) -> String {
match self {
Self::Ok => "ok".into(),
Self::Deferred(intervals) => format!("deferred_{intervals}"),
Self::Degraded => "degraded".into(),
}
}
}
pub fn phase_name(p: u8) -> &'static str {
match p {
PHASE_LOADING => "loading",
PHASE_IDLE => "idle",
PHASE_BUSY => "busy",
_ => "dead",
}
}
fn epoch() -> Instant {
static E: OnceLock<Instant> = OnceLock::new();
*E.get_or_init(Instant::now)
}
fn now_ms() -> u64 {
epoch().elapsed().as_millis() as u64
}
pub fn stall_threshold_ms() -> u64 {
std::env::var("MEMRA_HEALTH_STALL_S").ok()
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(120)
.max(1) * 1000
}
pub struct WorkerHealth {
beat_ms: AtomicU64,
phase: AtomicU8,
generation: AtomicU32,
tick_max_ms: AtomicU64,
faulted: AtomicBool,
fault_reason: Mutex<String>,
gpu_faulted: AtomicBool,
gpu_reason: Mutex<String>,
xid_warns: AtomicU64,
peer_probe_deferred_intervals: AtomicU64,
peer_probe_integrity_degraded: AtomicBool,
stall_ms: u64,
}
pub type SharedHealth = Arc<WorkerHealth>;
impl Default for WorkerHealth {
fn default() -> Self {
WorkerHealth {
beat_ms: AtomicU64::new(now_ms()),
phase: AtomicU8::new(PHASE_LOADING),
generation: AtomicU32::new(0),
tick_max_ms: AtomicU64::new(0),
faulted: AtomicBool::new(false),
fault_reason: Mutex::new(String::new()),
gpu_faulted: AtomicBool::new(false),
gpu_reason: Mutex::new(String::new()),
xid_warns: AtomicU64::new(0),
peer_probe_deferred_intervals: AtomicU64::new(0),
peer_probe_integrity_degraded: AtomicBool::new(false),
stall_ms: stall_threshold_ms(),
}
}
}
impl WorkerHealth {
pub fn new() -> SharedHealth {
Arc::new(Self::default())
}
#[cfg(test)]
pub(crate) fn with_stall_ms(stall_ms: u64) -> SharedHealth {
Arc::new(Self { stall_ms, ..Default::default() })
}
pub fn beat(&self) {
let t = now_ms();
let prev = self.beat_ms.swap(t, Ordering::Release);
let dt = t.saturating_sub(prev);
if dt > self.tick_max_ms.load(Ordering::Relaxed) {
self.tick_max_ms.store(dt, Ordering::Relaxed);
}
}
pub fn set_phase(&self, phase: u8) {
self.beat_ms.store(now_ms(), Ordering::Release);
self.phase.store(phase, Ordering::Release);
}
pub fn beat_busy(&self) {
self.phase.store(PHASE_BUSY, Ordering::Release);
self.beat();
}
pub fn mark_ready(&self) {
if let Ok(mut r) = self.fault_reason.lock() {
r.clear();
}
self.faulted.store(false, Ordering::Release);
self.set_phase(PHASE_IDLE);
}
pub fn mark_dead(&self, reason: impl Into<String>) {
let reason = reason.into();
if let Ok(mut r) = self.fault_reason.lock() {
*r = reason;
}
self.faulted.store(true, Ordering::Release);
self.phase.store(PHASE_DEAD, Ordering::Release);
}
pub fn mark_respawning(&self) {
self.generation.fetch_add(1, Ordering::Release);
self.set_phase(PHASE_LOADING);
}
pub fn generation(&self) -> u32 {
self.generation.load(Ordering::Acquire)
}
pub fn note_peer_probe_deferral(&self, consecutive_intervals: u64, degraded: bool) {
self.peer_probe_deferred_intervals
.store(consecutive_intervals, Ordering::Release);
if degraded {
self.peer_probe_integrity_degraded.store(true, Ordering::Release);
}
}
pub fn clear_peer_probe_deferral(&self) {
self.peer_probe_deferred_intervals.store(0, Ordering::Release);
self.peer_probe_integrity_degraded.store(false, Ordering::Release);
}
pub fn peer_probe_integrity(&self) -> PeerProbeIntegrity {
if self.peer_probe_integrity_degraded.load(Ordering::Acquire) {
PeerProbeIntegrity::Degraded
} else {
match self.peer_probe_deferred_intervals.load(Ordering::Acquire) {
0 => PeerProbeIntegrity::Ok,
intervals => PeerProbeIntegrity::Deferred(intervals),
}
}
}
pub fn peer_probe_allows_spec_admission(&self) -> bool {
!self.peer_probe_integrity_degraded.load(Ordering::Acquire)
}
pub fn mark_gpu_fault(&self, reason: impl Into<String>) {
let reason = reason.into();
eprintln!("[gpu-watch] CRITICAL: {reason}");
if let Ok(mut r) = self.gpu_reason.lock() {
if r.is_empty() {
*r = reason;
}
}
self.gpu_faulted.store(true, Ordering::Release);
}
pub fn note_xid_warn(&self, line: &str) {
self.xid_warns.fetch_add(1, Ordering::Relaxed);
eprintln!("[gpu-watch] WARN non-fatal Xid: {line}");
}
fn beat_age_ms(&self) -> u64 {
now_ms().saturating_sub(self.beat_ms.load(Ordering::Acquire))
}
fn stalled(&self) -> bool {
self.phase.load(Ordering::Acquire) == PHASE_BUSY
&& self.beat_age_ms() > self.stall_ms
}
pub fn live(&self) -> Result<(), String> {
if self.gpu_faulted.load(Ordering::Acquire) {
return Err(self.gpu_reason.try_lock().map(|r| r.clone())
.unwrap_or_else(|_| "gpu fault".into()));
}
if self.faulted.load(Ordering::Acquire) {
return Err(self.fault_reason.try_lock().map(|r| r.clone())
.unwrap_or_else(|_| "worker fault".into()));
}
match self.phase.load(Ordering::Acquire) {
PHASE_DEAD => Err("worker thread is gone".into()),
PHASE_LOADING => Err("worker is (re)loading weights".into()),
_ if self.stalled() => Err(format!(
"worker stalled: no scheduler-loop progress for {} ms (threshold {} ms)",
self.beat_age_ms(), self.stall_ms)),
_ => Ok(()),
}
}
pub fn ready(&self, draining: bool) -> Result<(), String> {
if draining {
return Err("draining (shutdown in progress)".into());
}
self.live()
}
pub fn snapshot(&self) -> HealthSnapshot {
HealthSnapshot {
phase: self.phase.load(Ordering::Acquire),
beat_age_ms: self.beat_age_ms(),
tick_max_ms: self.tick_max_ms.load(Ordering::Relaxed),
generation: self.generation(),
xid_warns: self.xid_warns.load(Ordering::Relaxed),
stall_threshold_ms: self.stall_ms,
}
}
}
pub struct HealthSnapshot {
pub phase: u8,
pub beat_age_ms: u64,
pub tick_max_ms: u64,
pub generation: u32,
pub xid_warns: u64,
pub stall_threshold_ms: u64,
}
fn gpu_watch_enabled() -> bool {
std::env::var("MEMRA_GPU_WATCH").as_deref() != Ok("0")
}
fn gpu_watch_interval_s() -> u64 {
std::env::var("MEMRA_GPU_WATCH_S").ok()
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(60)
.max(1)
}
fn gpu_probe_timeout_s() -> u64 {
std::env::var("MEMRA_GPU_PROBE_TIMEOUT_S").ok()
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(10)
.max(1)
}
const XID_FATAL: &[u32] = &[48, 64, 79, 94, 95, 119, 120];
pub fn classify_xid(line: &str) -> Option<(u32, bool)> {
let i = line.find("Xid")?;
let tail = &line[i + 3..];
let after = match tail.find(':') {
Some(_) if tail.trim_start().starts_with('(') => {
let close = tail.find(')')?;
let rest = &tail[close + 1..];
rest.trim_start().trim_start_matches(':')
}
_ => tail,
};
let digits: String = after.trim_start()
.chars().take_while(|c| c.is_ascii_digit()).collect();
let xid: u32 = digits.parse().ok()?;
Some((xid, XID_FATAL.contains(&xid)))
}
enum ProbeErr {
Hang,
Spawn(String),
Exit(String),
}
fn probe_smi(args: &[&str], deadline: Duration) -> Result<String, ProbeErr> {
use std::process::{Command, Stdio};
let mut child = Command::new("nvidia-smi")
.args(args)
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.map_err(|e| ProbeErr::Spawn(e.to_string()))?;
let t0 = Instant::now();
loop {
match child.try_wait() {
Ok(Some(status)) => {
let out = child.wait_with_output()
.map(|o| String::from_utf8_lossy(&o.stdout).to_string())
.unwrap_or_default();
if status.success() {
return Ok(out);
}
return Err(ProbeErr::Exit(format!("exit {status}")));
}
Ok(None) => {
if t0.elapsed() >= deadline {
let _ = child.kill();
let _ = child.wait();
return Err(ProbeErr::Hang);
}
std::thread::sleep(Duration::from_millis(100));
}
Err(e) => return Err(ProbeErr::Spawn(e.to_string())),
}
}
}
pub fn spawn_gpu_watch(health: SharedHealth) {
if !gpu_watch_enabled() {
eprintln!("[gpu-watch] disabled (MEMRA_GPU_WATCH=0)");
return;
}
spawn_xid_tail(health.clone());
let interval = Duration::from_secs(gpu_watch_interval_s());
let deadline = Duration::from_secs(gpu_probe_timeout_s());
let _ = std::thread::Builder::new().name("memra-gpu-watch".into()).spawn(move || {
const RICH: &[&str] = &["--query-gpu=timestamp,ecc.errors.uncorrected.volatile.total,\
retired_pages.pending,remapped_rows.failure",
"--format=csv,noheader"];
const MIN: &[&str] = &["--query-gpu=timestamp,memory.used", "--format=csv,noheader"];
let mut args: &[&str] = RICH;
match probe_smi(RICH, deadline) {
Ok(_) => {}
Err(ProbeErr::Hang) => {
health.mark_gpu_fault(format!(
"nvidia-smi did not answer within {}s at startup — GPU/driver wedge \
(the GSP-timeout class raises no Xid and hangs the query tools, so this \
timeout IS the fault)", deadline.as_secs()));
}
Err(ProbeErr::Spawn(e)) => {
eprintln!("[gpu-watch] nvidia-smi canary unavailable ({e}); Xid log watch only");
loop_xid_only();
return;
}
Err(ProbeErr::Exit(_)) => {
args = MIN;
eprintln!("[gpu-watch] rich ECC/Xid fields unsupported on this GPU; \
canary degraded to a driver-liveness query \
(the probe's own timeout stays the alarm)");
}
}
eprintln!("[gpu-watch] on: every {}s, probe deadline {}s, fatal Xid {:?}",
interval.as_secs(), deadline.as_secs(), XID_FATAL);
loop {
std::thread::sleep(interval);
match probe_smi(args, deadline) {
Ok(out) => {
if let Some(reason) = scan_smi_csv(&out) {
health.mark_gpu_fault(reason);
}
}
Err(ProbeErr::Hang) => health.mark_gpu_fault(format!(
"nvidia-smi did not answer within {}s — GPU/driver wedge",
deadline.as_secs())),
Err(ProbeErr::Spawn(e)) => {
eprintln!("[gpu-watch] canary spawn failed ({e}); continuing on Xid only");
}
Err(ProbeErr::Exit(e)) => {
eprintln!("[gpu-watch] canary exited nonzero ({e}) — not treated as a \
fault (a failing QUERY is not a failing card; the hang is)");
}
}
}
});
}
fn loop_xid_only() {
}
fn scan_smi_csv(out: &str) -> Option<String> {
let line = out.lines().next()?;
let fields: Vec<&str> = line.split(',').map(str::trim).collect();
if fields.len() >= 4 {
if let Ok(ecc) = fields[1].parse::<u64>() {
if ecc > 0 {
return Some(format!("uncorrected volatile ECC errors = {ecc} (nvidia-smi)"));
}
}
if fields[3].eq_ignore_ascii_case("yes") || fields[3] == "1" {
return Some("row-remap FAILURE reported by nvidia-smi (Xid 64 class)".into());
}
}
None
}
fn spawn_xid_tail(health: SharedHealth) {
let _ = std::thread::Builder::new().name("memra-xid-tail".into()).spawn(move || {
use std::io::{BufRead, BufReader};
if let Ok(f) = std::fs::File::open("/dev/kmsg") {
eprintln!("[gpu-watch] Xid source: /dev/kmsg");
let mut rd = BufReader::new(f);
let mut line = String::new();
loop {
line.clear();
match rd.read_line(&mut line) {
Ok(0) => break,
Ok(_) => handle_xid_line(&health, line.trim_end()),
Err(_) => break,
}
}
return;
}
let child = std::process::Command::new("journalctl")
.args(["-k", "-n", "0", "-f", "--no-pager"])
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::null())
.spawn();
match child {
Ok(mut c) => {
eprintln!("[gpu-watch] Xid source: journalctl -k -f \
(/dev/kmsg unreadable — kernel.dmesg_restrict)");
if let Some(out) = c.stdout.take() {
for line in BufReader::new(out).lines().map_while(Result::ok) {
handle_xid_line(&health, &line);
}
}
let _ = c.wait();
}
Err(e) => {
eprintln!("[gpu-watch] no Xid log source (/dev/kmsg unreadable, \
journalctl unavailable: {e}) — nvidia-smi canary only");
}
}
});
}
fn handle_xid_line(health: &SharedHealth, line: &str) {
if !line.contains("Xid") {
return;
}
match classify_xid(line) {
Some((xid, true)) => health.mark_gpu_fault(format!(
"NVRM Xid {xid} (fatal class) — {}", line.trim())),
Some((_, false)) => health.note_xid_warn(line.trim()),
None => {}
}
}
fn notify_socket() -> Option<&'static str> {
static S: OnceLock<Option<String>> = OnceLock::new();
S.get_or_init(|| {
let v = std::env::var("NOTIFY_SOCKET").ok()?;
if v.starts_with('@') {
eprintln!("[sd-notify] abstract socket {v:?} is not addressable from std — \
notifier disabled (use a path-form NOTIFY_SOCKET, i.e. a system unit)");
return None;
}
Some(v)
}).as_deref()
}
pub fn sd_notify(msg: &str) {
let Some(path) = notify_socket() else { return };
if let Ok(sock) = std::os::unix::net::UnixDatagram::unbound() {
let _ = sock.send_to(msg.as_bytes(), path);
}
}
pub fn spawn_sd_watchdog(health: SharedHealth) {
if notify_socket().is_none() {
return;
}
let usec: u64 = match std::env::var("WATCHDOG_USEC").ok().and_then(|v| v.parse().ok()) {
Some(v) if v > 0 => v,
_ => return, };
let every = Duration::from_micros(usec / 2);
eprintln!("[sd-notify] watchdog armed: WATCHDOG_USEC={usec}, pinging every {:.1}s while live",
every.as_secs_f64());
let _ = std::thread::Builder::new().name("memra-sd-watchdog".into()).spawn(move || {
loop {
std::thread::sleep(every);
match health.live() {
Ok(()) => sd_notify("WATCHDOG=1"),
Err(why) => eprintln!("[sd-notify] withholding WATCHDOG=1: {why}"),
}
}
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn xid_classification_covers_both_kernel_line_forms_and_the_fatal_set() {
assert_eq!(classify_xid("NVRM: Xid (PCI:0000:01:00): 119, pid=1234, GSP RPC timeout"),
Some((119, true)));
assert_eq!(classify_xid("NVRM: Xid 79, GPU has fallen off the bus"), Some((79, true)));
assert_eq!(classify_xid("Aug 06 20:00:00 box kernel: NVRM: Xid (PCI:0000:01:00): 48, \
Double-bit ECC"), Some((48, true)));
assert_eq!(classify_xid("NVRM: Xid (PCI:0000:01:00): 13, Graphics SM Warp Exception"),
Some((13, false)));
assert_eq!(classify_xid("NVRM: Xid (PCI:0000:01:00): 43, GPU stopped processing"),
Some((43, false)));
assert_eq!(classify_xid("NVRM: Xid (PCI:0000:01:00): 63, Row remap pending"),
Some((63, false)));
for id in [48u32, 64, 79, 94, 95, 119, 120] {
let line = format!("NVRM: Xid (PCI:0000:c1:00): {id}, something");
assert_eq!(classify_xid(&line), Some((id, true)), "xid {id} must be fatal");
}
assert_eq!(classify_xid("NVRM: GPU at PCI:0000:01:00 has been initialized"), None);
}
#[test]
fn idle_is_healthy_at_any_age_but_busy_stalls() {
let h = WorkerHealth::with_stall_ms(20);
h.mark_ready();
assert!(h.live().is_ok(), "a ready idle worker is live");
h.phase.store(PHASE_IDLE, Ordering::Release);
std::thread::sleep(Duration::from_millis(40));
assert!(h.beat_age_ms() > 20, "the beat must actually be stale for this to mean anything");
assert!(h.live().is_ok(), "idle staleness is meaningless");
h.phase.store(PHASE_BUSY, Ordering::Release);
let why = h.live().expect_err("a busy worker with a stale beat is not live");
assert!(why.contains("stalled"), "{why}");
h.beat();
assert!(h.live().is_ok());
}
#[test]
fn worker_death_and_gpu_fault_latch_immediately_without_a_threshold_wait() {
let h = WorkerHealth::new();
h.mark_ready();
h.beat(); h.mark_dead("worker thread panicked: test");
let why = h.live().expect_err("a dead worker is never live");
assert!(why.contains("panicked"), "{why}");
assert!(h.ready(false).is_err());
h.mark_ready();
assert!(h.live().is_ok());
h.mark_gpu_fault("NVRM Xid 119 (fatal class)");
h.mark_ready();
let why = h.live().expect_err("gpu fault must survive a worker respawn");
assert!(why.contains("Xid 119"), "{why}");
}
#[test]
fn readiness_is_off_while_draining_but_liveness_stays_on() {
let h = WorkerHealth::new();
h.mark_ready();
assert!(h.live().is_ok());
assert!(h.ready(false).is_ok());
assert!(h.ready(true).is_err());
assert!(h.live().is_ok());
}
#[test]
fn loading_is_not_live_and_not_ready() {
let h = WorkerHealth::new(); assert!(h.live().is_err(), "weights are not resident yet");
assert!(h.ready(false).is_err());
h.mark_ready();
assert!(h.live().is_ok());
h.mark_respawning();
assert!(h.live().is_err(), "a respawn load answers nothing");
assert_eq!(h.generation(), 1);
}
#[test]
fn tick_max_records_the_longest_iteration() {
let h = WorkerHealth::new();
h.mark_ready();
h.beat();
std::thread::sleep(Duration::from_millis(25));
h.beat();
let snap = h.snapshot();
assert!(snap.tick_max_ms >= 20, "tick_max_ms = {}", snap.tick_max_ms);
assert_eq!(snap.stall_threshold_ms, stall_threshold_ms());
}
#[test]
fn smi_csv_scan_faults_only_on_definite_values() {
assert!(scan_smi_csv("2026/08/06 20:00:00.000, [N/A], [N/A], [N/A]").is_none());
assert!(scan_smi_csv("2026/08/06 20:00:00.000, 0, 0, No").is_none());
assert!(scan_smi_csv("2026/08/06 20:00:00.000, 3, 0, No")
.is_some_and(|r| r.contains("ECC")));
assert!(scan_smi_csv("2026/08/06 20:00:00.000, 0, 0, Yes")
.is_some_and(|r| r.contains("row-remap")));
assert!(scan_smi_csv("2026/08/06 20:00:00.000, 512 MiB").is_none());
assert!(scan_smi_csv("").is_none());
}
}