Files
napalm-hpe-aruba-procurve/napalm_procurve/api_client.py
T
Christian ManivongandClaude Sonnet 4.6 d2915f5829 fix: extract J-code part numbers in REST API get_facts() path
api_client.get_facts() now applies the same J-code extraction as the
CLI path: product_model "HP2530-8G Switch(J9777A)" → model="HP2530-8G
Switch", part_number="J9777A".

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-06-22 23:10:07 +02:00

647 lines
25 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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 import helpers as napalm_helpers
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", ""))
model_raw = system.get("product_model", "")
# Extract J-code part numbers (e.g. J9777A) from the model string
import re as _re
_PN_RE = _re.compile(r"\b(J\d{4}[A-Z]{1,2})\b")
pn_match = _PN_RE.search(model_raw)
part_number = pn_match.group(1) if pn_match else ""
if pn_match:
clean = _PN_RE.sub("", model_raw)
clean = _re.sub(r"\(\s*\)", "", clean)
clean = _re.sub(r"\s{2,}", " ", clean).strip(" -()")
model = clean if clean else model_raw
else:
model = model_raw
return {
"vendor": "HPE Aruba",
"model": model,
"part_number": part_number,
"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] = {}
trunk_groups: Dict[str, list] = {}
trunk_modes: Dict[str, str] = {}
port_elements = list(ports_data.get("port_element", []))
# The /ports collection can under-report the switch's true port
# count (observed: a 2530-48G reports 51 of its 52 ports). Fill in
# any ports known to system/status/switch but missing from /ports.
seen_ids = {p.get("id") for p in port_elements}
try:
sw_status = self.get("system/status/switch")
for blade in sw_status.get("blades", []):
for p in blade.get("data_ports", []):
pid = p.get("port_name", "")
if pid and pid not in seen_ids:
port_elements.append({
"id": pid,
"name": "",
"is_port_up": p.get("operStatus") == "OPER_UP",
"is_port_enabled": p.get("adminStatus") == "ADMIN_UP",
"trunk_group": "",
})
seen_ids.add(pid)
except Exception:
pass
for port in port_elements:
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": "",
}
group = port.get("trunk_group") or ""
if group:
output[pid]["trunk_group"] = group
trunk_groups.setdefault(group, []).append(pid)
if port.get("trunk_mode"):
trunk_modes[group] = port["trunk_mode"]
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))
# Synthesize a logical interface entry for each configured LAG/trunk
# group so it shows up as its own row alongside its member ports.
for group, members in trunk_groups.items():
output[group] = {
"is_up": any(output[m]["is_up"] for m in members),
"is_enabled": any(output[m]["is_enabled"] for m in members),
"description": f"LAG ({', '.join(sorted(members, key=lambda s: int(s) if s.isdigit() else 0))})",
"last_flapped": -1.0,
"speed": sum(output[m]["speed"] for m in members),
"mtu": -1,
"mac_address": "",
"lag_members": members,
"lag_mode": "lacp" if trunk_modes.get(group) == "PTT_LACP" else "trunk",
}
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", []):
raw_mac = entry.get("mac_address", "")
try:
mac = napalm_helpers.mac(raw_mac)
except Exception:
mac = raw_mac
table.append(
{
"mac": mac,
"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,
}
}