From c11ea73eadc9edc83bd16975c0eb6ac838c314d4 Mon Sep 17 00:00:00 2001 From: ARIA Date: Sat, 3 Oct 2026 20:25:33 +0200 Subject: [PATCH] feat: GPU process list with per-process VRAM/utilization and kill Add a GPU process table to the Performance tab showing which processes use the selected GPU, with per-process VRAM, GPU utilization, CPU, host memory, and command. Includes a kill endpoint (SIGTERM/SIGKILL, with optional parent kill) behind a confirm dialog. - nvcurve/hal/processes.py: NVML + /proc process enumeration - server.py: GET /api/processes, POST /api/processes/kill (refuses to signal init/kthreadd) - frontend: ProcessList component, API client, types - tests: test_processes.py - docs: API reference for the new endpoints --- Makefile | 1 + docs/API.md | 2 + frontend/src/App.tsx | 22 +- frontend/src/api/client.ts | 12 + .../src/components/Monitor/ProcessList.tsx | 254 +++++++++++++ .../src/components/common/ConfirmDialog.tsx | 52 ++- frontend/src/types.ts | 20 + nvcurve/hal/processes.py | 237 ++++++++++++ nvcurve/server.py | 85 ++++- tests/test_processes.py | 346 ++++++++++++++++++ 10 files changed, 1008 insertions(+), 23 deletions(-) create mode 100644 frontend/src/components/Monitor/ProcessList.tsx create mode 100644 nvcurve/hal/processes.py create mode 100644 tests/test_processes.py diff --git a/Makefile b/Makefile index b66c615..7ecb162 100644 --- a/Makefile +++ b/Makefile @@ -26,6 +26,7 @@ test: ## Run the test suite $(UV) run python tests/test_security.py $(UV) run python tests/test_rm_power.py $(UV) run python tests/test_wireview.py + $(UV) run python tests/test_processes.py clean: ## Remove build artifacts rm -rf frontend/dist frontend/node_modules \ No newline at end of file diff --git a/docs/API.md b/docs/API.md index 16b187d..87d2338 100644 --- a/docs/API.md +++ b/docs/API.md @@ -125,6 +125,8 @@ Notes: | `GET /api/gpu?gpu_index=0` | Name, driver version, VRAM | | `GET /api/dashboard?gpu_index=0` | Static dashboard info (VBIOS, CUDA cores, PCIe, BAR1, …). Live values come from the monitor WebSocket | | `GET /api/monitor?gpu_index=0` | One-shot monitoring snapshot: voltage, clocks, temp, power, fans, p-state, VRAM, utilization, throttle reasons | +| `GET /api/processes?gpu_index=0` | Processes using this GPU, sorted by VRAM (descending): pid, user, dev, type (C/G/C+G), gpu_util_pct + mem_util_pct (last ~1 s only), vram_bytes, vram_pct, cpu_pct, mem_host_bytes, command | +| `POST /api/processes/kill` | Body `{pid, signal: "TERM" \| "KILL", parent: false}` (default TERM) → send signal. `parent: true` signals the process's parent instead (response reports the signalled pid). Errors: `400` bad pid/signal or no parent, `404` process not found, `403` permission denied. **Note:** the pid need not be a GPU process — this is an arbitrary-pid signal primitive. The server runs as root, so with no users configured (open API) any reachable client can signal any process; enable auth before exposing the port on a shared system | ### Curve diff --git a/frontend/src/App.tsx b/frontend/src/App.tsx index bed0e49..5bed3f4 100644 --- a/frontend/src/App.tsx +++ b/frontend/src/App.tsx @@ -9,6 +9,7 @@ import { CurveEditor } from "./components/CurveEditor/CurveEditor.js"; import { PointTable } from "./components/PointTable/PointTable.js"; import { PerformancePanel } from "./components/Limits/PerformancePanel.js"; import { PerformanceMonitor } from "./components/Monitor/PerformanceMonitor.js"; +import { ProcessList } from "./components/Monitor/ProcessList.js"; import { FanMonitor } from "./components/Monitor/FanMonitor.js"; import { FanCurveEditor } from "./components/Fans/FanCurveEditor.js"; import { WireViewPanel } from "./components/WireView/WireViewPanel.js"; @@ -290,16 +291,19 @@ function MainApp({ )} ) : activeTab === "performance" ? ( -
-
- -
-
- +
+
+
+ +
+
+ +
+
) : activeTab === "wireview" ? ( get<{ voltage_uv: number; voltage_mv: number }>("/voltage", gpuIndex), monitor: (gpuIndex: number) => get("/monitor", gpuIndex), + processes: (gpuIndex: number) => + get("/processes", gpuIndex), + killProcess: ( + pid: number, + signal: "TERM" | "KILL" = "TERM", + parent = false, + ) => + post<{ ok: boolean; pid: number; signal: string; parent: boolean }>( + "/processes/kill", + { pid, signal, parent }, + ), snapshots: (gpuIndex: number) => get("/snapshots", gpuIndex), /** Write per-point frequency deltas. deltas: { pointIndex: deltaKhz } */ diff --git a/frontend/src/components/Monitor/ProcessList.tsx b/frontend/src/components/Monitor/ProcessList.tsx new file mode 100644 index 0000000..a50d621 --- /dev/null +++ b/frontend/src/components/Monitor/ProcessList.tsx @@ -0,0 +1,254 @@ +import { useEffect, useRef, useState } from "react"; +import { Loader, X } from "lucide-react"; +import { api } from "../../api/client.js"; +import { useCurveStore } from "../../store/curveStore.js"; +import type { GpuProcess } from "../../types.js"; +import { fmt } from "../../utils/units.js"; +import { toast } from "sonner"; +import { ConfirmDialog } from "../common/ConfirmDialog.js"; + +const POLL_MS = 3000; + +const TYPE_STYLES: Record = + { + C: { label: "C", cls: "bg-cyan-500/15 text-cyan-400 border-cyan-500/30" }, + G: { + label: "G", + cls: "bg-emerald-500/15 text-emerald-400 border-emerald-500/30", + }, + "C+G": { + label: "C+G", + cls: "bg-violet-500/15 text-violet-400 border-violet-500/30", + }, + }; + +const TH = + "px-3 py-1.5 font-semibold text-left text-[10px] uppercase tracking-wider text-zinc-500 whitespace-nowrap"; +const TD = "px-3 py-1.5 whitespace-nowrap"; + +export function ProcessList() { + const { selectedGpuIndex } = useCurveStore(); + const [processes, setProcesses] = useState(null); + const [memTotal, setMemTotal] = useState(null); + const [error, setError] = useState(null); + const [killTarget, setKillTarget] = useState(null); + const [pendingAction, setPendingAction] = useState< + "TERM" | "KILL" | "KILL_PARENT" | null + >(null); + + // Ref so handleKill can trigger an immediate refresh after a kill. + const refreshRef = useRef<() => void>(() => {}); + + useEffect(() => { + let cancelled = false; + + async function fetchOnce() { + try { + const data = await api.processes(selectedGpuIndex); + if (cancelled) return; + setProcesses(data.processes); + setMemTotal(data.mem_total_bytes); + setError(null); + } catch (e) { + if (!cancelled) + setError(e instanceof Error ? e.message : String(e)); + } + } + + refreshRef.current = () => { + void fetchOnce(); + }; + fetchOnce(); + const timer = setInterval(fetchOnce, POLL_MS); + return () => { + cancelled = true; + clearInterval(timer); + }; + }, [selectedGpuIndex]); + + async function handleKill(action: "TERM" | "KILL" | "KILL_PARENT") { + if (!killTarget || pendingAction) return; + const pid = killTarget.pid; + setPendingAction(action); + try { + const signal = action === "TERM" ? "TERM" : "KILL"; + const parent = action === "KILL_PARENT"; + const res = await api.killProcess(pid, signal, parent); + toast.success( + parent + ? `Sent SIGKILL to parent process ${res.pid} of ${pid}` + : `Sent SIG${signal} to process ${pid}`, + ); + setKillTarget(null); + refreshRef.current(); + } catch (e) { + toast.error(e instanceof Error ? e.message : String(e)); + } finally { + setPendingAction(null); + } + } + + return ( +
+ {/* ── Header ─────────────────────────────────────────────────────── */} +
+ + GPU Processes + + {processes && processes.length > 0 && ( + {processes.length} + )} +
+ {memTotal != null && ( + + VRAM total {fmt.bytes(memTotal)} + + )} +
+
+ + {/* ── Error banner ────────────────────────────────────────────────── */} + {error && ( +
+ ⚠ {error} + +
+ )} + + {/* ── Table ───────────────────────────────────────────────────────── */} + {processes === null && !error ? ( +
+ + Loading processes… +
+ ) : processes && processes.length === 0 && !error ? ( +
+ {memTotal === null + ? "GPU process data unavailable (NVML not initialized)." + : "No processes using this GPU."} +
+ ) : ( +
+ + + + + + + + + + + + + + + + + + {processes?.map((p) => ( + + + + + + + + + + + + + + ))} + +
PIDUserDevType + GPU + VRAMVRAM %CPUMEMCommandKill
{p.pid}{p.user}{p.dev} + + {TYPE_STYLES[p.type].label} + + + {fmt.pct(p.gpu_util_pct ?? 0)} + + {fmt.bytes(p.vram_bytes)} + + {fmt.pct(p.vram_pct, 1)} + + {fmt.pct(p.cpu_pct, 1)} + + {fmt.bytes(p.mem_host_bytes)} + + + {p.command} + + + +
+
+ )} + + {/* ── Kill confirmation ───────────────────────────────────────────── */} + {killTarget && ( + setKillTarget(null)} + actions={[ + { + label: pendingAction === "TERM" ? "Sending…" : "SIGTERM", + onClick: () => handleKill("TERM"), + disabled: pendingAction != null, + className: + "px-3 py-1.5 rounded text-xs font-semibold transition-colors bg-zinc-800 hover:bg-zinc-700 text-zinc-200", + }, + { + label: pendingAction === "KILL" ? "Sending…" : "SIGKILL", + onClick: () => handleKill("KILL"), + disabled: pendingAction != null, + className: + "px-3 py-1.5 rounded text-xs font-semibold transition-colors bg-red-600 hover:bg-red-500 text-white", + }, + { + label: + pendingAction === "KILL_PARENT" ? "Sending…" : "SIGKILL Parent", + onClick: () => handleKill("KILL_PARENT"), + disabled: pendingAction != null, + className: + "px-3 py-1.5 rounded text-xs font-semibold transition-colors bg-red-900 hover:bg-red-800 text-red-100", + }, + ]} + /> + )} +
+ ); +} diff --git a/frontend/src/components/common/ConfirmDialog.tsx b/frontend/src/components/common/ConfirmDialog.tsx index 438a6a2..2b0384b 100644 --- a/frontend/src/components/common/ConfirmDialog.tsx +++ b/frontend/src/components/common/ConfirmDialog.tsx @@ -1,10 +1,20 @@ +export interface DialogAction { + label: string; + onClick: () => void; + className?: string; + disabled?: boolean; +} + interface Props { message: string; detail?: string; confirmLabel?: string; isDestructive?: boolean; - onConfirm: () => void; + /** Required unless `actions` is provided (which replaces the single confirm button). */ + onConfirm?: () => void; onCancel: () => void; + /** When provided, rendered after Cancel instead of the single confirm button. */ + actions?: DialogAction[]; } export function ConfirmDialog({ @@ -14,6 +24,7 @@ export function ConfirmDialog({ isDestructive = false, onConfirm, onCancel, + actions, }: Props) { return (
Cancel - + {actions ? ( + actions.map((a) => ( + + )) + ) : ( + + )}
diff --git a/frontend/src/types.ts b/frontend/src/types.ts index ae7226a..40dc8df 100644 --- a/frontend/src/types.ts +++ b/frontend/src/types.ts @@ -85,6 +85,26 @@ export interface GpuInfo { vram_gib: number | null; } +export interface GpuProcess { + pid: number; + user: string; + dev: number; + type: "C" | "G" | "C+G"; + // GPU utilization, only reported for processes active in the last ~1 s. + gpu_util_pct: number | null; + mem_util_pct: number | null; + vram_bytes: number | null; + vram_pct: number | null; + cpu_pct: number | null; + mem_host_bytes: number | null; + command: string; +} + +export interface ProcessListResponse { + processes: GpuProcess[]; + mem_total_bytes: number | null; +} + export interface SnapshotInfo { filepath: string; timestamp: string; diff --git a/nvcurve/hal/processes.py b/nvcurve/hal/processes.py new file mode 100644 index 0000000..6c9a285 --- /dev/null +++ b/nvcurve/hal/processes.py @@ -0,0 +1,237 @@ +"""GPU process list — which processes are using VRAM on a GPU. + +Per-PID VRAM and process type (compute/graphics) come from NVML; user, +CPU usage, host memory and the full command come from /proc. The server +runs as root, so it can read /proc entries of other users' processes. +""" + +import os +import pwd +import time +from typing import Any + +# NVML state lives in monitoring.py (nvmlInit is process-wide); import the +# module so attribute lookups see the current values, not import-time copies. +from . import monitoring as _mon + +_BOOT_TIME: float | None = None + + +def _get_boot_time() -> float: + """System boot time as a Unix timestamp (cached; read from /proc/stat).""" + global _BOOT_TIME + if _BOOT_TIME is None: + _BOOT_TIME = time.time() # fallback if /proc/stat is unreadable + try: + with open("/proc/stat") as f: + for line in f: + if line.startswith("btime "): + _BOOT_TIME = float(line.split()[1]) + break + except OSError: + pass + return _BOOT_TIME + + +def _proc_info(pid: int) -> dict[str, Any] | None: + """Read user, CPU usage, host memory and command for a pid from /proc. + + CPU% is the average over the process lifetime (total jiffies / age), + which needs a single /proc read — no sampling state. Returns None when + the process no longer exists. + """ + base = f"/proc/{pid}" + try: + with open(f"{base}/stat", "rb") as f: + stat_raw = f.read().decode("ascii", "replace") + with open(f"{base}/statm") as f: + statm = f.read().split() + with open(f"{base}/cmdline", "rb") as f: + cmdline = f.read() + except OSError: + return None + + # comm (field 2) may contain spaces and parentheses, so split at the + # last ')' — everything after it is fields 3..N. + lp = stat_raw.rfind(")") + if lp < 0: + return None + comm = stat_raw[stat_raw.index("(") + 1 : lp] + fields = stat_raw[lp + 2 :].split() + # field N -> fields[N-3]: utime=14, stime=15, starttime=22 + if len(fields) < 20: + return None + try: + utime = int(fields[11]) + stime = int(fields[12]) + starttime = int(fields[19]) + except ValueError: + return None + + uid: int | None = None + try: + with open(f"{base}/status") as f: + for line in f: + if line.startswith("Uid:"): + uid = int(line.split()[1]) + break + except OSError: + pass + try: + user = pwd.getpwuid(uid).pw_name if uid is not None else "unknown" + except KeyError: + user = str(uid) + + hertz = os.sysconf("SC_CLK_TCK") + age_s = ((time.time() - _get_boot_time()) * hertz - starttime) / hertz + cpu_pct = ( + round((utime + stime) / hertz / age_s * 100.0, 1) if age_s > 0 else None + ) + + page = os.sysconf("SC_PAGE_SIZE") + try: + mem_bytes = int(statm[1]) * page + except (IndexError, ValueError): + mem_bytes = None + + parts = [p for p in cmdline.split(b"\0") if p] + command = " ".join(p.decode("utf-8", "replace") for p in parts) or comm + + return { + "user": user, + "cpu_pct": cpu_pct, + "mem_bytes": mem_bytes, + "command": command, + } + + +def _nvml_process_lists(handle: Any) -> dict[int, dict[str, Any]]: + """Map pid -> {"type": "C"|"G"|"C+G", "vram_bytes": int} for one GPU.""" + out: dict[int, dict[str, Any]] = {} + pynvml = _mon._pynvml + try: + for p in pynvml.nvmlDeviceGetComputeRunningProcesses_v2(handle): + entry = out.setdefault(p.pid, {"type": "C", "vram_bytes": 0}) + entry["vram_bytes"] = max(entry["vram_bytes"], p.usedGpuMemory or 0) + except pynvml.NVMLError: + pass + try: + for p in pynvml.nvmlDeviceGetGraphicsRunningProcesses_v2(handle): + entry = out.setdefault(p.pid, {"type": "G", "vram_bytes": 0}) + entry["vram_bytes"] = max(entry["vram_bytes"], p.usedGpuMemory or 0) + if entry["type"] == "C": + entry["type"] = "C+G" + except pynvml.NVMLError: + pass + return out + + +def _nvml_process_utilization( + handle: Any, +) -> dict[int, tuple[float | None, float | None]]: + """Map pid -> (gpu_util_pct, mem_util_pct). + + Only processes that used the GPU within the last ~1 s are reported + (driver >= 525; older drivers raise NOT_SUPPORTED -> empty map). + """ + out: dict[int, tuple[float | None, float | None]] = {} + pynvml = _mon._pynvml + try: + for u in pynvml.nvmlDeviceGetProcessUtilization(handle, 0): + # Newer NVML renamed gpuUtil -> smUtil; accept both. + gpu_util = getattr(u, "smUtil", None) + if gpu_util is None: + gpu_util = getattr(u, "gpuUtil", None) + out[u.pid] = ( + float(gpu_util) if gpu_util is not None else None, + float(u.memUtil) if u.memUtil is not None else None, + ) + except pynvml.NVMLError: + pass + return out + + +def list_gpu_processes(gpu_index: int = 0) -> dict[str, Any]: + """List processes using the given GPU, with per-PID details. + + Returns {"processes": [...], "mem_total_bytes": int | None}. Processes + are sorted by VRAM usage (descending). Each process: + pid, user, dev (gpu_index), type ("C"/"G"/"C+G"), gpu_util_pct, + mem_util_pct, vram_bytes, vram_pct, cpu_pct, mem_host_bytes, command. + """ + from .monitoring import get_vram_total + + mem_total = get_vram_total(gpu_index) + processes: list[dict[str, Any]] = [] + + if _mon._NVML_AVAILABLE and _mon._nvml_initialized: + pynvml = _mon._pynvml + try: + handle = pynvml.nvmlDeviceGetHandleByIndex(gpu_index) + nvml_procs = _nvml_process_lists(handle) + util = _nvml_process_utilization(handle) + except pynvml.NVMLError: + nvml_procs, util = {}, {} + + for pid, info in sorted( + nvml_procs.items(), key=lambda kv: kv[1]["vram_bytes"], reverse=True + ): + proc = _proc_info(pid) + if proc is None: + continue # process exited between the NVML read and /proc read + gpu_util, mem_util = util.get(pid, (None, None)) + vram = info["vram_bytes"] + processes.append( + { + "pid": pid, + "user": proc["user"], + "dev": gpu_index, + "type": info["type"], + "gpu_util_pct": gpu_util, + "mem_util_pct": mem_util, + "vram_bytes": vram, + "vram_pct": ( + round(vram / mem_total * 100.0, 1) + if mem_total and vram is not None + else None + ), + "cpu_pct": proc["cpu_pct"], + "mem_host_bytes": proc["mem_bytes"], + "command": proc["command"], + } + ) + + return {"processes": processes, "mem_total_bytes": mem_total} + + +def kill_process(pid: int, sig: int) -> None: + """Send a signal to a pid. + + Raises ProcessLookupError if the pid does not exist, PermissionError if + the caller may not signal it. + """ + os.kill(pid, sig) + + +def get_parent_pid(pid: int) -> int | None: + """Return the parent pid of a process. + + Returns None when the process no longer exists, or 0 when it has no + parent (kernel threads, init). + """ + try: + with open(f"/proc/{pid}/stat") as f: + raw = f.read() + except OSError: + return None + lp = raw.rfind(")") + if lp < 0: + return None + fields = raw[lp + 2 :].split() + # field 4 (ppid) -> fields[1] + if len(fields) < 2: + return None + try: + return int(fields[1]) + except ValueError: + return None diff --git a/nvcurve/server.py b/nvcurve/server.py index 8fd4b40..78c9e97 100644 --- a/nvcurve/server.py +++ b/nvcurve/server.py @@ -9,6 +9,7 @@ Requires root (NvAPI needs it). import asyncio import logging import os +import signal from contextlib import asynccontextmanager, suppress from pathlib import Path from typing import Any, Protocol @@ -47,6 +48,7 @@ from .hal.monitoring import ( shutdown_nvml, throttle_reasons_label, ) +from .hal.processes import get_parent_pid, kill_process, list_gpu_processes from .hal.ranges import get_clock_ranges from .hal.snapshot import ( list_snapshots, @@ -72,8 +74,8 @@ from .profiles.native import ( rename_profile, save_profile, ) -from .wireview import create_device, find_wireview_ports from .safety import check_negative_freq_warnings, validate_write +from .wireview import create_device, find_wireview_ports log = logging.getLogger("nvcurve.server") @@ -837,6 +839,15 @@ class LoginRequest(BaseModel): password: str +class ProcessKillRequest(BaseModel): + pid: int + # "TERM" (default) or "KILL". + signal: str = "TERM" + # When true, the signal is sent to the process's parent instead of the + # process itself (e.g. to stop a parent that keeps respawning the child). + parent: bool = False + + # ── Helper: run blocking HAL call in thread pool ────────────────────────────── @@ -1107,6 +1118,77 @@ async def api_monitor(gpu_index: int = 0): return _sample_dict(sample) +@app.get("/api/processes", tags=["GPU"], responses={**R_AUTH, **R_GPU}) +async def api_processes(gpu_index: int = 0): + """Processes using this GPU, sorted by VRAM (descending). + + Per process: pid, user, dev (gpu_index), type (C/G/C+G), gpu_util_pct, + mem_util_pct, vram_bytes, vram_pct, cpu_pct, mem_host_bytes, command. + gpu_util_pct is only reported for processes active in the last ~1 s. + """ + _require_gpu(gpu_index) + return await _run(list_gpu_processes, gpu_index) + + +@app.post( + "/api/processes/kill", + tags=["GPU"], + responses={ + **R_AUTH, + 400: { + "description": "Invalid pid or signal (must be 'TERM' or 'KILL'), the process has no parent, or the target is init/kthreadd (pid 1/2)." + }, + 404: {"description": "Process (or its parent) not found."}, + 403: {"description": "Permission denied for the pid."}, + }, +) +async def api_process_kill(req: ProcessKillRequest): + """Send a signal to a process (or its parent). + + Body: {pid, signal: 'TERM' | 'KILL', parent: bool}. + parent=true signals the process's parent instead of the process itself. + The response reports the pid that was actually signalled. + """ + if req.pid <= 0: + raise HTTPException(status_code=400, detail="pid must be positive") + if req.signal not in ("TERM", "KILL"): + raise HTTPException(status_code=400, detail="signal must be 'TERM' or 'KILL'") + sig = signal.SIGKILL if req.signal == "KILL" else signal.SIGTERM + + target = req.pid + if req.parent: + ppid = await _run(get_parent_pid, req.pid) + if ppid is None: + raise HTTPException( + status_code=404, detail=f"Process {req.pid} not found" + ) + if ppid == 0: + raise HTTPException( + status_code=400, detail=f"Process {req.pid} has no parent" + ) + target = ppid + + # Never signal init (pid 1) or kthreadd (pid 2): SIGTERM to init is a + # system-shutdown vector, and kernel threads don't handle signals anyway. + if target in (1, 2): + raise HTTPException( + status_code=400, + detail=f"Refusing to signal pid {target} (init/kthreadd)", + ) + + try: + await _run(kill_process, target, sig) + except ProcessLookupError: + raise HTTPException( + status_code=404, detail=f"Process {target} not found" + ) from None + except PermissionError: + raise HTTPException( + status_code=403, detail=f"Permission denied for pid {target}" + ) from None + return {"ok": True, "pid": target, "signal": req.signal, "parent": req.parent} + + @app.get("/api/snapshots", tags=["Snapshots"], responses=R_AUTH) async def api_snapshots(): """List saved ClockBoostTable snapshots.""" @@ -2046,7 +2128,6 @@ async def api_shutdown(): systems stop the service via systemd instead. """ import os - import signal cfg: Config = _state["config"] if not cfg.allow_api_shutdown: diff --git a/tests/test_processes.py b/tests/test_processes.py new file mode 100644 index 0000000..9108d02 --- /dev/null +++ b/tests/test_processes.py @@ -0,0 +1,346 @@ +"""Tests for the GPU process list + kill endpoints. + +Standalone (no pytest required): + + python tests/test_processes.py + +Also works under pytest if available. Covers the /proc reader (including +comm values with spaces and parentheses), the kill endpoint's input +validation, a real SIGTERM round-trip, and the GET /api/processes payload. +""" + +import contextlib +import os +import subprocess +import sys +import time + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + +from nvcurve import server # noqa: E402 +from nvcurve.hal import processes # noqa: E402 + +PASS = 0 +FAIL = 0 + + +def check(name: str, cond: bool) -> None: + global PASS, FAIL + if cond: + PASS += 1 + print(f" PASS {name}") + else: + FAIL += 1 + print(f" FAIL {name}") + + +def test_proc_info_self() -> None: + """_proc_info must return sane data for our own process.""" + info = processes._proc_info(os.getpid()) + check("self info not None", info is not None) + if info is None: + return + check( + "user is non-empty string", + bool(isinstance(info["user"], str) and info["user"]), + ) + check( + "cpu_pct is None or non-negative number", + info["cpu_pct"] is None or (isinstance(info["cpu_pct"], float) and info["cpu_pct"] >= 0), + ) + check("mem_bytes > 0", info["mem_bytes"] is not None and info["mem_bytes"] > 0) + check("command non-empty", isinstance(info["command"], str) and len(info["command"]) > 0) + + +def test_proc_info_missing() -> None: + """A pid that does not exist must yield None (not an exception).""" + check("missing pid -> None", processes._proc_info(99999999) is None) + + +def test_proc_info_comm_with_parens() -> None: + """comm containing spaces and parentheses must not break /proc/stat parsing. + + The kernel comm is set via prctl(PR_SET_NAME); `exec -a` only changes + argv[0] and would leave comm as the executable name, so the parens path + would never be exercised. + """ + proc = subprocess.Popen( + [ + sys.executable, + "-c", + "import ctypes, time\n" + "ctypes.CDLL(None).prctl(15, b'proc (test) x', 0, 0, 0)\n" + "time.sleep(30)", + ] + ) + try: + # Wait until /proc//stat exists and comm actually carries the + # parenthesised name (prctl may take a moment after exec). + comm = "" + for _ in range(100): + try: + with open(f"/proc/{proc.pid}/stat") as f: + raw = f.read() + comm = raw[raw.index("(") + 1 : raw.rfind(")")] + except OSError: + pass + if "(" in comm: + break + time.sleep(0.05) + check("fixture comm contains parens", "(" in comm) + info = processes._proc_info(proc.pid) + check("weird-comm info not None", info is not None) + if info is not None: + check("weird-comm command parsed", "prctl" in info["command"]) + finally: + proc.kill() + proc.wait() + + +def test_kill_endpoint_validation() -> None: + """The kill endpoint must reject bad input and unknown pids.""" + from fastapi.testclient import TestClient + + client = TestClient(server.app) + r = client.post( + "/api/processes/kill", json={"pid": os.getpid(), "signal": "BOGUS"} + ) + check("invalid signal -> 400", r.status_code == 400) + r = client.post("/api/processes/kill", json={"pid": -5, "signal": "TERM"}) + check("negative pid -> 400", r.status_code == 400) + r = client.post("/api/processes/kill", json={"pid": 0, "signal": "TERM"}) + check("zero pid -> 400", r.status_code == 400) + r = client.post("/api/processes/kill", json={"pid": 99999999, "signal": "TERM"}) + check("unknown pid -> 404", r.status_code == 404) + + +def test_kill_endpoint_real() -> None: + """A real SIGTERM round-trip: the process must actually exit.""" + from fastapi.testclient import TestClient + + client = TestClient(server.app) + proc = subprocess.Popen(["sleep", "30"]) + try: + r = client.post( + "/api/processes/kill", json={"pid": proc.pid, "signal": "TERM"} + ) + check("kill -> 200", r.status_code == 200) + check("response ok", r.json().get("ok") is True) + try: + proc.wait(timeout=5) + check("process exited", True) + except subprocess.TimeoutExpired: + check("process exited", False) + finally: + proc.kill() + proc.wait() + + +def test_kill_endpoint_refuses_init_kthreadd() -> None: + """The endpoint must refuse to signal init (pid 1) or kthreadd (pid 2). + + SIGTERM to init is a system-shutdown vector; kernel threads don't handle + signals anyway. This fires before any NVML/state access, so no GPU needed. + """ + from fastapi.testclient import TestClient + + client = TestClient(server.app) + for pid in (1, 2): + r = client.post( + "/api/processes/kill", json={"pid": pid, "signal": "KILL"} + ) + check(f"pid {pid} refused -> 400", r.status_code == 400) + + +def test_get_parent_pid() -> None: + """get_parent_pid must return the kernel ppid and degrade gracefully.""" + check("self ppid matches", processes.get_parent_pid(os.getpid()) == os.getppid()) + check("missing pid -> None", processes.get_parent_pid(99999999) is None) + check("pid 1 has no parent", processes.get_parent_pid(1) == 0) + + +def test_kill_parent_endpoint() -> None: + """parent=true must signal the parent, leaving the child alive.""" + from fastapi.testclient import TestClient + + client = TestClient(server.app) + + # Validation cases first (no fixture needed). + r = client.post( + "/api/processes/kill", json={"pid": 99999999, "signal": "KILL", "parent": True} + ) + check("parent of missing pid -> 404", r.status_code == 404) + r = client.post( + "/api/processes/kill", json={"pid": 1, "signal": "KILL", "parent": True} + ) + check("pid 1 has no parent -> 400", r.status_code == 400) + + # bash forks a background sleep (child), then sleeps itself (parent). + parent = subprocess.Popen(["bash", "-c", "sleep 30 & sleep 30"]) + child: int | None = None + for _ in range(100): + try: + with open(f"/proc/{parent.pid}/task/{parent.pid}/children") as f: + kids = f.read().split() + if kids: + child = int(kids[0]) + break + except OSError: + pass + time.sleep(0.05) + try: + check("fixture has a child", child is not None) + if child is None: + return + r = client.post( + "/api/processes/kill", + json={"pid": child, "signal": "KILL", "parent": True}, + ) + check("kill parent -> 200", r.status_code == 200) + check("response targets the parent", r.json().get("pid") == parent.pid) + check("response reports parent flag", r.json().get("parent") is True) + try: + parent.wait(timeout=5) + check("parent died", True) + except subprocess.TimeoutExpired: + check("parent died", False) + time.sleep(0.2) + try: + os.kill(child, 0) + check("child survived (orphaned)", True) + except ProcessLookupError: + check("child survived (orphaned)", False) + finally: + for victim in (parent,): + with contextlib.suppress(ProcessLookupError, subprocess.TimeoutExpired): + victim.kill() + victim.wait(timeout=5) + if child is not None: + with contextlib.suppress(ProcessLookupError): + os.kill(child, 9) + + +def test_processes_endpoint() -> None: + """GET /api/processes must return the HAL payload for a known GPU.""" + from fastapi.testclient import TestClient + + fake = { + "processes": [ + { + "pid": 1, + "user": "root", + "dev": 3, + "type": "C", + "gpu_util_pct": 42.0, + "mem_util_pct": 1.0, + "vram_bytes": 1000, + "vram_pct": 0.1, + "cpu_pct": 1.5, + "mem_host_bytes": 2048, + "command": "/usr/bin/fake", + } + ], + "mem_total_bytes": 1000000, + } + orig = server.list_gpu_processes + server._state["gpus"][3] = {"gpu": object(), "gpu_name": "fake"} + server.list_gpu_processes = lambda idx: fake + try: + client = TestClient(server.app) + r = client.get("/api/processes", params={"gpu_index": 3}) + check("processes -> 200", r.status_code == 200) + check("payload matches HAL", r.json() == fake) + r = client.get("/api/processes", params={"gpu_index": 99}) + check("unknown gpu -> 404", r.status_code == 404) + finally: + server.list_gpu_processes = orig + server._state["gpus"].pop(3, None) + + +def test_list_gpu_processes_shape() -> None: + """With NVML available, list_gpu_processes must return a well-shaped payload.""" + from nvcurve.hal import monitoring + + if not monitoring._NVML_AVAILABLE: + print(" SKIP NVML not available") + return + if not monitoring.init_nvml(): + print(" SKIP NVML init failed (no driver?)") + return + try: + data = processes.list_gpu_processes(0) + check("has processes list", isinstance(data.get("processes"), list)) + check("has mem_total_bytes", "mem_total_bytes" in data) + keys = { + "pid", + "user", + "dev", + "type", + "gpu_util_pct", + "mem_util_pct", + "vram_bytes", + "vram_pct", + "cpu_pct", + "mem_host_bytes", + "command", + } + for p in data["processes"]: + if not keys.issubset(p.keys()): + check(f"process {p.get('pid')} has all keys", False) + break + else: + check("all processes have all keys", True) + # Sorted by VRAM descending. + vram = [p["vram_bytes"] or 0 for p in data["processes"]] + check("sorted by VRAM desc", vram == sorted(vram, reverse=True)) + finally: + monitoring.shutdown_nvml() + + +def test_list_gpu_processes_no_nvml() -> None: + """With NVML unavailable, list_gpu_processes returns an empty payload. + + This is the branch the shape test skips (no driver): the UI uses + mem_total_bytes is None to render 'data unavailable' instead of + 'no processes'. + """ + from nvcurve.hal import monitoring + + orig_avail = monitoring._NVML_AVAILABLE + orig_init = monitoring._nvml_initialized + monitoring._NVML_AVAILABLE = False + monitoring._nvml_initialized = False + try: + data = processes.list_gpu_processes(0) + check("empty processes list", data["processes"] == []) + check("mem_total_bytes is None", data["mem_total_bytes"] is None) + finally: + monitoring._NVML_AVAILABLE = orig_avail + monitoring._nvml_initialized = orig_init + + +def main() -> None: + tests = [ + test_proc_info_self, + test_proc_info_missing, + test_proc_info_comm_with_parens, + test_kill_endpoint_validation, + test_kill_endpoint_real, + test_kill_endpoint_refuses_init_kthreadd, + test_get_parent_pid, + test_kill_parent_endpoint, + test_processes_endpoint, + test_list_gpu_processes_shape, + test_list_gpu_processes_no_nvml, + ] + for t in tests: + print(f"== {t.__name__} ==") + t() + print() + print(f"{PASS} passed, {FAIL} failed") + if FAIL: + sys.exit(1) + + +if __name__ == "__main__": + main() -- 2.54.0