# -*- coding: utf-8 -*- # Licensed under the Apache License, Version 2.0 """REST API client for HPE/Aruba ProCurve switches (ArubaOS REST interface). Compatible with ArubaOS REST API versions v3, v6, and v7. Used automatically by ProcurveDriver when the REST API is reachable. """ import logging from typing import Any, Dict, Optional, Tuple import requests import urllib3 from napalm.base.exceptions import ConnectionException, ConnectAuthError logger = logging.getLogger("napalm_procurve.api") class ProcurveApiClient: """Thin wrapper around ArubaOS REST API sessions. Detection approach: probe ``GET /rest//system/status`` without authentication. A ``401`` response confirms the API is available. """ # Ordered list of API versions to probe (newest first) API_VERSIONS = ["v7", "v6", "v3"] def __init__( self, hostname: str, username: str, password: str, timeout: int = 60, ssl_verify: bool = False, api_version: Optional[str] = None, ) -> None: self.hostname = hostname self.username = username self.password = password self.timeout = timeout self.ssl_verify = ssl_verify self._api_version = api_version # None = auto-detect self._proto: str = "https" self._base_url: str = "" self._session: Optional[requests.Session] = None # ------------------------------------------------------------------ # Detection # ------------------------------------------------------------------ @classmethod def probe( cls, hostname: str, timeout: int = 5, ssl_verify: bool = False, ) -> Tuple[Optional[str], Optional[str]]: """Probe the device for an ArubaOS REST API. Returns ``(api_version, proto)`` on success, ``(None, None)`` if the API is unreachable or the device is not an ArubaOS switch. """ if not ssl_verify: urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) for proto in ("https", "http"): for ver in cls.API_VERSIONS: url = f"{proto}://{hostname}/rest/{ver}/system/status" try: resp = requests.get( url, verify=ssl_verify, timeout=timeout, allow_redirects=False, ) # 401 = API exists but not authenticated (most common) # 200 = API exists and accessible without auth (unusual) # 405 = wrong HTTP verb but API is there # 302/301/303 = redirect to login page (older firmware) # 403 = API exists but access forbidden if resp.status_code in (200, 301, 302, 303, 401, 403, 405): logger.debug( "API detected at %s (HTTP %s, %s %s)", hostname, resp.status_code, proto, ver, ) return ver, proto except requests.exceptions.SSLError: # SSL error on https → try http next continue except requests.exceptions.ConnectionError: # No route to host / port closed → next proto/ver break except requests.exceptions.Timeout: break return None, None # ------------------------------------------------------------------ # Session management # ------------------------------------------------------------------ def connect(self) -> None: """Open an authenticated REST API session.""" if not self._base_url: raise ConnectionException( "ProcurveApiClient.connect() called before api_version/proto set. " "Use ProcurveApiClient.probe() first." ) if not self.ssl_verify: urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) self._session = requests.Session() self._session.verify = self.ssl_verify self._session.headers.update( {"Content-Type": "application/json", "Connection": "close"} ) url = self._base_url + "login-sessions" resp = self._session.post( url, json={"userName": self.username, "password": self.password}, timeout=self.timeout, ) if resp.status_code != 201: raise ConnectAuthError( f"REST API login failed (HTTP {resp.status_code}) for {self.hostname}" ) # AOS-Switch returns the session cookie in the JSON body rather than # via Set-Cookie — it must be sent back as a Cookie header on every # subsequent request, otherwise writes (POST/PUT/DELETE) are rejected # with "Access is unauthorized" while reads still succeed. cookie = resp.json().get("cookie") if cookie: self._session.headers.update({"Cookie": cookie}) logger.debug("REST API login OK for %s", self.hostname) def setup(self, api_version: str, proto: str) -> None: """Set the API version and protocol (called after probe()).""" self._api_version = api_version self._proto = proto self._base_url = f"{proto}://{self.hostname}/rest/{api_version}/" def disconnect(self) -> None: """Log out from the REST API session.""" if self._session and self._base_url: try: self._session.delete( self._base_url + "login-sessions", timeout=self.timeout ) except Exception: pass if self._session: self._session.close() self._session = None def is_alive(self) -> bool: """Check if the API session is still usable.""" if not self._session: return False try: resp = self._session.get( self._base_url + "system/status", timeout=5 ) return resp.status_code == 200 except Exception: return False # ------------------------------------------------------------------ # HTTP helpers # ------------------------------------------------------------------ def get(self, endpoint: str) -> Dict[str, Any]: """GET ``{base_url}{endpoint}`` and return parsed JSON (or empty dict).""" url = self._base_url + endpoint try: resp = self._session.get(url, timeout=self.timeout) if resp.ok: return resp.json() if resp.status_code == 404: logger.debug("GET %s returned HTTP 404", url) else: logger.warning("GET %s returned HTTP %s: %s", url, resp.status_code, resp.text[:200]) except Exception as exc: logger.warning("GET %s error: %s", url, exc) return {} def post(self, endpoint: str, **kwargs: Any) -> requests.Response: """POST ``{base_url}{endpoint}``.""" return self._session.post(self._base_url + endpoint, timeout=self.timeout, **kwargs) def put(self, endpoint: str, **kwargs: Any) -> requests.Response: """PUT ``{base_url}{endpoint}``.""" return self._session.put(self._base_url + endpoint, timeout=self.timeout, **kwargs) def delete(self, endpoint: str, **kwargs: Any) -> requests.Response: """DELETE ``{base_url}{endpoint}``.""" return self._session.delete(self._base_url + endpoint, timeout=self.timeout, **kwargs) # ------------------------------------------------------------------ # NAPALM data getters (API-backed) # ------------------------------------------------------------------ def get_facts(self) -> Dict: """Return NAPALM facts from the REST API.""" system = self.get("system/status") if not system: # Fallback endpoint for stacked switches system = self.get("system/status/global_info") dns = self.get("dns") domain = "" if dns: domains = dns.get("dns_domain_names", []) domain = f".{domains[0]}" if domains else "" hostname = system.get("name", "") fqdn = f"{hostname}{domain}" if domain else hostname # Interface list from blades iface_list = [] sw_status = self.get("system/status/switch") for blade in sw_status.get("blades", []): for port in blade.get("data_ports", []): iface_list.append(port.get("port_name", "")) return { "vendor": "HPE Aruba", "model": system.get("product_model", ""), "hostname": hostname, "fqdn": fqdn, "os_version": system.get("firmware_version", ""), "serial_number": system.get("serial_number", ""), "uptime": float(system.get("uptime_seconds", -1)), "interface_list": iface_list, } def get_interfaces(self) -> Dict[str, Dict]: """Return NAPALM interfaces from the REST API.""" ports_data = self.get("ports") stats_data = self.get("port-statistics") output: Dict[str, Dict] = {} for port in ports_data.get("port_element", []): pid = port.get("id", "") output[pid] = { "is_up": bool(port.get("is_port_up", False)), "is_enabled": bool(port.get("is_port_enabled", False)), "description": port.get("name", ""), "last_flapped": -1.0, "speed": 0.0, "mtu": -1, "mac_address": "", } for stat in stats_data.get("port_statistics_element", []): pid = stat.get("id", "") if pid in output: output[pid]["speed"] = float(stat.get("port_speed_mbps", 0)) return output def get_interfaces_ip(self) -> Dict[str, Dict]: """Return NAPALM interfaces IP from the REST API.""" vlans_data = self.get("vlans") result: Dict[str, Dict] = {} for vlan in vlans_data.get("vlan_element", []): vlan_id = str(vlan.get("vlan_id", "")) vlan_name = f"VLAN{vlan_id}" # IPv4 from /vlans/{id}/ipv4 ipv4_data = self.get(f"vlans/{vlan_id}/ipv4") for entry in ipv4_data.get("ipv4_element", []): addr = entry.get("address", "") mask = entry.get("mask", "") if not addr or addr == "0.0.0.0": continue try: import netaddr plen = netaddr.IPNetwork(f"{addr}/{mask}").prefixlen except Exception: plen = 24 result.setdefault(vlan_name, {"ipv4": {}, "ipv6": {}}) result[vlan_name]["ipv4"][addr] = {"prefix_length": plen} return {k: v for k, v in result.items() if v["ipv4"] or v["ipv6"]} # ------------------------------------------------------------------ # PoE # ------------------------------------------------------------------ def get_poe_ports(self) -> Dict[str, Dict]: """Return PoE configuration for all ports, keyed by port id.""" data = self.get("poe/ports") return {p["port_id"]: p for p in data.get("port_poe", []) if p.get("port_id")} def set_port_poe(self, port_id: str, payload: Dict[str, Any]) -> requests.Response: """PUT PoE configuration changes for a single port.""" return self.put(f"ports/{port_id}/poe", json=payload) def get_arp_table(self, vrf: str = "") -> list: """Return NAPALM ARP table from the REST API.""" arp_data = self.get("arp") table = [] for entry in arp_data.get("arp_entry", []): table.append( { "interface": entry.get("port_id", ""), "mac": entry.get("mac_address", ""), "ip": entry.get("ip_address", ""), "age": float(entry.get("age", -1)), } ) return table def get_mac_address_table(self) -> list: """Return NAPALM MAC address table from the REST API.""" mac_data = self.get("mac-table") table = [] for entry in mac_data.get("mac_table_entry_element", []): table.append( { "mac": entry.get("mac_address", ""), "interface": entry.get("port_id", ""), "vlan": int(entry.get("vlan_id", 0)), "static": entry.get("mac_addr_type", "").lower() == "static", "active": True, "moves": -1, "last_move": -1.0, } ) return table def get_lldp_neighbors(self) -> Dict[str, list]: """Return NAPALM LLDP neighbors from the REST API.""" result: Dict[str, list] = {} data = self.get("lldp/remote-device") for nbr in data.get("lldp_remote_device_element", []): local_port = nbr.get("local_port", "") result.setdefault(local_port, []).append( { "hostname": nbr.get("system_name", ""), "port": nbr.get("port_id", ""), } ) return result def get_lldp_neighbors_detail(self) -> Dict[str, list]: """Return NAPALM LLDP neighbor details from the REST API.""" result: Dict[str, list] = {} data = self.get("lldp/remote-device") for nbr in data.get("lldp_remote_device_element", []): local_port = nbr.get("local_port", "") result.setdefault(local_port, []).append( { "parent_interface": "", "remote_chassis_id": nbr.get("chassis_id", ""), "remote_port": nbr.get("port_id", ""), "remote_port_description": nbr.get("port_description", ""), "remote_system_name": nbr.get("system_name", ""), "remote_system_description": nbr.get("system_description", ""), "remote_system_capab": [], "remote_system_enable_capab": [], } ) return result def get_config(self) -> Dict[str, str]: """Return running and startup configuration via REST API.""" running = self.get("running-config") startup = self.get("startup-config") return { "running": running.get("config", ""), "startup": startup.get("config", ""), "candidate": "", } def get_ntp_servers(self) -> Dict[str, Dict]: """Return NTP servers from REST API.""" data = self.get("ntp/server") servers: Dict[str, Dict] = {} for entry in data.get("ntp_server_element", []): addr = entry.get("server_address", "") if addr: servers[addr] = {} return servers def get_vlans(self) -> Dict[int, Dict]: """Return VLAN information with proper tagged/untagged separation. Fetches ``/vlans`` for VLAN names and ``/vlans-ports`` for port membership. The ``port_mode`` field in ``/vlans-ports`` varies by firmware: ======================== ======== Firmware ``port_mode`` Meaning ======================== ======== ``POM_UNTAGGED`` Untagged (access/native) ``POM_TAGGED_STATIC`` Tagged (trunk) ``POM_TAGGED`` Tagged (some older v3 API) ``POM_NATIVE_UNTAGGED`` Native untagged (also PVID) ``POM_NATIVE_TAGGED`` Native tagged (unusual) ======================== ======== Returns extended NAPALM-compatible dict:: { 1: {"name": "DEFAULT_VLAN", "interfaces": ["1","2","3"], "tagged": [], "untagged": ["1","2","3"]}, 10: {"name": "MGMT", "interfaces": ["1","2","3"], "tagged": ["1","2","3"], "untagged": []}, } """ vlans_data = self.get("vlans") ports_data = self.get("vlans-ports") # Build VLAN skeleton from /vlans result: Dict[int, Dict] = {} for vlan in vlans_data.get("vlan_element", []): try: vid = int(vlan["vlan_id"]) except (KeyError, ValueError, TypeError): continue result[vid] = { "name": vlan.get("name", f"VLAN{vid}"), "interfaces": [], "tagged": [], "untagged": [], } # Populate tagged/untagged from /vlans-ports for entry in ports_data.get("vlan_port_element", []): try: vid = int(entry["vlan_id"]) except (KeyError, ValueError, TypeError): continue port = entry.get("port_id", "") if not port: continue mode = entry.get("port_mode", "") # Normalise to lower-case for comparison mode_low = mode.lower() # Untagged variants: POM_UNTAGGED, POM_NATIVE_UNTAGGED, "Untagged" is_untagged = any( kw in mode_low for kw in ("untagged", "native_untagged", "pom_untagged") ) and "tagged" not in mode_low.replace("untagged", "") # More precise: contains "untagged" anywhere → untagged # Contains "tagged" but NOT "untagged" → tagged if "untagged" in mode_low: is_untagged = True elif "tagged" in mode_low: is_untagged = False else: # Unknown mode: skip logger.debug("Unknown port_mode %r for port %s VLAN %s", mode, port, vid) continue if vid not in result: result[vid] = { "name": f"VLAN{vid}", "interfaces": [], "tagged": [], "untagged": [], } if is_untagged: result[vid]["untagged"].append(port) else: result[vid]["tagged"].append(port) # Build "interfaces" as ordered union (untagged first, then tagged) for vdata in result.values(): seen: list = [] for p in vdata["untagged"] + vdata["tagged"]: if p not in seen: seen.append(p) vdata["interfaces"] = seen return result def delete_vlan(self, vlan_id: int) -> None: """Remove a VLAN via the REST API (``DELETE /vlans/{id}``). :param vlan_id: VLAN ID (1–4094) to delete. :raises ValueError: If *vlan_id* is out of range. :raises ConnectionException: If the API returns a non-success status. """ if not 1 <= vlan_id <= 4094: raise ValueError(f"VLAN ID {vlan_id} is out of range (1–4094)") resp = self.delete(f"vlans/{vlan_id}") if not resp.ok and resp.status_code != 404: raise ConnectionException( f"delete_vlan({vlan_id}): REST API returned HTTP {resp.status_code}" ) def get_snmp_communities(self) -> Dict[str, Dict]: """Return configured SNMP communities from REST API.""" data = self.get("snmp-server/community") result: Dict[str, Dict] = {} for entry in data.get("snmp_community_element", []): name = entry.get("community_name", "") if name: result[name] = {"access_type": entry.get("access_type", "")} return result def configure_snmp_community(self, community: str = "public") -> Tuple[bool, str]: """Ensure an SNMP read-only community exists via REST resource endpoint. Returns (success, message). """ existing = self.get_snmp_communities() if community in existing: return True, f"SNMP community '{community}' already exists." resp = self.post("snmp-server/community", json={ "community_name": community, "access_type": "MIB2", }) if resp.ok: return True, f"SNMP community '{community}' configured via REST API." if resp.status_code == 409: return True, f"SNMP community '{community}' already exists." return False, f"REST API returned HTTP {resp.status_code}: {resp.text[:200]}" def ping( self, destination: str, source: str = "", ttl: int = 255, timeout: int = 2, size: int = 100, count: int = 5, vrf: str = "", ) -> Dict: """Execute ping via REST API (if supported) or return not-implemented.""" payload: Dict = {"destination": destination, "count": count, "timeout": timeout} if source: payload["source"] = source resp = self.post("ping", json=payload) if not resp.ok: return {"error": f"Ping API returned HTTP {resp.status_code}"} data = resp.json() sent = data.get("sent", count) received = data.get("received", 0) rtt_min = float(data.get("min_rtt", 0)) rtt_max = float(data.get("max_rtt", 0)) rtt_avg = float(data.get("avg_rtt", 0)) results = [] for _ in range(received): results.append({"ip_address": destination, "rtt": rtt_avg}) return { "success": { "probes_sent": sent, "packet_loss": sent - received, "rtt_min": rtt_min, "rtt_max": rtt_max, "rtt_avg": rtt_avg, "rtt_stddev": 0.0, "results": results, } }