Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ fail-closed probe,也不提前声明 DP 或 task-per-GPU 支持。
- `stop` 只允许停止 UniLab record 证明拥有的 daemon;
- `env` 输出 launcher 环境而不修改当前 shell。
- 单 GPU 单 daemon 是默认路径:未传 `--gpus` 时选择唯一可见物理 GPU;未传
`--name` 时使用 `gpu-<uuid-prefix>`;未传 pipe/log 路径时使用用户 cache;
`--name` 时使用 `gpu-<full-uuid>`;未传 pipe/log 路径时使用 `/tmp/uni-cumps/<name>`(mode `0700`);
`env`/`stop` 未传 `--name` 时选择唯一 live record。多 GPU 或多 daemon 歧义必须
显式指定,不自动选择。
- Daemon record 以 `(UID, host, name)` 为 scope,持久化 canonical GPU UUID、绝对
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -117,9 +117,9 @@ Default values:
| Setting | Default |
| --- | --- |
| GPU | the sole visible physical GPU, resolved to its canonical UUID |
| daemon name | `gpu-<uuid-prefix>` |
| pipe directory | `~/.cache/unilab/cuda-mps/<name>/pipe` |
| log directory | `~/.cache/unilab/cuda-mps/<name>/log` |
| daemon name | `gpu-<full-uuid>` |
| pipe directory | `/tmp/uni-cumps/<name>/pipe` (mode `0700`) |
| log directory | `/tmp/uni-cumps/<name>/log` (mode `0700`) |
| `env`/`stop` target | the sole live UniLab-recorded daemon |

If more than one GPU is visible, `start` and `doctor` require
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -112,9 +112,9 @@ uv run uni-cumps stop
| 设置 | 默认值 |
| --- | --- |
| GPU | 唯一可见的物理 GPU,并解析为 canonical UUID |
| daemon 名称 | `gpu-<uuid-prefix>` |
| pipe 目录 | `~/.cache/unilab/cuda-mps/<name>/pipe` |
| log 目录 | `~/.cache/unilab/cuda-mps/<name>/log` |
| daemon 名称 | `gpu-<full-uuid>` |
| pipe 目录 | `/tmp/uni-cumps/<name>/pipe` (mode `0700`) |
| log 目录 | `/tmp/uni-cumps/<name>/log` (mode `0700`) |
| `env`/`stop` 目标 | 唯一 live 的 UniLab-recorded daemon |

如果可见 GPU 多于一张,`start` 与 `doctor` 必须显式传
Expand Down
112 changes: 106 additions & 6 deletions src/unilab/training/cuda_mps_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
import stat
import subprocess
import sys
import tempfile
import time
from collections.abc import Callable, Mapping, Sequence
from dataclasses import asdict, dataclass
Expand All @@ -41,6 +42,8 @@
_NAME_PATTERN = re.compile(r"^[a-z0-9][a-z0-9-]{0,63}$")
_DEFAULT_CONTROL_COMMAND = ("nvidia-cuda-mps-control", "-d")
_QUIT_COMMAND = ("nvidia-cuda-mps-control",)
_MAX_CONTROL_SOCKET_BYTES = 95
_INERT_CONTROL_ARTIFACTS = ("control_lock", "log")
_GPU_QUERY_COMMAND = (
"nvidia-smi",
"--query-gpu=index,uuid,compute_mode",
Expand Down Expand Up @@ -440,6 +443,17 @@ def _proc_stat(pid: int) -> tuple[int, int] | None:
return None


def _process_name_is_alive(pid: int, expected_name: str) -> bool:
"""Match a procfs comm name, allowing its conventional 15-character limit."""

try:
text = Path(f"/proc/{pid}/stat").read_text(encoding="utf-8")
process_name = text[text.index("(") + 1 : text.rindex(")")].strip()
except (OSError, ValueError):
return False
return process_name == expected_name or process_name.startswith(expected_name[:15])


def daemon_process_is_live(
daemon: DaemonRecord, *, proc_stat: Callable[[int], tuple[int, int] | None] = _proc_stat
) -> bool:
Expand Down Expand Up @@ -478,6 +492,54 @@ def _wait_for_control(pipe_directory: Path, timeout: float = 5.0) -> bool:
return False


def _remove_inert_control_artifacts(pipe_directory: Path) -> None:
"""Remove only artifacts left by a failed control launch; preserve sockets."""

if not pipe_directory.exists():
return
for name in _INERT_CONTROL_ARTIFACTS:
path = pipe_directory / name
try:
if path.is_fifo() or (name == "control_lock" and path.is_file()):
path.unlink()
except OSError as exc:
raise CudaMpsCliError(
f"Could not clean failed-control artifact {path}: {exc}."
) from exc


def _remove_orphan_default_control_sockets(pipe_directory: Path) -> None:
"""Remove UniLab-default sockets orphaned after a control process exit.

Explicit user-supplied directories are never modified. In the default
runtime root, NVIDIA leaves ``control`` and ``control_privileged`` sockets
after the control process exits, which otherwise blocks every restart. A
stale PID file alone is not ownership evidence: remove the sockets only
when its PID is absent or no longer has the expected control process name.
"""

identity = _control_daemon_identity(pipe_directory)
if identity is not None and _process_name_is_alive(identity[0], "nvidia-cuda-mps-control"):
return
for name in ("control", "control_privileged"):
path = pipe_directory / name
try:
if path.is_socket():
path.unlink()
except OSError as exc:
raise CudaMpsCliError(f"Could not clean orphan MPS socket {path}: {exc}.") from exc


def _control_log_tail(log_directory: Path, limit: int = 12) -> str | None:
path = log_directory / "control.log"
try:
lines = path.read_text(encoding="utf-8", errors="replace").splitlines()
except OSError:
return None
tail = "\n".join(lines[-limit:]).strip()
return tail or None


def _control_daemon_identity(pipe_directory: Path) -> tuple[int, int] | None:
pid: int | None = None
try:
Expand All @@ -490,8 +552,26 @@ def _control_daemon_identity(pipe_directory: Path) -> tuple[int, int] | None:


def _default_name(gpus: Sequence[GpuIdentity]) -> str:
"""Use the complete UUID so shortened prefixes cannot collide."""

uuid = gpus[0].uuid.removeprefix("GPU-").lower()
return f"gpu-{uuid[:12]}"
return f"gpu-{uuid}"


def _default_runtime_root(daemon_name: str) -> Path:
"""Return a short, user-private runtime root safe for Unix socket paths."""

root = Path(tempfile.gettempdir()) / "uni-cumps" / daemon_name
base = root.parent
try:
base.mkdir(mode=0o700, parents=True, exist_ok=True)
root.mkdir(mode=0o700, parents=True, exist_ok=True)
root.chmod(0o700)
except OSError as exc:
raise CudaMpsCliError(
f"Could not create user-owned MPS runtime directory {root}: {exc}."
) from exc
return root


def _write_daemon_record(root: Path, daemon: DaemonRecord) -> Path:
Expand Down Expand Up @@ -696,8 +776,10 @@ def start_daemon(
raise CudaMpsCliError(
f"Could not prepare daemon record for {daemon_name!r}: {exc}."
) from exc
pipe_value = pipe_dir or str(base / daemon_name / "pipe")
log_value = log_dir or str(base / daemon_name / "log")
runtime_base = _default_runtime_root(daemon_name)
using_default_runtime = pipe_dir is None
pipe_value = pipe_dir or str(runtime_base / "pipe")
log_value = log_dir or str(runtime_base / "log")
try:
pipe_directory = Path(pipe_value).expanduser().resolve(strict=False)
log_directory = Path(log_value).expanduser().resolve(strict=False)
Expand All @@ -706,6 +788,16 @@ def start_daemon(
if pipe_directory == log_directory:
raise CudaMpsCliError("MPS pipe and log directories must be different absolute paths.")
control = pipe_directory / "control"
control_bytes = len(str(control).encode())
if control_bytes > _MAX_CONTROL_SOCKET_BYTES:
raise CudaMpsCliError(
f"MPS control socket path is too long ({control_bytes} bytes; maximum "
f"{_MAX_CONTROL_SOCKET_BYTES}): {control}. Choose a shorter --pipe-dir, such as "
"/tmp/<short-name>/pipe."
)
_remove_inert_control_artifacts(pipe_directory)
if using_default_runtime:
_remove_orphan_default_control_sockets(pipe_directory)
if control.exists():
raise CudaMpsCliError(
f"Refusing to attach to an existing control path {control}; remove/rename it or choose a new daemon."
Expand All @@ -718,13 +810,21 @@ def start_daemon(
"CUDA_MPS_LOG_DIRECTORY": str(log_directory),
}
try:
pipe_directory.mkdir(parents=True, exist_ok=True)
log_directory.mkdir(parents=True, exist_ok=True)
pipe_directory.mkdir(mode=0o700, parents=True, exist_ok=True)
log_directory.mkdir(mode=0o700, parents=True, exist_ok=True)
pipe_directory.chmod(0o700)
log_directory.chmod(0o700)
except OSError as exc:
raise CudaMpsCliError(f"Could not create user-owned MPS directories: {exc}.") from exc
command = list(_DEFAULT_CONTROL_COMMAND)
if daemon:
_run_command(run_command, command, env=env) # type: ignore[call-overload]
try:
_run_command(run_command, command, env=env)
except CudaMpsCliError as exc:
_remove_inert_control_artifacts(pipe_directory)
tail = _control_log_tail(log_directory)
detail = f" Control log tail:\n{tail}" if tail is not None else ""
raise CudaMpsCliError(f"{exc}{detail}") from exc
else:
runner = foreground_runner if foreground_runner is not None else _foreground_control
return_code = runner(env)
Expand Down
138 changes: 138 additions & 0 deletions tests/training/test_cuda_mps_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
import platform
import socket
import stat
import tempfile
from pathlib import Path
from types import SimpleNamespace
from typing import Any
Expand Down Expand Up @@ -263,6 +264,143 @@ def run(command: list[str], **kwargs: Any) -> Completed:
assert (linux_host / "prod" / "daemon.json").is_file()


def test_default_daemon_name_uses_complete_uuid() -> None:
name = mps._default_name((mps.GpuIdentity(0, GPU_A),))

assert name == "gpu-aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"


def test_start_daemon_rejects_long_unix_socket_path(linux_host: Path) -> None:
long_dir = linux_host / ("x" * 96)

with pytest.raises(mps.CudaMpsCliError, match="MPS control socket path is too long"):
mps.start_daemon(
GPU_A,
name="prod",
pipe_dir=str(long_dir / "pipe"),
log_dir=str(linux_host / "log"),
root=linux_host,
run_command=fake_run(),
)


def test_start_daemon_recovers_only_inert_failed_control_artifacts(
linux_host: Path, monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
short_root = tmp_path.parent / "u"
monkeypatch.setattr(mps.tempfile, "gettempdir", lambda: str(short_root))
monkeypatch.setattr(mps.shutil, "which", lambda _name: "/fake/bin/nvidia-cuda-mps-control")
pipe_dir = short_root / "p" / "prod" / "pipe"
pipe_dir.mkdir(parents=True)
(pipe_dir / "control_lock").touch()
log_fifo = pipe_dir / "log"
os.mkfifo(log_fifo)
preserved = pipe_dir / "unrelated"
preserved.touch()
(pipe_dir / "control").touch()

with pytest.raises(mps.CudaMpsCliError, match="Refusing to attach"):
mps.start_daemon(
GPU_A,
name="prod",
pipe_dir=str(pipe_dir),
log_dir=str(short_root / "p" / "prod" / "log"),
root=linux_host,
run_command=fake_run(),
)

assert not (pipe_dir / "control_lock").exists()
assert not log_fifo.exists()
assert preserved.exists()


def test_start_daemon_failure_reports_control_log_and_cleans_artifacts(
linux_host: Path, monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
monkeypatch.setattr(mps.shutil, "which", lambda _name: "/fake/bin/nvidia-cuda-mps-control")
short_root = tmp_path.parent / "u2"
monkeypatch.setattr(mps.tempfile, "gettempdir", lambda: str(short_root))
pipe_dir = short_root / "p" / "prod" / "pipe"
log_dir = short_root / "p" / "prod" / "log"
log_dir.mkdir(parents=True)
(log_dir / "control.log").write_text("control failed\n", encoding="utf-8")

def run(command: list[str], **kwargs: Any) -> Completed:
if command[0] == "nvidia-smi":
return Completed(f"0, {GPU_A}, Default\n")
if command == list(mps._DEFAULT_CONTROL_COMMAND):
pipe_dir.mkdir(parents=True, exist_ok=True)
(pipe_dir / "control_lock").touch()
return Completed(returncode=1)
raise AssertionError(f"unexpected command {command}")

with pytest.raises(mps.CudaMpsCliError, match="control failed"):
mps.start_daemon(
GPU_A,
name="prod",
pipe_dir=str(pipe_dir),
log_dir=str(log_dir),
root=linux_host,
run_command=run,
)

assert not (pipe_dir / "control_lock").exists()


def test_default_runtime_recovers_orphan_control_sockets() -> None:
with tempfile.TemporaryDirectory() as raw:
root = Path(raw)
monkeypatch_root = root / "tmp"
runtime = monkeypatch_root / "uni-cumps" / "gpu-test" / "pipe"
runtime.mkdir(parents=True)
control = runtime / "control"
# Create actual socket nodes like NVIDIA leaves behind.
socket.socket(socket.AF_UNIX, socket.SOCK_STREAM).bind(str(control))
socket.socket(socket.AF_UNIX, socket.SOCK_STREAM).bind(str(runtime / "control_privileged"))

with pytest.MonkeyPatch.context() as monkeypatch:
monkeypatch.setattr(mps.tempfile, "gettempdir", lambda: str(monkeypatch_root))
root_path = mps._default_runtime_root("gpu-test")
mps._remove_orphan_default_control_sockets(root_path / "pipe")

assert not control.exists()
assert not (runtime / "control_privileged").exists()


def test_default_runtime_preserves_live_control_with_stale_pid_file(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with tempfile.TemporaryDirectory() as raw:
runtime = Path(raw) / "uni-cumps" / "gpu-test" / "pipe"
runtime.mkdir(parents=True)
(runtime / "nvidia-cuda-mps-control.pid").write_text("4242\n", encoding="utf-8")
control = runtime / "control"
socket.socket(socket.AF_UNIX, socket.SOCK_STREAM).bind(str(control))

monkeypatch.setattr(mps, "_proc_stat", live_proc_stat())
monkeypatch.setattr(
mps,
"_process_name_is_alive",
lambda pid, name: pid == 4242 and name.startswith("nvidia"),
)
mps._remove_orphan_default_control_sockets(runtime)

assert control.exists()


def test_default_runtime_removes_orphan_socket_with_stale_pid_file() -> None:
with tempfile.TemporaryDirectory() as raw:
runtime = Path(raw) / "uni-cumps" / "gpu-test" / "pipe"
runtime.mkdir(parents=True)
(runtime / "nvidia-cuda-mps-control.pid").write_text("99999999\n", encoding="utf-8")
control = runtime / "control"
socket.socket(socket.AF_UNIX, socket.SOCK_STREAM).bind(str(control))

mps._remove_orphan_default_control_sockets(runtime)

assert not control.exists()


def test_start_daemon_refuses_existing_unmanaged_control_path(
linux_host: Path,
) -> None:
Expand Down
Loading