"""GPU detection and Compose YAML injection helpers. Detection sources: * NVIDIA — `nvidia-smi` query (when the toolkit is present) * AMD/Intel — scan /dev/dri/renderD* + /sys/class/drm/*/device/vendor """ from __future__ import annotations import glob import os import subprocess from dataclasses import asdict, dataclass from typing import Literal, Optional Vendor = Literal["nvidia", "amd", "intel"] # PCI vendor IDs found in /sys/class/drm/*/device/vendor _PCI_VENDORS = { "0x10de": "nvidia", "0x1002": "amd", "0x8086": "intel", } @dataclass class GPUInfo: vendor: Vendor index: int name: str uuid: Optional[str] = None device_path: Optional[str] = None # e.g. /dev/dri/renderD128 driver: str = "" vram_mb: Optional[int] = None # Host group GIDs that own the DRI nodes — needed so a container can access # a passed-through iGPU (group_add). render = renderD* node, video = card* node. render_gid: Optional[int] = None video_gid: Optional[int] = None def to_dict(self) -> dict: return asdict(self) def _detect_nvidia() -> list[GPUInfo]: gpus: list[GPUInfo] = [] try: out = subprocess.run( [ "nvidia-smi", "--query-gpu=index,name,uuid,memory.total,driver_version", "--format=csv,noheader,nounits", ], capture_output=True, text=True, timeout=8, ) except (FileNotFoundError, subprocess.SubprocessError): return gpus if out.returncode != 0: return gpus for line in out.stdout.strip().splitlines(): parts = [p.strip() for p in line.split(",")] if len(parts) < 5: continue idx, name, uuid, mem, driver = parts[:5] try: index = int(idx) except ValueError: index = 0 try: vram = int(float(mem)) except ValueError: vram = None gpus.append( GPUInfo( vendor="nvidia", index=index, name=name, uuid=uuid, driver=driver, vram_mb=vram, ) ) return gpus def _read(path: str) -> str: try: with open(path, "r", encoding="utf-8") as fh: return fh.read().strip() except OSError: return "" def _stat_gid(path: Optional[str]) -> Optional[int]: """Group owner (GID) of a device node — the host's GID via the /dev bind.""" if not path: return None try: return os.stat(path).st_gid except OSError: return None def _paired_card(render_base: str) -> Optional[str]: """Find the /dev/dri/cardN node that belongs to the same PCI device.""" try: target = os.path.realpath(f"/sys/class/drm/{render_base}/device") except OSError: return None for card in sorted(glob.glob("/dev/dri/card*")): cb = os.path.basename(card) if os.path.realpath(f"/sys/class/drm/{cb}/device") == target: return card return None def _detect_dri() -> list[GPUInfo]: """Detect AMD/Intel render nodes via /dev/dri + sysfs.""" gpus: list[GPUInfo] = [] render_nodes = sorted(glob.glob("/dev/dri/renderD*")) for i, node in enumerate(render_nodes): base = os.path.basename(node) # renderD128 vendor_id = _read(f"/sys/class/drm/{base}/device/vendor").lower() vendor = _PCI_VENDORS.get(vendor_id) if vendor == "nvidia": # NVIDIA is reported via nvidia-smi; skip its DRI node here. continue if not vendor: continue # Try to read a human-ish name. device_id = _read(f"/sys/class/drm/{base}/device/device") name = { "amd": "AMD GPU", "intel": "Intel iGPU", }.get(vendor, "GPU") if device_id: name = f"{name} ({device_id})" card = _paired_card(base) gpus.append( GPUInfo( vendor=vendor, # type: ignore[arg-type] index=i, name=name, device_path=node, driver=vendor, render_gid=_stat_gid(node), video_gid=_stat_gid(card), ) ) return gpus def detect_gpus() -> list[GPUInfo]: return _detect_nvidia() + _detect_dri() def dri_group_gids() -> set[int]: """All render/video GIDs across detected DRI GPUs (for cleanup on remove).""" gids: set[int] = set() for g in _detect_dri(): if g.render_gid is not None: gids.add(g.render_gid) if g.video_gid is not None: gids.add(g.video_gid) return gids # --------------------------------------------------------------------------- # # Injection helpers — mutate a single service dict in place # --------------------------------------------------------------------------- # def inject_nvidia( service: dict, *, device_ids: Optional[list[str]] = None, count: Optional[int] = None, capabilities: Optional[list[str]] = None, ) -> dict: """Add an NVIDIA device reservation under deploy.resources.""" caps = capabilities or ["gpu"] reservation: dict = {"driver": "nvidia", "capabilities": caps} if device_ids: reservation["device_ids"] = device_ids elif count is not None: reservation["count"] = count else: reservation["count"] = "all" deploy = service.setdefault("deploy", {}) resources = deploy.setdefault("resources", {}) reservations = resources.setdefault("reservations", {}) devices = reservations.setdefault("devices", []) # Replace any existing nvidia reservation. devices[:] = [d for d in devices if d.get("driver") != "nvidia"] devices.append(reservation) return service def inject_dri( service: dict, *, vendor: Vendor, add_render_group: bool = True, add_video_group: bool = False, set_libva: bool = False, render_gid: Optional[int] = None, video_gid: Optional[int] = None, ) -> dict: """Pass through /dev/dri and (optionally) add render/video groups. When the host group GIDs are known they are used as numeric ``group_add`` entries (e.g. ``"993"``) — group *names* rarely resolve inside images, so the numeric GID is what actually grants access to a passed-through iGPU. Falls back to the group name when the GID is unknown. For Intel QSV, set LIBVA_DRIVER_NAME=iHD. """ devices = service.setdefault("devices", []) if "/dev/dri:/dev/dri" not in devices: devices.append("/dev/dri:/dev/dri") groups = service.setdefault("group_add", []) def _add(name: str, gid: Optional[int]) -> None: value = str(gid) if gid is not None else name existing = {str(g) for g in groups} if value not in existing and name not in existing: groups.append(value) if add_render_group: _add("render", render_gid) if add_video_group: _add("video", video_gid) if set_libva and vendor == "intel": _ensure_env(service, "LIBVA_DRIVER_NAME", "iHD") return service def remove_gpu(service: dict, dri_gids: Optional[set[int]] = None) -> dict: """Strip GPU-related config from a service. ``dri_gids`` are host render/video GIDs to also drop from ``group_add`` so switching a service away from iGPU passthrough cleans up numeric groups too. """ deploy = service.get("deploy", {}) resources = deploy.get("resources", {}) reservations = resources.get("reservations", {}) if "devices" in reservations: reservations["devices"] = [ d for d in reservations["devices"] if d.get("driver") != "nvidia" ] if not reservations["devices"]: reservations.pop("devices", None) if not reservations: resources.pop("reservations", None) if not resources: deploy.pop("resources", None) if not deploy: service.pop("deploy", None) # DRI passthrough if "devices" in service: service["devices"] = [ d for d in service["devices"] if d != "/dev/dri:/dev/dri" ] if not service["devices"]: service.pop("devices", None) if "group_add" in service: strip = {"render", "video"} | {str(g) for g in (dri_gids or set())} service["group_add"] = [ g for g in service["group_add"] if str(g) not in strip ] if not service["group_add"]: service.pop("group_add", None) _remove_env(service, "LIBVA_DRIVER_NAME") return service def _remove_env(service: dict, key: str) -> None: env = service.get("environment") if isinstance(env, dict): env.pop(key, None) if not env: service.pop("environment", None) elif isinstance(env, list): prefix = f"{key}=" service["environment"] = [ e for e in env if not (isinstance(e, str) and (e == key or e.startswith(prefix))) ] if not service["environment"]: service.pop("environment", None) def _ensure_env(service: dict, key: str, value: str) -> None: env = service.get("environment") if env is None: service["environment"] = [f"{key}={value}"] return if isinstance(env, dict): env[key] = value return # list form prefix = f"{key}=" env = [e for e in env if not (isinstance(e, str) and e.startswith(prefix))] env.append(f"{key}={value}") service["environment"] = env