Adds get_poe_status()/set_poe() to expose per-port PoE configuration (enable state, priority, allocation method, allocated power) and allow toggling it via the AOS-Switch REST API.
580 lines
22 KiB
Python
580 lines
22 KiB
Python
# -*- 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/<version>/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,
|
||
}
|
||
}
|