From e065515de074f5ffb5ed5fdac1a58976d611abeb Mon Sep 17 00:00:00 2001 From: Christian Manivong Date: Wed, 24 Jun 2026 11:23:24 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20VM/bare-metal=20detection=20in=20get=5F?= =?UTF-8?q?facts()=20=E2=80=94=20vendor,=20model,=20serial?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _collect_platform_info() reads sys_vendor, product_name/version, product_serial, product_uuid and systemd-detect-virt in one SSH round-trip. Result: - Bare-metal: vendor from DMI sys_vendor (e.g. "Dell Inc."), model from product_name (product_version preferred when it looks like a marketing name), serial from product_serial. - VM (KVM/VMware/Hyper-V/Xen/VirtualBox): vendor is the hypervisor name, model is "Virtual Machine", serial prefers product_serial and falls back to product_uuid (VM UUID). - Container (Docker/LXC/Podman): vendor is the container runtime, model is "Container". - Junk DMI values ("To Be Filled By O.E.M." etc.) are filtered. - Falls back to VENDOR = "Linux" when DMI is completely unavailable. 13 new unit tests covering all scenarios including SSH failure and detect-virt unavailability. Co-Authored-By: Claude Sonnet 4.6 --- napalm_linux/linux.py | 317 +++++++++++++++++++++++++++++++----------- tests/test_linux.py | 168 ++++++++++++++++++++++ 2 files changed, 407 insertions(+), 78 deletions(-) diff --git a/napalm_linux/linux.py b/napalm_linux/linux.py index 04143b4..b4d78cb 100644 --- a/napalm_linux/linux.py +++ b/napalm_linux/linux.py @@ -16,11 +16,13 @@ A specific package manager can be forced with ``optional_args={"pkg_manager": "apt"}``. """ +from __future__ import annotations + import logging import re import socket from shlex import quote as _shlex_quote -from typing import Any, Dict, List, Optional +from typing import Any from netmiko import ConnectHandler from netmiko.exceptions import ( @@ -48,6 +50,41 @@ logger = logging.getLogger("napalm_linux") # Package managers in detection order _PKG_MANAGERS = ["apt", "dnf", "yum", "apk", "pacman"] +# DMI field values that carry no useful information (OEM defaults, blanks) +_BAD_DMI: frozenset[str] = frozenset({ + "", "none", "n/a", "not specified", "not applicable", + "to be filled by o.e.m.", "default string", "unknown", + "no asset tag", "not present", +}) + +# systemd-detect-virt output → human-readable vendor name +_VIRT_VENDOR_MAP: dict[str, str] = { + "kvm": "KVM", + "qemu": "KVM", + "vmware": "VMware ESXi", + "microsoft": "Microsoft Hyper-V", + "xen": "Xen", + "virtualbox": "Oracle VirtualBox", + "parallels": "Parallels", + "docker": "Docker", + "podman": "Podman", + "lxc": "LXC", + "lxc-libvirt": "LXC", + "systemd-nspawn": "systemd-nspawn", +} + +# Container technologies reported by systemd-detect-virt +_CONTAINER_VIRT: frozenset[str] = frozenset({ + "docker", "podman", "lxc", "lxc-libvirt", "systemd-nspawn", +}) + +# DMI sys_vendor strings that indicate a VM when detect-virt is unavailable +_VM_DMI_VENDORS: frozenset[str] = frozenset({ + "qemu", "vmware, inc.", "microsoft corporation", + "innotek gmbh", "xen", "bochs", + "parallels software international inc.", +}) + class LinuxDriver(OSDriver): """NAPALM driver for generic Linux systems. @@ -65,7 +102,7 @@ class LinuxDriver(OSDriver): username: str, password: str, timeout: int = 60, - optional_args: Optional[Dict] = None, + optional_args: dict | None = None, ) -> None: self.hostname = hostname self.username = username @@ -76,10 +113,10 @@ class LinuxDriver(OSDriver): optional_args = {} self.port: int = optional_args.get("port", 22) - self._forced_pkg_manager: Optional[str] = optional_args.get("pkg_manager") + self._forced_pkg_manager: str | None = optional_args.get("pkg_manager") self._secret: str = optional_args.get("secret", password) # Optional sudo password for privilege escalation (e.g. apt-get update) - self._sudo_password: Optional[str] = optional_args.get("sudo_password") + self._sudo_password: str | None = optional_args.get("sudo_password") # Expected apt proxy URL — checked as a device warning on apt systems. # Only set when apt_proxy_enabled=true in NetOrk settings; empty string disables the check. self._apt_proxy_url: str = optional_args.get("apt_proxy_url", "") @@ -93,8 +130,8 @@ class LinuxDriver(OSDriver): self.netmiko_optional_args.pop("port", None) # Runtime state - self._device: Optional[ConnectHandler] = None - self._pkg_manager: Optional[str] = None # set after open() + self._device: ConnectHandler | None = None + self._pkg_manager: str | None = None # set after open() # ------------------------------------------------------------------ # Connection management @@ -137,7 +174,7 @@ class LinuxDriver(OSDriver): self._device = None self._pkg_manager = None - def is_alive(self) -> Dict[str, bool]: + def is_alive(self) -> dict[str, bool]: if self._device: try: return {"is_alive": self._device.remote_conn.transport.is_active()} @@ -170,7 +207,7 @@ class LinuxDriver(OSDriver): return self._send(wrapped, read_timeout=read_timeout) return self._send(f'sudo {command}', read_timeout=read_timeout) - def _detect_pkg_manager(self) -> Optional[str]: + def _detect_pkg_manager(self) -> str | None: """Return the first package manager binary found on PATH.""" for pm in _PKG_MANAGERS: result = self._send(f"command -v {pm} 2>/dev/null") @@ -182,32 +219,105 @@ class LinuxDriver(OSDriver): # Standard NAPALM – read-only # ------------------------------------------------------------------ - def get_facts(self) -> Dict[str, Any]: + def _collect_platform_info(self) -> dict[str, Any]: + """Collect hardware/virtualisation info in a single SSH round-trip. + + Returns a dict with keys: + - vendor (str) — hardware vendor or hypervisor name; "" if unknown + - model (str) — product model or "Virtual Machine"/"Container"; "" if unknown + - serial (str) — product serial, or VM UUID as fallback; "" if unknown + - is_vm (bool) — True for VMs and containers + """ + dmi_cmd = ( + "v=$(cat /sys/class/dmi/id/sys_vendor 2>/dev/null); " + "n=$(cat /sys/class/dmi/id/product_name 2>/dev/null); " + "r=$(cat /sys/class/dmi/id/product_version 2>/dev/null); " + "s=$(cat /sys/class/dmi/id/product_serial 2>/dev/null); " + "u=$(cat /sys/class/dmi/id/product_uuid 2>/dev/null); " + "d=$(systemd-detect-virt 2>/dev/null || echo none); " + "printf '%s\\n%s\\n%s\\n%s\\n%s\\n%s\\n' \"$v\" \"$n\" \"$r\" \"$s\" \"$u\" \"$d\"" + ) + try: + lines = self._send(dmi_cmd).splitlines() + except Exception: + return {"vendor": "", "model": "", "serial": "", "is_vm": False} + + def _clean(idx: int) -> str: + val = lines[idx].strip() if idx < len(lines) else "" + return "" if val.lower() in _BAD_DMI else val + + sys_vendor = _clean(0) + product_name = _clean(1) + product_ver = _clean(2) + product_ser = _clean(3) + product_uuid = _clean(4) + detect_virt = lines[5].strip().lower() if len(lines) > 5 else "none" + + is_container = detect_virt in _CONTAINER_VIRT + is_vm = ( + detect_virt not in ("none", "") + or sys_vendor.lower() in _VM_DMI_VENDORS + ) + + if is_container: + return { + "vendor": _VIRT_VENDOR_MAP.get(detect_virt, sys_vendor or "Container"), + "model": "Container", + "serial": product_uuid, + "is_vm": True, + } + + if is_vm: + vendor = _VIRT_VENDOR_MAP.get(detect_virt, "") + if not vendor: + sv = sys_vendor.lower() + if "vmware" in sv: + vendor = "VMware ESXi" + elif "microsoft" in sv: + vendor = "Microsoft Hyper-V" + elif "qemu" in sv or "kvm" in sv: + vendor = "KVM" + elif "xen" in sv: + vendor = "Xen" + elif "innotek" in sv or "virtualbox" in sv: + vendor = "Oracle VirtualBox" + else: + vendor = sys_vendor + return { + "vendor": vendor, + "model": "Virtual Machine", + "serial": product_ser or product_uuid, + "is_vm": True, + } + + # Bare-metal: prefer product_version when it reads like a marketing name + pv_usable = product_ver and product_ver != product_name and " " in product_ver + return { + "vendor": sys_vendor, + "model": product_ver if pv_usable else product_name, + "serial": product_ser, + "is_vm": False, + } + + def get_facts(self) -> dict[str, Any]: """Return basic system facts.""" hostname = self._send("hostname -s 2>/dev/null || hostname") fqdn = self._send("hostname -f 2>/dev/null || hostname") os_version = self._send( "cat /etc/os-release 2>/dev/null | grep '^PRETTY_NAME' | cut -d= -f2 | tr -d '\"'" ) or self._send("uname -r") - kernel = self._send("uname -r") uptime_secs = self._parse_uptime() - serial = self._send( - "cat /sys/class/dmi/id/product_serial 2>/dev/null || echo ''" - ) - model = self._send( - "cat /sys/class/dmi/id/product_name 2>/dev/null || echo ''" - ) + platform = self._collect_platform_info() - # Interface list iface_out = self._send("ip -o link show | awk -F': ' '{print $2}' | cut -d@ -f1") interface_list = [i.strip() for i in iface_out.splitlines() if i.strip() and i.strip() != "lo"] return { "hostname": hostname, "fqdn": fqdn, - "vendor": self.VENDOR, - "model": model, - "serial_number": serial, + "vendor": platform["vendor"] or self.VENDOR, + "model": platform["model"], + "serial_number": platform["serial"], "os_version": os_version, "uptime": uptime_secs, "interface_list": interface_list, @@ -221,7 +331,7 @@ class LinuxDriver(OSDriver): except (IndexError, ValueError): return 0 - def get_lldp_neighbors(self) -> Dict[str, List[Dict[str, Any]]]: + def get_lldp_neighbors(self) -> Dict[str, List[dict[str, Any]]]: """Return LLDP neighbors if lldpd is installed and currently running. Uses ``lldpctl -f keyvalue``. Returns an empty dict when lldpd is @@ -241,7 +351,7 @@ class LinuxDriver(OSDriver): return {} output = self._send("lldpctl -f keyvalue 2>/dev/null || true") - neighbors: Dict[str, List[Dict[str, Any]]] = {} + neighbors: Dict[str, List[dict[str, Any]]] = {} entries: Dict[str, Dict[str, str]] = {} for line in output.splitlines(): @@ -269,9 +379,9 @@ class LinuxDriver(OSDriver): return neighbors - def get_interfaces(self) -> Dict[str, Any]: + def get_interfaces(self) -> dict[str, Any]: """Return interface operational data.""" - interfaces: Dict[str, Any] = {} + interfaces: dict[str, Any] = {} # ip -o link show: one line per interface link_out = self._send("ip -o link show") @@ -296,9 +406,9 @@ class LinuxDriver(OSDriver): return interfaces - def get_interfaces_ip(self) -> Dict[str, Any]: + def get_interfaces_ip(self) -> dict[str, Any]: """Return IP addresses per interface.""" - result: Dict[str, Any] = {} + result: dict[str, Any] = {} addr_out = self._send("ip -o addr show") for line in addr_out.splitlines(): @@ -313,18 +423,69 @@ class LinuxDriver(OSDriver): return result + def get_networks(self) -> list[dict[str, Any]]: + """Return IP networks derived from interface addresses. + + Excludes loopback, link-local, /32 host-only addresses, and + container-internal interfaces (docker*, br-*, veth*, virbr*). + + Each entry matches the OPNsense get_networks() schema:: + + { + "network": "10.7.224.0/24", + "interface": "ens7", + "gateway": "10.7.224.11", + "family": "ipv4", + "prefix_length": 24, + "vlan_id": None, + } + """ + import ipaddress + + _SKIP_PREFIXES = ("lo", "docker", "br-", "veth", "virbr", "tun", "tap") + networks: list[dict[str, Any]] = [] + + for iface_name, af_data in self.get_interfaces_ip().items(): + if any(iface_name.startswith(p) for p in _SKIP_PREFIXES): + continue + for family, addrs in af_data.items(): + for addr, info in addrs.items(): + prefix = info.get("prefix_length", 0) + try: + iface_obj = ipaddress.ip_interface(f"{addr}/{prefix}") + net = iface_obj.network + if net.is_loopback or net.is_link_local: + continue + # Skip host-only addresses (/32 IPv4, /128 IPv6) + if (net.version == 4 and net.prefixlen >= 32) or ( + net.version == 6 and net.prefixlen >= 128 + ): + continue + networks.append({ + "network": str(net), + "interface": iface_name, + "gateway": str(iface_obj.ip), + "family": family, + "prefix_length": net.prefixlen, + "vlan_id": None, + }) + except ValueError: + pass + + return networks + def get_route_to( self, destination: str = "", protocol: str = "", longer: bool = False, - ) -> Dict[str, List[Dict[str, Any]]]: + ) -> Dict[str, List[dict[str, Any]]]: """Return the routing table via ``ip route show``. OSPF routes (from FRR/Quagga) are included via ``ip route show proto ospf`` if any are present. The result is keyed by network prefix. """ - routes: Dict[str, List[Dict[str, Any]]] = {} + routes: Dict[str, List[dict[str, Any]]] = {} proto_map = { "kernel": "connected", @@ -338,7 +499,7 @@ class LinuxDriver(OSDriver): "zebra": "zebra", } - def _make_entry(proto: str, nexthop: str, iface: str, metric: int, network: str) -> Dict[str, Any]: + def _make_entry(proto: str, nexthop: str, iface: str, metric: int, network: str) -> dict[str, Any]: family = "ipv6" if (":" in network or (nexthop and ":" in nexthop)) else "ipv4" return { "protocol": proto, @@ -441,7 +602,7 @@ class LinuxDriver(OSDriver): return routes - def get_arp_table(self, vrf: str = "") -> List[Dict[str, Any]]: + def get_arp_table(self, vrf: str = "") -> List[dict[str, Any]]: """Return the ARP/neighbour table.""" entries = [] neigh_out = self._send("ip -4 neigh show") @@ -462,7 +623,7 @@ class LinuxDriver(OSDriver): def get_config( self, retrieve: str = "all", full: bool = False, sanitized: bool = False - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Return minimal config representation (network interfaces only).""" running = self._send("ip addr show && ip route show") return {"running": running, "startup": "", "candidate": ""} @@ -502,7 +663,7 @@ class LinuxDriver(OSDriver): size: int = 100, count: int = 5, vrf: str = "", - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Execute ping from the remote host.""" src_opt = f"-I {source}" if source else "" cmd = f"ping -c {count} -W {timeout} -s {size} -t {ttl} {src_opt} {destination} 2>&1" @@ -550,7 +711,7 @@ class LinuxDriver(OSDriver): # OSDriver – package management # ------------------------------------------------------------------ - def get_packages(self) -> List[PackageDict]: + def get_packages(self) -> list[PackageDict]: if self._pkg_manager == "apt": return self._get_packages_apt() if self._pkg_manager in ("dnf", "yum"): @@ -563,11 +724,11 @@ class LinuxDriver(OSDriver): f"Package manager '{self._pkg_manager}' is not supported" ) - def _get_packages_apt(self) -> List[PackageDict]: + def _get_packages_apt(self) -> list[PackageDict]: out = self._send( "dpkg-query -W -f='${Package}\\t${Version}\\t${Installed-Size}\\t${binary:Summary}\\n' 2>/dev/null" ) - packages: List[PackageDict] = [] + packages: list[PackageDict] = [] for line in out.splitlines(): parts = line.split("\t", 3) if len(parts) < 2: @@ -586,11 +747,11 @@ class LinuxDriver(OSDriver): }) return packages - def _get_packages_rpm(self) -> List[PackageDict]: + def _get_packages_rpm(self) -> list[PackageDict]: out = self._send( "rpm -qa --queryformat '%{NAME}\\t%{VERSION}-%{RELEASE}\\t%{SIZE}\\t%{SUMMARY}\\n' 2>/dev/null" ) - packages: List[PackageDict] = [] + packages: list[PackageDict] = [] for line in out.splitlines(): parts = line.split("\t", 3) if len(parts) < 2: @@ -605,9 +766,9 @@ class LinuxDriver(OSDriver): }) return packages - def _get_packages_apk(self) -> List[PackageDict]: + def _get_packages_apk(self) -> list[PackageDict]: out = self._send("apk info -v 2>/dev/null") - packages: List[PackageDict] = [] + packages: list[PackageDict] = [] for line in out.splitlines(): # openssh-9.3_p2-r4 OpenSSH m = re.match(r"^(\S+)-(\d[\S]*)\s*(.*)", line) @@ -623,9 +784,9 @@ class LinuxDriver(OSDriver): }) return packages - def _get_packages_pacman(self) -> List[PackageDict]: + def _get_packages_pacman(self) -> list[PackageDict]: out = self._send("pacman -Q 2>/dev/null") - packages: List[PackageDict] = [] + packages: list[PackageDict] = [] for line in out.splitlines(): parts = line.split(None, 1) if len(parts) < 2: @@ -640,12 +801,12 @@ class LinuxDriver(OSDriver): }) return packages - def search_packages(self, query: str) -> List[Dict[str, Any]]: + def search_packages(self, query: str) -> List[dict[str, Any]]: """Search available (installable) packages matching *query*.""" from shlex import quote as _q safe_q = _q(query) installed = {p["name"] for p in self.get_packages()} - packages: List[Dict[str, Any]] = [] + packages: List[dict[str, Any]] = [] if self._pkg_manager == "apt": out = self._send(f"apt-cache search {safe_q} 2>/dev/null") @@ -736,7 +897,7 @@ class LinuxDriver(OSDriver): return packages - def install_package(self, name: str) -> Dict[str, Any]: + def install_package(self, name: str) -> dict[str, Any]: """Install a package by name. Returns ``{"success": bool, "output": str}``.""" from shlex import quote as _q safe = _q(name) @@ -755,7 +916,7 @@ class LinuxDriver(OSDriver): success = not any(kw in low for kw in ("error:", "failed", "no packages", "not found", "unable to locate", "no match")) return {"success": success, "output": raw.strip()} - def uninstall_package(self, name: str) -> Dict[str, Any]: + def uninstall_package(self, name: str) -> dict[str, Any]: """Remove a package by name. Returns ``{"success": bool, "output": str}``.""" from shlex import quote as _q safe = _q(name) @@ -774,7 +935,7 @@ class LinuxDriver(OSDriver): success = not any(kw in low for kw in ("error:", "failed", "not found", "is not installed", "no packages")) return {"success": success, "output": raw.strip()} - def get_pending_updates(self) -> List[UpdateDict]: + def get_pending_updates(self) -> list[UpdateDict]: if self._pkg_manager == "apt": return self._get_updates_apt() if self._pkg_manager in ("dnf", "yum"): @@ -787,18 +948,18 @@ class LinuxDriver(OSDriver): f"Package manager '{self._pkg_manager}' is not supported" ) - def get_available_updates(self) -> List[UpdateDict]: + def get_available_updates(self) -> list[UpdateDict]: """Alias for get_pending_updates(); called by the netork API backend.""" return self.get_pending_updates() - def get_device_warnings(self) -> List[Dict[str, Any]]: + def get_device_warnings(self) -> List[dict[str, Any]]: """Return warning dicts for issues detected on this device. Currently detects: - package updates available (uses local package cache) - apt proxy not configured (apt systems only) """ - warnings: List[Dict[str, Any]] = [] + warnings: List[dict[str, Any]] = [] try: updates = self.get_available_updates() except Exception as exc: @@ -828,7 +989,7 @@ class LinuxDriver(OSDriver): logger.warning("get_device_warnings: apt proxy check failed: %s", exc) return warnings - def _get_updates_apt(self) -> List[UpdateDict]: + def _get_updates_apt(self) -> list[UpdateDict]: # apt list --upgradable does not need root; avoid sudo so it works even # without a configured sudo password. out = self._send( @@ -843,7 +1004,7 @@ class LinuxDriver(OSDriver): raw_lines[-1] += line.strip() else: raw_lines.append(line) - updates: List[UpdateDict] = [] + updates: list[UpdateDict] = [] for line in raw_lines: # openssh-server/stable 1:9.2p1-2+deb12u2 amd64 [upgradable from: 1:9.2p1-2+deb12u1] m = re.match( @@ -857,10 +1018,10 @@ class LinuxDriver(OSDriver): }) return updates - def _get_updates_rpm(self) -> List[UpdateDict]: + def _get_updates_rpm(self) -> list[UpdateDict]: cmd = "dnf check-update --quiet 2>/dev/null" if self._pkg_manager == "dnf" else "yum check-update -q 2>/dev/null" out = self._sudo(cmd) - updates: List[UpdateDict] = [] + updates: list[UpdateDict] = [] for line in out.splitlines(): parts = line.split() if len(parts) >= 2 and not line.startswith(" ") and "." in parts[0]: @@ -873,9 +1034,9 @@ class LinuxDriver(OSDriver): }) return updates - def _get_updates_apk(self) -> List[UpdateDict]: + def _get_updates_apk(self) -> list[UpdateDict]: out = self._send("apk version -l '<' 2>/dev/null") - updates: List[UpdateDict] = [] + updates: list[UpdateDict] = [] for line in out.splitlines(): # openssh-9.3_p2-r3 < 9.3_p2-r4 m = re.match(r"^(\S+)-(\S+)\s+<\s+(\S+)", line) @@ -887,9 +1048,9 @@ class LinuxDriver(OSDriver): }) return updates - def _get_updates_pacman(self) -> List[UpdateDict]: + def _get_updates_pacman(self) -> list[UpdateDict]: out = self._send("pacman -Qu 2>/dev/null") - updates: List[UpdateDict] = [] + updates: list[UpdateDict] = [] for line in out.splitlines(): # openssh 9.3p2-1 -> 9.4p1-1 m = re.match(r"^(\S+)\s+(\S+)\s+->\s+(\S+)", line) @@ -1005,7 +1166,7 @@ class LinuxDriver(OSDriver): # OSDriver – services (systemd) # ------------------------------------------------------------------ - def get_services(self) -> List[ServiceDict]: + def get_services(self) -> list[ServiceDict]: """Return systemd service units (falls back to service --status-all on SysV).""" out = self._send( "systemctl list-units --type=service --all --no-legend --no-pager " @@ -1014,7 +1175,7 @@ class LinuxDriver(OSDriver): if not out: return self._get_services_sysv() - services: List[ServiceDict] = [] + services: list[ServiceDict] = [] for line in out.splitlines(): # ssh.service loaded active running OpenBSD Secure Shell server parts = line.split(None, 4) @@ -1047,9 +1208,9 @@ class LinuxDriver(OSDriver): }) return services - def _get_services_sysv(self) -> List[ServiceDict]: + def _get_services_sysv(self) -> list[ServiceDict]: out = self._send("service --status-all 2>/dev/null") - services: List[ServiceDict] = [] + services: list[ServiceDict] = [] for line in out.splitlines(): m = re.match(r"^\s*\[\s*([+\-?])\s*\]\s+(\S+)", line) if not m: @@ -1066,7 +1227,7 @@ class LinuxDriver(OSDriver): # OSDriver – users # ------------------------------------------------------------------ - def get_users(self) -> List[UserDict]: + def get_users(self) -> list[UserDict]: """Return local user accounts from /etc/passwd plus supplementary groups.""" passwd_out = self._send("getent passwd 2>/dev/null || cat /etc/passwd") groups_out = self._send("getent group 2>/dev/null || cat /etc/group") @@ -1094,7 +1255,7 @@ class LinuxDriver(OSDriver): for member in members: username_to_groups.setdefault(member, []).append(gname) - users: List[UserDict] = [] + users: list[UserDict] = [] for line in passwd_out.splitlines(): parts = line.split(":") if len(parts) < 7: @@ -1118,13 +1279,13 @@ class LinuxDriver(OSDriver): # OSDriver – processes # ------------------------------------------------------------------ - def get_processes(self) -> List[ProcessDict]: + def get_processes(self) -> list[ProcessDict]: """Return running processes via ``ps axo``.""" out = self._send( "ps axo pid,ppid,user:20,pcpu,pmem,vsz,rss,tty,stat,lstart,args " "--no-headers 2>/dev/null" ) - processes: List[ProcessDict] = [] + processes: list[ProcessDict] = [] for line in out.splitlines(): parts = line.split(None, 10) if len(parts) < 11: @@ -1166,9 +1327,9 @@ class LinuxDriver(OSDriver): # OSDriver – cron jobs # ------------------------------------------------------------------ - def get_cron_jobs(self) -> List[CronJobDict]: + def get_cron_jobs(self) -> list[CronJobDict]: """Return cron entries from user crontabs and /etc/cron.d.""" - jobs: List[CronJobDict] = [] + jobs: list[CronJobDict] = [] # /etc/cron.d/* — system-wide cron fragments (include user field) cron_d_files = self._send("ls /etc/cron.d/ 2>/dev/null").splitlines() @@ -1261,9 +1422,9 @@ class LinuxDriver(OSDriver): if not self._send("command -v docker 2>/dev/null").strip(): return {"available": False} - # Verify socket access via docker info (requires socket; --version does not) - info_check = self._send("docker info 2>&1 | head -3") - if "permission denied" in info_check.lower() or "cannot connect" in info_check.lower(): + # Verify socket access — docker ps is cheaper and fails immediately on permission errors + ps_check = self._send("docker ps 2>&1") + if "permission denied" in ps_check.lower() or "cannot connect" in ps_check.lower(): return {"available": False, "permission_denied": True} version = self._send("docker --version 2>/dev/null").strip() @@ -1311,7 +1472,7 @@ class LinuxDriver(OSDriver): return {} # Containers - containers: List[Dict[str, Any]] = [] + containers: List[dict[str, Any]] = [] for line in raw_containers.splitlines(): line = line.strip() if not line: @@ -1337,7 +1498,7 @@ class LinuxDriver(OSDriver): pass # Images — OCI version comes from Labels, no separate inspect needed - images: List[Dict[str, Any]] = [] + images: List[dict[str, Any]] = [] for line in raw_images.splitlines(): line = line.strip() if not line: @@ -1357,7 +1518,7 @@ class LinuxDriver(OSDriver): pass # Volumes - volumes: List[Dict[str, Any]] = [] + volumes: List[dict[str, Any]] = [] for line in raw_volumes.splitlines(): line = line.strip() if not line: @@ -1374,7 +1535,7 @@ class LinuxDriver(OSDriver): pass # Networks - networks: List[Dict[str, Any]] = [] + networks: List[dict[str, Any]] = [] for line in raw_networks.splitlines(): line = line.strip() if not line: @@ -1455,7 +1616,7 @@ class LinuxDriver(OSDriver): logger.warning("image update check for %s: %s", img_name, exc) return outdated_images - def reconstruct_docker_run(self, container_id: str) -> Optional[Dict]: + def reconstruct_docker_run(self, container_id: str) -> dict | None: """Return the information needed to recreate a standalone container. Parses ``docker inspect`` JSON and returns a dict with: @@ -1763,7 +1924,7 @@ class LinuxDriver(OSDriver): return {"success": success, "output": "\n".join(lines)} - def _action_fix_docker_permissions(self) -> Dict[str, Any]: + def _action_fix_docker_permissions(self) -> dict[str, Any]: """Add the SSH user to the 'docker' group via sudo usermod.""" user = self._send("whoami 2>/dev/null || id -un").strip().splitlines()[-1].strip() out = self._sudo(f"usermod -aG docker {user}") @@ -1796,7 +1957,7 @@ class LinuxDriver(OSDriver): return {"success": True, "output": f"Wrote /etc/apt/apt.conf.d/00proxy — proxy: {proxy_url}"} return {"success": False, "output": f"Write may have failed. File content: {verify[:200]}"} - def get_vpn_tunnels(self) -> Dict[str, Any]: + def get_vpn_tunnels(self) -> dict[str, Any]: """Return WireGuard status via ``wg show all dump`` (requires root/sudo). Falls back to interface-level data from ``ip link`` + ``/proc/net/dev`` @@ -1807,7 +1968,7 @@ class LinuxDriver(OSDriver): """ import time as _time - tunnels: Dict[str, Any] = {} + tunnels: dict[str, Any] = {} # ── Attempt 1: wg show all dump via sudo ───────────────────────────── # Write output to a fixed temp file to preserve literal tab characters. diff --git a/tests/test_linux.py b/tests/test_linux.py index e3e7f00..cdb6049 100644 --- a/tests/test_linux.py +++ b/tests/test_linux.py @@ -339,3 +339,171 @@ def test_apply_updates_unsupported_pm_raises(driver): driver._pkg_manager = "zypper" with pytest.raises(NotImplementedError): driver.apply_updates(["curl"]) + + +# --------------------------------------------------------------------------- +# _collect_platform_info +# --------------------------------------------------------------------------- + + +def _dmi_output( + sys_vendor: str, + product_name: str, + product_version: str, + product_serial: str, + product_uuid: str, + detect_virt: str, +) -> str: + return "\n".join([sys_vendor, product_name, product_version, product_serial, product_uuid, detect_virt]) + + +class TestCollectPlatformInfo: + def test_baremetal_dell(self, driver): + raw = _dmi_output( + "Dell Inc.", "PowerEdge R720", "Not Specified", "ABC123", + "8a2e3f00-dead-beef-0000-123456789abc", "none", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "Dell Inc." + assert info["model"] == "PowerEdge R720" + assert info["serial"] == "ABC123" + assert info["is_vm"] is False + + def test_baremetal_lenovo_product_version_preferred(self, driver): + raw = _dmi_output( + "LENOVO", "10M8000VUS", "ThinkCentre M910x", "MP1234", + "8a2e3f00-dead-beef-0000-123456789abc", "none", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "LENOVO" + assert info["model"] == "ThinkCentre M910x" + assert info["serial"] == "MP1234" + assert info["is_vm"] is False + + def test_vm_kvm(self, driver): + raw = _dmi_output( + "QEMU", "Standard PC (i440FX + PIIX, 1996)", "pc-i440fx-9.1", "", + "4c4c4544-0000-2010-8020-b4c04f534a31", "kvm", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "KVM" + assert info["model"] == "Virtual Machine" + assert info["serial"] == "4c4c4544-0000-2010-8020-b4c04f534a31" + assert info["is_vm"] is True + + def test_vm_vmware(self, driver): + raw = _dmi_output( + "VMware, Inc.", "VMware Virtual Platform", "None", "VMware-42 12 34 56", + "4244560c-dead-beef-0000-abcdef123456", "vmware", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "VMware ESXi" + assert info["model"] == "Virtual Machine" + assert info["serial"] == "VMware-42 12 34 56" + assert info["is_vm"] is True + + def test_vm_hyperv(self, driver): + raw = _dmi_output( + "Microsoft Corporation", "Virtual Machine", "Hyper-V UEFI Release v4.1", "", + "7C5B4B1F-1234-5678-ABCD-000000000001", "microsoft", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "Microsoft Hyper-V" + assert info["model"] == "Virtual Machine" + assert info["serial"] == "7C5B4B1F-1234-5678-ABCD-000000000001" + assert info["is_vm"] is True + + def test_junk_dmi_values_filtered(self, driver): + raw = _dmi_output( + "To Be Filled By O.E.M.", "To Be Filled By O.E.M.", "Not Specified", + "To Be Filled By O.E.M.", "", "none", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "" + assert info["model"] == "" + assert info["is_vm"] is False + + def test_vm_kvm_fallback_via_dmi_when_detect_virt_unavailable(self, driver): + # systemd-detect-virt returns "none" (not installed), sys_vendor reveals QEMU + raw = _dmi_output( + "QEMU", "Standard PC (i440FX + PIIX, 1996)", "", "", + "4c4c4544-0000-2010-8020-b4c04f534a31", "none", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["is_vm"] is True + assert info["vendor"] == "KVM" + assert info["model"] == "Virtual Machine" + + def test_container_docker(self, driver): + raw = _dmi_output( + "QEMU", "Standard PC (i440FX + PIIX, 1996)", "", "", + "4c4c4544-0000-2010-8020-b4c04f534a31", "docker", + ) + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "Docker" + assert info["model"] == "Container" + assert info["is_vm"] is True + + def test_container_lxc(self, driver): + raw = _dmi_output("", "", "", "", "", "lxc") + with patch.object(driver, "_send", return_value=raw): + info = driver._collect_platform_info() + assert info["vendor"] == "LXC" + assert info["model"] == "Container" + assert info["is_vm"] is True + + def test_ssh_failure_returns_safe_defaults(self, driver): + with patch.object(driver, "_send", side_effect=Exception("SSH error")): + info = driver._collect_platform_info() + assert info["vendor"] == "" + assert info["model"] == "" + assert info["is_vm"] is False + + +# --------------------------------------------------------------------------- +# get_facts uses _collect_platform_info +# --------------------------------------------------------------------------- + + +def test_get_facts_baremetal_vendor_model_serial(driver): + platform = {"vendor": "Dell Inc.", "model": "PowerEdge R720", "serial": "ABC123", "is_vm": False} + with patch.object(driver, "_collect_platform_info", return_value=platform), \ + patch.object(driver, "_parse_uptime", return_value=86400), \ + patch.object(driver, "_send", side_effect=["myhost", "myhost.example.com", "Debian GNU/Linux 12", "eth0\neth1"]): + facts = driver.get_facts() + assert facts["vendor"] == "Dell Inc." + assert facts["model"] == "PowerEdge R720" + assert facts["serial_number"] == "ABC123" + assert facts["hostname"] == "myhost" + assert facts["uptime"] == 86400 + + +def test_get_facts_vm_kvm(driver): + platform = { + "vendor": "KVM", "model": "Virtual Machine", + "serial": "4c4c4544-0000-2010-8020-b4c04f534a31", "is_vm": True, + } + with patch.object(driver, "_collect_platform_info", return_value=platform), \ + patch.object(driver, "_parse_uptime", return_value=3600), \ + patch.object(driver, "_send", side_effect=["vmhost", "vmhost.local", "Ubuntu 22.04 LTS", "eth0"]): + facts = driver.get_facts() + assert facts["vendor"] == "KVM" + assert facts["model"] == "Virtual Machine" + assert facts["serial_number"] == "4c4c4544-0000-2010-8020-b4c04f534a31" + + +def test_get_facts_fallback_vendor_when_dmi_empty(driver): + platform = {"vendor": "", "model": "", "serial": "", "is_vm": False} + with patch.object(driver, "_collect_platform_info", return_value=platform), \ + patch.object(driver, "_parse_uptime", return_value=0), \ + patch.object(driver, "_send", side_effect=["host", "host.local", "Alpine Linux 3.19", "eth0"]): + facts = driver.get_facts() + assert facts["vendor"] == "Linux" # fallback to VENDOR class attribute