From ae27cd5469159f855f848170584d35f27fb764d0 Mon Sep 17 00:00:00 2001 From: Christian Manivong Date: Fri, 29 May 2026 09:22:10 +0200 Subject: [PATCH] initial commit --- .github/workflows/ci.yml | 44 ++ .gitignore | 42 ++ README.md | 193 +++++ examples/test_driver.py | 66 ++ napalm_opnsense/__init__.py | 3 + napalm_opnsense/opnsense.py | 1390 +++++++++++++++++++++++++++++++++++ pyproject.toml | 55 ++ tests/__init__.py | 0 tests/unit/__init__.py | 0 tests/unit/test_driver.py | 1245 +++++++++++++++++++++++++++++++ 10 files changed, 3038 insertions(+) create mode 100644 .github/workflows/ci.yml create mode 100644 .gitignore create mode 100644 README.md create mode 100644 examples/test_driver.py create mode 100644 napalm_opnsense/__init__.py create mode 100644 napalm_opnsense/opnsense.py create mode 100644 pyproject.toml create mode 100644 tests/__init__.py create mode 100644 tests/unit/__init__.py create mode 100644 tests/unit/test_driver.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..2207b4c --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,44 @@ +name: CI + +on: + push: + branches: ["**"] + pull_request: + branches: ["**"] + +jobs: + test: + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + python-version: ["3.9", "3.10", "3.11", "3.12"] + steps: + - name: Checkout + uses: actions/checkout@v4 + + - name: Setup Python + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + cache: pip + + - name: Install package with dev extras + run: | + python -m pip install --upgrade pip + python -m pip install -e ".[dev]" + + - name: Run unit tests + run: | + python -m pytest -q --tb=short + + - name: Build wheel and sdist + run: | + python -m pip install build + python -m build + + - name: Upload dist artifacts + uses: actions/upload-artifact@v4 + with: + name: dist-${{ matrix.python-version }} + path: dist/* \ No newline at end of file diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..80bbc57 --- /dev/null +++ b/.gitignore @@ -0,0 +1,42 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# Virtual environments +.venv/ +venv/ +ENV/ +env/ + +# Distribution / packaging +.Python +build/ +dist/ +*.egg-info/ +*.egg +*.whl +pip-wheel-metadata/ + +# Testing / coverage +.pytest_cache/ +.coverage +.coverage.* +coverage.xml +htmlcov/ + +# Type checking / lint +.mypy_cache/ +.ruff_cache/ +.tox/ + +# IDE / editor +.vscode/ +.idea/ + +# OS / misc +.DS_Store +Thumbs.db + +# Logs +*.log diff --git a/README.md b/README.md new file mode 100644 index 0000000..ce8683d --- /dev/null +++ b/README.md @@ -0,0 +1,193 @@ +# napalm-opnsense + +NAPALM community driver for **OPNsense** firewalls (read-only, via REST API). + +## Tested devices + +| Model | OPNsense Version | Tested | +|---|---|---| +| OPNsense (virtual/bare metal) | 24.7 | ✅ | + +> Additional OPNsense versions should work — contributions welcome. + +## Requirements + +| Dependency | Minimum version | +|---|---| +| Python | 3.9 | +| NAPALM | 4.0 | +| requests | 2.28 | + +## Installation + +```bash +pip install napalm-opnsense +``` + +Or from source: + +```bash +git clone https://github.com/napalm-automation-community/napalm-opnsense +cd napalm-opnsense +pip install -e . +``` + +## Quick start + +```python +from napalm import get_network_driver + +driver = get_network_driver("opnsense") +with driver( + "192.168.1.1", + "my_api_key", + "my_api_secret", + optional_args={"verify": False}, +) as device: + facts = device.get_facts() + print(facts) +``` + +## Authentication + +OPNsense uses API key/secret pairs instead of username/password. Generate +a key pair in the OPNsense GUI under **System → Access → Users → edit user → +API keys**. + +Pass the credentials via `optional_args`: + +```python +optional_args={ + "api_key": "your-api-key", + "api_secret": "your-api-secret", + "verify": False, # set to a CA bundle path or True in production +} +``` + +Alternatively, pass the key/secret as the positional `username`/`password` +arguments. + +## Implemented getters + +| Getter | Status | OPNsense API Endpoint | +|---|---|---| +| `get_facts` | ✅ | `GET /api/core/system/status` | +| `get_interfaces` | ✅ | `GET /api/interfaces/overview/export` | +| `get_interfaces_ip` | ✅ | `GET /api/interfaces/addresses/export` | +| `get_interfaces_counters` | ✅ | `GET /api/diagnostics/interface/get_interface_statistics` | +| `get_arp_table` | ✅ | `GET /api/diagnostics/interface/get_arp` | +| `get_ipv6_neighbors_table` | ✅ | `GET /api/diagnostics/interface/get_ndp` | +| `get_route_to` | ✅ | `GET /api/diagnostics/interface/get_routes` | +| `get_environment` | ✅ | `GET /api/diagnostics/system/system_resources` + `system_temperature` | +| `get_lldp_neighbors` | ✅ ¹ | `GET /api/lldpd/service/neighbor` | +| `get_lldp_neighbors_detail` | ✅ ¹ | `GET /api/lldpd/service/neighbor` | +| `get_ntp_servers` | ✅ | `GET /api/ntpd/service/status` | +| `get_config` | ✅ | `GET /api/core/backup/download/this` (XML) | +| `is_alive` | ✅ | TCP socket check | +| `get_bgp_neighbors` | ✅ ² | `GET /api/quagga/bgp/get` + `GET /api/quagga/diagnostics/bgpneighbors` | +| `get_vlans` | ✅ | `GET /api/interfaces/vlan_settings/search_item` | +| `get_mac_address_table` | ❌ | Not applicable (firewall, no L2 switching) | + +> ¹ Requires the `os-lldpd` plugin. Returns empty dict if the plugin is not installed. +> ² Requires the `os-frr` (FRR/Quagga) plugin. Returns empty dict if the plugin is not installed or FRR is not running. + +## Config management + +OPNsense does not expose a single generic "push config" endpoint. Config is +managed per-module via separate API controllers. This driver implements config +management for **static routes** via `/api/routes/routes/`. + +> **Why routes?** Routes are the most common network-automation target on a +> firewall, and the OPNsense routes API provides full CRUD operations. + +### Supported methods + +| Method | OPNsense API | +|---|---| +| `load_merge_candidate(config=...)` | stages routes in memory | +| `compare_config()` | diffs against `GET /api/routes/routes/searchroute` | +| `commit_config()` | snapshots backup → `POST /api/routes/routes/addroute` × n → `reconfigure` | +| `discard_config()` | clears staged candidate | +| `rollback()` | `POST /api/core/backup/revert_backup/{id}` — restores the config.xml snapshot taken before the last commit | + +### Config format + +The `load_merge_candidate` `config` parameter must be a **JSON array** of route +objects, each with `network` and `gateway` keys. `gateway` must be the +**name** of an existing OPNsense gateway (as configured under +*System → Gateways → Configuration*), not an IP address. + +```json +[ + { + "network": "10.0.0.0/8", + "gateway": "WAN_GW", + "descr": "Corporate internal", + "disabled": "0" + }, + { + "network": "0.0.0.0/0", + "gateway": "WAN_GW", + "descr": "Default route" + } +] +``` + +### Example + +```python +import json + +driver = get_network_driver("opnsense") +d = driver("192.168.1.1", "user", "pass", optional_args={"api_key": "k", "api_secret": "s"}) +d.open() + +routes = json.dumps([{"network": "10.0.0.0/8", "gateway": "WAN_GW"}]) +d.load_merge_candidate(config=routes) + +print(d.compare_config()) # unified diff +d.commit_config() # applies routes and calls reconfigure +d.rollback() # removes the routes just added +d.close() +``` + +### Limitations + +- `load_replace_candidate` is not supported (no XML upload endpoint in the JSON API). +- Non-route config (firewall rules, DHCP, DNS, interfaces, …) must be managed + via OPNsense module-specific controllers — outside the scope of this driver. +- `rollback()` reverts the entire `config.xml` to the pre-commit state (not just + the routes). OPNsense's backup API has no partial-restore capability. +- Rollback uses the backup snapshot taken at `commit_config()` time. If no + commit was made in the current session, the most recent available backup is + used as a fallback. + +## Optional arguments + +| Argument | Default | Description | +|---|---|---| +| `api_key` | `username` | OPNsense API key | +| `api_secret` | `password` | OPNsense API secret | +| `base_url` | `https://` | Override the base URL | +| `verify` | `True` | TLS certificate verification (path or bool) | + +## Development + +```bash +pip install -e ".[dev]" +pytest tests/ +``` + +## CI + +This project includes a GitHub Actions workflow that: +- Runs unit tests across Python 3.9–3.12 +- Builds sdist and wheel +- Uploads build artifacts + +See [.github/workflows/ci.yml](.github/workflows/ci.yml). + +## License + +Apache 2.0 — see [LICENSE](LICENSE). + diff --git a/examples/test_driver.py b/examples/test_driver.py new file mode 100644 index 0000000..ea92dd1 --- /dev/null +++ b/examples/test_driver.py @@ -0,0 +1,66 @@ +from napalm_opnsense.opnsense import OPNsenseDriver + + +def main(): + driver = OPNsenseDriver( + hostname="test.local", + username="api_key", + password="api_secret", + optional_args={ + "base_url": "https://test.local", + "verify": False, + }, + ) + + # Mock API responses to avoid real network calls + def fake_get(path: str): + if path == "/api/core/system/status": + return { + "hostname": "opnsense01", + "version": "24.7", + "model": "OPNsense", + "serial": "ABC123", + "uptime": 12345, + } + if path == "/api/interfaces/overview/export": + return { + "interfaces": [ + { + "name": "em0", + "up": True, + "enabled": True, + "descr": "LAN", + "mac": "AA:BB:CC:DD:EE:FF", + "speed_mbps": 1000, + "mtu": 1500, + } + ] + } + if path == "/api/interfaces/addresses/export": + return { + "items": [ + {"interface": "em0", "address": "192.0.2.10", "prefix": 24}, + {"interface": "em0", "address": "2001:db8::1", "prefix": 64}, + ] + } + return {} + + # Monkeypatch before open() + driver._get = fake_get # type: ignore + + driver.open() + + print("Facts:") + print(driver.get_facts()) + + print("\nInterfaces:") + print(driver.get_interfaces()) + + print("\nInterfaces IP:") + print(driver.get_interfaces_ip()) + + driver.close() + + +if __name__ == "__main__": + main() diff --git a/napalm_opnsense/__init__.py b/napalm_opnsense/__init__.py new file mode 100644 index 0000000..d8ba7c0 --- /dev/null +++ b/napalm_opnsense/__init__.py @@ -0,0 +1,3 @@ +from napalm_opnsense.opnsense import OPNsenseDriver + +__all__ = ["OPNsenseDriver"] diff --git a/napalm_opnsense/opnsense.py b/napalm_opnsense/opnsense.py new file mode 100644 index 0000000..52a1888 --- /dev/null +++ b/napalm_opnsense/opnsense.py @@ -0,0 +1,1390 @@ +# -*- 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 OPNsense firewalls (REST API). + +OPNsense exposes a JSON REST API at ``/api/``. Authentication uses an +API key/secret pair generated in the OPNsense GUI (System → Access → +Users → edit user → API keys). Pass them via *optional_args*: + + optional_args={ + "api_key": "", + "api_secret": "", + "verify": False, # disable TLS verification for self-signed certs + } + +If *api_key* / *api_secret* are omitted the driver falls back to the +positional *username* / *password* arguments. + +Config management (load_merge_candidate / commit_config / rollback) +operates on **static routes** via ``/api/routes/routes/``. Other parts +of the OPNsense configuration (firewall rules, interface settings, …) are +not exposed through a generic config-push endpoint and must be managed +per-module through their respective API controllers. +""" + +from __future__ import annotations + +import difflib +import json +import socket +from typing import Any, Dict, List, Optional + +import requests +from requests.exceptions import RequestException + +from napalm_device_types import FirewallDriver +from napalm.base.exceptions import ConnectionException, ConnectionClosedException, MergeConfigException + + +class OPNsenseDriver(FirewallDriver): + """NAPALM driver for OPNsense (read-only, REST API).""" + + VENDOR = "OPNsense" + + def __init__( + self, + hostname: str, + username: str, + password: str, + timeout: int = 60, + optional_args: Optional[Dict[str, Any]] = None, + ) -> None: + self.hostname = hostname + self.username = username + self.password = password + self.timeout = timeout + self.optional_args = optional_args or {} + + # NAPALM standard attributes + self.force_no_enable = True + self.use_canonical_interface = False + + # OPNsense REST API settings + self.base_url = self.optional_args.get("base_url") or f"https://{hostname}" + self.verify = self.optional_args.get("verify", True) + self.api_key = self.optional_args.get("api_key") or username + self.api_secret = self.optional_args.get("api_secret") or password + + self.session: Optional[requests.Session] = None + + # Config-management state + self._candidate_config: Optional[List[Dict[str, Any]]] = None + # Backup ID of the config snapshot taken just before commit_config(). + # Used by rollback() to restore the exact pre-commit state. + self._pre_commit_backup_id: Optional[str] = None + + # ------------------------------------------------------------------ + # Connection management + # ------------------------------------------------------------------ + + def open(self) -> None: + """Open an HTTPS session to the OPNsense API and validate credentials.""" + try: + s = requests.Session() + s.verify = self.verify + s.headers.update({"Accept": "application/json"}) + s.auth = (self.api_key, self.api_secret) + self.session = s + # Lightweight connectivity and auth check + self._get("/api/core/system/status") + except RequestException as exc: + self.session = None + raise ConnectionException( + f"Cannot connect to OPNsense API at {self.base_url}: {exc}" + ) from exc + + def close(self) -> None: + """Close the HTTPS session.""" + if self.session is not None: + self.session.close() + self.session = None + + def is_alive(self) -> Dict[str, bool]: + """Return whether the session is usable. + + Performs a lightweight socket-level check without sending a full + HTTP request to avoid polluting API logs. + """ + if self.session is None: + return {"is_alive": False} + try: + host = self.hostname + port = 443 + with socket.create_connection((host, port), timeout=5): + pass + return {"is_alive": True} + except (socket.error, OSError): + return {"is_alive": False} + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + def _get(self, path: str) -> Dict[str, Any]: + """Perform a GET request against the OPNsense REST API. + + :raises ConnectionClosedException: if called before :meth:`open`. + :raises RequestException: on HTTP-level errors. + """ + if self.session is None: + raise ConnectionClosedException("Not connected – call open() first.") + url = self.base_url.rstrip("/") + path + response = self.session.get(url, timeout=self.timeout) + response.raise_for_status() + return response.json() + + def _post(self, path: str, data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: + """Perform a POST request against the OPNsense REST API. + + :param path: API path, e.g. ``/api/routes/routes/addroute``. + :param data: JSON-serialisable payload (sent as ``application/json``). + :raises ConnectionClosedException: if called before :meth:`open`. + :raises RequestException: on HTTP-level errors. + """ + if self.session is None: + raise ConnectionClosedException("Not connected – call open() first.") + url = self.base_url.rstrip("/") + path + response = self.session.post(url, json=data or {}, timeout=self.timeout) + response.raise_for_status() + return response.json() + + # ------------------------------------------------------------------ + # NAPALM getters + # ------------------------------------------------------------------ + + def get_facts(self) -> Dict[str, Any]: + """Return a dictionary of general device facts. + + Calls ``GET /api/core/system/status``. + + Returned keys (NAPALM standard): + ``vendor``, ``model``, ``hostname``, ``fqdn``, ``os_version``, + ``serial_number``, ``uptime``, ``interface_list``. + """ + status = self._get("/api/core/system/status") + + hostname = status.get("hostname") or status.get("name") or self.hostname + version = status.get("version") or status.get("product_version") or "unknown" + + try: + interface_list = list(self.get_interfaces().keys()) + except Exception: + interface_list = [] + + return { + "vendor": self.VENDOR, + "model": status.get("model") or self.VENDOR, + "hostname": hostname, + "fqdn": hostname, + "os_version": version, + "serial_number": status.get("serial") or "", + "uptime": status.get("uptime", -1), + "interface_list": interface_list, + } + + def get_interfaces(self) -> Dict[str, Dict[str, Any]]: + """Return interface details keyed by interface name. + + Calls ``GET /api/interfaces/overview/export``. + + Each entry contains NAPALM standard keys: + ``is_up``, ``is_enabled``, ``description``, ``last_flapped``, + ``mac_address``, ``speed``, ``mtu``. + """ + data = self._get("/api/interfaces/overview/export") + interfaces: Dict[str, Dict[str, Any]] = {} + + # API returns a bare list in newer OPNsense versions; + # older/wrapped format uses {"interfaces": [...]} + items: list = data if isinstance(data, list) else data.get("interfaces", []) + + for iface in items: + name = iface.get("device") or iface.get("name", "") + if not name: + continue + interfaces[name] = { + "is_up": iface.get("status", "") == "up" if "status" in iface else bool(iface.get("up", False)), + "is_enabled": bool(iface.get("enabled", True)), + "description": iface.get("description") or iface.get("descr") or "", + "last_flapped": -1.0, + "mac_address": (iface.get("macaddr") or iface.get("mac") or "").lower(), + "speed": float(iface["speed_mbps"]) if iface.get("speed_mbps") else 0.0, + "mtu": int(iface["mtu"]) if iface.get("mtu") else 0, + } + + return interfaces + + def get_interfaces_ip(self) -> Dict[str, Dict[str, Any]]: + """Return IP addresses grouped by interface name. + + Calls ``GET /api/interfaces/addresses/export``. + + Structure:: + + { + "em0": { + "ipv4": {"192.0.2.10": {"prefix_length": 24}}, + "ipv6": {"2001:db8::1": {"prefix_length": 64}}, + } + } + """ + data = self._get("/api/interfaces/addresses/export") + result: Dict[str, Dict[str, Any]] = {} + + for item in data.get("items", []): + ifname: str = item["interface"] + ip: str = item["address"] + prefix: int = int(item["prefix"]) + family = "ipv6" if ":" in ip else "ipv4" + + result.setdefault(ifname, {"ipv4": {}, "ipv6": {}}) + result[ifname][family][ip] = {"prefix_length": prefix} + + return result + + def get_networks(self) -> List[Dict[str, Any]]: + """Return the IP networks this firewall is authoritative for. + + Derived from the interface overview (``GET /api/interfaces/overview/export``). + Loopback and link-local addresses are excluded. + + Each entry:: + + { + "network": "192.168.1.0/24", + "interface": "em0", + "gateway": "192.168.1.1", + "family": "ipv4", + "prefix_length": 24, + } + """ + import ipaddress + + data = self._get("/api/interfaces/overview/export") + items: list = data if isinstance(data, list) else data.get("interfaces", []) + + networks: List[Dict[str, Any]] = [] + + def _add(ifname: str, cidr: str, family: str) -> None: + """Parse a CIDR string (e.g. '10.0.0.1/24') and append to networks.""" + try: + iface_obj = ipaddress.ip_interface(cidr) + net = iface_obj.network + if net.is_loopback or net.is_link_local: + return + networks.append({ + "network": str(net), + "interface": ifname, + "gateway": str(iface_obj.ip), + "family": family, + "prefix_length": net.prefixlen, + }) + except ValueError: + pass + + for iface in items: + name: str = iface.get("device") or iface.get("name", "") + if not name: + continue + + # Primary format: flat CIDR strings in addr4/addr6 + addr4: str = iface.get("addr4", "") + addr6: str = iface.get("addr6", "") + if addr4: + _add(name, addr4, "ipv4") + if addr6: + _add(name, addr6, "ipv6") + + # Fallback: ipv4/ipv6 arrays where ipaddr may include prefix + if not addr4: + for entry in iface.get("ipv4") or []: + ip_field = entry.get("ipaddr") or entry.get("ip", "") + subnet = entry.get("subnetbits") or entry.get("prefix_length") + cidr = f"{ip_field}/{subnet}" if subnet and "/" not in ip_field else ip_field + if cidr: + _add(name, cidr, "ipv4") + if not addr6: + for entry in iface.get("ipv6") or []: + ip_field = entry.get("ipaddr") or entry.get("ip", "") + prefix = entry.get("prefixlen") or entry.get("prefix_length") + cidr = f"{ip_field}/{prefix}" if prefix and "/" not in ip_field else ip_field + if cidr: + _add(name, cidr, "ipv6") + + return networks + + def get_arp_table(self, vrf: str = "") -> List[Dict[str, Any]]: + """Return the ARP table. + + Calls ``GET /api/diagnostics/interface/get_arp``. + + Each entry contains: ``interface``, ``mac``, ``ip``, ``age``. + """ + data = self._get("/api/diagnostics/interface/get_arp") + arp_table: List[Dict[str, Any]] = [] + + for entry in data if isinstance(data, list) else data.get("arp", []): + arp_table.append( + { + "interface": entry.get("intf") or entry.get("interface", ""), + "mac": (entry.get("mac") or "").lower(), + "ip": entry.get("ip") or entry.get("address", ""), + "age": float(entry.get("expires", 0)), + } + ) + + return arp_table + + def get_interfaces_counters(self) -> Dict[str, Dict[str, Any]]: + """Return per-interface packet and byte counters. + + Calls ``GET /api/diagnostics/interface/get_interface_statistics``. + + Each entry contains NAPALM standard keys: + ``tx_errors``, ``rx_errors``, ``tx_discards``, ``rx_discards``, + ``tx_octets``, ``rx_octets``, + ``tx_unicast_packets``, ``rx_unicast_packets``, + ``tx_multicast_packets``, ``rx_multicast_packets``, + ``tx_broadcast_packets``, ``rx_broadcast_packets``. + """ + data = self._get("/api/diagnostics/interface/get_interface_statistics") + counters: Dict[str, Dict[str, Any]] = {} + + for iface, stats in data.get("statistics", {}).items(): + counters[iface] = { + "tx_errors": int(stats.get("output-errors", 0)), + "rx_errors": int(stats.get("input-errors", 0)), + "tx_discards": int(stats.get("output-drops", 0)), + "rx_discards": int(stats.get("input-drops", 0)), + "tx_octets": int(stats.get("output-bytes", 0)), + "rx_octets": int(stats.get("input-bytes", 0)), + "tx_unicast_packets": int(stats.get("output-packets", 0)), + "rx_unicast_packets": int(stats.get("input-packets", 0)), + "tx_multicast_packets": int(stats.get("output-multicasts", 0)), + "rx_multicast_packets": int(stats.get("input-multicasts", 0)), + "tx_broadcast_packets": int(stats.get("output-broadcasts", 0)), + "rx_broadcast_packets": int(stats.get("input-broadcasts", 0)), + } + + return counters + + def get_environment(self) -> Dict[str, Any]: + """Return device environment data (CPU, memory, temperature). + + Calls: + - ``GET /api/diagnostics/system/system_resources`` for CPU and memory. + - ``GET /api/diagnostics/system/system_temperature`` for temperature sensors. + + ``fan`` and ``power`` fields are not exposed by OPNsense and are + returned with assumed-healthy placeholder values. + """ + resources = self._get("/api/diagnostics/system/system_resources") + + cpu_pct = float(resources.get("cpu", {}).get("used", 0)) + mem_total = int(resources.get("memory", {}).get("total", 0) or 0) + mem_used = int(resources.get("memory", {}).get("used", 0) or 0) + + try: + temp_data = self._get("/api/diagnostics/system/system_temperature") + except Exception: + temp_data = {} + + temperature: Dict[str, Any] = {} + for sensor in temp_data.get("data", []): + name = sensor.get("device") or sensor.get("name", "") + temp_val = float(sensor.get("temperature", 0)) + temperature[name] = { + "temperature": temp_val, + "is_alert": temp_val > 80.0, + "is_critical": temp_val > 95.0, + } + + return { + "fans": {}, + "temperature": temperature, + "power": {}, + "cpu": {0: {"%usage": cpu_pct}}, + "memory": { + "available_ram": mem_total - mem_used, + "used_ram": mem_used, + }, + } + + def get_route_to( + self, + destination: str = "", + protocol: str = "", + longer: bool = False, + ) -> Dict[str, List[Dict[str, Any]]]: + """Return routing table entries. + + Calls ``GET /api/diagnostics/interface/get_routes``. + + :param destination: Filter by exact prefix (e.g. ``"192.0.2.0/24"``). + :param protocol: Filter by protocol name (``"static"``, ``"connected"``). + :param longer: Ignored (OPNsense does not support longer-prefixes filter). + + Returns a NAPALM-standard route dict keyed by network prefix. + """ + data = self._get("/api/diagnostics/interface/get_routes") + routes: Dict[str, List[Dict[str, Any]]] = {} + + proto_map = { + "static": "static", + "ospf": "ospf", + "bgp": "bgp", + "rip": "rip", + "kernel": "connected", + "connected": "connected", + "local": "connected", + } + + for route in data.get("route", []): + network = route.get("network") or route.get("destination", "") + if not network: + continue + + if destination and network != destination: + continue + + flags = route.get("flags", "").upper() + proto_raw = route.get("proto", "").lower() + proto = proto_map.get(proto_raw, proto_raw) + + if protocol and proto != protocol.lower(): + continue + + gateway = route.get("gateway") or route.get("nexthop", "") + iface = route.get("netif") or route.get("interface", "") + + entry: Dict[str, Any] = { + "protocol": proto, + "current_active": "U" in flags, + "last_active": False, + "age": -1, + "next_hop": gateway if gateway not in ("link#", "0.0.0.0", "") else "", + "outgoing_interface": iface, + "selected_next_hop": True, + "preference": int(route.get("priority", 0)), + "inactive_reason": "", + "routing_table": "global", + "protocol_attributes": {}, + } + routes.setdefault(network, []).append(entry) + + return routes + + def get_ipv6_neighbors_table(self) -> List[Dict[str, Any]]: + """Return the IPv6 Neighbor Discovery (NDP) table. + + Calls ``GET /api/diagnostics/interface/get_ndp``. + + Each entry contains: ``interface``, ``mac``, ``ip``, ``age``, + ``state`` (best-effort from NDP flags). + """ + data = self._get("/api/diagnostics/interface/get_ndp") + neighbors: List[Dict[str, Any]] = [] + + rows = data if isinstance(data, list) else data.get("rows", []) + for entry in rows: + neighbors.append( + { + "interface": entry.get("intf") or entry.get("interface", ""), + "mac": (entry.get("mac") or "").lower(), + "ip": entry.get("ip") or entry.get("address", ""), + "age": float(entry.get("expires", 0)), + "state": entry.get("state", ""), + } + ) + + return neighbors + + def get_lldp_neighbors(self) -> Dict[str, List[Dict[str, Any]]]: + """Return LLDP neighbors grouped by local port. + + Calls ``GET /api/lldpd/service/neighbor``. + + Requires the ``os-lldpd`` plugin to be installed on OPNsense. + Returns an empty dict if the plugin is not present. + """ + try: + data = self._get("/api/lldpd/service/neighbor") + except Exception: + return {} + + neighbors: Dict[str, List[Dict[str, Any]]] = {} + for row in data.get("rows", []): + port = row.get("local_port") or row.get("port", "") + neighbors.setdefault(port, []).append( + { + "hostname": row.get("system_name") or row.get("chassis", ""), + "port": row.get("port_id") or row.get("port", ""), + } + ) + return neighbors + + def get_lldp_neighbors_detail( + self, interface: str = "" + ) -> Dict[str, List[Dict[str, Any]]]: + """Return detailed LLDP neighbor information. + + Calls ``GET /api/lldpd/service/neighbor``. + + Requires the ``os-lldpd`` plugin. Returns an empty dict if the + plugin is not available. + + :param interface: If set, filter results to this local port. + """ + try: + data = self._get("/api/lldpd/service/neighbor") + except Exception: + return {} + + details: Dict[str, List[Dict[str, Any]]] = {} + for row in data.get("rows", []): + port = row.get("local_port") or row.get("port", "") + if interface and port != interface: + continue + details.setdefault(port, []).append( + { + "parent_interface": "", + "remote_port": row.get("port_id") or row.get("port", ""), + "remote_port_description": row.get("port_description", ""), + "remote_chassis_id": row.get("chassis_id") or row.get("chassis", ""), + "remote_system_name": row.get("system_name", ""), + "remote_system_description": row.get("system_description", ""), + "remote_system_capab": [ + c.strip().lower() + for c in row.get("system_capabilities", "").split(",") + if c.strip() + ], + "remote_system_enable_capab": [ + c.strip().lower() + for c in row.get("enabled_capabilities", "").split(",") + if c.strip() + ], + } + ) + return details + + def get_ntp_servers(self) -> Dict[str, Dict[str, Any]]: + """Return configured NTP servers. + + Calls ``GET /api/ntpd/service/status`` which includes the list of + configured peer addresses in the ``peers`` field. + + Returns a dict keyed by server address with an empty value dict + (NAPALM standard format). + """ + try: + data = self._get("/api/ntpd/service/status") + except Exception: + return {} + + servers: Dict[str, Dict[str, Any]] = {} + for peer in data.get("peers", []): + addr = peer.get("address") or peer.get("remote", "") + if addr: + servers[addr] = {} + return servers + + def get_vlans(self) -> Dict[str, Dict[str, Any]]: + """Return configured VLANs. + + Calls ``GET /api/interfaces/vlan_settings/search_item``. + + Each VLAN device configured under *Interfaces → Other Types → VLAN* + becomes one entry, keyed by VLAN tag (as a string). + + The ``interfaces`` list contains the VLAN device name (e.g. + ``em0_vlan10``). When a device is already assigned to a logical + interface OPNsense appends the logical name in brackets + (``"em0_vlan10 [LAN]"``); this driver strips that annotation and + stores the bare device name. + + Returned keys (NAPALM standard): + ``name`` (description, or device name if no description is set), + ``interfaces`` (list with the VLAN device name). + """ + data = self._get("/api/interfaces/vlan_settings/search_item") + vlans: Dict[str, Dict[str, Any]] = {} + + for row in data.get("rows", []): + tag = str(row.get("tag", "")).strip() + if not tag: + continue + + # vlanif may be "em0_vlan10 [LAN]" when assigned to a logical interface + vlanif_raw = str(row.get("vlanif", "")) + vlanif = vlanif_raw.split(" [")[0].strip() + + descr = str(row.get("descr", "")).strip() + vlans[tag] = { + "name": descr or vlanif, + "interfaces": [vlanif] if vlanif else [], + } + + return vlans + + def get_bgp_neighbors(self) -> Dict[str, Any]: + """Return BGP neighbor state. + + Requires the FRR plugin (``os-frr``) to be installed on OPNsense. + Returns an empty dict if the plugin is absent or FRR is not running. + + Calls: + + - ``GET /api/quagga/bgp/get`` — local AS number and router-id from + the BGP configuration model. + - ``GET /api/quagga/diagnostics/bgpneighbors`` — live neighbor state + as returned by FRR (``vtysh -c "show bgp neighbors json"``). + + Returns a NAPALM-standard dict keyed by VRF name. Only the default + VRF (``"global"``) is populated; per-VRF BGP is not yet mapped. + + Each peer entry contains: + ``local_as``, ``remote_as``, ``remote_id``, ``is_up``, + ``is_enabled``, ``description``, ``uptime``, + ``address_family`` (``ipv4`` and/or ``ipv6`` with prefix counters). + """ + try: + bgp_cfg = self._get("/api/quagga/bgp/get") + neighbors_data = self._get("/api/quagga/diagnostics/bgpneighbors") + except Exception: + return {} + + bgp = bgp_cfg.get("bgp", {}) + local_as_default = int(bgp.get("asnumber", 0) or 0) + router_id = bgp.get("routerid", "") + + raw_neighbors = neighbors_data.get("response", {}) + if not isinstance(raw_neighbors, dict): + return {} + + # FRR address-family key → NAPALM address-family key + _AF_MAP = { + "ipv4Unicast": "ipv4", + "ipv6Unicast": "ipv6", + } + + peers: Dict[str, Any] = {} + for peer_ip, nbr in raw_neighbors.items(): + if not isinstance(nbr, dict): + continue + + is_up = nbr.get("bgpState", "").lower() == "established" + uptime_msec = int(nbr.get("bgpTimerUpMsec", 0) or 0) + uptime = uptime_msec // 1000 if is_up else -1 + + address_family: Dict[str, Any] = {} + for frr_af, napalm_af in _AF_MAP.items(): + af = nbr.get("addressFamilyInfo", {}).get(frr_af) + if af is not None: + address_family[napalm_af] = { + "sent_prefixes": int(af.get("sentPrefixCounter", 0) or 0), + "received_prefixes": int(af.get("prefixReceivedCount", 0) or 0), + "accepted_prefixes": int(af.get("acceptedPrefixCounter", 0) or 0), + } + + if not address_family: + # FRR did not report AF info (session not yet established) + address_family["ipv4"] = { + "sent_prefixes": -1, + "received_prefixes": -1, + "accepted_prefixes": -1, + } + + peers[peer_ip] = { + "local_as": int(nbr.get("localAs", local_as_default) or local_as_default), + "remote_as": int(nbr.get("remoteAs", 0) or 0), + "remote_id": nbr.get("remoteRouterId", ""), + "is_up": is_up, + "is_enabled": not bool(nbr.get("adminShutdown", False)), + "description": nbr.get("nbrDesc", ""), + "uptime": uptime, + "address_family": address_family, + } + + return { + "global": { + "router_id": router_id, + "peers": peers, + } + } + + def get_config( + self, + retrieve: str = "all", + full: bool = False, + sanitized: bool = False, + format: str = "text", + ) -> Dict[str, str]: + """Return device configuration. + + OPNsense stores its configuration as XML. This getter returns the + raw XML text in the ``running`` slot. ``startup`` mirrors ``running`` + (OPNsense applies config immediately). The ``candidate`` slot shows + the JSON-serialised staged routes when a candidate has been loaded, + or an empty string otherwise. + + Calls ``GET /api/core/backup/download/this``. + """ + configs: Dict[str, str] = {"running": "", "startup": "", "candidate": ""} + + if retrieve in ("all", "running", "startup"): + try: + response = self._get("/api/core/backup/download/this") + xml_text: str = ( + response if isinstance(response, str) else str(response) + ) + if retrieve in ("all", "running"): + configs["running"] = xml_text + if retrieve in ("all", "startup"): + configs["startup"] = xml_text + except Exception: + pass + + if retrieve in ("all", "candidate") and self._candidate_config is not None: + configs["candidate"] = json.dumps(self._candidate_config, indent=2) + + return configs + + # ------------------------------------------------------------------ + # Config management (static routes) + # ------------------------------------------------------------------ + + def load_merge_candidate(self, filename: Optional[str] = None, config: Optional[str] = None) -> None: + """Stage a set of static-route additions as a candidate config. + + OPNsense does not offer a single generic config-push endpoint. + This method targets the **routes** subsystem + (``/api/routes/routes/``) and accepts a JSON list of route objects. + + The *config* parameter must be a JSON string containing a list of + route dicts. Each dict may contain the following keys: + + .. code-block:: json + + [ + { + "network": "10.0.0.0/8", + "gateway": "WAN_GW", + "descr": "optional description", + "disabled": "0" + } + ] + + ``gateway`` must be the **name** of an existing OPNsense gateway + (not an IP address) as shown in + *System → Gateways → Configuration*. + + :param filename: Path to a JSON file containing the route list. + :param config: JSON string containing the route list. + :raises MergeConfigException: if neither or both arguments are given, + or if the JSON is malformed. + """ + if filename is None and config is None: + raise MergeConfigException("Provide either 'filename' or 'config'.") + if filename is not None and config is not None: + raise MergeConfigException("Provide either 'filename' or 'config', not both.") + + if filename is not None: + try: + with open(filename, "r", encoding="utf-8") as fh: + config = fh.read() + except OSError as exc: + raise MergeConfigException(f"Cannot read file {filename!r}: {exc}") from exc + + try: + routes = json.loads(config) # type: ignore[arg-type] + except json.JSONDecodeError as exc: + raise MergeConfigException(f"Invalid JSON in candidate config: {exc}") from exc + + if not isinstance(routes, list): + raise MergeConfigException( + "Candidate config must be a JSON array of route objects." + ) + for i, route in enumerate(routes): + if not isinstance(route, dict): + raise MergeConfigException(f"Route at index {i} must be a JSON object.") + if "network" not in route or "gateway" not in route: + raise MergeConfigException( + f"Route at index {i} is missing 'network' or 'gateway'." + ) + + self._candidate_config = routes + + def compare_config(self) -> str: + """Return a unified diff of the candidate vs the current routes. + + Fetches the live route list from ``GET /api/routes/routes/searchroute`` + and diffs it against the staged candidate. + + Returns an empty string if no candidate has been loaded. + """ + if self._candidate_config is None: + return "" + + current_routes = self._fetch_current_routes() + current_text = json.dumps(current_routes, indent=2, sort_keys=True) + candidate_text = json.dumps(self._candidate_config, indent=2, sort_keys=True) + + diff = difflib.unified_diff( + current_text.splitlines(keepends=True), + candidate_text.splitlines(keepends=True), + fromfile="current", + tofile="candidate", + ) + return "".join(diff) + + def commit_config(self, message: str = "", revert_in: Optional[int] = None) -> None: + """Apply the staged candidate routes to the device. + + Each route in the candidate is submitted via + ``POST /api/routes/routes/addroute``. After all routes are added, + ``POST /api/routes/routes/reconfigure`` is called to activate them. + + The UUIDs returned by OPNsense are stored internally so that + :meth:`rollback` can remove exactly these routes. + + :param message: Ignored (OPNsense has no commit-message concept). + :param revert_in: Ignored (auto-rollback not supported via API). + :raises MergeConfigException: if no candidate is staged, or if the + API returns an error for any route. + """ + if self._candidate_config is None: + raise MergeConfigException("No candidate config loaded. Call load_merge_candidate() first.") + + # Record the current backup ID so rollback() can restore exactly + # this state after OPNsense writes the new config to disk. + self._pre_commit_backup_id = self._get_latest_backup_id() + + try: + for route in self._candidate_config: + payload = { + "route": { + "network": route["network"], + "gateway": route["gateway"], + "descr": route.get("descr", ""), + "disabled": route.get("disabled", "0"), + } + } + self._post("/api/routes/routes/addroute", payload) + + self._post("/api/routes/routes/reconfigure") + except RequestException as exc: + raise MergeConfigException(f"Failed to apply route config: {exc}") from exc + + self._candidate_config = None + + def discard_config(self) -> None: + """Discard the staged candidate config without applying it.""" + self._candidate_config = None + + def rollback(self) -> None: + """Revert the device to the configuration state captured before the last commit. + + OPNsense automatically saves a backup of ``config.xml`` before applying + configuration changes. :meth:`commit_config` records the ID of the + most recent backup at the time of the commit so that this method can + restore exactly the pre-commit state via + ``POST /api/core/backup/revert_backup/{backup_id}``. + + If no backup ID is available (i.e. :meth:`commit_config` was never + called in this session, or the backup list was empty at commit time), + the most recent backup from the device is used as a fallback. + + If no backups exist at all, this is a no-op. + """ + backup_id = self._pre_commit_backup_id or self._get_latest_backup_id() + if backup_id is None: + return + + self._post(f"/api/core/backup/revert_backup/{backup_id}") + self._pre_commit_backup_id = None + + # ------------------------------------------------------------------ + # Internal helpers (config management) + # ------------------------------------------------------------------ + + def _get_latest_backup_id(self) -> Optional[str]: + """Return the ID of the most recent server-side config backup, or ``None``. + + Calls ``GET /api/core/backup/backups/this``. The response is sorted + newest-first by OPNsense. Returns ``None`` when no backups exist or + the request fails. + """ + try: + data = self._get("/api/core/backup/backups/this") + items = data.get("items", []) + return items[0]["id"] if items else None + except Exception: + return None + + def _fetch_current_routes(self) -> List[Dict[str, Any]]: + """Return the current static routes from the OPNsense API. + + Calls ``GET /api/routes/routes/searchroute`` and normalises the + response to the same keys used by :meth:`load_merge_candidate`. + """ + data = self._get("/api/routes/routes/searchroute") + routes: List[Dict[str, Any]] = [] + for row in data.get("rows", []): + routes.append( + { + "network": row.get("network", ""), + "gateway": row.get("gateway", ""), + "descr": row.get("descr", ""), + "disabled": row.get("disabled", "0"), + } + ) + return routes + + # ------------------------------------------------------------------ + # NetOrch extensions: packages, services, updates + # ------------------------------------------------------------------ + + def get_packages(self) -> List[Dict[str, Any]]: + """Return installed OPNsense plugins. + + Calls ``GET /api/core/firmware/info`` and returns the ``plugin`` + list filtered to entries where ``installed == "1"``. The format + mirrors the OpenWrt NAPALM driver so NetOrch can render both + drivers with the same UI component: + ``{name, version, installed, description, size, source}`` + """ + info = self._get("/api/core/firmware/info") + result: List[Dict[str, Any]] = [] + for p in info.get("plugin", []): + if p.get("installed") != "1": + continue + result.append({ + "name": p.get("name", ""), + "version": p.get("version", ""), + "installed": True, + "description": p.get("comment", ""), + "size": 0, + "source": "opnsense-plugins", + }) + return sorted(result, key=lambda x: x["name"].lower()) + + def get_dhcp_leases(self) -> List[Dict[str, Any]]: + """Return active DHCP leases from OPNsense. + + Tries all known DHCP backends in order: + + 1. **Kea DHCPv4** (os-kea plugin) — ``GET /api/kea/leases4/search`` + Fields: ``hw-address``, ``ip-address``, ``hostname``, ``expire`` + 2. **ISC DHCP** (legacy) — ``POST /api/dhcpv4/leases/searchlease`` + Fields: ``mac``, ``address``, ``hostname``, ``ends`` + 3. **ARP table** fallback — ``GET /api/diagnostics/interface/get_arp`` + Provides IP only, hostname will be empty. + + Each returned entry contains: + + * ``mac`` — lower-case MAC address + * ``ip`` — assigned IP address + * ``hostname`` — client hostname (may be empty) + * ``lease_end`` — Unix timestamp when the lease expires (0 if unknown) + * ``state`` — raw state string + """ + def _parse_kea(rows: List[Dict]) -> List[Dict[str, Any]]: + result = [] + for row in rows: + # OPNsense Kea uses "hwaddr"; standard Kea uses "hw-address"; ISC DHCP uses "mac" + mac = ( + row.get("hwaddr") or row.get("hw-address") or row.get("mac") or "" + ).lower().strip() + if not mac: + continue + raw_end = row.get("expire") or row.get("ends") or 0 + try: + lease_end = int(raw_end) + except (ValueError, TypeError): + lease_end = 0 + result.append({ + "mac": mac, + "ip": (row.get("address") or row.get("ip-address") or row.get("ip") or "").strip(), + "hostname": (row.get("hostname") or "").strip(), + "lease_end": lease_end, + "state": str(row.get("state") or ""), + }) + return result + + # 1. Kea DHCPv4 plugin + try: + data = self._get("/api/kea/leases4/search") + rows = data.get("rows") or data.get("leases") or (data if isinstance(data, list) else []) + leases = _parse_kea(rows) + if leases: + return leases + except Exception: + pass + + # 2. ISC DHCP (legacy) + try: + data = self._post( + "/api/dhcpv4/leases/searchlease", + {"current": 1, "rowCount": -1, "searchPhrase": "", "sort": {}}, + ) + leases = _parse_kea(data.get("rows", [])) + if leases: + return leases + except Exception: + pass + + # 3. ARP table fallback (IP only, no hostname) + try: + arp = self.get_arp_table() + return [ + {"mac": e["mac"], "ip": e["ip"], "hostname": "", "lease_end": 0, "state": "arp"} + for e in arp + if e.get("mac") and e.get("ip") + ] + except Exception: + return [] + + def get_services(self) -> List[Dict[str, Any]]: + """Return running services from OPNsense. + + Calls ``GET /api/core/service/search`` and normalises the rows to + ``{name, running, enabled, pid}`` — the same format used by the + OpenWrt driver so the UI can render them identically. + """ + data = self._get("/api/core/service/search") + result: List[Dict[str, Any]] = [] + for row in data.get("rows", []): + result.append({ + "name": row.get("name") or row.get("id", ""), + "running": bool(row.get("running", 0)), + "enabled": True, # OPNsense has no separate enabled/disabled state + "pid": 0, + }) + return sorted(result, key=lambda x: x["name"].lower()) + + def manage_service(self, name: str, action: str) -> Dict[str, Any]: + """Execute a lifecycle action on an OPNsense service. + + OPNsense exposes per-plugin service endpoints at + ``/api//service/``. The ``name`` parameter must + match the service ``id`` returned by :meth:`get_services`. + + Supported actions: ``start``, ``stop``, ``restart``. + ``enable`` and ``disable`` are not supported on OPNsense (plugins + are enabled/disabled by installing or removing them). + """ + 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'): + return {"success": False, "output": f"Action '{action}' is not supported on OPNsense services"} + try: + result = self._post(f"/api/{name}/service/{action}") + return {"success": True, "output": str(result)} + except Exception as exc: + return {"success": False, "output": str(exc)} + + def get_vpn_tunnels(self) -> Dict[str, Dict[str, Any]]: + """Return status of all configured VPN tunnels. + + Queries IPsec, OpenVPN, and WireGuard in order and merges results + into a single dict keyed by a unique tunnel identifier. + + Each entry follows the ``VPNTunnelDict`` schema: + ``type``, ``local_endpoint``, ``remote_endpoint``, ``is_up``, + ``uptime``, ``bytes_in``, ``bytes_out``, ``description``. + + * **IPsec** — ``GET /api/ipsec/sessions`` + Falls back to ``GET /api/ipsec/leases/searchPhase2`` on older + OPNsense releases. + * **OpenVPN** — ``GET /api/openvpn/instances/search`` + Instance status obtained from ``GET /api/openvpn/service/show``. + * **WireGuard** — ``GET /api/wireguard/service/show`` + """ + tunnels: Dict[str, Dict[str, Any]] = {} + + # ── IPsec ────────────────────────────────────────────────────────── + try: + data = self._get("/api/ipsec/sessions") + # OPNsense 24.x returns {"response": [...]} or a bare list + sessions = data.get("response") if isinstance(data, dict) else data + if isinstance(sessions, list): + for session in sessions: + # Each session object may contain multiple child SAs; we + # report one entry per IKE peer. + remote_host = ( + session.get("remote-host") + or session.get("remote_host") + or session.get("remote-id") + or "" + ) + local_host = ( + session.get("local-host") + or session.get("local_host") + or session.get("local-id") + or "" + ) + name = ( + session.get("uniqueid") + or session.get("con-id") + or session.get("name") + or remote_host + or f"ipsec-{len(tunnels)}" + ) + state = str(session.get("state") or session.get("ikey-state") or "").lower() + is_up = state in ("established", "up", "installed") + uptime_raw = session.get("established") or session.get("uptime") or 0 + try: + uptime = int(uptime_raw) + except (ValueError, TypeError): + uptime = 0 + # Byte counters may live in child-sa entries + bytes_in = 0 + bytes_out = 0 + for child in session.get("child-sas", {}).values() if isinstance(session.get("child-sas"), dict) else []: + try: + bytes_in += int(child.get("bytes-in", 0) or 0) + bytes_out += int(child.get("bytes-out", 0) or 0) + except (ValueError, TypeError): + pass + description = ( + session.get("local-id") + or session.get("description") + or "" + ) + key = f"ipsec-{name}" + tunnels[key] = { + "type": "IPsec", + "local_endpoint": local_host, + "remote_endpoint": remote_host, + "is_up": is_up, + "uptime": uptime, + "bytes_in": bytes_in, + "bytes_out": bytes_out, + "description": description, + } + except Exception: + # Try legacy Phase-2 leases endpoint (OPNsense < 23.x) + try: + data = self._post( + "/api/ipsec/leases/searchPhase2", + {"current": 1, "rowCount": -1, "searchPhrase": "", "sort": {}}, + ) + for row in data.get("rows", []): + name = row.get("id") or row.get("con") or f"ipsec-{len(tunnels)}" + key = f"ipsec-{name}" + state = str(row.get("state") or "").lower() + is_up = state in ("established", "installed", "up") + tunnels[key] = { + "type": "IPsec", + "local_endpoint": row.get("local-ts", ""), + "remote_endpoint": row.get("remote-ts", ""), + "is_up": is_up, + "uptime": 0, + "bytes_in": int(row.get("bytes-in", 0) or 0), + "bytes_out": int(row.get("bytes-out", 0) or 0), + "description": row.get("con", ""), + } + except Exception: + pass + + # ── OpenVPN ──────────────────────────────────────────────────────── + try: + data = self._get("/api/openvpn/instances/search") + instances = data.get("rows", []) + # Fetch live service status to determine up/down state + try: + show = self._get("/api/openvpn/service/show") + # show is a dict of {instance_id: {status, ...}} + status_map: Dict[str, Any] = show if isinstance(show, dict) else {} + except Exception: + status_map = {} + for inst in instances: + iid = inst.get("id") or inst.get("vpnid") or f"ovpn-{len(tunnels)}" + name = inst.get("description") or inst.get("dev") or str(iid) + key = f"openvpn-{iid}" + svc = status_map.get(str(iid), {}) + is_up = str(svc.get("running", False)).lower() in ("true", "1", "yes") + local_addr = inst.get("local") or inst.get("interface") or "" + remote_addr = inst.get("server") or inst.get("remote") or "" + tunnels[key] = { + "type": "SSL", + "local_endpoint": local_addr, + "remote_endpoint": remote_addr, + "is_up": is_up, + "uptime": 0, + "bytes_in": 0, + "bytes_out": 0, + "description": name, + } + except Exception: + pass + + # ── WireGuard ────────────────────────────────────────────────────── + try: + # /api/wireguard/service/show returns structured JSON: + # {"total": N, "rows": [ + # {"if":"wg0", "type":"interface", ...}, + # {"if":"wg0", "type":"peer", "public-key":"...", + # "endpoint":"1.2.3.4:51820", "transfer-rx":N, "transfer-tx":N, + # "name":"HCQ", "peer-status":"online", + # "latest-handshake-age":45, "ifname":"HCQ"}, ... + # ]} + data = self._get("/api/wireguard/service/show") + rows = data.get("rows", []) if isinstance(data, dict) else [] + + wg_idx = 0 + for row in rows: + if row.get("type") != "peer": + continue + wg_idx += 1 + + endpoint = (row.get("endpoint") or "").strip() + if endpoint and endpoint != "(none)": + remote_ip = endpoint.rsplit(":", 1)[0].strip("[]") + else: + remote_ip = "" + + is_up = row.get("peer-status") == "online" + bytes_in = int(row.get("transfer-rx") or 0) + bytes_out = int(row.get("transfer-tx") or 0) + hs_age = row.get("latest-handshake-age") + uptime = int(hs_age) if hs_age else 0 + + description = ( + (row.get("name") or "").strip() + or (row.get("ifname") or "").strip() + or f"wg-peer-{wg_idx}" + ) + + tunnels[f"wireguard-{wg_idx}"] = { + "type": "WireGuard", + "local_endpoint": "", + "remote_endpoint": remote_ip, + "is_up": is_up, + "uptime": uptime, + "bytes_in": bytes_in, + "bytes_out": bytes_out, + "description": description, + } + except Exception: + pass + + return tunnels + + def get_available_updates(self) -> List[Dict[str, Any]]: + """Return available firmware and package updates. + + Triggers an async update-check on OPNsense via + ``POST /api/core/firmware/check``, then polls + ``GET /api/core/firmware/status`` for up to 15 seconds. + Returns a list of ``{name, current_version, new_version}`` dicts, + or an empty list when everything is up to date or the check has + not yet finished. + """ + import time + try: + self._post("/api/core/firmware/check") + except Exception: + pass + + for _ in range(5): + time.sleep(3) + try: + status = self._get("/api/core/firmware/status") + state = status.get("status", "none") + if state in ("update", "upgrade"): + updates = status.get("updates") or [] + return [ + { + "name": u.get("name", ""), + "current_version": u.get("current_version", u.get("version", "")), + "new_version": u.get("new_version", u.get("version", "")), + } + for u in updates + ] + if state == "latest": + return [] + except Exception: + pass + return [] + + def apply_updates(self, packages: List[str]) -> Dict[str, Any]: + """Trigger a full firmware upgrade on OPNsense. + + Note: OPNsense upgrades the entire system at once rather than + individual packages. The ``packages`` parameter is accepted for + API compatibility but is ignored — the full upgrade is always + applied. + + Calls ``POST /api/core/firmware/upgrade``. + """ + try: + result = self._post("/api/core/firmware/upgrade") + return {"success": True, "output": str(result)} + except Exception as exc: + return {"success": False, "output": str(exc)} + + def search_packages(self, query: str) -> List[Dict[str, Any]]: + """Search available OPNsense plugins by name or description. + + Filters the full plugin list from ``GET /api/core/firmware/info`` + against *query* (case-insensitive substring match on name and + comment). Returns all matching plugins with an ``installed`` + flag so the UI can show which ones are already active. + """ + q = query.lower() + info = self._get("/api/core/firmware/info") + result: List[Dict[str, Any]] = [] + for p in info.get("plugin", []): + name = p.get("name", "") + comment = p.get("comment", "") + if q in name.lower() or q in comment.lower(): + result.append({ + "name": name, + "version": p.get("version", ""), + "installed": p.get("installed") == "1", + "description": comment, + "size": 0, + "source": "opnsense-plugins", + }) + return sorted(result, key=lambda x: x["name"].lower()) + + def install_package(self, name: str) -> Dict[str, Any]: + """Install an OPNsense plugin by name. + + Calls ``POST /api/core/firmware/install/{name}``. + """ + import re as _re + if not _re.match(r'^[a-zA-Z0-9_\-\.]+$', name): + raise ValueError(f"Invalid package name: {name!r}") + try: + result = self._post(f"/api/core/firmware/install/{name}") + return {"success": True, "output": str(result)} + except Exception as exc: + return {"success": False, "output": str(exc)} + + def uninstall_package(self, name: str) -> Dict[str, Any]: + """Remove an OPNsense plugin by name. + + Calls ``POST /api/core/firmware/remove/{name}``. + """ + import re as _re + if not _re.match(r'^[a-zA-Z0-9_\-\.]+$', name): + raise ValueError(f"Invalid package name: {name!r}") + try: + result = self._post(f"/api/core/firmware/remove/{name}") + return {"success": True, "output": str(result)} + except Exception as exc: + return {"success": False, "output": str(exc)} diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..dfd351c --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,55 @@ +[build-system] +requires = ["setuptools>=68", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "napalm-opnsense" +version = "0.1.0" +description = "NAPALM driver for OPNsense (read-only via REST API)." +readme = "README.md" +license = { text = "Apache-2.0" } +requires-python = ">=3.9" +authors = [ + { name = "Christian Manivong" }, +] +classifiers = [ + "Topic :: Utilities", + "License :: OSI Approved :: Apache Software License", + "Programming Language :: Python :: 3", + "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.2.0", + "requests>=2.28.0", +] + +[project.optional-dependencies] +dev = [ + "pytest", + "pytest-cov", + "black", + "ruff", +] + +[project.entry-points."napalm.drivers"] +opnsense = "napalm_opnsense.opnsense:OPNsenseDriver" + +[project.urls] +Repository = "https://github.com/napalm-automation-community/napalm-opnsense" + +[tool.setuptools.packages.find] +where = ["."] +include = ["napalm_opnsense*"] + +[tool.ruff] +line-length = 100 +target-version = "py39" + +[tool.pytest.ini_options] +testpaths = ["tests"] 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..c928af1 --- /dev/null +++ b/tests/unit/test_driver.py @@ -0,0 +1,1245 @@ +"""Unit tests for OPNsenseDriver — no real device required.""" + +import json +import pytest +from unittest.mock import MagicMock, patch + +from napalm_opnsense.opnsense import OPNsenseDriver +from napalm.base.exceptions import ConnectionException, ConnectionClosedException + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture +def driver(): + """Return a driver instance with a mocked requests.Session.""" + with patch("napalm_opnsense.opnsense.requests.Session"): + drv = OPNsenseDriver( + hostname="opnsense.example.com", + username="api_key", + password="api_secret", + optional_args={"verify": False}, + ) + drv.session = MagicMock() + yield drv + + +# --------------------------------------------------------------------------- +# Sample API responses +# --------------------------------------------------------------------------- + +STATUS_RESPONSE = { + "hostname": "opnsense01", + "version": "24.7", + "model": "OPNsense", + "serial": "ABC123", + "uptime": 12345, +} + +INTERFACES_RESPONSE = { + "interfaces": [ + { + "name": "em0", + "up": True, + "enabled": True, + "descr": "LAN", + "mac": "AA:BB:CC:DD:EE:FF", + "speed_mbps": 1000, + "mtu": 1500, + }, + { + "name": "em1", + "up": False, + "enabled": True, + "descr": "WAN", + "mac": "AA:BB:CC:DD:EE:00", + "speed_mbps": None, + "mtu": 0, + }, + ] +} + +ADDRESSES_RESPONSE = { + "items": [ + {"interface": "em0", "address": "192.0.2.10", "prefix": 24}, + {"interface": "em0", "address": "2001:db8::1", "prefix": 64}, + {"interface": "em1", "address": "203.0.113.5", "prefix": 30}, + ] +} + +ARP_RESPONSE = { + "arp": [ + {"intf": "em0", "mac": "AA:BB:CC:DD:EE:01", "ip": "192.0.2.1", "expires": 900}, + {"intf": "em0", "mac": "AA:BB:CC:DD:EE:02", "ip": "192.0.2.2", "expires": 600}, + ] +} + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_json_response(data): + mock_resp = MagicMock() + mock_resp.json.return_value = data + mock_resp.raise_for_status.return_value = None + return mock_resp + + +# --------------------------------------------------------------------------- +# open() / close() / is_alive() +# --------------------------------------------------------------------------- + +class TestOpenClose: + def test_open_raises_connection_exception_on_error(self): + drv = OPNsenseDriver( + hostname="unreachable.invalid", + username="k", + password="s", + optional_args={"verify": False}, + ) + with pytest.raises(ConnectionException): + drv.open() + + def test_close_clears_session(self, driver): + driver.close() + assert driver.session is None + + def test_close_is_idempotent(self, driver): + driver.close() + driver.close() # second call must not raise + + +class TestIsAlive: + def test_returns_false_when_no_session(self): + drv = OPNsenseDriver("host", "u", "p") + assert drv.is_alive() == {"is_alive": False} + + def test_returns_true_on_successful_connection(self, driver): + with patch("napalm_opnsense.opnsense.socket.create_connection") as mock_conn: + mock_conn.return_value.__enter__ = MagicMock(return_value=None) + mock_conn.return_value.__exit__ = MagicMock(return_value=False) + result = driver.is_alive() + assert result == {"is_alive": True} + + def test_returns_false_on_socket_error(self, driver): + with patch( + "napalm_opnsense.opnsense.socket.create_connection", + side_effect=OSError("refused"), + ): + result = driver.is_alive() + assert result == {"is_alive": False} + + +# --------------------------------------------------------------------------- +# _get() +# --------------------------------------------------------------------------- + +class TestInternalGet: + def test_raises_when_no_session(self): + drv = OPNsenseDriver("host", "u", "p") + with pytest.raises(ConnectionClosedException): + drv._get("/api/core/system/status") + + def test_calls_correct_url(self, driver): + driver.session.get.return_value = _make_json_response(STATUS_RESPONSE) + driver._get("/api/core/system/status") + driver.session.get.assert_called_once_with( + "https://opnsense.example.com/api/core/system/status", + timeout=60, + ) + + +# --------------------------------------------------------------------------- +# get_facts() +# --------------------------------------------------------------------------- + +class TestGetFacts: + def test_returns_required_keys(self, driver): + driver._get = lambda path: ( + STATUS_RESPONSE if "status" in path else INTERFACES_RESPONSE + ) + facts = driver.get_facts() + for key in ("vendor", "model", "hostname", "fqdn", "os_version", + "serial_number", "uptime", "interface_list"): + assert key in facts + + def test_vendor_constant(self, driver): + driver._get = lambda path: ( + STATUS_RESPONSE if "status" in path else INTERFACES_RESPONSE + ) + facts = driver.get_facts() + assert facts["vendor"] == "OPNsense" + + def test_hostname_parsed(self, driver): + driver._get = lambda path: ( + STATUS_RESPONSE if "status" in path else INTERFACES_RESPONSE + ) + facts = driver.get_facts() + assert facts["hostname"] == "opnsense01" + + def test_os_version_parsed(self, driver): + driver._get = lambda path: ( + STATUS_RESPONSE if "status" in path else INTERFACES_RESPONSE + ) + facts = driver.get_facts() + assert facts["os_version"] == "24.7" + + def test_interface_list_populated(self, driver): + driver._get = lambda path: ( + STATUS_RESPONSE if "status" in path else INTERFACES_RESPONSE + ) + facts = driver.get_facts() + assert "em0" in facts["interface_list"] + assert "em1" in facts["interface_list"] + + def test_interface_list_empty_on_getter_failure(self, driver): + def fail_on_interfaces(path): + if "overview" in path: + raise RuntimeError("no endpoint") + return STATUS_RESPONSE + + driver._get = fail_on_interfaces + facts = driver.get_facts() + assert facts["interface_list"] == [] + + def test_serial_number(self, driver): + driver._get = lambda path: ( + STATUS_RESPONSE if "status" in path else INTERFACES_RESPONSE + ) + facts = driver.get_facts() + assert facts["serial_number"] == "ABC123" + + def test_uptime(self, driver): + driver._get = lambda path: ( + STATUS_RESPONSE if "status" in path else INTERFACES_RESPONSE + ) + facts = driver.get_facts() + assert facts["uptime"] == 12345 + + +# --------------------------------------------------------------------------- +# get_interfaces() +# --------------------------------------------------------------------------- + +class TestGetInterfaces: + def test_interface_count(self, driver): + driver._get = lambda path: INTERFACES_RESPONSE + ifaces = driver.get_interfaces() + assert len(ifaces) == 2 + + def test_is_up_and_enabled(self, driver): + driver._get = lambda path: INTERFACES_RESPONSE + ifaces = driver.get_interfaces() + assert ifaces["em0"]["is_up"] is True + assert ifaces["em0"]["is_enabled"] is True + assert ifaces["em1"]["is_up"] is False + + def test_speed(self, driver): + driver._get = lambda path: INTERFACES_RESPONSE + ifaces = driver.get_interfaces() + assert ifaces["em0"]["speed"] == 1000.0 + assert ifaces["em1"]["speed"] == 0.0 + + def test_mac_address_lowercase(self, driver): + driver._get = lambda path: INTERFACES_RESPONSE + ifaces = driver.get_interfaces() + assert ifaces["em0"]["mac_address"] == "aa:bb:cc:dd:ee:ff" + + def test_description(self, driver): + driver._get = lambda path: INTERFACES_RESPONSE + ifaces = driver.get_interfaces() + assert ifaces["em0"]["description"] == "LAN" + assert ifaces["em1"]["description"] == "WAN" + + def test_last_flapped_is_negative_one(self, driver): + driver._get = lambda path: INTERFACES_RESPONSE + ifaces = driver.get_interfaces() + assert ifaces["em0"]["last_flapped"] == -1.0 + + def test_empty_response(self, driver): + driver._get = lambda path: {"interfaces": []} + assert driver.get_interfaces() == {} + + +# --------------------------------------------------------------------------- +# get_interfaces_ip() +# --------------------------------------------------------------------------- + +class TestGetInterfacesIp: + def test_entry_count(self, driver): + driver._get = lambda path: ADDRESSES_RESPONSE + result = driver.get_interfaces_ip() + # em0 has 2 addresses, em1 has 1 + assert len(result) == 2 + + def test_ipv4_entry(self, driver): + driver._get = lambda path: ADDRESSES_RESPONSE + result = driver.get_interfaces_ip() + assert "192.0.2.10" in result["em0"]["ipv4"] + assert result["em0"]["ipv4"]["192.0.2.10"]["prefix_length"] == 24 + + def test_ipv6_entry(self, driver): + driver._get = lambda path: ADDRESSES_RESPONSE + result = driver.get_interfaces_ip() + assert "2001:db8::1" in result["em0"]["ipv6"] + assert result["em0"]["ipv6"]["2001:db8::1"]["prefix_length"] == 64 + + def test_second_interface(self, driver): + driver._get = lambda path: ADDRESSES_RESPONSE + result = driver.get_interfaces_ip() + assert "203.0.113.5" in result["em1"]["ipv4"] + + +# --------------------------------------------------------------------------- +# get_arp_table() +# --------------------------------------------------------------------------- + +class TestGetArpTable: + def test_entry_count(self, driver): + driver._get = lambda path: ARP_RESPONSE + table = driver.get_arp_table() + assert len(table) == 2 + + def test_entry_structure(self, driver): + driver._get = lambda path: ARP_RESPONSE + entry = driver.get_arp_table()[0] + for key in ("interface", "mac", "ip", "age"): + assert key in entry + + def test_ip_values(self, driver): + driver._get = lambda path: ARP_RESPONSE + ips = {e["ip"] for e in driver.get_arp_table()} + assert "192.0.2.1" in ips + assert "192.0.2.2" in ips + + def test_mac_lowercase(self, driver): + driver._get = lambda path: ARP_RESPONSE + macs = {e["mac"] for e in driver.get_arp_table()} + assert all(m == m.lower() for m in macs) + + def test_list_response_format(self, driver): + """ARP endpoint may return a bare list instead of dict.""" + bare_list = ARP_RESPONSE["arp"] + driver._get = lambda path: bare_list + table = driver.get_arp_table() + assert len(table) == 2 + + +# --------------------------------------------------------------------------- +# get_interfaces_counters() +# --------------------------------------------------------------------------- + +INTERFACE_STATISTICS_RESPONSE = { + "statistics": { + "em0": { + "input-packets": 10000, + "output-packets": 8000, + "input-bytes": 1000000, + "output-bytes": 800000, + "input-errors": 5, + "output-errors": 2, + "input-drops": 3, + "output-drops": 1, + "input-multicasts": 100, + "output-multicasts": 50, + "input-broadcasts": 20, + "output-broadcasts": 10, + }, + "em1": { + "input-packets": 500, + "output-packets": 300, + "input-bytes": 50000, + "output-bytes": 30000, + "input-errors": 0, + "output-errors": 0, + "input-drops": 0, + "output-drops": 0, + "input-multicasts": 0, + "output-multicasts": 0, + "input-broadcasts": 0, + "output-broadcasts": 0, + }, + } +} + + +class TestGetInterfacesCounters: + def test_interface_count(self, driver): + driver._get = lambda path: INTERFACE_STATISTICS_RESPONSE + counters = driver.get_interfaces_counters() + assert len(counters) == 2 + + def test_required_keys(self, driver): + driver._get = lambda path: INTERFACE_STATISTICS_RESPONSE + entry = driver.get_interfaces_counters()["em0"] + for key in ( + "tx_errors", "rx_errors", "tx_discards", "rx_discards", + "tx_octets", "rx_octets", + "tx_unicast_packets", "rx_unicast_packets", + "tx_multicast_packets", "rx_multicast_packets", + "tx_broadcast_packets", "rx_broadcast_packets", + ): + assert key in entry + + def test_rx_octets(self, driver): + driver._get = lambda path: INTERFACE_STATISTICS_RESPONSE + assert driver.get_interfaces_counters()["em0"]["rx_octets"] == 1000000 + + def test_tx_errors(self, driver): + driver._get = lambda path: INTERFACE_STATISTICS_RESPONSE + assert driver.get_interfaces_counters()["em0"]["tx_errors"] == 2 + + def test_rx_errors(self, driver): + driver._get = lambda path: INTERFACE_STATISTICS_RESPONSE + assert driver.get_interfaces_counters()["em0"]["rx_errors"] == 5 + + def test_zero_counters(self, driver): + driver._get = lambda path: INTERFACE_STATISTICS_RESPONSE + em1 = driver.get_interfaces_counters()["em1"] + assert em1["tx_errors"] == 0 + assert em1["rx_errors"] == 0 + + def test_empty_statistics(self, driver): + driver._get = lambda path: {"statistics": {}} + assert driver.get_interfaces_counters() == {} + + +# --------------------------------------------------------------------------- +# get_environment() +# --------------------------------------------------------------------------- + +SYSTEM_RESOURCES_RESPONSE = { + "cpu": {"used": "25"}, + "memory": {"total": "4096000000", "used": "2048000000"}, +} + +SYSTEM_TEMP_RESPONSE = { + "data": [ + {"device": "cpu0", "temperature": "52.5"}, + {"device": "cpu1", "temperature": "48.0"}, + ] +} + + +class TestGetEnvironment: + def test_required_top_keys(self, driver): + def fake_get(path): + if "temperature" in path: + return SYSTEM_TEMP_RESPONSE + return SYSTEM_RESOURCES_RESPONSE + + driver._get = fake_get + env = driver.get_environment() + for key in ("fans", "temperature", "power", "cpu", "memory"): + assert key in env + + def test_cpu_usage(self, driver): + driver._get = lambda path: ( + SYSTEM_TEMP_RESPONSE if "temperature" in path else SYSTEM_RESOURCES_RESPONSE + ) + env = driver.get_environment() + assert env["cpu"][0]["%usage"] == 25.0 + + def test_memory_values(self, driver): + driver._get = lambda path: ( + SYSTEM_TEMP_RESPONSE if "temperature" in path else SYSTEM_RESOURCES_RESPONSE + ) + env = driver.get_environment() + assert env["memory"]["used_ram"] == 2048000000 + assert env["memory"]["available_ram"] == 2048000000 + + def test_temperature_sensors(self, driver): + driver._get = lambda path: ( + SYSTEM_TEMP_RESPONSE if "temperature" in path else SYSTEM_RESOURCES_RESPONSE + ) + env = driver.get_environment() + assert "cpu0" in env["temperature"] + assert env["temperature"]["cpu0"]["temperature"] == 52.5 + + def test_temperature_alert_thresholds(self, driver): + driver._get = lambda path: ( + {"data": [{"device": "cpu0", "temperature": "85.0"}]} + if "temperature" in path else SYSTEM_RESOURCES_RESPONSE + ) + env = driver.get_environment() + assert env["temperature"]["cpu0"]["is_alert"] is True + assert env["temperature"]["cpu0"]["is_critical"] is False + + def test_temperature_endpoint_failing_gracefully(self, driver): + """Driver must not raise if temperature endpoint is unavailable.""" + def fake_get(path): + if "temperature" in path: + raise Exception("no sensor data") + return SYSTEM_RESOURCES_RESPONSE + + driver._get = fake_get + env = driver.get_environment() + assert env["temperature"] == {} + + +# --------------------------------------------------------------------------- +# get_route_to() +# --------------------------------------------------------------------------- + +ROUTES_RESPONSE = { + "route": [ + { + "network": "0.0.0.0/0", + "gateway": "192.0.2.1", + "flags": "UGS", + "netif": "em1", + "proto": "static", + "priority": 1, + }, + { + "network": "192.0.2.0/24", + "gateway": "", + "flags": "U", + "netif": "em0", + "proto": "kernel", + "priority": 0, + }, + { + "network": "198.51.100.0/24", + "gateway": "192.0.2.5", + "flags": "UGS", + "netif": "em0", + "proto": "static", + "priority": 1, + }, + ] +} + + +class TestGetRouteTo: + def test_returns_all_routes_without_filter(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + routes = driver.get_route_to() + assert len(routes) == 3 + + def test_default_route_present(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + assert "0.0.0.0/0" in driver.get_route_to() + + def test_next_hop(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + entry = driver.get_route_to()["0.0.0.0/0"][0] + assert entry["next_hop"] == "192.0.2.1" + + def test_outgoing_interface(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + entry = driver.get_route_to()["0.0.0.0/0"][0] + assert entry["outgoing_interface"] == "em1" + + def test_protocol_static(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + assert driver.get_route_to()["0.0.0.0/0"][0]["protocol"] == "static" + + def test_protocol_kernel_mapped_to_connected(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + assert driver.get_route_to()["192.0.2.0/24"][0]["protocol"] == "connected" + + def test_filter_by_destination(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + routes = driver.get_route_to(destination="0.0.0.0/0") + assert "0.0.0.0/0" in routes + assert "192.0.2.0/24" not in routes + + def test_filter_by_protocol(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + routes = driver.get_route_to(protocol="static") + assert all( + e["protocol"] == "static" + for entries in routes.values() + for e in entries + ) + + def test_required_keys_in_entry(self, driver): + driver._get = lambda path: ROUTES_RESPONSE + entry = driver.get_route_to()["0.0.0.0/0"][0] + for key in ( + "protocol", "current_active", "last_active", "age", + "next_hop", "outgoing_interface", "selected_next_hop", + "preference", "inactive_reason", "routing_table", + "protocol_attributes", + ): + assert key in entry + + +# --------------------------------------------------------------------------- +# get_ipv6_neighbors_table() +# --------------------------------------------------------------------------- + +NDP_RESPONSE = { + "rows": [ + { + "intf": "em0", + "mac": "aa:bb:cc:dd:ee:01", + "ip": "fe80::1", + "expires": 120, + "state": "REACHABLE", + }, + { + "intf": "em0", + "mac": "aa:bb:cc:dd:ee:02", + "ip": "2001:db8::1", + "expires": 60, + "state": "STALE", + }, + ] +} + + +class TestGetIpv6NeighborsTable: + def test_entry_count(self, driver): + driver._get = lambda path: NDP_RESPONSE + assert len(driver.get_ipv6_neighbors_table()) == 2 + + def test_required_keys(self, driver): + driver._get = lambda path: NDP_RESPONSE + entry = driver.get_ipv6_neighbors_table()[0] + for key in ("interface", "mac", "ip", "age", "state"): + assert key in entry + + def test_ip_values(self, driver): + driver._get = lambda path: NDP_RESPONSE + ips = {e["ip"] for e in driver.get_ipv6_neighbors_table()} + assert "fe80::1" in ips + assert "2001:db8::1" in ips + + def test_mac_lowercase(self, driver): + driver._get = lambda path: NDP_RESPONSE + macs = {e["mac"] for e in driver.get_ipv6_neighbors_table()} + assert all(m == m.lower() for m in macs) + + def test_list_format_response(self, driver): + """Endpoint may return a bare list.""" + driver._get = lambda path: NDP_RESPONSE["rows"] + assert len(driver.get_ipv6_neighbors_table()) == 2 + + def test_empty_response(self, driver): + driver._get = lambda path: {"rows": []} + assert driver.get_ipv6_neighbors_table() == [] + + +# --------------------------------------------------------------------------- +# get_lldp_neighbors() / get_lldp_neighbors_detail() +# --------------------------------------------------------------------------- + +LLDP_RESPONSE = { + "rows": [ + { + "local_port": "em0", + "port_id": "eth1", + "chassis_id": "aa:bb:cc:dd:ee:ff", + "system_name": "core-sw-01", + "port_description": "uplink", + "system_description": "Cisco IOS", + "system_capabilities": "bridge, router", + "enabled_capabilities": "bridge", + } + ] +} + + +class TestGetLldpNeighbors: + def test_returns_neighbor(self, driver): + driver._get = lambda path: LLDP_RESPONSE + neighbors = driver.get_lldp_neighbors() + assert "em0" in neighbors + assert neighbors["em0"][0]["hostname"] == "core-sw-01" + assert neighbors["em0"][0]["port"] == "eth1" + + def test_plugin_not_installed_returns_empty(self, driver): + driver._get = lambda path: (_ for _ in ()).throw(Exception("404")) + assert driver.get_lldp_neighbors() == {} + + +class TestGetLldpNeighborsDetail: + def test_required_keys(self, driver): + driver._get = lambda path: LLDP_RESPONSE + detail = driver.get_lldp_neighbors_detail() + entry = detail["em0"][0] + for key in ( + "remote_chassis_id", "remote_system_name", "remote_port", + "remote_port_description", "remote_system_description", + "remote_system_capab", "remote_system_enable_capab", + ): + assert key in entry + + def test_chassis_id(self, driver): + driver._get = lambda path: LLDP_RESPONSE + assert driver.get_lldp_neighbors_detail()["em0"][0]["remote_chassis_id"] == "aa:bb:cc:dd:ee:ff" + + def test_system_name(self, driver): + driver._get = lambda path: LLDP_RESPONSE + assert driver.get_lldp_neighbors_detail()["em0"][0]["remote_system_name"] == "core-sw-01" + + def test_capabilities_parsed(self, driver): + driver._get = lambda path: LLDP_RESPONSE + capab = driver.get_lldp_neighbors_detail()["em0"][0]["remote_system_capab"] + assert "bridge" in capab + assert "router" in capab + + def test_interface_filter(self, driver): + driver._get = lambda path: LLDP_RESPONSE + detail = driver.get_lldp_neighbors_detail(interface="em99") + assert detail == {} + + def test_plugin_not_installed_returns_empty(self, driver): + driver._get = lambda path: (_ for _ in ()).throw(Exception("404")) + assert driver.get_lldp_neighbors_detail() == {} + + +# --------------------------------------------------------------------------- +# get_ntp_servers() +# --------------------------------------------------------------------------- + +NTP_STATUS_RESPONSE = { + "peers": [ + {"address": "pool.ntp.org", "state": "synced"}, + {"address": "time.cloudflare.com", "state": "candidate"}, + ] +} + + +class TestGetNtpServers: + def test_entry_count(self, driver): + driver._get = lambda path: NTP_STATUS_RESPONSE + servers = driver.get_ntp_servers() + assert len(servers) == 2 + + def test_server_addresses(self, driver): + driver._get = lambda path: NTP_STATUS_RESPONSE + servers = driver.get_ntp_servers() + assert "pool.ntp.org" in servers + assert "time.cloudflare.com" in servers + + def test_value_is_empty_dict(self, driver): + driver._get = lambda path: NTP_STATUS_RESPONSE + for v in driver.get_ntp_servers().values(): + assert v == {} + + def test_endpoint_failure_returns_empty(self, driver): + driver._get = lambda path: (_ for _ in ()).throw(Exception("service not running")) + assert driver.get_ntp_servers() == {} + + def test_empty_peers(self, driver): + driver._get = lambda path: {"peers": []} + assert driver.get_ntp_servers() == {} + + +# --------------------------------------------------------------------------- +# get_vlans() +# --------------------------------------------------------------------------- + +VLAN_SEARCH_RESPONSE = { + "rows": [ + {"tag": "10", "vlanif": "em0_vlan10", "if": "em0", "descr": "Management", "pcp": "0"}, + {"tag": "20", "vlanif": "em0_vlan20 [LAN]", "if": "em0", "descr": "", "pcp": "0"}, + {"tag": "100", "vlanif": "em1_vlan100", "if": "em1", "descr": "Guest WiFi", "pcp": "0"}, + ], + "rowCount": 3, + "total": 3, + "current": 1, +} + + +class TestGetVlans: + def test_returns_all_vlans(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + assert len(driver.get_vlans()) == 3 + + def test_keyed_by_tag_string(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + vlans = driver.get_vlans() + assert "10" in vlans + assert "20" in vlans + assert "100" in vlans + + def test_name_from_descr(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + assert driver.get_vlans()["10"]["name"] == "Management" + + def test_name_falls_back_to_vlanif_when_no_descr(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + # tag 20 has no descr — should use stripped vlanif + assert driver.get_vlans()["20"]["name"] == "em0_vlan20" + + def test_interface_in_list(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + assert driver.get_vlans()["10"]["interfaces"] == ["em0_vlan10"] + + def test_bracket_annotation_stripped_from_vlanif(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + # "em0_vlan20 [LAN]" must be stored as "em0_vlan20" + assert driver.get_vlans()["20"]["interfaces"] == ["em0_vlan20"] + + def test_required_keys_present(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + for vlan in driver.get_vlans().values(): + assert "name" in vlan + assert "interfaces" in vlan + + def test_empty_response_returns_empty_dict(self, driver): + driver._get = lambda path: {"rows": [], "rowCount": 0, "total": 0} + assert driver.get_vlans() == {} + + def test_multiple_vlans_on_different_parents(self, driver): + driver._get = lambda path: VLAN_SEARCH_RESPONSE + vlans = driver.get_vlans() + assert vlans["100"]["interfaces"] == ["em1_vlan100"] + + +# --------------------------------------------------------------------------- +# get_bgp_neighbors() +# --------------------------------------------------------------------------- + +BGP_CFG_RESPONSE = { + "bgp": { + "asnumber": "65000", + "routerid": "1.2.3.4", + "enabled": "1", + } +} + +BGP_NEIGHBORS_RESPONSE = { + "response": { + "10.0.0.1": { + "remoteAs": 65001, + "localAs": 65000, + "nbrDesc": "upstream-peer", + "bgpState": "Established", + "bgpTimerUpMsec": 3723000, + "remoteRouterId": "10.0.0.1", + "adminShutdown": False, + "addressFamilyInfo": { + "ipv4Unicast": { + "sentPrefixCounter": 5, + "prefixReceivedCount": 20, + "acceptedPrefixCounter": 18, + } + }, + }, + "10.0.0.2": { + "remoteAs": 65002, + "localAs": 65000, + "nbrDesc": "", + "bgpState": "Active", + "bgpTimerUpMsec": 0, + "remoteRouterId": "", + "adminShutdown": True, + "addressFamilyInfo": {}, + }, + } +} + + +class TestGetBgpNeighbors: + def _fake_get(self, path): + if "diagnostics/bgpneighbors" in path: + return BGP_NEIGHBORS_RESPONSE + return BGP_CFG_RESPONSE + + def test_returns_global_vrf(self, driver): + driver._get = self._fake_get + result = driver.get_bgp_neighbors() + assert "global" in result + + def test_router_id(self, driver): + driver._get = self._fake_get + assert driver.get_bgp_neighbors()["global"]["router_id"] == "1.2.3.4" + + def test_peer_count(self, driver): + driver._get = self._fake_get + peers = driver.get_bgp_neighbors()["global"]["peers"] + assert len(peers) == 2 + + def test_established_peer_is_up(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + assert peer["is_up"] is True + + def test_active_peer_is_not_up(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.2"] + assert peer["is_up"] is False + + def test_admin_shutdown_peer_is_disabled(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.2"] + assert peer["is_enabled"] is False + + def test_non_shutdown_peer_is_enabled(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + assert peer["is_enabled"] is True + + def test_uptime_established(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + assert peer["uptime"] == 3723 # 3723000 ms → 3723 s + + def test_uptime_not_established(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.2"] + assert peer["uptime"] == -1 + + def test_remote_as(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + assert peer["remote_as"] == 65001 + + def test_local_as(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + assert peer["local_as"] == 65000 + + def test_remote_id(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + assert peer["remote_id"] == "10.0.0.1" + + def test_description(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + assert peer["description"] == "upstream-peer" + + def test_ipv4_prefix_counters(self, driver): + driver._get = self._fake_get + af = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"]["address_family"]["ipv4"] + assert af["sent_prefixes"] == 5 + assert af["received_prefixes"] == 20 + assert af["accepted_prefixes"] == 18 + + def test_no_af_info_falls_back_to_ipv4_minus_one(self, driver): + driver._get = self._fake_get + af = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.2"]["address_family"] + assert "ipv4" in af + assert af["ipv4"]["sent_prefixes"] == -1 + + def test_ipv6_af_populated_when_present(self, driver): + def fake_get(path): + if "diagnostics/bgpneighbors" in path: + return { + "response": { + "2001:db8::1": { + "remoteAs": 65010, + "localAs": 65000, + "nbrDesc": "", + "bgpState": "Established", + "bgpTimerUpMsec": 1000, + "remoteRouterId": "2001:db8::1", + "adminShutdown": False, + "addressFamilyInfo": { + "ipv6Unicast": { + "sentPrefixCounter": 3, + "prefixReceivedCount": 7, + "acceptedPrefixCounter": 7, + } + }, + } + } + } + return BGP_CFG_RESPONSE + + driver._get = fake_get + af = driver.get_bgp_neighbors()["global"]["peers"]["2001:db8::1"]["address_family"] + assert "ipv6" in af + assert af["ipv6"]["sent_prefixes"] == 3 + + def test_plugin_absent_returns_empty(self, driver): + driver._get = lambda path: (_ for _ in ()).throw(Exception("404")) + assert driver.get_bgp_neighbors() == {} + + def test_frr_not_running_returns_empty(self, driver): + def fake_get(path): + if "diagnostics" in path: + return {"response": "error"} # non-dict response + return BGP_CFG_RESPONSE + + driver._get = fake_get + assert driver.get_bgp_neighbors() == {} + + def test_required_peer_keys(self, driver): + driver._get = self._fake_get + peer = driver.get_bgp_neighbors()["global"]["peers"]["10.0.0.1"] + for key in ("local_as", "remote_as", "remote_id", "is_up", "is_enabled", + "description", "uptime", "address_family"): + assert key in peer + + +# --------------------------------------------------------------------------- +# _post() +# --------------------------------------------------------------------------- + +class TestInternalPost: + def test_raises_when_no_session(self): + drv = OPNsenseDriver("host", "u", "p") + with pytest.raises(ConnectionClosedException): + drv._post("/api/routes/routes/reconfigure") + + def test_calls_correct_url(self, driver): + driver.session.post.return_value = _make_json_response({"result": "ok"}) + driver._post("/api/routes/routes/reconfigure") + driver.session.post.assert_called_once_with( + "https://opnsense.example.com/api/routes/routes/reconfigure", + json={}, + timeout=60, + ) + + def test_sends_json_payload(self, driver): + driver.session.post.return_value = _make_json_response({"uuid": "abc-123"}) + payload = {"route": {"network": "10.0.0.0/8", "gateway": "WAN_GW"}} + driver._post("/api/routes/routes/addroute", payload) + driver.session.post.assert_called_once_with( + "https://opnsense.example.com/api/routes/routes/addroute", + json=payload, + timeout=60, + ) + + def test_returns_parsed_json(self, driver): + driver.session.post.return_value = _make_json_response({"result": "saved"}) + result = driver._post("/api/routes/routes/reconfigure") + assert result == {"result": "saved"} + + +# --------------------------------------------------------------------------- +# load_merge_candidate() +# --------------------------------------------------------------------------- + +VALID_ROUTES_CONFIG = json.dumps([ + {"network": "10.0.0.0/8", "gateway": "WAN_GW", "descr": "internal"}, + {"network": "0.0.0.0/0", "gateway": "WAN_GW", "descr": "default"}, +]) + + +class TestLoadMergeCandidate: + def test_accepts_valid_json_string(self, driver): + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + assert driver._candidate_config is not None + assert len(driver._candidate_config) == 2 + + def test_parses_required_keys(self, driver): + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + first = driver._candidate_config[0] + assert first["network"] == "10.0.0.0/8" + assert first["gateway"] == "WAN_GW" + + def test_reads_from_file(self, driver, tmp_path): + cfg_file = tmp_path / "routes.json" + cfg_file.write_text(VALID_ROUTES_CONFIG) + driver.load_merge_candidate(filename=str(cfg_file)) + assert len(driver._candidate_config) == 2 + + def test_raises_if_both_args_given(self, driver): + from napalm.base.exceptions import MergeConfigException + with pytest.raises(MergeConfigException): + driver.load_merge_candidate(filename="f.json", config="{}") + + def test_raises_if_no_args_given(self, driver): + from napalm.base.exceptions import MergeConfigException + with pytest.raises(MergeConfigException): + driver.load_merge_candidate() + + def test_raises_on_invalid_json(self, driver): + from napalm.base.exceptions import MergeConfigException + with pytest.raises(MergeConfigException, match="Invalid JSON"): + driver.load_merge_candidate(config="not json {{{") + + def test_raises_if_not_a_list(self, driver): + from napalm.base.exceptions import MergeConfigException + with pytest.raises(MergeConfigException, match="array"): + driver.load_merge_candidate(config='{"network": "1.0.0.0/8"}') + + def test_raises_if_route_missing_network(self, driver): + from napalm.base.exceptions import MergeConfigException + bad = json.dumps([{"gateway": "GW1"}]) + with pytest.raises(MergeConfigException, match="'network'"): + driver.load_merge_candidate(config=bad) + + def test_raises_if_route_missing_gateway(self, driver): + from napalm.base.exceptions import MergeConfigException + bad = json.dumps([{"network": "10.0.0.0/8"}]) + with pytest.raises(MergeConfigException, match="'gateway'"): + driver.load_merge_candidate(config=bad) + + def test_raises_on_missing_file(self, driver): + from napalm.base.exceptions import MergeConfigException + with pytest.raises(MergeConfigException, match="Cannot read"): + driver.load_merge_candidate(filename="/nonexistent/path/routes.json") + + +# --------------------------------------------------------------------------- +# compare_config() +# --------------------------------------------------------------------------- + +SEARCH_ROUTE_RESPONSE = { + "rows": [ + {"network": "192.168.1.0/24", "gateway": "LAN_GW", "descr": "lan", "disabled": "0"}, + ] +} + + +class TestCompareConfig: + def test_returns_empty_string_without_candidate(self, driver): + assert driver.compare_config() == "" + + def test_returns_diff_string(self, driver): + driver._get = lambda path: SEARCH_ROUTE_RESPONSE + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + diff = driver.compare_config() + assert "---" in diff + assert "+++" in diff + + def test_diff_shows_added_routes(self, driver): + driver._get = lambda path: SEARCH_ROUTE_RESPONSE + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + diff = driver.compare_config() + assert "WAN_GW" in diff + + def test_empty_diff_when_config_matches(self, driver): + same_config = json.dumps([ + {"network": "192.168.1.0/24", "gateway": "LAN_GW", "descr": "lan", "disabled": "0"} + ]) + driver._get = lambda path: SEARCH_ROUTE_RESPONSE + driver.load_merge_candidate(config=same_config) + diff = driver.compare_config() + assert diff == "" + + +# --------------------------------------------------------------------------- +# commit_config() +# --------------------------------------------------------------------------- + +BACKUPS_RESPONSE = { + "items": [ + {"id": "config-opnsense01-1234567890.xml", "time": "1234567890", "description": "before change"}, + {"id": "config-opnsense01-1234567800.xml", "time": "1234567800", "description": "initial"}, + ] +} + + +class TestCommitConfig: + def test_raises_without_candidate(self, driver): + from napalm.base.exceptions import MergeConfigException + with pytest.raises(MergeConfigException, match="No candidate"): + driver.commit_config() + + def test_posts_each_route_and_reconfigure(self, driver): + # _get for backup list + POST for 2 routes + POST reconfigure + driver._get = lambda path: BACKUPS_RESPONSE + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + driver.session.post.return_value = _make_json_response({}) + driver.commit_config() + assert driver.session.post.call_count == 3 # addroute×2 + reconfigure + + def test_records_pre_commit_backup_id(self, driver): + driver._get = lambda path: BACKUPS_RESPONSE + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + driver.session.post.return_value = _make_json_response({}) + driver.commit_config() + assert driver._pre_commit_backup_id == "config-opnsense01-1234567890.xml" + + def test_clears_candidate_after_commit(self, driver): + driver._get = lambda path: BACKUPS_RESPONSE + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + driver.session.post.return_value = _make_json_response({}) + driver.commit_config() + assert driver._candidate_config is None + + def test_records_none_backup_when_no_backups_exist(self, driver): + driver._get = lambda path: {"items": []} + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + driver.session.post.return_value = _make_json_response({}) + driver.commit_config() + assert driver._pre_commit_backup_id is None + + def test_raises_on_api_error(self, driver): + from napalm.base.exceptions import MergeConfigException + import requests as _req + driver._get = lambda path: BACKUPS_RESPONSE + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + driver.session.post.side_effect = _req.exceptions.RequestException("timeout") + with pytest.raises(MergeConfigException, match="Failed to apply"): + driver.commit_config() + + +# --------------------------------------------------------------------------- +# discard_config() +# --------------------------------------------------------------------------- + +class TestDiscardConfig: + def test_clears_candidate(self, driver): + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + driver.discard_config() + assert driver._candidate_config is None + + def test_idempotent_when_no_candidate(self, driver): + driver.discard_config() # must not raise + assert driver._candidate_config is None + + +# --------------------------------------------------------------------------- +# rollback() +# --------------------------------------------------------------------------- + +class TestRollback: + def test_noop_when_no_backups_and_no_commit(self, driver): + driver._get = lambda path: {"items": []} + driver.rollback() # must not raise + driver.session.post.assert_not_called() + + def test_uses_pre_commit_backup_id(self, driver): + driver._pre_commit_backup_id = "config-opnsense01-1234567890.xml" + driver.session.post.return_value = _make_json_response({"status": "ok"}) + driver.rollback() + url = driver.session.post.call_args[0][0] + assert "config-opnsense01-1234567890.xml" in url + assert "revert_backup" in url + + def test_falls_back_to_latest_backup_without_commit(self, driver): + driver._get = lambda path: BACKUPS_RESPONSE + driver.session.post.return_value = _make_json_response({"status": "ok"}) + driver.rollback() + url = driver.session.post.call_args[0][0] + assert "config-opnsense01-1234567890.xml" in url + + def test_clears_pre_commit_backup_id_after_rollback(self, driver): + driver._pre_commit_backup_id = "config-opnsense01-1234567890.xml" + driver.session.post.return_value = _make_json_response({}) + driver.rollback() + assert driver._pre_commit_backup_id is None + + def test_calls_revert_backup_exactly_once(self, driver): + driver._pre_commit_backup_id = "config-opnsense01-1234567890.xml" + driver.session.post.return_value = _make_json_response({}) + driver.rollback() + assert driver.session.post.call_count == 1 + + +# --------------------------------------------------------------------------- +# get_config() — candidate slot +# --------------------------------------------------------------------------- + +class TestGetConfigCandidate: + def test_candidate_empty_without_staged_config(self, driver): + driver._get = lambda path: "" + result = driver.get_config() + assert result["candidate"] == "" + + def test_candidate_contains_staged_routes(self, driver): + driver._get = lambda path: "" + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + result = driver.get_config() + assert "WAN_GW" in result["candidate"] + + def test_candidate_is_valid_json(self, driver): + driver._get = lambda path: "" + driver.load_merge_candidate(config=VALID_ROUTES_CONFIG) + result = driver.get_config() + parsed = json.loads(result["candidate"]) + assert isinstance(parsed, list)