From 1f0349aee99ba8b4cde9589ed7b326c16165ed85 Mon Sep 17 00:00:00 2001 From: Christian Manivong Date: Fri, 29 May 2026 09:10:40 +0200 Subject: [PATCH] initial commit --- .gitignore | 45 + README.md | 117 ++ napalm_openwrt/__init__.py | 5 + napalm_openwrt/openwrt.py | 2604 ++++++++++++++++++++++++++++++++++++ pyproject.toml | 54 + requirements.txt | 4 + tests/__init__.py | 0 tests/unit/__init__.py | 0 tests/unit/test_driver.py | 835 ++++++++++++ 9 files changed, 3664 insertions(+) create mode 100644 .gitignore create mode 100644 README.md create mode 100644 napalm_openwrt/__init__.py create mode 100644 napalm_openwrt/openwrt.py create mode 100644 pyproject.toml create mode 100644 requirements.txt create mode 100644 tests/__init__.py create mode 100644 tests/unit/__init__.py create mode 100644 tests/unit/test_driver.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..d75d676 --- /dev/null +++ b/.gitignore @@ -0,0 +1,45 @@ +# Python +__pycache__/ +*.py[cod] +*.pyo +*.pyd +*.so +*.egg +*.egg-info/ +dist/ +build/ +.eggs/ +wheels/ + +# Virtual environments +.venv/ +venv/ +env/ +.env + +# Packaging +*.tar.gz +*.whl +MANIFEST + +# Testing +.pytest_cache/ +.coverage +.coverage.* +htmlcov/ +coverage.xml + +# Type checking +.mypy_cache/ +.ruff_cache/ + +# IDEs +.vscode/ +.idea/ +*.swp +*~ +napalm-openwrt.code-workspace + +# OS +.DS_Store +Thumbs.db diff --git a/README.md b/README.md new file mode 100644 index 0000000..d89cbc6 --- /dev/null +++ b/README.md @@ -0,0 +1,117 @@ +# napalm-openwrt + +NAPALM community driver for **OpenWrt** routers and access-points. + +Communicates over SSH using [Netmiko](https://github.com/ktbyers/netmiko) (`linux` device type). +Requires OpenWrt **19.07** or newer. + +## Tested devices + +| Model | OpenWrt version | Tested | +|---|---|---| +| TP-Link TL-WR1043N/ND v5 | 23.05.3 | ✅ | + +> Contributions for additional devices and firmware versions are welcome. + +## Requirements + +| Dependency | Minimum version | +|---|---| +| Python | 3.8 | +| NAPALM | 4.0 | +| Netmiko | 4.0 | + +## Installation + +From source: + +```bash +git clone https://github.com/napalm-automation-community/napalm-openwrt +cd napalm-openwrt +pip install -e . +``` + +## Quick start + +```python +from napalm import get_network_driver + +driver = get_network_driver("openwrt") +device = driver("192.168.1.1", "root", "") + +device.open() + +facts = device.get_facts() +print(facts) + +interfaces = device.get_interfaces() +print(interfaces) + +device.close() +``` + +## Supported NAPALM getters + +| Getter | Supported | Notes | +|---|---|---| +| `get_facts` | ✅ | Uses `/etc/openwrt_release`, `/tmp/sysinfo/model`, `/proc/uptime` | +| `get_interfaces` | ✅ | Uses `ip link show` | +| `get_interfaces_ip` | ✅ | Uses `ip addr show` | +| `get_interfaces_counters` | ✅ | Uses `/proc/net/dev` | +| `get_arp_table` | ✅ | Uses `ip neigh show` | +| `get_mac_address_table` | ✅ | Uses `bridge fdb show` | +| `get_config` | ✅ | Uses `uci export` | +| `get_environment` | ✅ | CPU from `/proc/stat`, memory from `/proc/meminfo` | +| `get_lldp_neighbors` | ✅ | Requires `lldpd` package installed on device | +| `get_lldp_neighbors_detail` | ✅ | Requires `lldpd` package installed on device | + +## Configuration management + +Configuration is managed via [UCI](https://openwrt.org/docs/guide-user/base-system/uci) +(Unified Configuration Interface). + +### Merge candidate + +```python +device.load_merge_candidate(config=""" +uci set system.@system[0].hostname='my-router' +uci set network.lan.ipaddr='10.0.0.1' +""") + +print(device.compare_config()) +device.commit_config() +``` + +### Replace candidate + +```python +with open("full-config.uci") as f: + device.load_replace_candidate(config=f.read()) + +print(device.compare_config()) +device.commit_config() +``` + +### Rollback + +```python +# Reverts to the config state before the last commit_config call +device.rollback() +``` + +## Development + +```bash +# Create and activate a virtual environment +python -m venv .venv +source .venv/bin/activate + +# Install with dev dependencies +pip install -e ".[dev]" + +# Run tests +pytest + +# Lint +ruff check napalm_openwrt/ +``` diff --git a/napalm_openwrt/__init__.py b/napalm_openwrt/__init__.py new file mode 100644 index 0000000..54d7f90 --- /dev/null +++ b/napalm_openwrt/__init__.py @@ -0,0 +1,5 @@ +"""NAPALM driver for OpenWrt routers/access-points.""" + +from napalm_openwrt.openwrt import OpenWrtDriver + +__all__ = ["OpenWrtDriver"] diff --git a/napalm_openwrt/openwrt.py b/napalm_openwrt/openwrt.py new file mode 100644 index 0000000..11f6f2c --- /dev/null +++ b/napalm_openwrt/openwrt.py @@ -0,0 +1,2604 @@ +# -*- coding: utf-8 -*- +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""NAPALM driver for OpenWrt routers and access-points. + +Communicates via SSH. The device must be running OpenWrt 19.07 or newer. +Netmiko device_type: ``linux`` +""" + +import re +import socket +from typing import Dict, List, Optional, Union + +import netaddr +from netmiko import ConnectHandler +from netmiko.exceptions import NetmikoTimeoutException, NetmikoAuthenticationException + +from napalm_device_types import AccessPointDriver +from napalm.base import helpers as napalm_helpers +from napalm.base.exceptions import ( + ConnectionException, + ConnectionClosedException, + CommandErrorException, + MergeConfigException, + ReplaceConfigException, +) +from napalm.base.netmiko_helpers import netmiko_args + + +class OpenWrtDriver(AccessPointDriver): + """NAPALM driver for OpenWrt routers and access-points.""" + + VENDOR = "OpenWrt" + NETMIKO_DEVICE_TYPE = "linux" + + def __init__( + self, + hostname: str, + username: str, + password: str, + timeout: int = 60, + optional_args: Optional[Dict] = None, + ) -> None: + self.hostname = hostname + self.username = username + self.password = password + self.timeout = timeout + self.device: Optional[ConnectHandler] = None + + if optional_args is None: + optional_args = {} + + self.port = optional_args.pop("port", 22) + self.netmiko_optional_args = netmiko_args(optional_args) + + # Config management state + self._candidate_config: Optional[str] = None + self._candidate_mode: Optional[str] = None # 'merge' or 'replace' + self._backup_config: Optional[str] = None + + # ------------------------------------------------------------------ + # Connection management + # ------------------------------------------------------------------ + + def open(self) -> None: + """Open an SSH connection to the device.""" + try: + self.device = ConnectHandler( + device_type=self.NETMIKO_DEVICE_TYPE, + host=self.hostname, + username=self.username, + password=self.password, + timeout=self.timeout, + port=self.port, + **self.netmiko_optional_args, + ) + except NetmikoTimeoutException as exc: + raise ConnectionException( + f"Cannot connect to {self.hostname}: {exc}" + ) from exc + except NetmikoAuthenticationException as exc: + raise ConnectionException( + f"Authentication failed for {self.hostname}: {exc}" + ) from exc + + def close(self) -> None: + """Close the SSH connection.""" + if self.device: + self.device.disconnect() + self.device = None + + def is_alive(self) -> Dict[str, bool]: + """Return connection liveness.""" + if self.device is None: + return {"is_alive": False} + try: + return {"is_alive": self.device.remote_conn.transport.is_active()} + except (socket.error, EOFError, AttributeError): + return {"is_alive": False} + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + def _send_command(self, command: Union[str, List[str]]) -> str: + """Send a shell command (or list of fallback commands) to the device. + + When a list is supplied, commands are tried in order and the first + one that does not return an error indicator is returned. + """ + def _do_send(cmd: str) -> str: + return self.device.send_command( + cmd, + read_timeout=self.timeout, + ).strip() + + try: + if isinstance(command, list): + output = "" + for cmd in command: + output = _do_send(cmd) + if not output.startswith(("sh: ", "ash: ", "-ash: ", "command not found")): + return output + return output + return _do_send(command) + except (socket.error, EOFError) as exc: + raise ConnectionClosedException(str(exc)) from exc + + @staticmethod + def _parse_openwrt_release(output: str) -> Dict[str, str]: + """Parse ``/etc/openwrt_release`` key=value pairs.""" + result: Dict[str, str] = {} + for line in output.splitlines(): + m = re.match(r'^(\w+)=["\']?([^"\']*)["\']?$', line.strip()) + if m: + result[m.group(1)] = m.group(2) + return result + + @staticmethod + def _parse_uptime_seconds(uptime_str: str) -> float: + """Convert ``/proc/uptime`` first field (seconds.hundredths) to float.""" + try: + return float(uptime_str.split()[0]) + except (IndexError, ValueError): + return 0.0 + + # ------------------------------------------------------------------ + # NAPALM getters + # ------------------------------------------------------------------ + + def get_facts(self) -> Dict: + """Return a dictionary of general device facts. + + Retrieves data from: + - ``/etc/openwrt_release`` → os_version, model + - ``/proc/uptime`` → uptime + - ``uci get system.@system[0].hostname`` or ``hostname`` → hostname + - ``cat /tmp/sysinfo/model`` → model (preferred) + - ``ip link show`` → interface_list + """ + release_out = self._send_command("cat /etc/openwrt_release") + release = self._parse_openwrt_release(release_out) + + os_version = release.get("DISTRIB_RELEASE", "") + model = self._send_command("cat /tmp/sysinfo/model") + if not model or model.startswith("cat:"): + model = release.get("DISTRIB_TARGET", "") + + uptime_out = self._send_command("cat /proc/uptime") + uptime = self._parse_uptime_seconds(uptime_out) + + hostname = self._send_command( + ["uci get system.@system[0].hostname", "hostname"] + ) + + interface_list = self._get_interface_list() + + return { + "vendor": self.VENDOR, + "model": model, + "hostname": hostname, + "fqdn": hostname, + "os_version": os_version, + "serial_number": "", + "uptime": uptime, + "interface_list": interface_list, + } + + def _get_interface_list(self) -> List[str]: + """Return a sorted list of interface names from ``ip link show``.""" + output = self._send_command("ip link show") + interfaces = [] + for line in output.splitlines(): + m = re.match(r"^\d+:\s+(\S+?)[@:]", line) + if m: + name = m.group(1) + if ( + name not in self._EXCLUDED_INTERFACES + and not name.startswith(self._EXCLUDED_INTERFACE_PREFIXES) + ): + interfaces.append(name) + return sorted(set(interfaces)) + + def get_interfaces(self) -> Dict[str, Dict]: + """Return interface details, excluding loopback and raw radio (phy*) interfaces.""" + output = self._send_command("ip link show") + return self._filter_interfaces(self._parse_ip_link(output)) + + def _parse_ip_link(self, output: str) -> Dict[str, Dict]: + """Parse ``ip link show`` output into NAPALM interface dicts.""" + interfaces: Dict[str, Dict] = {} + current: Optional[str] = None + + for line in output.splitlines(): + # New interface block: "2: eth0: mtu 1500 ..." + m = re.match( + r"^\d+:\s+(\S+?)[@:].*<([^>]*)>.*\bmtu\s+(\d+).*\bstate\s+(\S+)", + line, + ) + if m: + name = m.group(1) + flags = m.group(2).upper() + mtu = int(m.group(3)) + state = m.group(4).upper() + + is_up = state in ("UP", "UNKNOWN") and "UP" in flags.split(",") + is_enabled = "UP" in flags.split(",") + + interfaces[name] = { + "is_up": is_up, + "is_enabled": is_enabled, + "description": "", + "last_flapped": -1.0, + "speed": -1.0, + "mtu": mtu, + "mac_address": "", + } + current = name + continue + + # MAC address line: " link/ether aa:bb:cc:dd:ee:ff ..." + if current and "link/ether" in line: + m2 = re.search(r"link/ether\s+(\S+)", line) + if m2: + try: + interfaces[current]["mac_address"] = napalm_helpers.mac(m2.group(1)) + except Exception: + interfaces[current]["mac_address"] = m2.group(1) + + return interfaces + + def get_interfaces_ip(self) -> Dict[str, Dict]: + """Return all configured IP addresses grouped by interface. + + Uses ``ip addr show``. + + Example output:: + + 2: eth0: mtu 1500 ... + inet 192.168.1.1/24 brd 192.168.1.255 scope global eth0 + inet6 fd00::1/64 scope global + """ + output = self._send_command("ip addr show") + interfaces_ip: Dict[str, Dict] = {} + current_iface: Optional[str] = None + + for line in output.splitlines(): + # Interface line + m = re.match(r"^\d+:\s+(\S+?)[@:]", line) + if m: + current_iface = m.group(1) + continue + + if current_iface is None: + continue + + # IPv4 + m = re.match(r"^\s+inet\s+(\S+)", line) + if m: + cidr = m.group(1) + try: + ip_net = netaddr.IPNetwork(cidr) + except (netaddr.AddrFormatError, ValueError): + continue + if current_iface not in interfaces_ip: + interfaces_ip[current_iface] = {} + interfaces_ip[current_iface].setdefault("ipv4", {})[str(ip_net.ip)] = { + "prefix_length": ip_net.prefixlen + } + continue + + # IPv6 + m = re.match(r"^\s+inet6\s+(\S+)", line) + if m: + cidr = m.group(1) + try: + ip_net = netaddr.IPNetwork(cidr) + except (netaddr.AddrFormatError, ValueError): + continue + if current_iface not in interfaces_ip: + interfaces_ip[current_iface] = {} + interfaces_ip[current_iface].setdefault("ipv6", {})[str(ip_net.ip)] = { + "prefix_length": ip_net.prefixlen + } + + return interfaces_ip + + def get_config( + self, + retrieve: str = "all", + full: bool = False, + sanitized: bool = False, + format: str = "text", + ) -> Dict[str, str]: + """Return the device configuration via ``uci export``. + + OpenWrt does not have a distinct startup/candidate config concept. + ``running`` and ``startup`` both return ``uci export`` output. + ``candidate`` is always empty. + """ + configs = {"running": "", "startup": "", "candidate": ""} + + if retrieve in ("all", "running"): + configs["running"] = self._send_command("uci export") + + if retrieve in ("all", "startup"): + configs["startup"] = self._send_command("uci export") + + return configs + + def get_arp_table(self, vrf: str = "") -> List[Dict]: + """Return the ARP/neighbour table. + + Uses ``ip neigh show`` (preferred) which produces:: + + 192.168.1.100 dev br-lan lladdr aa:bb:cc:dd:ee:ff REACHABLE + 192.168.1.1 dev br-lan lladdr 00:11:22:33:44:55 STALE + """ + output = self._send_command(["ip neigh show", "cat /proc/net/arp"]) + arp_table = [] + + # ip neigh show format + for line in output.splitlines(): + line_s = line.strip() + if not line_s: + continue + + # Skip incomplete/failed entries + if "FAILED" in line_s or "INCOMPLETE" in line_s: + continue + + # ip neigh show: "192.168.1.100 dev br-lan lladdr aa:bb:cc:dd:ee:ff REACHABLE" + m = re.match( + r"^(\S+)\s+dev\s+(\S+)\s+lladdr\s+(\S+)", + line_s, + re.I, + ) + if m: + ip_addr = m.group(1) + interface = m.group(2) + mac_raw = m.group(3) + + try: + netaddr.IPAddress(ip_addr) + except (netaddr.AddrFormatError, ValueError): + continue + + try: + mac_addr = napalm_helpers.mac(mac_raw) + except Exception: + mac_addr = mac_raw + + arp_table.append( + { + "interface": interface, + "mac": mac_addr, + "ip": ip_addr, + "age": 0.0, + } + ) + continue + + # /proc/net/arp fallback: "IP address HW type Flags HW address Mask Device" + # skip header + if line_s.startswith("IP address"): + continue + parts = line_s.split() + if len(parts) >= 6: + ip_addr = parts[0] + mac_raw = parts[3] + interface = parts[5] + + try: + netaddr.IPAddress(ip_addr) + except (netaddr.AddrFormatError, ValueError): + continue + + if mac_raw in ("00:00:00:00:00:00", ""): + continue + + try: + mac_addr = napalm_helpers.mac(mac_raw) + except Exception: + mac_addr = mac_raw + + arp_table.append( + { + "interface": interface, + "mac": mac_addr, + "ip": ip_addr, + "age": 0.0, + } + ) + + return arp_table + + def get_mac_address_table(self) -> List[Dict]: + """Return the bridge forwarding database (MAC address table). + + Uses ``bridge fdb show`` which produces:: + + aa:bb:cc:dd:ee:ff dev br-lan master br-lan permanent + 11:22:33:44:55:66 dev eth0.1 vlan 1 master br-lan + """ + output = self._send_command("bridge fdb show") + mac_table = [] + + for line in output.splitlines(): + line_s = line.strip() + if not line_s: + continue + + m = re.match(r"^(\S+)\s+dev\s+(\S+)", line_s) + if not m: + continue + + mac_raw = m.group(1) + interface = m.group(2) + + # Skip broadcast/multicast self-entries that are always present + if mac_raw.lower() in ("ff:ff:ff:ff:ff:ff", "33:33:00:00:00:01"): + continue + + static = "permanent" in line_s or "static" in line_s + + # Extract VLAN if present: "vlan 10" + vlan = 0 + vlan_m = re.search(r"\bvlan\s+(\d+)", line_s) + if vlan_m: + vlan = int(vlan_m.group(1)) + + try: + mac_addr = napalm_helpers.mac(mac_raw) + except Exception: + mac_addr = mac_raw + + mac_table.append( + { + "mac": mac_addr, + "interface": interface, + "vlan": vlan, + "static": static, + "active": True, + "moves": None, + "last_move": None, + } + ) + + return mac_table + + def get_lldp_neighbors(self) -> Dict[str, List[Dict]]: + """Return LLDP neighbors (requires ``lldpd`` to be installed on the device). + + Uses ``lldpctl -f keyvalue`` output:: + + lldp.eth0.port.ifname=eth1 + lldp.eth0.chassis.name=router-core + """ + neighbors: Dict[str, List[Dict]] = {} + for row in self._get_lldp_table(): + neighbors.setdefault(row["local_port"], []).append( + {"hostname": row["system_name"], "port": row["port_id"]} + ) + return neighbors + + def _get_lldp_table(self) -> List[Dict]: + """Parse ``lldpctl -f keyvalue`` into a list of row dicts. + + Ensures ``lldpd`` is enabled and running before querying; if it was + not already running the daemon needs time to discover neighbors so + the first call after a fresh install will return an empty list. + """ + # Enable + start lldpd if not already running (idempotent / silent) + self._send_command( + "pgrep lldpd >/dev/null 2>&1 || " + "(/etc/init.d/lldpd enable 2>/dev/null; /etc/init.d/lldpd start 2>/dev/null)" + ) + output = self._send_command("lldpctl -f keyvalue") + rows: List[Dict] = [] + + # Group by local interface prefix: lldp..* + entries: Dict[str, Dict[str, str]] = {} + for line in output.splitlines(): + line_s = line.strip() + if "=" not in line_s: + continue + key, _, value = line_s.partition("=") + parts = key.split(".") + # parts: ['lldp', '', , , ...] + if len(parts) < 3 or parts[0] != "lldp": + continue + iface = parts[1] + subkey = ".".join(parts[2:]) + entries.setdefault(iface, {})[subkey] = value + + for iface, data in entries.items(): + rows.append( + { + "local_port": iface, + "remote_chassis_id": data.get("chassis.mac", data.get("chassis.id.value", "")), + "port_id": data.get("port.ifname", data.get("port.id.value", "")), + "mgmt_address": data.get("chassis.mgmt-ip", ""), + "port_description": data.get("port.descr", ""), + "system_name": data.get("chassis.name", ""), + } + ) + + return rows + + def get_lldp_neighbors_detail(self, interface: str = "") -> Dict[str, List[Dict]]: + """Return detailed LLDP neighbor info.""" + details: Dict[str, List[Dict]] = {} + + for row in self._get_lldp_table(): + if interface and row["local_port"] != interface: + continue + details.setdefault(row["local_port"], []).append( + { + "parent_interface": "", + "remote_port": row["port_id"], + "remote_port_description": row["port_description"], + "remote_chassis_id": row["remote_chassis_id"], + "remote_system_name": row["system_name"], + "remote_system_description": "", + "remote_system_capab": [], + "remote_system_enable_capab": [], + } + ) + + return details + + def get_environment(self) -> Dict: + """Return device environment data (CPU, memory). + + CPU usage from ``/proc/stat`` (two samples 1 second apart via ``awk``). + Memory from ``/proc/meminfo``. + """ + cpu_out = self._send_command( + "awk '/^cpu /{idle1=$5; total1=$2+$3+$4+$5+$6+$7+$8} END{print (1-(idle1/total1))*100}' /proc/stat" + ) + mem_out = self._send_command("cat /proc/meminfo") + + cpu_pct = 0.0 + try: + cpu_pct = float(cpu_out.strip()) + except (ValueError, AttributeError): + pass + + mem_total = 0 + mem_available = 0 + for line in mem_out.splitlines(): + if line.startswith("MemTotal:"): + try: + mem_total = int(line.split()[1]) + except (IndexError, ValueError): + pass + elif line.startswith("MemAvailable:"): + try: + mem_available = int(line.split()[1]) + except (IndexError, ValueError): + pass + + return { + "fans": {}, + "temperature": {}, + "power": {}, + "cpu": {0: {"%usage": round(cpu_pct, 1)}}, + "memory": { + "available_ram": mem_available * 1024, + "used_ram": (mem_total - mem_available) * 1024, + }, + } + + def get_interfaces_counters(self) -> Dict[str, Dict]: + """Return per-interface packet and byte counters from ``/proc/net/dev``. + + ``/proc/net/dev`` columns (Receive | Transmit):: + + face |bytes packets errs drop fifo frame compressed multicast| \ + bytes packets errs drop fifo colls carrier compressed + """ + output = self._send_command("cat /proc/net/dev") + counters: Dict[str, Dict] = {} + + for line in output.splitlines(): + # Skip header lines + if "|" in line or "Inter" in line: + continue + line_s = line.strip() + if not line_s: + continue + + parts = line_s.replace(":", " ").split() + if len(parts) < 17: + continue + + iface = parts[0] + try: + counters[iface] = { + "tx_errors": int(parts[10]), + "rx_errors": int(parts[3]), + "tx_discards": int(parts[11]), + "rx_discards": int(parts[4]), + "tx_octets": int(parts[9]), + "rx_octets": int(parts[1]), + "tx_unicast_packets": int(parts[10 - 1]), # packets field + "rx_unicast_packets": int(parts[2]), + "tx_multicast_packets": 0, + "rx_multicast_packets": int(parts[8]), + "tx_broadcast_packets": 0, + "rx_broadcast_packets": 0, + } + except (IndexError, ValueError): + continue + + return counters + + # ------------------------------------------------------------------ + # NAPALM configuration management + # ------------------------------------------------------------------ + + def load_merge_candidate( + self, filename: Optional[str] = None, config: Optional[str] = None + ) -> None: + """Stage a set of UCI commands to be applied to the running config. + + *config* is a plain-text string of UCI commands (``uci set``, + ``uci add``, ``uci del``, etc.) – one command per line. Blank lines + and lines starting with ``#`` are ignored. + + The configuration is **not** applied until :meth:`commit_config` is + called. + + :raises MergeConfigException: on invalid input. + """ + if filename is not None: + try: + with open(filename) as fh: + config = fh.read() + except OSError as exc: + raise MergeConfigException(str(exc)) from exc + if config is None: + raise MergeConfigException("Either 'filename' or 'config' must be provided.") + self._candidate_config = config + self._candidate_mode = "merge" + + def load_replace_candidate( + self, filename: Optional[str] = None, config: Optional[str] = None + ) -> None: + """Stage a full ``uci export`` replacement candidate. + + The candidate should be a complete ``uci export`` output. + :meth:`compare_config` shows a unified diff against the current config. + :meth:`commit_config` imports the candidate via ``uci import`` and + commits all affected packages. + + :raises ReplaceConfigException: on invalid input. + """ + if filename is not None: + try: + with open(filename) as fh: + config = fh.read() + except OSError as exc: + raise ReplaceConfigException(str(exc)) from exc + if config is None: + raise ReplaceConfigException("Either 'filename' or 'config' must be provided.") + self._candidate_config = config + self._candidate_mode = "replace" + + def compare_config(self) -> str: + """Return a human-readable diff of the pending candidate vs running config. + + For a **merge** candidate: returns the staged UCI commands prefixed + with ``+``. + + For a **replace** candidate: returns a unified diff between the current + ``uci export`` output and the candidate text. + + Returns an empty string when no candidate is staged. + """ + if self._candidate_config is None: + return "" + + if self._candidate_mode == "merge": + lines = [] + for line in self._candidate_config.splitlines(): + if line.strip() and not line.strip().startswith("#"): + lines.append(f"+{line}") + return "\n".join(lines) + + # replace mode – unified diff + import difflib + running = self._send_command("uci export") + diff = difflib.unified_diff( + running.splitlines(), + self._candidate_config.splitlines(), + fromfile="running-config", + tofile="candidate-config", + lineterm="", + ) + return "\n".join(diff) + + def commit_config(self, message: str = "", revert_in: Optional[int] = None) -> None: + """Apply the staged candidate configuration and commit it. + + **Merge mode**: each UCI command line is sent to the device shell, then + ``uci commit`` is called to persist the changes. + + **Replace mode**: the candidate is piped through ``uci import`` and + then ``uci commit`` is called for every affected package. + + :raises MergeConfigException: if no candidate is staged or if + commands are rejected. + :raises ReplaceConfigException: same, for replace candidates. + """ + if self._candidate_config is None: + raise MergeConfigException("No candidate configuration is staged.") + + ex_cls = ReplaceConfigException if self._candidate_mode == "replace" else MergeConfigException + + # Save backup for potential rollback + self._backup_config = self._send_command("uci export") + + errors: List[str] = [] + try: + if self._candidate_mode == "merge": + for line in self._candidate_config.splitlines(): + stripped = line.strip() + if not stripped or stripped.startswith("#"): + continue + out = self._send_command(stripped) + if out and ("uci: " in out.lower() or "error" in out.lower()): + errors.append(f" {stripped!r}: {out}") + self._send_command("uci commit") + else: + # Replace: pipe candidate through uci import + # Write to a temp file and import it + escaped = self._candidate_config.replace("'", "'\\''") + self._send_command(f"printf '%s' '{escaped}' > /tmp/napalm_candidate.uci") + out = self._send_command("uci import < /tmp/napalm_candidate.uci && uci commit") + self._send_command("rm -f /tmp/napalm_candidate.uci") + if out and "error" in out.lower(): + errors.append(out) + except Exception as exc: + raise ex_cls(str(exc)) from exc + + if errors: + raise ex_cls("The following commands were rejected:\n" + "\n".join(errors)) + + self._candidate_config = None + self._candidate_mode = None + + def discard_config(self) -> None: + """Discard the staged candidate configuration without applying it.""" + self._candidate_config = None + self._candidate_mode = None + + def rollback(self) -> None: + """Restore the UCI configuration to the state before the last :meth:`commit_config`. + + Pipes the saved backup through ``uci import`` and then commits. + + :raises CommandErrorException: if no backup is available. + """ + if self._backup_config is None: + raise CommandErrorException( + "No backup configuration available – commit_config has not been called in this session." + ) + + escaped = self._backup_config.replace("'", "'\\''") + self._send_command(f"printf '%s' '{escaped}' > /tmp/napalm_rollback.uci") + self._send_command("uci import < /tmp/napalm_rollback.uci && uci commit") + self._send_command("rm -f /tmp/napalm_rollback.uci") + + self._backup_config = None + + def has_pending_commit(self) -> bool: + """Return True when a candidate configuration is staged but not yet committed.""" + return self._candidate_config is not None + + def get_vlans(self) -> Dict[str, Dict]: + """Return VLAN information with proper tagged/untagged separation. + + Uses ``bridge vlan show`` (DSA-based OpenWrt ≥21.02) for VLAN/port + membership and ``uci show network`` for VLAN names. + + A port marked *PVID Egress Untagged* is an untagged member. + All other VLAN memberships for the same port are tagged. + + Also detects legacy 802.1q sub-interfaces (``eth0.10`` etc.) from + ``ip link show``. The parent interface (``eth0``) is added as a + tagged member for every such VLAN. + """ + bridge_out = self._send_command("bridge vlan show") + uci_out = self._send_command("uci show network") + + # vlan_id → {name, tagged: [], untagged: []} + vlans: Dict[str, Dict] = {} + current_port: Optional[str] = None + + for line in bridge_out.splitlines(): + line_s = line.strip() + if not line_s or line_s.lower().startswith("port"): + continue + + # Port line: "eth0 1 PVID Egress Untagged" + m = re.match(r"^(\S+)\s+(\d+)(.*)", line) + if m: + current_port = m.group(1) + vlan_id = str(int(m.group(2))) + flags = m.group(3).upper() + vlans.setdefault(vlan_id, {"name": "", "tagged": [], "untagged": []}) + if "PVID" in flags or "UNTAGGED" in flags: + if current_port not in vlans[vlan_id]["untagged"]: + vlans[vlan_id]["untagged"].append(current_port) + else: + if current_port not in vlans[vlan_id]["tagged"]: + vlans[vlan_id]["tagged"].append(current_port) + continue + + # Continuation line with only a VLAN ID (tagged for current_port) + m = re.match(r"^(\d+)", line_s) + if m and current_port: + vlan_id = str(int(m.group(1))) + vlans.setdefault(vlan_id, {"name": "", "tagged": [], "untagged": []}) + if current_port not in vlans[vlan_id]["tagged"]: + vlans[vlan_id]["tagged"].append(current_port) + + # Enrich with UCI VLAN names from explicit bridge-vlan sections + uci_entries: Dict[str, Dict[str, str]] = {} + for line in uci_out.splitlines(): + m = re.match(r"network\.@bridge-vlan\[(\d+)\]\.(\w+)='([^']*)'", line.strip()) + if m: + idx, key, value = m.group(1), m.group(2), m.group(3) + uci_entries.setdefault(idx, {})[key] = value + + for entry in uci_entries.values(): + if "vlan" in entry and "name" in entry: + vlan_id = str(int(entry["vlan"])) + if vlan_id in vlans: + vlans[vlan_id]["name"] = entry["name"] + + # Also derive VLAN names from UCI network interface sections that + # reference subinterfaces like eth0.N or br-ap.N: + # network.guest.device='eth0.8' → VLAN 8 name = "guest" + # network.ap_v8.device='br-ap.8' → VLAN 8 name = "ap_v8" + # Only fills in names that are still empty after bridge-vlan lookup. + for line in uci_out.splitlines(): + m = re.match(r"network\.(\w+)\.device='[\w-]+\.(\d+)'", line.strip()) + if m: + section_name, vid_str = m.group(1), m.group(2) + vlan_id = str(int(vid_str)) + if vlan_id in vlans and not vlans[vlan_id]["name"]: + vlans[vlan_id]["name"] = section_name + + # Also detect VLAN sub-interfaces (eth0.10, br-ap.8, …) from ip link show. + # The sub-interface is the untagged egress point; its parent is tagged. + link_out = self._send_command("ip link show") + for line in link_out.splitlines(): + lm = re.match(r"^\d+:\s+(\S+?)[@:]", line) + if not lm: + continue + iface = lm.group(1) + vm = re.match(r"^([\w-]+)\.(\d+)$", iface) # allow hyphens (br-ap) + if not vm: + continue + parent = vm.group(1) # e.g. "eth0" or "br-ap" + vlan_id = str(int(vm.group(2))) # e.g. "10" + vlans.setdefault(vlan_id, {"name": "", "tagged": [], "untagged": []}) + # sub-interface itself → untagged egress + if iface not in vlans[vlan_id]["untagged"] and iface not in vlans[vlan_id]["tagged"]: + vlans[vlan_id]["untagged"].append(iface) + # parent → tagged trunk + if parent not in vlans[vlan_id]["tagged"] and parent not in vlans[vlan_id]["untagged"]: + vlans[vlan_id]["tagged"].append(parent) + + return vlans + + def delete_vlan(self, vlan_id: int) -> None: + """Remove a VLAN from the device by deleting its UCI bridge-vlan section. + + Finds the ``network.@bridge-vlan[N]`` section whose ``.vlan`` matches + *vlan_id*, deletes it and commits. If no matching section is found the + method is a no-op (the VLAN may only exist as an eth0.N sub-interface, + which cannot be deleted via UCI alone). + + :param vlan_id: VLAN ID to remove. + :raises ValueError: If *vlan_id* is out of the valid range. + """ + if not 1 <= vlan_id <= 4094: + raise ValueError(f"VLAN ID {vlan_id} is out of range (1–4094)") + + uci_out = self.cli(["uci show network"]).get("uci show network", "") + idx = None + for line in uci_out.splitlines(): + m = re.match(r"network\.@bridge-vlan\[(\d+)\]\.vlan='(\d+)'", line.strip()) + if m and int(m.group(2)) == vlan_id: + idx = m.group(1) + break + + if idx is None: + # No explicit bridge-vlan section — nothing to delete via UCI + return + + self.cli([ + f"uci delete network.@bridge-vlan[{idx}]", + "uci commit network", + "/etc/init.d/network reload", + ]) + + # ------------------------------------------------------------------ + # CLI pass-through + # ------------------------------------------------------------------ + + def get_ssids(self) -> Dict[str, Dict]: + """Return configured SSIDs from UCI wireless configuration. + + Parses ``uci show wireless`` for ``wifi-iface`` entries and enriches + each entry with: + + * ``band`` — human-readable frequency band ("2.4 GHz", "5 GHz", "6 GHz") + derived from the radio's ``band`` or ``hwmode`` UCI key. + * ``encryption`` — human-readable security mode ("WPA2-PSK", "Open", …). + + When the same SSID name is broadcast on multiple radios, the keys in + the returned dict are disambiguated as ``"ssid (2.4 GHz)"`` / + ``"ssid (5 GHz)"``. + """ + uci_out = self._send_command("uci show wireless") + + # Collect radio band info: radio0 → "2g", radio1 → "5g", … + radio_bands: Dict[str, str] = {} + iface_entries: Dict[str, Dict[str, str]] = {} + + # First pass: identify named sections that are wifi-iface types and + # collect radio band info. + named_iface_sections: set = set() + for line in uci_out.splitlines(): + line_s = line.strip() + # named wifi-iface declaration: wireless.managed_family_2g=wifi-iface + nm = re.match(r"wireless\.(\w+)=wifi-iface", line_s) + if nm: + named_iface_sections.add(nm.group(1)) + continue + # radio device config: wireless.radio0.band='2g' + rm = re.match(r"wireless\.(radio\d+)\.(band|hwmode)='([^']*)'", line_s) + if rm: + radio, key, val = rm.group(1), rm.group(2), rm.group(3) + if key == "band" or radio not in radio_bands: + radio_bands[radio] = val + + # Second pass: collect iface properties (both anonymous and named sections) + for line in uci_out.splitlines(): + line_s = line.strip() + # radio device config (already handled above) + rm = re.match(r"wireless\.(radio\d+)\.(band|hwmode)='([^']*)'", line_s) + if rm: + radio, key, val = rm.group(1), rm.group(2), rm.group(3) + # Prefer 'band' over 'hwmode' when both present + if key == "band" or radio not in radio_bands: + radio_bands[radio] = val + continue + # anonymous wifi-iface values: wireless.@wifi-iface[0].ssid='MyNet' + im = re.match(r"wireless\.@wifi-iface\[(\d+)\]\.(\w+)='([^']*)'", line_s) + if im: + idx, key, val = im.group(1), im.group(2), im.group(3) + iface_entries.setdefault(idx, {})[key] = val + continue + # named wifi-iface values: wireless.managed_family_2g.ssid='manivong' + nm = re.match(r"wireless\.(\w+)\.(\w+)='([^']*)'", line_s) + if nm and nm.group(1) in named_iface_sections: + section, key, val = nm.group(1), nm.group(2), nm.group(3) + iface_entries.setdefault(section, {})[key] = val + + def _band_label(radio: str) -> str: + raw = radio_bands.get(radio, "").lower() + if raw in ("2g", "11g", "b", "g", "bg", "bgn", "b/g", "b/g/n"): + return "2.4 GHz" + if raw in ("5g", "11a", "a", "ac", "ax5", "a/n", "a/n/ac"): + return "5 GHz" + if raw in ("6g", "ax6"): + return "6 GHz" + return "" + + _ENC_MAP = { + "": "Open", "none": "Open", "0": "Open", + "wep": "WEP", "wep-open": "WEP (Open)", "wep-shared": "WEP (Shared)", + "psk": "WPA-PSK", + "psk+ccmp": "WPA-PSK", + "psk-mixed": "WPA/WPA2-PSK", + "psk2": "WPA2-PSK", + "psk2+ccmp": "WPA2-PSK", + "psk2+aes": "WPA2-PSK", + "psk3": "WPA3-SAE", + "psk2+psk3": "WPA2/WPA3", + "sae": "WPA3-SAE", + "sae-mixed": "WPA2/WPA3", + "wpa": "WPA-Enterprise", + "wpa2": "WPA2-Enterprise", + "wpa3": "WPA3-Enterprise", + "ccmp": "WPA2-PSK", + } + + def _enc_label(enc_raw: str) -> str: + return _ENC_MAP.get(enc_raw.lower(), enc_raw.upper() or "Open") + + # Build network→vlan_id map from UCI network config. + # A wifi-iface has option network='ap_7'; the corresponding UCI network + # interface has either an explicit vid ('7') or a bridge device whose + # name encodes the VLAN, e.g. br-ap.7 → VLAN 7. + def _vlan_from_device(dev: str) -> Optional[int]: + m = re.search(r"\.(\d+)$", dev) + if m: + return int(m.group(1)) + return None + + net_vlan: Dict[str, int] = {} + try: + net_out = self._send_command("uci show network 2>/dev/null || true") + net_entries: Dict[str, Dict[str, str]] = {} + for line in net_out.splitlines(): + line_s = line.strip() + m = re.match(r"network\.(\w+)\.(\w+)='([^']*)'", line_s) + if m: + iface, key, val = m.group(1), m.group(2), m.group(3) + net_entries.setdefault(iface, {})[key] = val + for iface, props in net_entries.items(): + vid_str = props.get("vid") or props.get("vlan") + if vid_str and vid_str.isdigit(): + net_vlan[iface] = int(vid_str) + continue + dev = props.get("device", "") + vlan = _vlan_from_device(dev) + if vlan is not None: + net_vlan[iface] = vlan + except Exception: + pass # Non-fatal: VLAN info is optional enrichment + + # Build result; group entries with the same SSID name, merging bands + result: Dict[str, Dict] = {} + # Intermediate: ssid -> list of bands seen + ssid_bands: Dict[str, List[str]] = {} + for entry in iface_entries.values(): + ssid = entry.get("ssid") + if not ssid: + continue + radio = entry.get("device", "") + band = _band_label(radio) + disabled = entry.get("disabled", "0") == "1" + enc_raw = entry.get("encryption", "") or "" + encryption = _enc_label(enc_raw) + hidden = entry.get("hidden", "0") == "1" + network_name = entry.get("network", "") + vlan_id: Optional[int] = net_vlan.get(network_name) + ft_enabled = entry.get("ieee80211r", "0") == "1" + ft_mobility_domain = entry.get("mobility_domain", "") + ft_over_ds = entry.get("ft_over_ds", "1") == "1" + client_isolation = entry.get("isolate", "0") == "1" + _max_raw = entry.get("maxassoc") + max_clients: Optional[int] = int(_max_raw) if _max_raw and str(_max_raw).isdigit() else None + + if ssid in result: + # Merge: append band if not already present + if band and band not in ssid_bands[ssid]: + ssid_bands[ssid].append(band) + # Keep alphabetical order so 2.4 GHz comes before 5 GHz + ssid_bands[ssid].sort() + result[ssid]["band"] = " + ".join(ssid_bands[ssid]) + result[ssid]["bands_list"] = list(ssid_bands[ssid]) + # If one radio is enabled, the SSID counts as enabled + if not disabled: + result[ssid]["enabled"] = True + # Keep vlan_id if not yet set + if result[ssid].get("vlan_id") is None and vlan_id is not None: + result[ssid]["vlan_id"] = vlan_id + # FT: if any radio has ieee80211r enabled, mark the SSID as FT-enabled + if ft_enabled: + result[ssid]["ieee80211r"] = True + result[ssid]["mobility_domain"] = ft_mobility_domain + result[ssid]["ft_over_ds"] = ft_over_ds + # Client isolation: if any iface has it, mark True + if client_isolation: + result[ssid]["client_isolation"] = True + # Max clients: keep first non-None value + if max_clients is not None and result[ssid].get("max_clients") is None: + result[ssid]["max_clients"] = max_clients + else: + ssid_bands[ssid] = [band] if band else [] + result[ssid] = { + "enabled": not disabled, + "radio": radio, + "band": band, + "bands_list": list(ssid_bands[ssid]), + "bssid": "", + "encryption": encryption, + "encryption_uci": enc_raw, + "hidden": hidden, + "client_isolation": client_isolation, + "max_clients": max_clients, + "clients": 0, + "vlan_id": vlan_id, + "ieee80211r": ft_enabled, + "mobility_domain": ft_mobility_domain, + "ft_over_ds": ft_over_ds, + } + return result + + def get_wireless_clients(self) -> List[Dict]: + """Return currently associated wireless clients from all AP interfaces. + + Uses ``iw dev`` to discover AP-mode interfaces and then + ``iw dev station dump`` to collect per-client statistics. + """ + from napalm_device_types.models import WirelessClientDict + + # Step 1: discover interfaces and their SSIDs / radio mappings + iw_out = self._send_command("iw dev 2>/dev/null || true") + + iface_info: Dict[str, Dict[str, str]] = {} + current_phy: str = "" + current_iface: str = "" + + for line in iw_out.splitlines(): + stripped = line.strip() + phy_m = re.match(r"^phy#(\d+)$", stripped) + if phy_m: + current_phy = f"radio{phy_m.group(1)}" + current_iface = "" + continue + + iface_m = re.match(r"^Interface\s+(\S+)$", stripped) + if iface_m: + current_iface = iface_m.group(1) + iface_info[current_iface] = {"ssid": "", "radio": current_phy, "type": ""} + continue + + if not current_iface: + continue + + ssid_m = re.match(r"^ssid\s+(.+)$", stripped) + if ssid_m: + iface_info[current_iface]["ssid"] = ssid_m.group(1) + continue + + type_m = re.match(r"^type\s+(\S+)$", stripped) + if type_m: + iface_info[current_iface]["type"] = type_m.group(1) + continue + + # channel 6 (2437 MHz), width: 20 MHz, ... + chan_m = re.match(r"^channel\s+\d+\s+\((\d+)\s+MHz\)", stripped) + if chan_m: + try: + freq = int(chan_m.group(1)) + if freq < 3000: + iface_info[current_iface]["band"] = "2.4 GHz" + elif freq < 6000: + iface_info[current_iface]["band"] = "5 GHz" + else: + iface_info[current_iface]["band"] = "6 GHz" + except ValueError: + pass + + # Filter to AP-mode interfaces only + ap_ifaces = { + name: info + for name, info in iface_info.items() + if info.get("type", "").upper() in ("AP", "AP/VLAN") + } + + if not ap_ifaces: + return [] + + # Step 2: fetch station dumps for all AP interfaces in one SSH call + dump_cmd = " ; ".join( + f"echo '=== {name} ===' && iw dev {name} station dump 2>/dev/null || true" + for name in ap_ifaces + ) + station_out = self._send_command(dump_cmd) + + # Step 3: parse station dump output + results: List[Dict] = [] + active_iface: str = "" + current_station: Optional[Dict] = None + + def _flush() -> None: + if current_station and current_station.get("mac"): + info = ap_ifaces.get(active_iface, {}) + results.append(WirelessClientDict( + mac=current_station["mac"], + ssid=info.get("ssid", ""), + radio=info.get("band") or info.get("radio", ""), + signal=current_station.get("signal", 0), + noise=0, + tx_rate=current_station.get("tx_rate", 0.0), + rx_rate=current_station.get("rx_rate", 0.0), + uptime=current_station.get("uptime", 0), + )) + + for line in station_out.splitlines(): + stripped = line.strip() + + # Section header injected above: === wlan0 === + hdr_m = re.match(r"^=== (\S+) ===$", stripped) + if hdr_m: + _flush() + active_iface = hdr_m.group(1) + current_station = None + continue + + # Station aa:bb:cc:dd:ee:ff (on wlan0) + sta_m = re.match(r"^Station\s+([\da-fA-F:]{17})\s+\(", stripped) + if sta_m: + _flush() + current_station = {"mac": sta_m.group(1)} + continue + + if current_station is None: + continue + + # signal: -65 dBm (may be "signal: -65 [-65] dBm") + sig_m = re.match(r"^signal:\s+([-\d]+)", stripped) + if sig_m: + try: + current_station["signal"] = int(sig_m.group(1)) + except ValueError: + pass + continue + + # tx bitrate: 54.0 MBit/s + tx_m = re.match(r"^tx bitrate:\s+([\d.]+)", stripped) + if tx_m: + try: + current_station["tx_rate"] = float(tx_m.group(1)) + except ValueError: + pass + continue + + # rx bitrate: 72.2 MBit/s + rx_m = re.match(r"^rx bitrate:\s+([\d.]+)", stripped) + if rx_m: + try: + current_station["rx_rate"] = float(rx_m.group(1)) + except ValueError: + pass + continue + + # connected time: 3600 seconds + uptime_m = re.match(r"^connected time:\s+(\d+)", stripped) + if uptime_m: + try: + current_station["uptime"] = int(uptime_m.group(1)) + except ValueError: + pass + + _flush() + return results + + def get_radio_status(self) -> Dict[str, Dict]: + """Return radio status from UCI and iwinfo. + + Combines ``uci show wireless`` for static config with ``iwinfo`` + output for runtime channel/frequency and tx-power data. + + Returns a dict keyed by radio name (e.g. ``"radio0"``) with: + + * enabled (bool) + * band (str) — ``"2.4GHz"``, ``"5GHz"``, ``"6GHz"`` + * channel (int) — 0 means auto + * channel_width (int) — channel bandwidth in MHz (0 if unknown) + * tx_power (int) — TX power in dBm (0 if unknown) + * frequency (float) — centre frequency in MHz (0 if unknown) + * htmode (str) — e.g. ``"HT20"``, ``"VHT80"``, ``"HE80"`` + * country (str) — regulatory country code, e.g. ``"DE"`` + """ + from napalm_device_types.models import RadioStatusDict + + uci_out = self._send_command("uci show wireless") + radios: Dict[str, Dict[str, str]] = {} + + for line in uci_out.splitlines(): + # wifi-device section: wireless.radio0.band='2g' + m = re.match(r"wireless\.(radio\d+)\.(\w+)='([^']*)'", line.strip()) + if m: + radio, key, value = m.group(1), m.group(2), m.group(3) + radios.setdefault(radio, {})[key] = value + + result: Dict[str, Dict] = {} + for radio, cfg in sorted(radios.items()): + band_raw = cfg.get("band", cfg.get("hwmode", "")) + # Normalise band: '2g'/'11g' → '2.4GHz', '5g'/'11a' → '5GHz', '6g' → '6GHz' + if band_raw in ("2g", "11g", "b", "g", "bg", "bgn"): + band = "2.4GHz" + elif band_raw in ("5g", "11a", "a", "ac", "ax5"): + band = "5GHz" + elif band_raw in ("6g", "ax6"): + band = "6GHz" + else: + band = band_raw or "unknown" + + try: + channel = int(cfg.get("channel", 0)) + except (ValueError, TypeError): + channel = 0 # 'auto' + + try: + tx_power = int(cfg.get("txpower", 0)) + except (ValueError, TypeError): + tx_power = 0 + + disabled = cfg.get("disabled", "0") == "1" + htmode = cfg.get("htmode", "") + country = cfg.get("country", "") + + # Derive channel_width from htmode string (e.g. VHT80 → 80 MHz) + _HTMODE_WIDTH = { + "HT20": 20, "HT40": 40, + "VHT20": 20, "VHT40": 40, "VHT80": 80, "VHT80+80": 80, "VHT160": 160, + "HE20": 20, "HE40": 40, "HE80": 80, "HE160": 160, + "EHT20": 20, "EHT40": 40, "EHT80": 80, "EHT160": 160, "EHT320": 320, + } + channel_width = _HTMODE_WIDTH.get(htmode.upper(), 0) + + result[radio] = { + **RadioStatusDict( + enabled=not disabled, + band=band, + channel=channel, + channel_width=channel_width, + tx_power=tx_power, + frequency=0.0, # enriched below via iwinfo + ), + "htmode": htmode, + "country": country, + } + + # Enrich with iwinfo runtime data (channel, frequency, tx_power, channel_width) + # iwinfo groups output per interface; we need to map interface → radio. + # "phy0-ap0 ESSID: "MyNet"" → radio0 + # " Tx-Power: 23 dBm" + # " Channel: 44 (5.220 GHz), Width: 80 MHz" + try: + iwinfo_out = self._send_command("iwinfo 2>/dev/null || true") + except Exception: + iwinfo_out = "" + + current_radio: Optional[str] = None + for line in iwinfo_out.splitlines(): + # Interface header line: "phy0-ap0 ESSID: ..." + iface_m = re.match(r"^(\S+)\s+ESSID:", line) + if iface_m: + iface_name = iface_m.group(1) + phy_m = re.match(r"^phy(\d+)", iface_name) + if phy_m: + current_radio = f"radio{phy_m.group(1)}" + else: + current_radio = None + continue + + if current_radio is None or current_radio not in result: + continue + + # Channel and frequency: "Channel: 44 (5.220 GHz), Width: 80 MHz" + ch_m = re.search(r"Channel:\s+(\d+)\s+\(([\d.]+)\s+GHz\)", line) + if ch_m: + result[current_radio]["channel"] = int(ch_m.group(1)) + result[current_radio]["frequency"] = float(ch_m.group(2)) * 1000 + + # Width (MHz): "Width: 80 MHz" or ", Width: 80 MHz" + width_m = re.search(r"Width:\s+(\d+)\s+MHz", line) + if width_m: + result[current_radio]["channel_width"] = int(width_m.group(1)) + + # Tx-Power: "Tx-Power: 23 dBm" + pwr_m = re.search(r"Tx-Power:\s+(\d+)\s+dBm", line) + if pwr_m: + result[current_radio]["tx_power"] = int(pwr_m.group(1)) + + return result + + def get_system_config(self) -> Dict: + """Return system-level configuration from UCI. + + Reads ``uci show system`` and ``uci show dropbear`` to collect: + + * hostname (str) + * timezone (str) — POSIX TZ string, e.g. ``"CET-1CEST,M3.5.0,M10.5.0/3"`` + * zonename (str) — human-readable name, e.g. ``"Europe/Berlin"`` + * ntp_servers (list[str]) + * dropbear_port (int) — SSH port + * dropbear_password_auth (bool) — whether password login is allowed + * dropbear_root_password_auth (bool) + """ + sys_out = self._send_command("uci show system 2>/dev/null || true") + db_out = self._send_command("uci show dropbear 2>/dev/null || true") + + sys_cfg: Dict[str, str] = {} + for line in sys_out.splitlines(): + m = re.match(r"system\.@system\[0\]\.(\w+)='([^']*)'", line.strip()) + if m: + sys_cfg[m.group(1)] = m.group(2) + + # NTP server list: all servers on a single line, space-separated quoted values + # e.g. system.ntp.server='0.openwrt.pool.ntp.org' '1.openwrt.pool.ntp.org' ... + ntp_servers: List[str] = [] + for line in sys_out.splitlines(): + if re.match(r"system\.ntp\.server=", line.strip()): + ntp_servers = re.findall(r"'([^']+)'", line) + break + + # Dropbear settings + db_cfg: Dict[str, str] = {} + for line in db_out.splitlines(): + # May be @dropbear[0] or named section + m = re.match(r"dropbear\.[@\w]+\.(\w+)='([^']*)'", line.strip()) + if m: + db_cfg.setdefault(m.group(1), m.group(2)) + + try: + ssh_port = int(db_cfg.get("Port", "22")) + except (ValueError, TypeError): + ssh_port = 22 + + def _bool_uci(val: str, default: bool = True) -> bool: + return val.lower() not in ("0", "off", "false", "no") if val else default + + return { + "hostname": sys_cfg.get("hostname", ""), + "timezone": sys_cfg.get("timezone", ""), + "zonename": sys_cfg.get("zonename", ""), + "ntp_servers": ntp_servers, + "dropbear_port": ssh_port, + "dropbear_password_auth": _bool_uci(db_cfg.get("PasswordAuth", "on")), + "dropbear_root_password_auth": _bool_uci(db_cfg.get("RootPasswordAuth", "on")), + } + + # ------------------------------------------------------------------ + # CLI pass-through + # ------------------------------------------------------------------ + + def cli( + self, commands: List[str], encoding: str = "text" + ) -> Dict[str, Union[str, Dict]]: + """Execute a list of shell commands and return their output. + + Each command is run via SSH. The key in the returned dictionary is + the command string; the value is the raw text output. + + Example:: + + device.cli(["uname -a", "cat /etc/openwrt_release"]) + """ + result: Dict[str, Union[str, Dict]] = {} + for cmd in commands: + result[cmd] = self._send_command(cmd) + return result + + # ------------------------------------------------------------------ + # Package management (opkg ≤ OpenWrt 23 / apk ≥ OpenWrt 24) + # ------------------------------------------------------------------ + + def get_packages(self) -> List[Dict]: + """Return installed packages from the device's package manager. + + Automatically detects whether to use ``apk`` (OpenWrt 24+, Alpine + APK) or ``opkg`` (older OpenWrt releases). Returns one entry per + installed package. + """ + pm = self._send_command("command -v apk 2>/dev/null || echo __no_apk__").strip() + if "__no_apk__" not in pm and pm: + return self._get_packages_apk() + return self._get_packages_opkg() + + def _get_packages_opkg(self) -> List[Dict]: + """Parse ``opkg status`` (dpkg-style stanzas).""" + out = self._send_command("opkg status") + packages: List[Dict] = [] + stanza: Dict[str, str] = {} + for raw in out.splitlines(): + line = raw.rstrip() + if line == "": + if stanza.get("Package"): + packages.append(self._opkg_stanza_to_dict(stanza)) + stanza = {} + elif line[:1] in (" ", "\t"): + # Continuation of previous field (e.g. multi-line Description) + last_key = list(stanza)[-1] if stanza else None + if last_key: + stanza[last_key] += " " + line.strip() + elif ":" in line: + key, _, val = line.partition(":") + stanza[key.strip()] = val.strip() + if stanza.get("Package"): + packages.append(self._opkg_stanza_to_dict(stanza)) + return sorted(packages, key=lambda p: p["name"].lower()) + + @staticmethod + def _opkg_stanza_to_dict(stanza: Dict[str, str]) -> Dict: + status = stanza.get("Status", "") + try: + size = int(stanza.get("Installed-Size", 0) or 0) + except ValueError: + size = 0 + return { + "name": stanza["Package"], + "version": stanza.get("Version", ""), + "installed": "installed" in status.lower(), + "description": stanza.get("Description", ""), + "size": size, + "source": stanza.get("Section", ""), + } + + def _get_packages_apk(self) -> List[Dict]: + """Parse ``apk list --installed`` output. + + Line format:: + + busybox-1.37.0-r0 x86_64 {busybox} (GPL-2.0-only) [installed] + kmod-nft-bridge-6.6.75-r0 mips_24kc {kmod-nft-bridge} (GPL-2.0-only) [installed] + """ + out = self._send_command("apk list --installed 2>/dev/null") + packages: List[Dict] = [] + for line in out.splitlines(): + line = line.strip() + if not line or "[installed]" not in line: + continue + # Split name from version: version always starts with a digit after '-' + m = re.match(r"^(.*?)-(\d\S*)\s+\S+\s+\{(\S+)\}", line) + if m: + name, version, origin = m.group(1), m.group(2), m.group(3) + else: + # Minimal fallback: first token only + token = line.split()[0] + vm = re.search(r"-(\d\S*)$", token) + name = token[: vm.start()] if vm else token + version = vm.group(1) if vm else "" + origin = "" + packages.append({ + "name": name, + "version": version, + "installed": True, + "description": "", + "size": 0, + "source": origin, + }) + return sorted(packages, key=lambda p: p["name"].lower()) + + def _pm_type(self) -> str: + """Return ``'apk'`` if device has apk (OpenWrt 24+), otherwise ``'opkg'``.""" + out = self._send_command("command -v apk 2>/dev/null || echo __no_apk__").strip() + return "apk" if ("__no_apk__" not in out and out) else "opkg" + + def search_packages(self, query: str) -> List[Dict]: + """Search available packages matching *query* (name or description). + + Runs ``opkg update`` / ``apk update`` first to ensure the package + index is populated (OpenWrt stores it in RAM and loses it on reboot). + """ + import shlex + safe_q = shlex.quote(query) + if self._pm_type() == "apk": + # Refresh index (no-ops if already current, safe to run every time) + self._send_command("apk update 2>/dev/null || true") + out = self._send_command(f"apk search {safe_q} 2>/dev/null") + installed = {p["name"] for p in self._get_packages_apk()} + packages: List[Dict] = [] + for line in out.splitlines(): + line = line.strip() + if not line: + continue + m = re.match(r"^(.*?)-(\d\S*)(?:\s+(.*))?$", line) + if m: + name, version, description = m.group(1), m.group(2), (m.group(3) or "") + else: + name, version, description = line, "", "" + packages.append({ + "name": name, + "version": version, + "installed": name in installed, + "description": description, + "size": 0, + "source": "", + }) + else: + # opkg lists live in /var/opkg-lists/ (RAM) — cleared on reboot + self._send_command("opkg update 2>/dev/null || true") + out = self._send_command(f"opkg list 2>/dev/null | grep -i {safe_q}") + installed = {p["name"] for p in self._get_packages_opkg()} + packages = [] + for line in out.splitlines(): + line = line.strip() + if not line: + continue + parts = line.split(" - ", 2) + name = parts[0].strip() + version = parts[1].strip() if len(parts) > 1 else "" + description = parts[2].strip() if len(parts) > 2 else "" + packages.append({ + "name": name, + "version": version, + "installed": name in installed, + "description": description, + "size": 0, + "source": "", + }) + return packages + + @staticmethod + def _clean_pkg_output(raw: str) -> str: + """Strip ANSI/VT100 escape sequences and progress-bar lines.""" + # Strip CSI sequences (\x1b[...X), OSC, charset designations, and + # 2-byte DEC private sequences like ESC 7 (cursor save) / ESC 8 (restore) + cleaned = re.sub( + r'\x1b(?:\[[0-9;?]*[a-zA-Z]|\][^\x07]*\x07|[()][0-9A-Za-z]|[\x30-\x7e])', + '', raw, + ) + # After stripping cursor-save/restore sequences, apk progress updates + # end up concatenated on a single line. Strip those inline patterns. + cleaned = re.sub(r'\s*\d{1,3}%\s*#*', ' ', cleaned) + lines = [] + for segment in cleaned.split('\n'): + # \r overwrites the line; keep only the portion after the last \r + part = segment.split('\r')[-1].strip() + if not part: + continue + # Drop pure progress-bar lines (only #, spaces, digits, %) + if re.match(r'^[#\s\d%]*$', part): + continue + lines.append(part) + return '\n'.join(lines) + + def install_package(self, name: str) -> Dict: + """Install a package by name. Returns ``{"success": bool, "output": str}``.""" + import shlex + safe_name = shlex.quote(name) + if self._pm_type() == "apk": + raw = self._send_command(f"apk add {safe_name} 2>&1") + else: + raw = self._send_command(f"opkg install {safe_name} 2>&1") + out = self._clean_pkg_output(raw) + low = out.lower() + success = not any(kw in low for kw in ("error:", "failed", "not found", "unknown package")) + return {"success": success, "output": out} + + def uninstall_package(self, name: str) -> Dict: + """Remove a package by name. Returns ``{"success": bool, "output": str}``.""" + import shlex + safe_name = shlex.quote(name) + if self._pm_type() == "apk": + raw = self._send_command(f"apk del {safe_name} 2>&1") + else: + raw = self._send_command(f"opkg remove {safe_name} 2>&1") + out = self._clean_pkg_output(raw) + low = out.lower() + success = not any(kw in low for kw in ("error:", "failed", "not found", "unknown package")) + return {"success": success, "output": out} + + # ------------------------------------------------------------------ + # Device warnings & generic actions + # ------------------------------------------------------------------ + + def get_device_warnings(self) -> List[Dict]: + """Return a list of warning dicts for issues detected on this device. + + Currently detects: + - lldpd not installed (LLDP neighbor discovery unavailable) + - package updates available (uses local package cache, no network call) + - update notifications not configured (auc not installed, opkg only) + """ + warnings: List[Dict] = [] + + # 1. LLDP daemon + lldpd_path = self._send_command("which lldpd 2>/dev/null").strip() + if not lldpd_path: + warnings.append({ + "code": "lldpd_not_installed", + "severity": "warning", + "action": "install_lldpd", + }) + + pm = self._pm_type() + + # 2. Package updates available (local cache only – no opkg update) + try: + if pm == "apk": + raw_upg = self._send_command("apk version 2>/dev/null | grep '<'") + else: + raw_upg = self._send_command("opkg list-upgradable 2>/dev/null") + upgradable = [ln.strip() for ln in raw_upg.splitlines() if ln.strip()] + except Exception: + upgradable = [] + + if upgradable: + warnings.append({ + "code": "updates_available", + "severity": "info", + "action": None, + "meta": { + "count": len(upgradable), + "packages": upgradable[:10], + }, + }) + + # 3. Attended sysupgrade client not installed (opkg systems only) + if pm == "opkg": + auc_path = self._send_command("which auc 2>/dev/null").strip() + if not auc_path: + warnings.append({ + "code": "update_notifications_disabled", + "severity": "warning", + "action": "install_auc", + }) + + # 4. base64 not available — needed for efficient config apply + b64_path = self._send_command("command -v base64 2>/dev/null").strip() + if not b64_path: + warnings.append({ + "code": "no_base64", + "severity": "warning", + "action": "install_coreutils_base64", + }) + + return warnings + + def get_services(self) -> List[Dict]: + """Return all system services with their running and enabled state. + + Uses ``ubus call service list`` for running/PID info and + ``/etc/rc.d/S*`` symlinks for enabled-at-boot state. + """ + import json as _json + + # -- enabled set: names from /etc/rc.d/S symlinks ---- + rc_out = self._send_command( + "ls /etc/rc.d/ 2>/dev/null | grep '^S' | sed 's/^S[0-9]*//'" + ) + enabled: set[str] = {s.strip() for s in rc_out.splitlines() if s.strip()} + + # -- running info from procd via ubus --------------------------------- + ubus_raw = self._send_command("ubus call service list 2>/dev/null") + ubus_data: dict = {} + try: + ubus_data = _json.loads(ubus_raw) + except (ValueError, TypeError): + pass + + # Build index from ubus data + service_map: dict[str, dict] = {} + for svc_name, svc_info in ubus_data.items(): + if not isinstance(svc_info, dict): + continue + instances = svc_info.get("instances", {}) + running = any( + inst.get("running", False) + for inst in instances.values() + if isinstance(inst, dict) + ) + pid = next( + ( + inst.get("pid", 0) + for inst in instances.values() + if isinstance(inst, dict) and inst.get("running") + ), + 0, + ) + service_map[svc_name] = {"running": running, "pid": pid} + + # -- all init scripts ------------------------------------------------- + init_raw = self._send_command("ls -1 /etc/init.d/ 2>/dev/null") + init_scripts: set[str] = {s.strip() for s in init_raw.splitlines() if s.strip()} + + # Merge: all known services (from init.d + ubus) + all_names = init_scripts | set(service_map.keys()) + # Exclude procd internal pseudo-service + all_names.discard("") + + result: List[Dict] = [] + for name in sorted(all_names): + info = service_map.get(name, {}) + result.append({ + "name": name, + "running": info.get("running", False), + "enabled": name in enabled, + "pid": info.get("pid", 0), + }) + + return result + + def get_available_updates(self) -> List[Dict]: + """Return list of upgradable packages from the local package manager cache.""" + import re as _re + pm = self._pm_type() + updates: list[dict] = [] + + if pm == "apk": + # Output format: "pkgname-current_ver < new_ver" + raw = self._send_command("apk version 2>/dev/null | grep '<'") + for line in raw.splitlines(): + line = line.strip() + m = _re.match(r'^(.+)-(\d\S*)\s+<\s+(\S+)', line) + if m: + updates.append({ + "name": m.group(1), + "current_version": m.group(2), + "new_version": m.group(3), + }) + else: + # opkg output: "pkgname - current_ver - new_ver" + raw = self._send_command("opkg list-upgradable 2>/dev/null") + for line in raw.splitlines(): + parts = [p.strip() for p in line.split(" - ")] + if len(parts) == 3: + updates.append({ + "name": parts[0], + "current_version": parts[1], + "new_version": parts[2], + }) + + return sorted(updates, key=lambda u: u["name"]) + + def apply_updates(self, packages: List[str]) -> Dict: + """Upgrade the given packages using the device's package manager.""" + import re as _re + for pkg in packages: + if not _re.match(r'^[a-zA-Z0-9_\-\+\.]+$', pkg): + raise ValueError(f"Invalid package name: {pkg!r}") + pm = self._pm_type() + pkg_args = " ".join(packages) + if pm == "apk": + cmd = f"apk upgrade {pkg_args} 2>&1" + else: + cmd = f"opkg upgrade {pkg_args} 2>&1" + output = self._send_command(cmd) + return {"success": True, "output": output} + + def manage_service(self, name: str, action: str) -> Dict: + """Execute a lifecycle action (start/stop/restart/enable/disable) on a service.""" + import re as _re + if not _re.match(r'^[a-zA-Z0-9_\-]+$', name): + raise ValueError(f"Invalid service name: {name!r}") + if action not in ('start', 'stop', 'restart', 'enable', 'disable'): + raise ValueError(f"Invalid action: {action!r}") + output = self._send_command(f"/etc/init.d/{name} {action} 2>&1") + return {"success": True, "output": output} + + def run_device_action(self, action: str) -> Dict: + """Execute a named action on the device.""" + if action == "install_lldpd": + return self._action_install_lldpd() + if action == "install_auc": + return self._action_install_auc() + if action == "install_coreutils_base64": + return self._action_install_coreutils_base64() + raise NotImplementedError(f"Unknown action: {action!r}") + + def _action_install_coreutils_base64(self) -> Dict: + """Install coreutils-base64 via the device package manager.""" + pm = self._pm_type() + if pm == "apk": + raw = self._send_command("apk add coreutils-base64 2>&1") + else: + self._send_command("opkg update 2>&1") + raw = self._send_command("opkg install coreutils-base64 2>&1") + out = self._clean_pkg_output(raw) + low = out.lower() + success = not any(kw in low for kw in ("error:", "failed", "not found", "unknown package")) + return {"success": success, "output": out} + + def _action_install_auc(self) -> Dict: + """Install the attended sysupgrade client (auc) via opkg.""" + self._send_command("opkg update 2>&1") + raw = self._send_command("opkg install auc 2>&1") + out = self._clean_pkg_output(raw) + low = out.lower() + success = not any(kw in low for kw in ("error:", "failed", "not found", "unknown package")) + return {"success": success, "output": out} + + def _action_install_lldpd(self) -> Dict: + """Install lldpd, add eth0 to its interface list and start the service.""" + pm = self._pm_type() + if pm == "apk": + raw = self._send_command("apk add lldpd 2>&1") + else: + self._send_command("opkg update 2>&1") + raw = self._send_command("opkg install lldpd 2>&1") + out = self._clean_pkg_output(raw) + + # Add eth0 to lldpd UCI interface list (idempotent) + current_ifaces = self._send_command("uci get lldpd.config.interface 2>/dev/null").strip() + if "eth0" not in current_ifaces: + self._send_command( + "uci add_list lldpd.config.interface='eth0' 2>/dev/null; " + "uci commit lldpd 2>/dev/null" + ) + + # Enable and start the service + self._send_command( + "/etc/init.d/lldpd enable 2>/dev/null; " + "/etc/init.d/lldpd start 2>/dev/null" + ) + + low = out.lower() + success = not any(kw in low for kw in ("error:", "failed", "not found", "unknown package")) + return {"success": success, "output": out} + + # ------------------------------------------------------------------ + # Hostname configuration + # ------------------------------------------------------------------ + + def set_hostname(self, new_hostname: str) -> None: + """Set the system hostname via UCI and reload the system service.""" + self._send_command( + f"uci set system.@system[0].hostname='{new_hostname}' && " + f"uci commit system && " + f"/etc/init.d/system reload" + ) + + # ------------------------------------------------------------------ + # Users + # ------------------------------------------------------------------ + + def get_users(self) -> Dict[str, Dict]: + """Return users configured on the device. + + Parses ``/etc/passwd`` for accounts with a valid login shell. + SSH public keys are read from ``~/.ssh/authorized_keys`` + (Dropbear also stores root keys at ``/etc/dropbear/authorized_keys``). + + Level mapping: + - UID 0 (root) → 15 (full access) + - all other users → 1 + """ + passwd_out = self._send_command("cat /etc/passwd") + # Root authorized_keys locations on OpenWrt + root_keys_out = self._send_command( + ["cat /root/.ssh/authorized_keys", "cat /etc/dropbear/authorized_keys"] + ) + + users: Dict[str, Dict] = {} + valid_shells = {"/bin/sh", "/bin/ash", "/bin/bash", "/usr/bin/fish"} + + for line in passwd_out.splitlines(): + parts = line.strip().split(":") + if len(parts) < 7: + continue + username, password_hash, uid_str, _, _, home, shell = ( + parts[0], parts[1], parts[2], parts[3], parts[4], parts[5], parts[6], + ) + if shell not in valid_shells: + continue + try: + uid = int(uid_str) + except ValueError: + continue + + level = 15 if uid == 0 else 1 + + # Collect SSH keys for this user + sshkeys: List[str] = [] + if uid == 0: + for line_k in root_keys_out.splitlines(): + line_k = line_k.strip() + if line_k and not line_k.startswith("#"): + sshkeys.append(line_k) + else: + # Try reading per-user authorized_keys + keys_out = self._send_command(f"cat {home}/.ssh/authorized_keys 2>/dev/null") + for line_k in keys_out.splitlines(): + line_k = line_k.strip() + if line_k and not line_k.startswith("#"): + sshkeys.append(line_k) + + users[username] = { + "level": level, + "password": password_hash, + "sshkeys": sshkeys, + } + + return users + + # ------------------------------------------------------------------ + # NTP + # ------------------------------------------------------------------ + + def get_ntp_servers(self) -> Dict[str, Dict]: + """Return configured NTP servers from ``uci show system``. + + UCI example:: + + system.ntp.server='0.openwrt.pool.ntp.org 1.openwrt.pool.ntp.org' + """ + uci_out = self._send_command("uci show system") + servers: Dict[str, Dict] = {} + + for line in uci_out.splitlines(): + # Handles both list and single-value UCI representations + m = re.match(r"system\.ntp\.server(?:\[\d+\])?='([^']*)'", line.strip()) + if m: + for srv in m.group(1).split(): + srv = srv.strip() + if srv: + servers[srv] = {} + + return servers + + def get_ntp_peers(self) -> Dict[str, Dict]: + """Return NTP peers from ``uci show system``. + + OpenWrt's busybox ntpd does not differentiate peers from servers; + the same UCI ``ntp.server`` list is returned. + """ + return self.get_ntp_servers() + + def get_ntp_stats(self) -> List[Dict]: + """Return NTP synchronisation statistics. + + Tries ``ntpq -pn`` first (ntpd), then ``chronyc sources -v`` (chrony). + Returns an empty list when neither tool is available. + + ``ntpq -pn`` example line:: + + *188.114.101.4 188.114.100.1 4 u 107 256 377 164.228 -13.866 2.695 + + ``chronyc sources -v`` example line:: + + ^* 192.168.1.1 2 6 17 8 +2345us[ 0ns] +/- 15ms + """ + ntpq_out = self._send_command("ntpq -pn") + if ntpq_out and not ntpq_out.startswith(("ntpq: ", "sh: ", "ash: ", "command not found")): + return self._parse_ntpq(ntpq_out) + + chrony_out = self._send_command("chronyc sources -v") + if chrony_out and not chrony_out.startswith(("sh: ", "ash: ", "command not found")): + return self._parse_chronyc(chrony_out) + + return [] + + @staticmethod + def _parse_ntpq(output: str) -> List[Dict]: + """Parse ``ntpq -pn`` tabular output.""" + stats = [] + for line in output.splitlines(): + line_s = line.strip() + if not line_s or line_s.startswith(("remote", "=")): + continue + # First char is the tally code (* = synchronized, + = candidate, etc.) + tally = line_s[0] if line_s[0] in "* +-x.o#" else " " + parts = line_s[1:].split() + if len(parts) < 10: + continue + try: + stats.append({ + "remote": parts[0], + "referenceid": parts[1], + "synchronized": tally == "*", + "stratum": int(parts[2]), + "type": parts[3], + "when": parts[4], + "hostpoll": int(parts[5]), + "reachability": int(parts[6], 8), # octal + "delay": float(parts[7]), + "offset": float(parts[8]), + "jitter": float(parts[9]), + }) + except (ValueError, IndexError): + continue + return stats + + @staticmethod + def _parse_chronyc(output: str) -> List[Dict]: + """Parse ``chronyc sources -v`` tabular output.""" + stats = [] + for line in output.splitlines(): + line_s = line.strip() + # Data lines start with ^* ^+ ^- ^? + m = re.match(r"^(\^[*+\-?])\s+(\S+)\s+(\d+)\s+(\d+)\s+(\d+)\s+(\S+)\s+(.*)", line_s) + if not m: + continue + tally = m.group(1) + try: + stats.append({ + "remote": m.group(2), + "referenceid": "", + "synchronized": tally == "^*", + "stratum": int(m.group(3)), + "type": "u", + "when": m.group(6), + "hostpoll": int(m.group(4)), + "reachability": int(m.group(5), 8) if re.match(r"^[0-7]+$", m.group(5)) else 0, + "delay": 0.0, + "offset": 0.0, + "jitter": 0.0, + }) + except (ValueError, IndexError): + continue + return stats + + # ------------------------------------------------------------------ + # SNMP + # ------------------------------------------------------------------ + + def get_snmp_information(self) -> Dict: + """Return SNMP configuration from ``uci show snmpd``. + + UCI example:: + + snmpd.@com2sec[0].community='public' + snmpd.@com2sec[0].secname='public' + snmpd.@system[0].sysContact='root@localhost' + snmpd.@system[0].sysLocation='Unknown' + """ + uci_out = self._send_command("uci show snmpd") + + contact = "" + location = "" + chassis_id = "" + community: Dict[str, Dict] = {} + + # Track com2sec entries by index + com2sec: Dict[str, Dict[str, str]] = {} + + for line in uci_out.splitlines(): + line_s = line.strip() + m = re.match(r"snmpd\.@com2sec\[(\d+)\]\.(\w+)='([^']*)'", line_s) + if m: + idx, key, val = m.group(1), m.group(2), m.group(3) + com2sec.setdefault(idx, {})[key] = val + continue + m = re.match(r"snmpd\.@system\[0\]\.sys(\w+)='([^']*)'", line_s) + if m: + key, val = m.group(1).lower(), m.group(2) + if key == "contact": + contact = val + elif key == "location": + location = val + elif key == "name": + chassis_id = val + + for entry in com2sec.values(): + name = entry.get("community", entry.get("secname", "")) + if not name: + continue + # OpenWrt snmpd doesn't distinguish rw/ro per community via UCI by default + mode = "ro" + if entry.get("secname", "").lower() in ("private", "readwrite", "rw"): + mode = "rw" + community[name] = { + "acl": entry.get("source", "N/A"), + "mode": mode, + } + + return { + "chassis_id": chassis_id, + "community": community, + "contact": contact, + "location": location, + } + + # ------------------------------------------------------------------ + # Ping + # ------------------------------------------------------------------ + + def ping( + self, + destination: str, + source: str = "", + ttl: int = 255, + timeout: int = 2, + size: int = 56, + count: int = 5, + vrf: str = "", + source_interface: str = "", + ) -> Dict: + """Execute ping on the device and return statistics. + + Builds a ``ping`` command with standard BusyBox/iputils flags:: + + ping -c -W -s [-t ] [-I ] + + Returns ``{'success': {...}}`` or ``{'error': ''}``. + """ + cmd_parts = ["ping", "-c", str(count), "-W", str(timeout), "-s", str(size)] + if ttl != 255: + cmd_parts += ["-t", str(ttl)] + if source_interface: + cmd_parts += ["-I", source_interface] + elif source: + cmd_parts += ["-I", source] + cmd_parts.append(destination) + + output = self._send_command(" ".join(cmd_parts)) + + # Check for hard failure before parsing + if re.search(r"unknown host|bad address|Network unreachable|not reachable", output, re.I): + m = re.search(r"(unknown host.*|bad address.*|Network unreachable)", output, re.I) + return {"error": m.group(0) if m else output.strip()} + + return self._parse_ping_output(output, destination) + + @staticmethod + def _parse_ping_output(output: str, destination: str) -> Dict: + """Parse BusyBox/iputils ping output into NAPALM format.""" + # "2 packets transmitted, 2 packets received, 0% packet loss" + summary_m = re.search( + r"(\d+)\s+packets?\s+transmitted.*?(\d+)\s+(?:packets?\s+)?received.*?(\d+)%\s+packet\s+loss", + output, + re.S | re.I, + ) + if not summary_m: + return {"error": output.strip() or f"No response from {destination}"} + + sent = int(summary_m.group(1)) + received = int(summary_m.group(2)) + loss = sent - received + + # "round-trip min/avg/max = 6.987/7.055/7.123 ms" (BusyBox) + # "rtt min/avg/max/mdev = 6.987/7.055/7.123/0.094 ms" (iputils) + rtt_m = re.search( + r"(?:round-trip|rtt)\s+min/avg/max(?:/(?:mdev|stddev))?\s*=\s*([\d.]+)/([\d.]+)/([\d.]+)(?:/([\d.]+))?", + output, + re.I, + ) + rtt_min = rtt_avg = rtt_max = rtt_stddev = 0.0 + if rtt_m: + rtt_min = float(rtt_m.group(1)) + rtt_avg = float(rtt_m.group(2)) + rtt_max = float(rtt_m.group(3)) + rtt_stddev = float(rtt_m.group(4)) if rtt_m.group(4) else 0.0 + + # Individual probe results + results = [] + for m in re.finditer( + r"(\d+)\s+bytes\s+from\s+(\S+?):\s+(?:icmp_seq|seq)=\d+\s+.*?time=([\d.]+)\s*ms", + output, + re.I, + ): + ip = m.group(2).rstrip(":") + results.append({"ip_address": ip, "rtt": float(m.group(3))}) + + return { + "success": { + "probes_sent": sent, + "packet_loss": loss, + "rtt_min": rtt_min, + "rtt_avg": rtt_avg, + "rtt_max": rtt_max, + "rtt_stddev": rtt_stddev, + "results": results, + } + } + + # ------------------------------------------------------------------ + # IPv6 neighbours + # ------------------------------------------------------------------ + + def get_ipv6_neighbors_table(self) -> List[Dict]: + """Return the IPv6 neighbour table from ``ip -6 neigh show``. + + Example output:: + + 2001:db8::1 dev eth0 lladdr aa:bb:cc:dd:ee:ff REACHABLE + fe80::1 dev br-lan lladdr 11:22:33:44:55:66 STALE + """ + output = self._send_command("ip -6 neigh show") + table = [] + + for line in output.splitlines(): + line_s = line.strip() + if not line_s or "FAILED" in line_s or "INCOMPLETE" in line_s: + continue + + m = re.match( + r"^(\S+)\s+dev\s+(\S+)\s+lladdr\s+(\S+)\s+(\S+)", + line_s, + re.I, + ) + if not m: + continue + + ip_addr = m.group(1) + interface = m.group(2) + mac_raw = m.group(3) + state = m.group(4) + + try: + netaddr.IPAddress(ip_addr, version=6) + except (netaddr.AddrFormatError, ValueError): + continue + + try: + mac_addr = napalm_helpers.mac(mac_raw) + except Exception: + mac_addr = mac_raw + + table.append({ + "interface": interface, + "mac": mac_addr, + "ip": ip_addr, + "age": -1.0, + "state": state, + }) + + return table + + # ------------------------------------------------------------------ + # Routing + # ------------------------------------------------------------------ + + def get_route_to( + self, destination: str = "", protocol: str = "", longer: bool = False + ) -> Dict[str, List[Dict]]: + """Return routes to *destination* from the kernel routing table. + + Uses ``ip route show`` (optionally filtered by prefix/match) and + ``ip route get `` for the best-path lookup. + + Protocol filter is applied post-parse (kernel proto names: + ``kernel``, ``static``, ``dhcp``, ``bird``, ``zebra``, …). + + Example ``ip route show`` output:: + + default via 192.168.1.1 dev br-wan proto dhcp src 203.0.113.1 metric 100 + 192.168.1.0/24 dev br-lan proto kernel scope link src 192.168.1.1 + """ + if destination: + cmd = f"ip route show {'match ' if longer else ''}{destination}" + else: + cmd = "ip route show" + + output = self._send_command(cmd) + routes: Dict[str, List[Dict]] = {} + + for line in output.splitlines(): + line_s = line.strip() + if not line_s: + continue + + # Determine the prefix + # "default via ..." → prefix = "0.0.0.0/0" + # "192.168.1.0/24 dev ..." → prefix as-is + if line_s.startswith("default"): + prefix = "0.0.0.0/0" + rest = line_s[len("default"):].strip() + else: + parts = line_s.split() + prefix = parts[0] + rest = " ".join(parts[1:]) + + # Extract fields + next_hop = "" + outgoing_iface = "" + proto_raw = "kernel" + metric = 0 + + m = re.search(r"\bvia\s+(\S+)", rest) + if m: + next_hop = m.group(1) + + m = re.search(r"\bdev\s+(\S+)", rest) + if m: + outgoing_iface = m.group(1) + + m = re.search(r"\bproto\s+(\S+)", rest) + if m: + proto_raw = m.group(1) + + m = re.search(r"\bmetric\s+(\d+)", rest) + if m: + metric = int(m.group(1)) + + # Map proto to NAPALM-style name + proto_map = { + "kernel": "connected", + "static": "static", + "dhcp": "static", + "bird": "bgp", + "zebra": "ospf", + } + napalm_proto = proto_map.get(proto_raw.lower(), proto_raw) + + if protocol and napalm_proto.lower() != protocol.lower(): + continue + + entry = { + "protocol": napalm_proto, + "current_active": True, + "last_active": True, + "age": 0, + "next_hop": next_hop, + "outgoing_interface": outgoing_iface, + "selected_next_hop": True, + "preference": metric, + "inactive_reason": "", + "routing_table": "default", + "protocol_attributes": {}, + } + routes.setdefault(prefix, []).append(entry) + + return routes + + # ------------------------------------------------------------------ + # Traceroute + # ------------------------------------------------------------------ + + def traceroute( + self, + destination: str, + source: str = "", + ttl: int = 30, + timeout: int = 3, + vrf: str = "", + ) -> Dict: + """Execute traceroute on the device. + + Uses ``traceroute -m -w `` (BusyBox-compatible). + Falls back to ``traceroute6`` for IPv6 destinations. + + Returns ``{'success': {hop: {'probes': {probe: {rtt, ip_address, host_name}}}}}`` + or ``{'error': ''}``. + """ + # Detect IPv6 destination + try: + is_ipv6 = netaddr.IPAddress(destination).version == 6 + except (netaddr.AddrFormatError, ValueError): + is_ipv6 = ":" in destination + + cmd_base = "traceroute6" if is_ipv6 else "traceroute" + cmd_parts = [cmd_base, "-m", str(ttl), "-w", str(timeout)] + if source: + cmd_parts += ["-s", source] + cmd_parts.append(destination) + + output = self._send_command(" ".join(cmd_parts)) + + if re.search(r"unknown host|bad address|not reachable|cannot resolve", output, re.I): + m = re.search(r"(unknown host.*|bad address.*|cannot resolve.*)", output, re.I) + return {"error": m.group(0) if m else output.strip()} + + return self._parse_traceroute_output(output) + + @staticmethod + def _parse_traceroute_output(output: str) -> Dict: + """Parse BusyBox traceroute output into NAPALM format. + + Example lines:: + + 1 192.168.1.1 (192.168.1.1) 1.123 ms 1.456 ms 1.789 ms + 2 * * * + """ + hops: Dict[int, Dict] = {} + + for line in output.splitlines(): + line_s = line.strip() + # Hop line starts with an integer + m = re.match(r"^(\d+)\s+(.*)", line_s) + if not m: + continue + + hop_id = int(m.group(1)) + rest = m.group(2).strip() + + # All-star line: no response + if re.match(r"^\*[\s*]*$", rest): + hops[hop_id] = { + "probes": { + 1: {"rtt": -1.0, "ip_address": "*", "host_name": "*"}, + 2: {"rtt": -1.0, "ip_address": "*", "host_name": "*"}, + 3: {"rtt": -1.0, "ip_address": "*", "host_name": "*"}, + } + } + continue + + # Extract host/IP and RTT values + # Format: "hostname (ip) 1.1 ms 2.2 ms 3.3 ms" + # or: "ip 1.1 ms 2.2 ms 3.3 ms" + host_m = re.match(r"^(\S+)\s+\((\S+)\)", rest) + if host_m: + host_name = host_m.group(1) + ip_address = host_m.group(2) + else: + # IP only + ip_m = re.match(r"^(\d[\d.]+|[0-9a-f:]+)", rest) + if ip_m: + ip_address = ip_m.group(1) + host_name = ip_address + else: + continue + + rtt_values = [float(x) for x in re.findall(r"([\d.]+)\s+ms", rest)] + + probes: Dict[int, Dict] = {} + for i, rtt in enumerate(rtt_values[:3], start=1): + probes[i] = { + "rtt": rtt, + "ip_address": ip_address, + "host_name": host_name, + } + # Fill missing probes with star entries + for i in range(len(rtt_values) + 1, 4): + probes[i] = {"rtt": -1.0, "ip_address": "*", "host_name": "*"} + + if probes: + hops[hop_id] = {"probes": probes} + + if not hops: + return {"error": output.strip() or "No traceroute output received"} + + return {"success": hops} + + # ------------------------------------------------------------------ + # Network instances (namespaces / default VRF) + # ------------------------------------------------------------------ + + def get_network_instances(self, name: str = "") -> Dict[str, Dict]: + """Return network instances (Linux network namespaces + default). + + The ``default`` instance contains all interfaces not assigned to a + named namespace. Named namespaces are discovered via ``ip netns list``. + + Example:: + + { + 'default': { + 'name': 'default', + 'type': 'DEFAULT_INSTANCE', + 'state': {'route_distinguisher': None}, + 'interfaces': {'interface': {'br-lan': {}, 'eth0': {}}} + } + } + """ + netns_out = self._send_command("ip netns list") + iface_list = self._get_interface_list() + + instances: Dict[str, Dict] = {} + + # Named namespaces + netns_names: List[str] = [] + for line in netns_out.splitlines(): + line_s = line.strip() + if not line_s: + continue + # "myns (id: 3)" or just "myns" + ns_name = line_s.split()[0] + netns_names.append(ns_name) + + # Interfaces inside the namespace + ns_ifaces_out = self._send_command(f"ip netns exec {ns_name} ip link show") + ns_ifaces: Dict[str, Dict] = {} + for iline in ns_ifaces_out.splitlines(): + im = re.match(r"^\d+:\s+(\S+?)[@:]", iline) + if im and im.group(1) != "lo": + ns_ifaces[im.group(1)] = {} + + instances[ns_name] = { + "name": ns_name, + "type": "L3VRF", + "state": {"route_distinguisher": None}, + "interfaces": {"interface": ns_ifaces}, + } + + # Default instance: interfaces NOT in any named namespace + # (on most OpenWrt devices there are no named namespaces) + default_ifaces = {iface: {} for iface in iface_list} + instances["default"] = { + "name": "default", + "type": "DEFAULT_INSTANCE", + "state": {"route_distinguisher": None}, + "interfaces": {"interface": default_ifaces}, + } + + if name: + return {k: v for k, v in instances.items() if k == name} + + return instances diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..2e8704d --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,54 @@ +[build-system] +requires = ["setuptools>=64", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "napalm-openwrt" +version = "0.1.0" +description = "NAPALM driver for OpenWrt routers/access-points" +readme = "README.md" +license = { text = "Apache-2.0" } +requires-python = ">=3.8" +authors = [ + { name = "Christian Manivong" }, +] +classifiers = [ + "Topic :: Utilities", + "License :: OSI Approved :: Apache Software License", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.8", + "Programming Language :: Python :: 3.9", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Operating System :: POSIX :: Linux", + "Operating System :: MacOS", +] +dependencies = [ + "napalm>=4.0.0", + "napalm_device_types>=0.1.0", + "netmiko>=4.0.0", + "netaddr", +] + +[project.optional-dependencies] +dev = [ + "pytest", + "pytest-cov", + "black", + "ruff", +] + +[project.entry-points."napalm.drivers"] +openwrt = "napalm_openwrt:OpenWrtDriver" + +[project.urls] +Repository = "https://github.com/napalm-automation-community/napalm-openwrt" + +[tool.setuptools.packages.find] +where = ["."] +include = ["napalm_openwrt*"] + +[tool.ruff] +line-length = 100 +target-version = "py38" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..a6587eb --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +napalm>=4.0.0 +napalm_device_types>=0.1.0 +netmiko>=4.0.0 +netaddr diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/unit/__init__.py b/tests/unit/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/unit/test_driver.py b/tests/unit/test_driver.py new file mode 100644 index 0000000..467910f --- /dev/null +++ b/tests/unit/test_driver.py @@ -0,0 +1,835 @@ +"""Unit tests for OpenWrtDriver — no real device required.""" + +import pytest +from unittest.mock import MagicMock, patch + +from napalm_openwrt.openwrt import OpenWrtDriver + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture +def driver(): + """Return a driver instance with a mocked Netmiko connection.""" + with patch("napalm_openwrt.openwrt.ConnectHandler"): + drv = OpenWrtDriver( + hostname="192.168.1.1", + username="root", + password="", + ) + drv.device = MagicMock() + yield drv + + +# --------------------------------------------------------------------------- +# Sample command output (as would be returned by the device) +# --------------------------------------------------------------------------- + +OPENWRT_RELEASE = """\ +DISTRIB_ID="OpenWrt" +DISTRIB_RELEASE="23.05.3" +DISTRIB_REVISION="r23809-234f1a2efa" +DISTRIB_TARGET="ath79/generic" +DISTRIB_ARCH="mips_24kc" +DISTRIB_CODENAME="Restoring Earth" +DISTRIB_TAINTS="" +""" + +SYSINFO_MODEL = "TP-Link TL-WR1043N/ND v5" + +UPTIME = "352467.12 345823.44" + +IP_LINK_SHOW = """\ +1: lo: mtu 65536 qdisc noqueue state UNKNOWN mode DEFAULT group default qlen 1000 + link/loopback 00:00:00:00:00:00 brd 00:00:00:00:00:00 +2: eth0: mtu 1500 qdisc fq_codel state UP mode DEFAULT group default qlen 1000 + link/ether b0:95:75:aa:bb:cc brd ff:ff:ff:ff:ff:ff +3: eth1: mtu 1500 qdisc noop state DOWN mode DEFAULT group default qlen 1000 + link/ether b0:95:75:aa:bb:dd brd ff:ff:ff:ff:ff:ff +4: br-lan: mtu 1500 qdisc noqueue state UP mode DEFAULT group default qlen 1000 + link/ether b0:95:75:aa:bb:cc brd ff:ff:ff:ff:ff:ff +""" + +IP_ADDR_SHOW = """\ +1: lo: mtu 65536 qdisc noqueue state UNKNOWN group default qlen 1000 + link/loopback 00:00:00:00:00:00 brd 00:00:00:00:00:00 + inet 127.0.0.1/8 scope host lo + valid_lft forever preferred_lft forever +2: eth0: mtu 1500 qdisc fq_codel state UP group default qlen 1000 + link/ether b0:95:75:aa:bb:cc brd ff:ff:ff:ff:ff:ff +4: br-lan: mtu 1500 qdisc noqueue state UP group default qlen 1000 + link/ether b0:95:75:aa:bb:cc brd ff:ff:ff:ff:ff:ff + inet 192.168.1.1/24 brd 192.168.1.255 scope global br-lan + valid_lft forever preferred_lft forever + inet6 fd00::1/64 scope global + valid_lft forever preferred_lft forever +""" + +IP_NEIGH_SHOW = """\ +192.168.1.100 dev br-lan lladdr aa:bb:cc:dd:ee:ff REACHABLE +192.168.1.101 dev br-lan lladdr 11:22:33:44:55:66 STALE +192.168.1.102 dev br-lan FAILED +""" + +BRIDGE_FDB = """\ +aa:bb:cc:dd:ee:ff dev br-lan master br-lan permanent +11:22:33:44:55:66 dev eth0 vlan 1 master br-lan +33:33:00:00:00:01 dev br-lan self permanent +""" + +UCI_EXPORT = """\ +package system + +config system +\toption hostname 'OpenWrt' +\toption timezone 'UTC' + +package network + +config interface 'loopback' +\toption device 'lo' +\toption proto 'static' +\toption ipaddr '127.0.0.1' +\toption netmask '255.0.0.0' + +config interface 'lan' +\toption device 'br-lan' +\toption proto 'static' +\toption ipaddr '192.168.1.1' +\toption netmask '255.255.255.0' +""" + +LLDPCTL_KV = """\ +lldp.eth0.via=LLDP +lldp.eth0.rid=1 +lldp.eth0.age=0 day, 01:23:45 +lldp.eth0.chassis.mac=00:aa:bb:cc:dd:ee +lldp.eth0.chassis.name=core-router +lldp.eth0.chassis.descr=RouterOS 7.x +lldp.eth0.chassis.mgmt-ip=10.0.0.1 +lldp.eth0.chassis.cap.available=Router, Bridge +lldp.eth0.chassis.cap.enabled=Router +lldp.eth0.port.ifname=ether1 +lldp.eth0.port.descr=uplink +""" + +PROC_NET_DEV = """\ +Inter-| Receive | Transmit + face |bytes packets errs drop fifo frame compressed multicast|bytes packets errs drop fifo colls carrier compressed + lo: 1234 12 0 0 0 0 0 0 1234 12 0 0 0 0 0 0 + eth0: 9876543 12345 0 0 0 0 0 100 1234567 9876 0 0 0 0 0 0 +br-lan: 8765432 11234 0 0 0 0 0 50 1123456 8765 0 0 0 0 0 0 +""" + +HOSTNAME = "OpenWrt" + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + +class TestGetFacts: + def _make_send(self): + """Return a _send_command mock that handles both str and list args.""" + def _send(cmd, **kw): + key = cmd[0] if isinstance(cmd, list) else cmd + if "openwrt_release" in key: + return OPENWRT_RELEASE + if "sysinfo/model" in key: + return SYSINFO_MODEL + if "uptime" in key: + return UPTIME + if "hostname" in key.lower() or "system.@system" in key: + return HOSTNAME + if "ip link" in key: + return IP_LINK_SHOW + return "" + return _send + + def test_returns_required_keys(self, driver): + driver._send_command = self._make_send() + + facts = driver.get_facts() + assert set(facts.keys()) == { + "vendor", "model", "hostname", "fqdn", "os_version", + "serial_number", "uptime", "interface_list", + } + + def test_vendor(self, driver): + driver._send_command = self._make_send() + assert driver.get_facts()["vendor"] == "OpenWrt" + + def test_os_version(self, driver): + driver._send_command = self._make_send() + assert driver.get_facts()["os_version"] == "23.05.3" + + def test_uptime_parsed(self, driver): + driver._send_command = self._make_send() + facts = driver.get_facts() + assert facts["uptime"] == pytest.approx(352467.12) + + +class TestParseOpenwrtRelease: + def test_parses_release(self): + result = OpenWrtDriver._parse_openwrt_release(OPENWRT_RELEASE) + assert result["DISTRIB_RELEASE"] == "23.05.3" + assert result["DISTRIB_ID"] == "OpenWrt" + assert result["DISTRIB_TARGET"] == "ath79/generic" + + +class TestGetInterfaces: + def test_returns_dict(self, driver): + driver._send_command = lambda cmd, **kw: IP_LINK_SHOW + result = driver.get_interfaces() + assert isinstance(result, dict) + + def test_eth0_is_up(self, driver): + driver._send_command = lambda cmd, **kw: IP_LINK_SHOW + result = driver.get_interfaces() + assert "eth0" in result + assert result["eth0"]["is_up"] is True + assert result["eth0"]["is_enabled"] is True + + def test_eth1_is_down(self, driver): + driver._send_command = lambda cmd, **kw: IP_LINK_SHOW + result = driver.get_interfaces() + assert "eth1" in result + assert result["eth1"]["is_up"] is False + + def test_mac_address_populated(self, driver): + driver._send_command = lambda cmd, **kw: IP_LINK_SHOW + result = driver.get_interfaces() + assert result["eth0"]["mac_address"] != "" + + def test_required_keys(self, driver): + driver._send_command = lambda cmd, **kw: IP_LINK_SHOW + result = driver.get_interfaces() + for iface_data in result.values(): + assert set(iface_data.keys()) >= { + "is_up", "is_enabled", "description", + "last_flapped", "speed", "mtu", "mac_address", + } + + +class TestGetInterfacesIP: + def test_br_lan_ipv4(self, driver): + driver._send_command = lambda cmd, **kw: IP_ADDR_SHOW + result = driver.get_interfaces_ip() + assert "br-lan" in result + assert "192.168.1.1" in result["br-lan"]["ipv4"] + assert result["br-lan"]["ipv4"]["192.168.1.1"]["prefix_length"] == 24 + + def test_br_lan_ipv6(self, driver): + driver._send_command = lambda cmd, **kw: IP_ADDR_SHOW + result = driver.get_interfaces_ip() + assert "ipv6" in result["br-lan"] + assert "fd00::1" in result["br-lan"]["ipv6"] + + +class TestGetArpTable: + def test_returns_list(self, driver): + driver._send_command = lambda cmd, **kw: IP_NEIGH_SHOW + result = driver.get_arp_table() + assert isinstance(result, list) + + def test_entries_count(self, driver): + driver._send_command = lambda cmd, **kw: IP_NEIGH_SHOW + result = driver.get_arp_table() + # FAILED entry should be skipped + assert len(result) == 2 + + def test_entry_keys(self, driver): + driver._send_command = lambda cmd, **kw: IP_NEIGH_SHOW + result = driver.get_arp_table() + for entry in result: + assert set(entry.keys()) >= {"interface", "mac", "ip", "age"} + + def test_ip_value(self, driver): + driver._send_command = lambda cmd, **kw: IP_NEIGH_SHOW + result = driver.get_arp_table() + ips = {e["ip"] for e in result} + assert "192.168.1.100" in ips + assert "192.168.1.101" in ips + + +class TestGetMacAddressTable: + def test_returns_list(self, driver): + driver._send_command = lambda cmd, **kw: BRIDGE_FDB + result = driver.get_mac_address_table() + assert isinstance(result, list) + + def test_multicast_skipped(self, driver): + driver._send_command = lambda cmd, **kw: BRIDGE_FDB + result = driver.get_mac_address_table() + macs = [e["mac"] for e in result] + assert not any("33:33" in m for m in macs) + + def test_vlan_parsed(self, driver): + driver._send_command = lambda cmd, **kw: BRIDGE_FDB + result = driver.get_mac_address_table() + vlan1_entries = [e for e in result if e["vlan"] == 1] + assert len(vlan1_entries) >= 1 + + +class TestGetConfig: + def test_returns_uci_export(self, driver): + driver._send_command = lambda cmd, **kw: UCI_EXPORT + result = driver.get_config() + assert "running" in result + assert "startup" in result + assert "candidate" in result + assert "package" in result["running"] + + def test_candidate_always_empty(self, driver): + driver._send_command = lambda cmd, **kw: UCI_EXPORT + result = driver.get_config() + assert result["candidate"] == "" + + +class TestGetLldpNeighbors: + def test_returns_dict(self, driver): + driver._send_command = lambda cmd, **kw: LLDPCTL_KV + result = driver.get_lldp_neighbors() + assert isinstance(result, dict) + + def test_eth0_neighbor(self, driver): + driver._send_command = lambda cmd, **kw: LLDPCTL_KV + result = driver.get_lldp_neighbors() + assert "eth0" in result + assert result["eth0"][0]["hostname"] == "core-router" + assert result["eth0"][0]["port"] == "ether1" + + +BRIDGE_VLAN_SHOW = """\ +port vlan-id +eth0 1 PVID Egress Untagged + 10 + 20 +br-lan 1 PVID Egress Untagged + 10 + 20 +eth1 20 PVID Egress Untagged +""" + +UCI_NETWORK_VLANS = """\ +network.@bridge-vlan[0]=bridge-vlan +network.@bridge-vlan[0].device='br-lan' +network.@bridge-vlan[0].vlan='10' +network.@bridge-vlan[0].name='management' +network.@bridge-vlan[1]=bridge-vlan +network.@bridge-vlan[1].device='br-lan' +network.@bridge-vlan[1].vlan='20' +network.@bridge-vlan[1].name='iot' +""" + + +class TestGetVlans: + def _send(self, cmd, **kw): + if "bridge vlan" in cmd: + return BRIDGE_VLAN_SHOW + if "uci show network" in cmd: + return UCI_NETWORK_VLANS + return "" + + def test_returns_dict(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + assert isinstance(result, dict) + + def test_vlan_ids_present(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + assert "1" in result + assert "10" in result + assert "20" in result + + def test_required_keys(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + for vlan_data in result.values(): + assert "name" in vlan_data + assert "interfaces" in vlan_data + + def test_interfaces_for_vlan10(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + assert set(result["10"]["interfaces"]) == {"eth0", "br-lan"} + + def test_interfaces_for_vlan20(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + assert set(result["20"]["interfaces"]) == {"eth0", "br-lan", "eth1"} + + def test_uci_names_applied(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + assert result["10"]["name"] == "management" + assert result["20"]["name"] == "iot" + + def test_vlan_without_uci_name_is_empty_string(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + assert result["1"]["name"] == "" + + def test_no_duplicate_interfaces(self, driver): + driver._send_command = self._send + result = driver.get_vlans() + for vlan_data in result.values(): + assert len(vlan_data["interfaces"]) == len(set(vlan_data["interfaces"])) + + def test_empty_bridge_output(self, driver): + driver._send_command = lambda cmd, **kw: "" + result = driver.get_vlans() + assert result == {} + + +class TestConfigManagement: + def test_load_merge_candidate(self, driver): + driver.load_merge_candidate(config="uci set system.@system[0].hostname='MyRouter'") + assert driver._candidate_config is not None + assert driver._candidate_mode == "merge" + + def test_load_replace_candidate(self, driver): + driver.load_replace_candidate(config=UCI_EXPORT) + assert driver._candidate_config is not None + assert driver._candidate_mode == "replace" + + def test_discard_config(self, driver): + driver.load_merge_candidate(config="uci set system.@system[0].hostname='test'") + driver.discard_config() + assert driver._candidate_config is None + assert not driver.has_pending_commit() + + def test_compare_merge_candidate(self, driver): + driver._send_command = lambda cmd, **kw: UCI_EXPORT + driver.load_merge_candidate(config="uci set system.@system[0].hostname='test'") + diff = driver.compare_config() + assert diff.startswith("+") + + def test_compare_no_candidate(self, driver): + assert driver.compare_config() == "" + + def test_has_pending_commit_false_initially(self, driver): + assert not driver.has_pending_commit() + + def test_has_pending_commit_true_after_load(self, driver): + driver.load_merge_candidate(config="uci set system.@system[0].hostname='test'") + assert driver.has_pending_commit() + + +# --------------------------------------------------------------------------- +# Sample data for new methods +# --------------------------------------------------------------------------- + +PASSWD = """\ +root:$1$xyz:0:0:root:/root:/bin/ash +daemon:*:1:1:daemon:/var:/bin/false +nobody:*:65534:65534:nobody:/var:/bin/false +alice:$6$abc:1001:1001:Alice:/home/alice:/bin/ash +""" + +ROOT_AUTHORIZED_KEYS = "ssh-rsa AAAAB3NzaC1yc2EAAAADAQABroot@host" +ALICE_AUTHORIZED_KEYS = "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5alice@host" + +UCI_SYSTEM_NTP = """\ +system.@system[0]=system +system.@system[0].hostname='OpenWrt' +system.@system[0].timezone='UTC' +system.ntp=timeserver +system.ntp.server='0.openwrt.pool.ntp.org 1.openwrt.pool.ntp.org 2.openwrt.pool.ntp.org' +system.ntp.enabled='1' +system.ntp.enable_server='0' +""" + +NTPQ_OUTPUT = """\ + remote refid st t when poll reach delay offset jitter +============================================================================== +*188.114.101.4 188.114.100.1 4 u 107 256 377 164.228 -13.866 2.695 ++37.187.56.220 192.53.103.108 2 u 22 64 377 30.112 5.123 1.100 +""" + +UCI_SNMPD = """\ +snmpd.@agent[0]=agent +snmpd.@agent[0].agentaddress='UDP:161' +snmpd.@com2sec[0]=com2sec +snmpd.@com2sec[0].secname='public' +snmpd.@com2sec[0].source='default' +snmpd.@com2sec[0].community='public' +snmpd.@com2sec[1]=com2sec +snmpd.@com2sec[1].secname='private' +snmpd.@com2sec[1].source='10.0.0.0/8' +snmpd.@com2sec[1].community='private' +snmpd.@system[0]=system +snmpd.@system[0].sysContact='admin@example.com' +snmpd.@system[0].sysLocation='Server Room' +snmpd.@system[0].sysName='MyRouter' +""" + +PING_SUCCESS = """\ +PING 8.8.8.8 (8.8.8.8): 56 data bytes +64 bytes from 8.8.8.8: seq=0 ttl=120 time=7.123 ms +64 bytes from 8.8.8.8: seq=1 ttl=120 time=6.987 ms +64 bytes from 8.8.8.8: seq=2 ttl=120 time=7.234 ms +--- 8.8.8.8 ping statistics --- +3 packets transmitted, 3 packets received, 0% packet loss +round-trip min/avg/max = 6.987/7.115/7.234 ms +""" + +PING_LOSS = """\ +PING 10.0.0.99 (10.0.0.99): 56 data bytes +--- 10.0.0.99 ping statistics --- +3 packets transmitted, 0 packets received, 100% packet loss +""" + +PING_ERROR = "ping: bad address 'invalid.host'" + +IPV6_NEIGH = """\ +2001:db8::1 dev eth0 lladdr aa:bb:cc:dd:ee:ff REACHABLE +fe80::1 dev br-lan lladdr 11:22:33:44:55:66 STALE +2001:db8::2 dev eth0 FAILED +""" + +IP_ROUTE_SHOW = """\ +default via 192.168.1.1 dev br-wan proto dhcp src 203.0.113.1 metric 100 +192.168.1.0/24 dev br-lan proto kernel scope link src 192.168.1.1 +10.0.0.0/8 via 192.168.1.254 dev br-wan proto static metric 50 +""" + +TRACEROUTE_OUTPUT = """\ +traceroute to 8.8.8.8 (8.8.8.8), 30 hops max, 38 byte packets + 1 192.168.1.1 (192.168.1.1) 1.123 ms 1.456 ms 1.789 ms + 2 10.0.0.1 (10.0.0.1) 5.123 ms 5.456 ms 5.789 ms + 3 * * * + 4 8.8.8.8 (8.8.8.8) 7.001 ms 6.999 ms 7.100 ms +""" + +TRACEROUTE_ERROR = "traceroute: unknown host invalid.host" + +IP_NETNS_LIST = """\ +vpn (id: 1) +mgmt (id: 2) +""" + + +# --------------------------------------------------------------------------- +# Tests for new methods +# --------------------------------------------------------------------------- + +class TestCli: + def test_returns_dict_keyed_by_command(self, driver): + driver._send_command = lambda cmd, **kw: f"output of {cmd}" + result = driver.cli(["uname -a", "uptime"]) + assert set(result.keys()) == {"uname -a", "uptime"} + + def test_output_content(self, driver): + driver._send_command = lambda cmd, **kw: "Linux OpenWrt" + result = driver.cli(["uname -a"]) + assert result["uname -a"] == "Linux OpenWrt" + + def test_empty_command_list(self, driver): + result = driver.cli([]) + assert result == {} + + +class TestGetUsers: + def _send(self, cmd, **kw): + key = cmd[0] if isinstance(cmd, list) else cmd + if "/etc/passwd" in key: + return PASSWD + if "root/.ssh/authorized_keys" in key or "dropbear/authorized_keys" in key: + return ROOT_AUTHORIZED_KEYS + if "alice" in key: + return ALICE_AUTHORIZED_KEYS + return "" + + def test_returns_dict(self, driver): + driver._send_command = self._send + assert isinstance(driver.get_users(), dict) + + def test_root_level_15(self, driver): + driver._send_command = self._send + users = driver.get_users() + assert "root" in users + assert users["root"]["level"] == 15 + + def test_regular_user_level_1(self, driver): + driver._send_command = self._send + users = driver.get_users() + assert "alice" in users + assert users["alice"]["level"] == 1 + + def test_system_accounts_excluded(self, driver): + driver._send_command = self._send + users = driver.get_users() + assert "daemon" not in users + assert "nobody" not in users + + def test_root_has_ssh_key(self, driver): + driver._send_command = self._send + users = driver.get_users() + assert len(users["root"]["sshkeys"]) >= 1 + assert users["root"]["sshkeys"][0].startswith("ssh-rsa") + + def test_required_keys(self, driver): + driver._send_command = self._send + users = driver.get_users() + for data in users.values(): + assert "level" in data + assert "password" in data + assert "sshkeys" in data + + +class TestGetNtpServers: + def test_returns_dict(self, driver): + driver._send_command = lambda cmd, **kw: UCI_SYSTEM_NTP + result = driver.get_ntp_servers() + assert isinstance(result, dict) + + def test_servers_present(self, driver): + driver._send_command = lambda cmd, **kw: UCI_SYSTEM_NTP + result = driver.get_ntp_servers() + assert "0.openwrt.pool.ntp.org" in result + assert "1.openwrt.pool.ntp.org" in result + assert "2.openwrt.pool.ntp.org" in result + + def test_empty_when_no_ntp(self, driver): + driver._send_command = lambda cmd, **kw: "" + assert driver.get_ntp_servers() == {} + + +class TestGetNtpStats: + def test_returns_list(self, driver): + driver._send_command = lambda cmd, **kw: NTPQ_OUTPUT + result = driver.get_ntp_stats() + assert isinstance(result, list) + + def test_synchronized_entry(self, driver): + driver._send_command = lambda cmd, **kw: NTPQ_OUTPUT + result = driver.get_ntp_stats() + synced = [e for e in result if e["synchronized"]] + assert len(synced) == 1 + assert synced[0]["remote"] == "188.114.101.4" + + def test_required_keys(self, driver): + driver._send_command = lambda cmd, **kw: NTPQ_OUTPUT + result = driver.get_ntp_stats() + for entry in result: + assert set(entry.keys()) >= { + "remote", "referenceid", "synchronized", "stratum", + "type", "when", "hostpoll", "reachability", "delay", "offset", "jitter", + } + + def test_empty_on_no_tool(self, driver): + driver._send_command = lambda cmd, **kw: "sh: ntpq: not found" + result = driver.get_ntp_stats() + assert result == [] + + +class TestGetSnmpInformation: + def test_returns_dict(self, driver): + driver._send_command = lambda cmd, **kw: UCI_SNMPD + result = driver.get_snmp_information() + assert isinstance(result, dict) + + def test_required_keys(self, driver): + driver._send_command = lambda cmd, **kw: UCI_SNMPD + result = driver.get_snmp_information() + assert set(result.keys()) >= {"chassis_id", "community", "contact", "location"} + + def test_contact_and_location(self, driver): + driver._send_command = lambda cmd, **kw: UCI_SNMPD + result = driver.get_snmp_information() + assert result["contact"] == "admin@example.com" + assert result["location"] == "Server Room" + + def test_community_entries(self, driver): + driver._send_command = lambda cmd, **kw: UCI_SNMPD + result = driver.get_snmp_information() + assert "public" in result["community"] + assert "private" in result["community"] + + def test_community_mode(self, driver): + driver._send_command = lambda cmd, **kw: UCI_SNMPD + result = driver.get_snmp_information() + assert result["community"]["public"]["mode"] == "ro" + assert result["community"]["private"]["mode"] == "rw" + + +class TestPing: + def test_success_result(self, driver): + driver._send_command = lambda cmd, **kw: PING_SUCCESS + result = driver.ping("8.8.8.8") + assert "success" in result + assert result["success"]["probes_sent"] == 3 + assert result["success"]["packet_loss"] == 0 + + def test_rtt_values(self, driver): + driver._send_command = lambda cmd, **kw: PING_SUCCESS + result = driver.ping("8.8.8.8") + s = result["success"] + assert s["rtt_min"] == 6.987 + assert s["rtt_max"] == 7.234 + assert s["rtt_avg"] == 7.115 + + def test_probe_results(self, driver): + driver._send_command = lambda cmd, **kw: PING_SUCCESS + result = driver.ping("8.8.8.8") + assert len(result["success"]["results"]) == 3 + assert result["success"]["results"][0]["ip_address"] == "8.8.8.8" + + def test_100_percent_loss(self, driver): + driver._send_command = lambda cmd, **kw: PING_LOSS + result = driver.ping("10.0.0.99") + assert "success" in result + assert result["success"]["packet_loss"] == 3 + + def test_error_on_bad_host(self, driver): + driver._send_command = lambda cmd, **kw: PING_ERROR + result = driver.ping("invalid.host") + assert "error" in result + + +class TestGetIpv6NeighborsTable: + def test_returns_list(self, driver): + driver._send_command = lambda cmd, **kw: IPV6_NEIGH + assert isinstance(driver.get_ipv6_neighbors_table(), list) + + def test_failed_entries_excluded(self, driver): + driver._send_command = lambda cmd, **kw: IPV6_NEIGH + result = driver.get_ipv6_neighbors_table() + assert len(result) == 2 + + def test_required_keys(self, driver): + driver._send_command = lambda cmd, **kw: IPV6_NEIGH + for entry in driver.get_ipv6_neighbors_table(): + assert set(entry.keys()) >= {"interface", "mac", "ip", "age", "state"} + + def test_state_values(self, driver): + driver._send_command = lambda cmd, **kw: IPV6_NEIGH + states = {e["state"] for e in driver.get_ipv6_neighbors_table()} + assert "REACHABLE" in states + assert "STALE" in states + + +class TestGetRouteTo: + def test_returns_dict(self, driver): + driver._send_command = lambda cmd, **kw: IP_ROUTE_SHOW + assert isinstance(driver.get_route_to(), dict) + + def test_default_route_present(self, driver): + driver._send_command = lambda cmd, **kw: IP_ROUTE_SHOW + result = driver.get_route_to() + assert "0.0.0.0/0" in result + + def test_next_hop(self, driver): + driver._send_command = lambda cmd, **kw: IP_ROUTE_SHOW + result = driver.get_route_to() + default = result["0.0.0.0/0"][0] + assert default["next_hop"] == "192.168.1.1" + assert default["outgoing_interface"] == "br-wan" + + def test_connected_route(self, driver): + driver._send_command = lambda cmd, **kw: IP_ROUTE_SHOW + result = driver.get_route_to() + assert "192.168.1.0/24" in result + assert result["192.168.1.0/24"][0]["protocol"] == "connected" + + def test_protocol_filter(self, driver): + driver._send_command = lambda cmd, **kw: IP_ROUTE_SHOW + result = driver.get_route_to(protocol="static") + assert "10.0.0.0/8" in result + assert "192.168.1.0/24" not in result + + def test_required_keys(self, driver): + driver._send_command = lambda cmd, **kw: IP_ROUTE_SHOW + for prefix, routes in driver.get_route_to().items(): + for route in routes: + assert set(route.keys()) >= { + "protocol", "current_active", "next_hop", + "outgoing_interface", "preference", "routing_table", + } + + +class TestTraceroute: + def test_success_result(self, driver): + driver._send_command = lambda cmd, **kw: TRACEROUTE_OUTPUT + result = driver.traceroute("8.8.8.8") + assert "success" in result + + def test_hop_count(self, driver): + driver._send_command = lambda cmd, **kw: TRACEROUTE_OUTPUT + result = driver.traceroute("8.8.8.8") + assert len(result["success"]) == 4 + + def test_hop_1_rtt(self, driver): + driver._send_command = lambda cmd, **kw: TRACEROUTE_OUTPUT + result = driver.traceroute("8.8.8.8") + hop1 = result["success"][1]["probes"][1] + assert hop1["rtt"] == 1.123 + assert hop1["ip_address"] == "192.168.1.1" + + def test_star_hop(self, driver): + driver._send_command = lambda cmd, **kw: TRACEROUTE_OUTPUT + result = driver.traceroute("8.8.8.8") + hop3 = result["success"][3]["probes"][1] + assert hop3["ip_address"] == "*" + + def test_error_on_unknown_host(self, driver): + driver._send_command = lambda cmd, **kw: TRACEROUTE_ERROR + result = driver.traceroute("invalid.host") + assert "error" in result + + +class TestGetNetworkInstances: + def test_default_instance_always_present(self, driver): + driver._send_command = lambda cmd, **kw: ( + IP_LINK_SHOW if "ip link" in cmd else "" + ) + result = driver.get_network_instances() + assert "default" in result + + def test_default_instance_type(self, driver): + driver._send_command = lambda cmd, **kw: ( + IP_LINK_SHOW if "ip link" in cmd else "" + ) + result = driver.get_network_instances() + assert result["default"]["type"] == "DEFAULT_INSTANCE" + + def test_default_interfaces_populated(self, driver): + driver._send_command = lambda cmd, **kw: ( + IP_LINK_SHOW if "ip link" in cmd else "" + ) + result = driver.get_network_instances() + ifaces = result["default"]["interfaces"]["interface"] + assert "eth0" in ifaces + + def test_named_namespaces(self, driver): + def _send(cmd, **kw): + if "netns list" in cmd: + return IP_NETNS_LIST + if "netns exec" in cmd: + return "" # empty namespace + if "ip link" in cmd: + return IP_LINK_SHOW + return "" + driver._send_command = _send + result = driver.get_network_instances() + assert "vpn" in result + assert "mgmt" in result + assert result["vpn"]["type"] == "L3VRF" + + def test_name_filter(self, driver): + driver._send_command = lambda cmd, **kw: ( + IP_LINK_SHOW if "ip link" in cmd else "" + ) + result = driver.get_network_instances(name="default") + assert list(result.keys()) == ["default"] + + def test_required_keys(self, driver): + driver._send_command = lambda cmd, **kw: ( + IP_LINK_SHOW if "ip link" in cmd else "" + ) + for inst in driver.get_network_instances().values(): + assert set(inst.keys()) >= {"name", "type", "state", "interfaces"}