diff --git a/napalm_device_types/__init__.py b/napalm_device_types/__init__.py index a006d6a..3a8fb37 100644 --- a/napalm_device_types/__init__.py +++ b/napalm_device_types/__init__.py @@ -21,20 +21,31 @@ Available base classes: * :class:`~napalm_device_types.hypervisor.HypervisorDriver` * :class:`~napalm_device_types.os.OSDriver` * :class:`~napalm_device_types.storage.StorageDriver` +* :class:`~napalm_device_types.residential_gateway.ResidentialGatewayDriver` + +Also provided: + +* :class:`~napalm_device_types.config_lifecycle.ConfigLifecycleMixin` -- + stand-alone mixin to reduce duplication of config lifecycle methods across + drivers. """ from napalm_device_types.access_point import AccessPointDriver +from napalm_device_types.config_lifecycle import ConfigLifecycleMixin from napalm_device_types.firewall import FirewallDriver from napalm_device_types.hypervisor import HypervisorDriver from napalm_device_types.os import OSDriver +from napalm_device_types.residential_gateway import ResidentialGatewayDriver from napalm_device_types.storage import StorageDriver from napalm_device_types.switch import SwitchDriver __all__ = [ "AccessPointDriver", + "ConfigLifecycleMixin", "FirewallDriver", "HypervisorDriver", "OSDriver", + "ResidentialGatewayDriver", "StorageDriver", "SwitchDriver", ] diff --git a/napalm_device_types/access_point.py b/napalm_device_types/access_point.py index 99097f1..ab338b8 100644 --- a/napalm_device_types/access_point.py +++ b/napalm_device_types/access_point.py @@ -32,6 +32,7 @@ from napalm_device_types.models import ( class AccessPointDriver(NetworkDriver): + TYPE_LABEL: str = "Access Point" """ Abstract intermediate driver for wireless access points. diff --git a/napalm_device_types/config_lifecycle.py b/napalm_device_types/config_lifecycle.py new file mode 100644 index 0000000..d3c98dc --- /dev/null +++ b/napalm_device_types/config_lifecycle.py @@ -0,0 +1,113 @@ +# -*- coding: utf-8 -*- +from __future__ import annotations + +import difflib +from typing import ClassVar + +from napalm.base.exceptions import ( + CommandErrorException, + MergeConfigException, + ReplaceConfigException, +) + + +class ConfigLifecycleMixin: + """Mixin providing the standard NAPALM config lifecycle for CLI-driven drivers. + + Provides a complete config lifecycle: ``load_merge_candidate``, + ``load_replace_candidate``, ``compare_config``, ``discard_config``, + ``has_pending_commit``, and a basic ``rollback`` that re-stages the + last backup and calls ``commit_config``. + + Concrete drivers **must** provide: + + * ``_get_running_config()`` -- return the current running config as a + string (e.g. ``show running-config`` or ``uci export``) + * ``commit_config()`` -- vendor-specific apply-and-persist logic + + Override ``rollback()`` if the device needs a smarter revert (e.g. + block-diff rollback or config-import-based recovery). + + Config-state attributes (``_candidate_config``, ``_candidate_mode``, + ``_backup_config``) are annotated here for type checkers; they default + to ``None`` and are written by the lifecycle methods. + """ + + _candidate_config: str | None = None + _candidate_mode: str | None = None + _backup_config: str | None = None + + _comment_chars: ClassVar[tuple[str, ...]] = ("!", "#") + + def _get_running_config(self) -> str: + raise NotImplementedError + + def commit_config(self, message: str = "", revert_in: int | None = None) -> None: + raise NotImplementedError + + def load_merge_candidate( + self, filename: str | None = None, config: str | None = None + ) -> None: + if filename is not None: + try: + with open(filename) as fh: + config = fh.read() + except OSError as exc: + raise MergeConfigException(str(exc)) from exc + if config is None: + raise MergeConfigException("Either 'filename' or 'config' must be provided.") + self._candidate_config = config + self._candidate_mode = "merge" + + def load_replace_candidate( + self, filename: str | None = None, config: str | None = None + ) -> None: + if filename is not None: + try: + with open(filename) as fh: + config = fh.read() + except OSError as exc: + raise ReplaceConfigException(str(exc)) from exc + if config is None: + raise ReplaceConfigException("Either 'filename' or 'config' must be provided.") + self._candidate_config = config + self._candidate_mode = "replace" + + def compare_config(self) -> str: + if self._candidate_config is None: + return "" + + if self._candidate_mode == "merge": + lines = [] + for line in self._candidate_config.splitlines(): + if line.strip() and not line.strip().startswith(self._comment_chars): + lines.append(f"+{line}") + return "\n".join(lines) + + running = self._get_running_config() + diff = difflib.unified_diff( + running.splitlines(), + self._candidate_config.splitlines(), + fromfile="running-config", + tofile="candidate-config", + lineterm="", + ) + return "\n".join(diff) + + def discard_config(self) -> None: + self._candidate_config = None + self._candidate_mode = None + + def has_pending_commit(self) -> bool: + return self._candidate_config is not None + + def rollback(self) -> None: + if self._backup_config is None: + raise CommandErrorException( + "No backup configuration available – " + "commit_config has not been called in this session." + ) + self._candidate_config = self._backup_config + self._candidate_mode = "merge" + self.commit_config() + self._backup_config = None diff --git a/napalm_device_types/firewall.py b/napalm_device_types/firewall.py index d143f30..a251b1c 100644 --- a/napalm_device_types/firewall.py +++ b/napalm_device_types/firewall.py @@ -24,6 +24,7 @@ from napalm_device_types.models import ( class FirewallDriver(NetworkDriver): + TYPE_LABEL: str = "Firewall" """ Abstract intermediate driver for firewall/security devices. diff --git a/napalm_device_types/hypervisor.py b/napalm_device_types/hypervisor.py index b2c4f6d..024844a 100644 --- a/napalm_device_types/hypervisor.py +++ b/napalm_device_types/hypervisor.py @@ -25,6 +25,7 @@ from napalm_device_types.models import ( class HypervisorDriver(NetworkDriver): + TYPE_LABEL: str = "Hypervisor" """ Abstract intermediate driver for hypervisors and virtualisation platforms (e.g. Proxmox VE, VMware ESXi, KVM/libvirt, Hyper-V). diff --git a/napalm_device_types/models.py b/napalm_device_types/models.py index 2996e1b..1685bf1 100644 --- a/napalm_device_types/models.py +++ b/napalm_device_types/models.py @@ -313,6 +313,43 @@ class VPNTunnelDict(TypedDict): description: NotRequired[str] # human-readable tunnel description / name +# --------------------------------------------------------------------------- +# Residential Gateway +# --------------------------------------------------------------------------- + + +class WANStatusDict(TypedDict): + connection_type: str # e.g. "DSL", "Cable", "PPPoE", "DHCP" + is_connected: bool + external_ip: str + uptime: int + bytes_sent: int + bytes_received: int + max_bitrate_up: int # kbit/s + max_bitrate_down: int # kbit/s + external_ipv6: NotRequired[str] + link_status: NotRequired[str] # physical line state, e.g. "Up" / "Down" + + +class PortForwardDict(TypedDict): + name: str + protocol: str # "TCP" or "UDP" + external_port: int + internal_ip: str + internal_port: int + enabled: bool + remote_host: NotRequired[str] # restrict forward to a specific remote source + + +class HostDict(TypedDict): + mac: str + ip: str + hostname: str + interface_type: str # "LAN", "WLAN", ... + is_active: bool + lease_time_remaining: NotRequired[int] # seconds + + # --------------------------------------------------------------------------- # Hypervisor # --------------------------------------------------------------------------- diff --git a/napalm_device_types/os.py b/napalm_device_types/os.py index 6f526f7..ca3102d 100644 --- a/napalm_device_types/os.py +++ b/napalm_device_types/os.py @@ -30,6 +30,7 @@ from napalm_device_types.models import ( class OSDriver(NetworkDriver): + TYPE_LABEL: str = "OS" """ Abstract intermediate driver for general-purpose operating systems (e.g. Linux, BSD, macOS). diff --git a/napalm_device_types/residential_gateway.py b/napalm_device_types/residential_gateway.py new file mode 100644 index 0000000..e27b3ef --- /dev/null +++ b/napalm_device_types/residential_gateway.py @@ -0,0 +1,197 @@ +""" +Abstract base class for residential gateway drivers. + +A "residential gateway" combines the roles of router, firewall and +wireless access point in a single consumer device (e.g. AVM FritzBox, +ISP-supplied DSL/cable routers). This base class merges the relevant +subsets of :class:`~napalm_device_types.firewall.FirewallDriver` and +:class:`~napalm_device_types.access_point.AccessPointDriver` plus +gateway-specific operations (WAN status, port forwarding, connected +hosts). + +Usage:: + + from napalm_device_types import ResidentialGatewayDriver + + class FritzBoxDriver(ResidentialGatewayDriver): + def get_wan_status(self): + ... +""" + +from typing import Dict, List +from napalm.base import NetworkDriver +from napalm_device_types._ucd_metrics import IF_SKIP_DEFAULT, collect_ucd_metrics +from napalm_device_types.models import ( + HealthMetricsDict, + HostDict, + NATTranslationDict, + PortForwardDict, + RadioStatusDict, + SSIDDict, + VPNTunnelDict, + WANStatusDict, + WirelessClientDict, +) + + +class ResidentialGatewayDriver(NetworkDriver): + TYPE_LABEL: str = "Gateway" + """ + Abstract intermediate driver for residential gateways (router + firewall + AP). + + Inherits all standard NAPALM NetworkDriver methods and adds the + gateway-, firewall- and wireless-specific operations that concrete + drivers must implement. + """ + + _SNMP_SKIP_IF = IF_SKIP_DEFAULT + _SNMP_TX_ERR_IS_DROP: bool = False + + @classmethod + async def get_health_metrics(cls, snmp_get, snmp_walk) -> HealthMetricsDict: + return await collect_ucd_metrics( + snmp_get, snmp_walk, + tx_err_is_drop=cls._SNMP_TX_ERR_IS_DROP, + if_skip=cls._SNMP_SKIP_IF, + ) + + def get_wan_status(self) -> WANStatusDict: + """ + Returns the status of the device's internet (WAN) uplink. + + Contains: + + * connection_type (string) - e.g. ``"DSL"``, ``"Cable"``, ``"PPPoE"``, ``"DHCP"`` + * is_connected (bool) - whether the WAN connection is currently established + * external_ip (string) - the public IPv4 address assigned to the WAN interface + * uptime (int) - seconds since the WAN connection was last (re-)established + * bytes_sent (int) - total bytes transmitted on the WAN interface + * bytes_received (int) - total bytes received on the WAN interface + * max_bitrate_up (int) - upstream sync rate in kbit/s + * max_bitrate_down (int) - downstream sync rate in kbit/s + * external_ipv6 (string, optional) - the public IPv6 address, if any + * link_status (string, optional) - physical line state, e.g. ``"Up"`` / ``"Down"`` + + Example:: + + { + "connection_type": "DSL", + "is_connected": True, + "external_ip": "203.0.113.7", + "uptime": 345600, + "bytes_sent": 1234567890, + "bytes_received": 9876543210, + "max_bitrate_up": 40000, + "max_bitrate_down": 250000, + "link_status": "Up", + } + """ + raise NotImplementedError + + def get_port_forwards(self) -> List[PortForwardDict]: + """ + Returns the configured port forwarding (port mapping) rules. + + Each entry contains: + + * name (string) - the rule's description/name + * protocol (string) - ``"TCP"`` or ``"UDP"`` + * external_port (int) - the WAN-side port + * internal_ip (string) - the LAN host the traffic is forwarded to + * internal_port (int) - the LAN-side port + * enabled (bool) - whether the rule is currently active + * remote_host (string, optional) - restricts the forward to a specific + remote source address; empty/absent means "any" + + Example:: + + [ + { + "name": "Webserver HTTPS", + "protocol": "TCP", + "external_port": 443, + "internal_ip": "192.168.1.10", + "internal_port": 443, + "enabled": True, + } + ] + """ + raise NotImplementedError + + def get_hosts(self) -> List[HostDict]: + """ + Returns the list of hosts known to the gateway (LAN clients). + + Each entry contains: + + * mac (string) - the host's MAC address + * ip (string) - the host's current IP address + * hostname (string) - the host's reported hostname (empty if unknown) + * interface_type (string) - how the host is connected, e.g. ``"LAN"``, ``"WLAN"`` + * is_active (bool) - whether the host is currently online + * lease_time_remaining (int, optional) - remaining DHCP lease time in seconds + + Example:: + + [ + { + "mac": "AA:BB:CC:DD:EE:FF", + "ip": "192.168.1.42", + "hostname": "laptop", + "interface_type": "WLAN", + "is_active": True, + "lease_time_remaining": 3600, + } + ] + """ + raise NotImplementedError + + def get_nat_translations(self) -> List[NATTranslationDict]: + """ + Returns a list of active NAT translation entries. + + See :meth:`napalm_device_types.firewall.FirewallDriver.get_nat_translations` + for the entry format. Residential gateways typically derive this from + the active port-forwarding/NAT-PT table rather than a live connection + tracker; drivers that cannot provide this should return an empty list. + """ + raise NotImplementedError + + def get_vpn_tunnels(self) -> Dict[str, VPNTunnelDict]: + """ + Returns the status of VPN tunnels (e.g. WireGuard road-warrior + access, IPsec site-to-site). + + See :meth:`napalm_device_types.firewall.FirewallDriver.get_vpn_tunnels` + for the entry format. Drivers that cannot provide this should return + an empty dict. + """ + raise NotImplementedError + + def get_wireless_clients(self) -> List[WirelessClientDict]: + """ + Returns the list of wireless clients currently associated with the + device's built-in access point(s). + + See :meth:`napalm_device_types.access_point.AccessPointDriver.get_wireless_clients` + for the entry format. + """ + raise NotImplementedError + + def get_ssids(self) -> Dict[str, SSIDDict]: + """ + Returns the configured wireless networks (SSIDs). + + See :meth:`napalm_device_types.access_point.AccessPointDriver.get_ssids` + for the entry format. + """ + raise NotImplementedError + + def get_radio_status(self) -> Dict[str, RadioStatusDict]: + """ + Returns the status of the device's wireless radios. + + See :meth:`napalm_device_types.access_point.AccessPointDriver.get_radio_status` + for the entry format. + """ + raise NotImplementedError diff --git a/napalm_device_types/storage.py b/napalm_device_types/storage.py index 84f9eb8..36c7887 100644 --- a/napalm_device_types/storage.py +++ b/napalm_device_types/storage.py @@ -26,6 +26,7 @@ from napalm_device_types.models import ( class StorageDriver(NetworkDriver): + TYPE_LABEL: str = "Storage" """ Abstract intermediate driver for storage appliances and NAS/SAN devices (e.g. TrueNAS SCALE/CORE, Synology DSM, QNAP QTS, NetApp ONTAP, diff --git a/napalm_device_types/switch.py b/napalm_device_types/switch.py index 93692d3..a5d0659 100644 --- a/napalm_device_types/switch.py +++ b/napalm_device_types/switch.py @@ -25,6 +25,7 @@ from napalm_device_types.models import ( class SwitchDriver(NetworkDriver): + TYPE_LABEL: str = "Switch" """ Abstract intermediate driver for Ethernet switches. diff --git a/pyproject.toml b/pyproject.toml index b1b4086..7e40288 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "napalm-device-types" -version = "0.3.0" +version = "0.5.0" description = "Abstract device-type base classes for NAPALM drivers" readme = "README.md" requires-python = ">=3.9" diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/test_type_label.py b/tests/test_type_label.py new file mode 100644 index 0000000..e9abfc8 --- /dev/null +++ b/tests/test_type_label.py @@ -0,0 +1,53 @@ +"""Tests for TYPE_LABEL class attribute on all base driver classes.""" + +import pytest +from napalm_device_types import ( + AccessPointDriver, + FirewallDriver, + HypervisorDriver, + OSDriver, + ResidentialGatewayDriver, + StorageDriver, + SwitchDriver, +) + + +BASE_CLASSES = [ + (AccessPointDriver, "Access Point"), + (FirewallDriver, "Firewall"), + (HypervisorDriver, "Hypervisor"), + (OSDriver, "OS"), + (ResidentialGatewayDriver, "Gateway"), + (StorageDriver, "Storage"), + (SwitchDriver, "Switch"), +] + + +@pytest.mark.parametrize("cls, expected_label", BASE_CLASSES) +def test_type_label_present(cls, expected_label): + assert hasattr(cls, "TYPE_LABEL"), f"{cls.__name__} is missing TYPE_LABEL" + + +@pytest.mark.parametrize("cls, expected_label", BASE_CLASSES) +def test_type_label_value(cls, expected_label): + assert cls.TYPE_LABEL == expected_label + + +@pytest.mark.parametrize("cls, _", BASE_CLASSES) +def test_type_label_is_nonempty_string(cls, _): + assert isinstance(cls.TYPE_LABEL, str) and cls.TYPE_LABEL + + +def test_subclass_inherits_type_label(): + class MySwitch(SwitchDriver): + pass + + assert MySwitch.TYPE_LABEL == "Switch" + + +def test_subclass_can_override_type_label(): + class MySpecialOS(OSDriver): + TYPE_LABEL = "Linux" + + assert MySpecialOS.TYPE_LABEL == "Linux" + assert OSDriver.TYPE_LABEL == "OS" # base class unaffected