initial commit

This commit is contained in:
Christian Manivong
2026-05-29 09:30:53 +02:00
commit 35032fd034
9 changed files with 3301 additions and 0 deletions
+524
View File
@@ -0,0 +1,524 @@
# -*- 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
if resp.status_code in (200, 401, 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}"
)
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()
logger.debug("GET %s returned HTTP %s", url, resp.status_code)
except Exception as exc:
logger.debug("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"]}
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 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,
}
}