From 6dc3946bc620aefba9a1d0d90237df14a64646fb Mon Sep 17 00:00:00 2001 From: oliver Date: Sun, 2 Aug 2026 17:03:42 +0800 Subject: [PATCH] Split UME routes and managed-NE service by domain. Keep thin facades for existing imports, tighten CSV import required columns, and restore Chinese WebCRT connection labels in status responses. Co-authored-by: Cursor --- netx_api/ne_service.py | 1028 ++------------------------ netx_api/ne_service_common.py | 252 +++++++ netx_api/ne_service_crud.py | 358 +++++++++ netx_api/ne_service_import.py | 247 +++++++ netx_api/ne_service_webcrt.py | 210 ++++++ netx_api/ume_alarm_ws.py | 2 +- netx_api/ume_alarms_router.py | 509 +++++++++++++ netx_api/ume_inventory_router.py | 177 +++++ netx_api/ume_key_alert_router.py | 339 +++++++++ netx_api/ume_router.py | 1175 +----------------------------- netx_api/ume_sync_router.py | 247 +++++++ netx_api/ume_token_router.py | 164 +++++ tests/test_managed_ne.py | 9 + 13 files changed, 2596 insertions(+), 2121 deletions(-) create mode 100644 netx_api/ne_service_common.py create mode 100644 netx_api/ne_service_crud.py create mode 100644 netx_api/ne_service_import.py create mode 100644 netx_api/ne_service_webcrt.py create mode 100644 netx_api/ume_alarms_router.py create mode 100644 netx_api/ume_inventory_router.py create mode 100644 netx_api/ume_key_alert_router.py create mode 100644 netx_api/ume_sync_router.py create mode 100644 netx_api/ume_token_router.py diff --git a/netx_api/ne_service.py b/netx_api/ne_service.py index 2361c26..3c0ccbf 100644 --- a/netx_api/ne_service.py +++ b/netx_api/ne_service.py @@ -1,976 +1,66 @@ +"""Managed NE service facade — CRUD / WebCRT hosts / import / credentials.""" from __future__ import annotations -from datetime import datetime -from io import BytesIO -import re -from typing import Any - -import pandas as pd -from fastapi import HTTPException -from sqlalchemy import or_ -from sqlalchemy.orm import Session - -from .device_types import ( - SUPPORTED_DEVICE_TYPES, - SUPPORTED_VENDORS, - WEBCRT_DEVICE_TYPES, - WEBCRT_NE_SOURCE, +from .ne_service_common import ( + IMPORT_COLUMNS, + UME_SYNC_SOURCE, + UME_SYNC_TAG, + WEBCRT_SOURCE, + _normalize_hop_target_auth_mode, + _normalize_hop_vendor, + _normalize_protocol, + _require_crypto, + get_device_credentials, + row_to_out, ) -from .models import ManagedNE, UmeInventoryNE -from .ne_crypto import CredentialCryptoError, credentials_configured, decrypt_secret, encrypt_secret -from .ne_schemas import ( - BatchAccountConfig, - HopProxyConfig, - ImportFailure, - ImportResult, - ManagedNeCreate, - ManagedNeOut, - ManagedNeUpdate, - UmeManagedDeleteResult, - UmeManagedSyncResult, +from .ne_service_crud import ( + batch_apply_account, + batch_apply_hop_proxy, + batch_delete_managed_ne, + create_managed_ne, + delete_managed_ne, + get_ids_by_tag, + get_managed_ne, + get_managed_ne_stats, + list_managed_ne, + update_managed_ne, ) -from .ne_session_factory import default_bastion_username_template, default_hop_command_template - -IMPORT_COLUMNS = ( - "device_type", - "ip", - "username", - "password", - "port", - "protocol", - "name", - "vendor", - "tags", - "remark", +from .ne_service_import import ( + build_managed_ne_import_template, + delete_ume_synced_managed_ne, + import_managed_ne, + sync_ume_inventory_to_managed_ne, +) +from .ne_service_webcrt import ( + upsert_webcrt_managed_ne, + upsert_webcrt_session_host, ) -UME_SYNC_SOURCE = "ume_sync" -UME_SYNC_TAG = "UME" -# Re-export for callers (WebCRT quick-connect). -WEBCRT_SOURCE = WEBCRT_NE_SOURCE -# Re-export for callers (WebCRT Quick Connect). -WEBCRT_SOURCE = WEBCRT_NE_SOURCE -_BUILTIN_NE_TYPE_RULES: list[tuple[re.Pattern[str], str, str]] = [ - (re.compile(r"ZXR|ZXCTN|M6000|\bBN\b", re.I), "zte_zxros", "ZTE"), - (re.compile(r"NE40|CE\b|ATN|MA5800|OptiX", re.I), "huawei", "Huawei"), - (re.compile(r"ASR|NCS|IOS.?XR|XR\b", re.I), "cisco_xr", "Cisco"), - (re.compile(r"Catalyst|Nexus|C9[0-9]{3}|ISR", re.I), "cisco_ios", "Cisco"), +__all__ = [ + "IMPORT_COLUMNS", + "UME_SYNC_SOURCE", + "UME_SYNC_TAG", + "WEBCRT_SOURCE", + "_normalize_hop_target_auth_mode", + "_normalize_hop_vendor", + "_normalize_protocol", + "_require_crypto", + "batch_apply_account", + "batch_apply_hop_proxy", + "batch_delete_managed_ne", + "build_managed_ne_import_template", + "create_managed_ne", + "delete_managed_ne", + "delete_ume_synced_managed_ne", + "get_device_credentials", + "get_ids_by_tag", + "get_managed_ne", + "get_managed_ne_stats", + "import_managed_ne", + "list_managed_ne", + "row_to_out", + "sync_ume_inventory_to_managed_ne", + "update_managed_ne", + "upsert_webcrt_managed_ne", + "upsert_webcrt_session_host", ] - - -def _now() -> datetime: - return datetime.utcnow() - - -def _require_crypto() -> None: - if not credentials_configured(): - raise HTTPException(status_code=503, detail="credential_secret_key_not_configured") - - -def _normalize_ip(ip: str) -> str: - return str(ip or "").strip() - - -def _normalize_protocol(protocol: str) -> str: - p = str(protocol or "ssh").strip().lower() - return p if p in ("ssh", "telnet") else "ssh" - - -def _normalize_hop_vendor(vendor: str) -> str: - v = str(vendor or "zte").strip().lower() - return v if v in ("zte", "linux", "huawei", "cisco", "bastion") else "zte" - - -def _normalize_hop_target_auth_mode(mode: str) -> str: - m = str(mode or "bastion_managed").strip().lower() - return m if m in ("bastion_managed", "manual") else "bastion_managed" - - -def _normalize_vendor(vendor: str) -> str: - raw = str(vendor or "").strip() - if not raw: - return "Other" - for item in SUPPORTED_VENDORS: - if item.lower() == raw.lower(): - return item - return "Other" - - -def _merge_tags(tags: str, *extras: str) -> str: - seen: set[str] = set() - out: list[str] = [] - for token in str(tags or "").split(): - t = token.strip() - if t and t not in seen: - seen.add(t) - out.append(t) - for extra in extras: - t = str(extra or "").strip() - if t and t not in seen: - seen.add(t) - out.append(t) - return " ".join(out) - - -def _infer_managed_ne_type_vendor(ne_type: str, vendor: str) -> tuple[str, str]: - raw_vendor = _normalize_vendor(vendor) - text = str(ne_type or "").strip() - for pattern, device_type, inferred_vendor in _BUILTIN_NE_TYPE_RULES: - if pattern.search(text): - dt = device_type if device_type in SUPPORTED_DEVICE_TYPES else "zte_zxros" - return dt, _normalize_vendor(inferred_vendor or raw_vendor) - if raw_vendor == "Huawei": - return "huawei", "Huawei" - if raw_vendor == "Cisco": - return "cisco_ios", "Cisco" - if raw_vendor == "ZTE": - return "zte_zxros", "ZTE" - return "zte_zxros", raw_vendor - - -def _validate_hop_on_create(body: ManagedNeCreate) -> None: - if not body.hop_enabled: - return - if not str(body.hop_host or "").strip(): - raise HTTPException(status_code=400, detail="hop_host_required") - if not str(body.hop_username or "").strip(): - raise HTTPException(status_code=400, detail="hop_username_required") - if not str(body.hop_password or "").strip(): - raise HTTPException(status_code=400, detail="hop_password_required") - hop_vendor = _normalize_hop_vendor(body.hop_vendor) - if hop_vendor == "bastion" and _normalize_hop_target_auth_mode(body.hop_target_auth_mode) == "manual": - if not str(body.password or "").strip(): - raise HTTPException(status_code=400, detail="password_required") - - -def _parse_import_bool(value: Any) -> bool: - return str(value or "").strip().lower() in ("1", "true", "yes", "y", "on") - - -def _import_cell_str(value: Any) -> str: - if value is None or (isinstance(value, float) and pd.isna(value)): - return "" - text = str(value).strip() - return "" if text.lower() == "nan" else text - - -def _apply_hop_create(row: ManagedNE, body: ManagedNeCreate) -> None: - row.hop_enabled = bool(body.hop_enabled) - row.hop_vendor = _normalize_hop_vendor(body.hop_vendor) - row.hop_host = str(body.hop_host or "").strip() - row.hop_port = int(body.hop_port or 22) - row.hop_protocol = _normalize_protocol(body.hop_protocol) - row.hop_username = str(body.hop_username or "").strip() - row.hop_password_enc = encrypt_secret(body.hop_password) if body.hop_enabled else "" - row.hop_command_template = str(body.hop_command_template or "").strip() - row.hop_vrf = str(body.hop_vrf or "").strip() - row.hop_target_auth_mode = _normalize_hop_target_auth_mode(body.hop_target_auth_mode) - - -def _apply_hop_update(row: ManagedNE, data: dict[str, Any]) -> None: - if "hop_enabled" in data and data["hop_enabled"] is not None: - row.hop_enabled = bool(data["hop_enabled"]) - if "hop_vendor" in data and data["hop_vendor"] is not None: - row.hop_vendor = _normalize_hop_vendor(data["hop_vendor"]) - if "hop_host" in data and data["hop_host"] is not None: - row.hop_host = str(data["hop_host"]).strip() - if "hop_port" in data and data["hop_port"] is not None: - row.hop_port = int(data["hop_port"]) - if "hop_protocol" in data and data["hop_protocol"] is not None: - row.hop_protocol = _normalize_protocol(data["hop_protocol"]) - if "hop_username" in data and data["hop_username"] is not None: - row.hop_username = str(data["hop_username"]).strip() - if "hop_password" in data and data["hop_password"]: - _require_crypto() - row.hop_password_enc = encrypt_secret(str(data["hop_password"])) - if "hop_command_template" in data and data["hop_command_template"] is not None: - row.hop_command_template = str(data["hop_command_template"]).strip() - if "hop_vrf" in data and data["hop_vrf"] is not None: - row.hop_vrf = str(data["hop_vrf"]).strip() - if "hop_target_auth_mode" in data and data["hop_target_auth_mode"] is not None: - row.hop_target_auth_mode = _normalize_hop_target_auth_mode(data["hop_target_auth_mode"]) - if row.hop_enabled: - if not str(row.hop_host or "").strip(): - raise HTTPException(status_code=400, detail="hop_host_required") - if not str(row.hop_username or "").strip(): - raise HTTPException(status_code=400, detail="hop_username_required") - if ( - not str(row.hop_password_enc or "").strip() - and _normalize_hop_target_auth_mode(row.hop_target_auth_mode) != "bastion_managed" - ): - raise HTTPException(status_code=400, detail="hop_password_required") - - -def row_to_out(row: ManagedNE) -> ManagedNeOut: - status = str(row.connect_status or "unknown") - if status not in ("unknown", "testing", "pass", "fail"): - status = "unknown" - return ManagedNeOut( - id=str(row.id), - name=str(row.name or ""), - vendor=str(row.vendor or "Other"), - device_type=str(row.device_type or ""), - ip_address=str(row.ip_address or ""), - port=int(row.port or 22), - protocol=str(row.protocol or "ssh"), - username=str(row.username or ""), - connect_status=status, # type: ignore[arg-type] - connect_message=str(row.connect_message or "")[:500], - connect_detail=str(row.connect_detail or "")[:8000], - connect_tested_at=row.connect_tested_at, - tags=str(row.tags or ""), - remark=str(row.remark or ""), - source=str(row.source or ""), - source_ref=str(row.source_ref or ""), - hop_enabled=bool(row.hop_enabled), - hop_vendor=str(row.hop_vendor or "zte"), - hop_host=str(row.hop_host or ""), - hop_port=int(row.hop_port or 22), - hop_protocol=str(row.hop_protocol or "ssh"), - hop_username=str(row.hop_username or ""), - hop_command_template=str(row.hop_command_template or ""), - hop_vrf=str(row.hop_vrf or ""), - hop_target_auth_mode=str(row.hop_target_auth_mode or "bastion_managed"), - created_at=row.created_at, - updated_at=row.updated_at, - ) - - -def list_managed_ne( - db: Session, - *, - keyword: str | None = None, - vendor: str | None = None, - connect_status: str | None = None, - page: int = 1, - page_size: int = 50, -) -> dict[str, Any]: - stmt = db.query(ManagedNE) - kw = str(keyword or "").strip() - v = str(vendor or "").strip() - cs = str(connect_status or "").strip() - if kw: - like = f"%{kw}%" - stmt = stmt.filter( - or_( - ManagedNE.name.ilike(like), - ManagedNE.ip_address.ilike(like), - ManagedNE.username.ilike(like), - ManagedNE.tags.ilike(like), - ManagedNE.vendor.ilike(like), - ManagedNE.device_type.ilike(like), - ) - ) - if v: - stmt = stmt.filter(ManagedNE.vendor == v) - if cs: - stmt = stmt.filter(ManagedNE.connect_status == cs) - total = int(stmt.count()) - rows = ( - stmt.order_by(ManagedNE.updated_at.desc()) - .offset((page - 1) * page_size) - .limit(page_size) - .all() - ) - return { - "total": total, - "page": page, - "page_size": page_size, - "items": [row_to_out(x).model_dump() for x in rows], - } - - -def get_managed_ne(db: Session, ne_id: str) -> ManagedNeOut: - row = db.get(ManagedNE, ne_id) - if not row: - raise HTTPException(status_code=404, detail="managed_ne_not_found") - return row_to_out(row) - - -def create_managed_ne(db: Session, body: ManagedNeCreate) -> ManagedNeOut: - _require_crypto() - _validate_hop_on_create(body) - ip = _normalize_ip(body.ip_address) - if not ip: - raise HTTPException(status_code=400, detail="ip_address_required") - if body.device_type not in SUPPORTED_DEVICE_TYPES: - raise HTTPException(status_code=400, detail="unsupported_device_type") - existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() - if existing: - raise HTTPException(status_code=400, detail="ip_address_exists") - now = _now() - row = ManagedNE( - name=str(body.name or "").strip() or ip, - vendor=body.vendor, - device_type=body.device_type, - ip_address=ip, - port=int(body.port or 22), - protocol=_normalize_protocol(body.protocol), - username=str(body.username or "").strip(), - password_enc=encrypt_secret(body.password) if str(body.password or "").strip() else "", - enable_secret_enc="", - connect_status="unknown", - tags=str(body.tags or "").strip(), - remark=str(body.remark or "").strip(), - source="", - source_ref="", - created_at=now, - updated_at=now, - ) - _apply_hop_create(row, body) - db.add(row) - db.commit() - db.refresh(row) - return row_to_out(row) - - -def _normalize_webcrt_device_type(device_type: str) -> str: - dt = str(device_type or "").strip() - low = dt.lower() - if low in ("linux", "linux_ssh", "linux_telnet"): - return "linux" - if low in ("generic", "generic_ssh", "generic_telnet", "terminal_server", "generic_termserver"): - return "generic" - return dt - - -def upsert_webcrt_managed_ne(db: Session, body: ManagedNeCreate) -> tuple[ManagedNeOut, str]: - """Create/update a WebCRT-origin NE, or reuse an existing inventory NE by IP. - - Returns ``(ne_out, action)`` where action is ``created`` | ``updated`` | ``reused``. - """ - _require_crypto() - _validate_hop_on_create(body) - ip = _normalize_ip(body.ip_address) - if not ip: - raise HTTPException(status_code=400, detail="ip_address_required") - if not str(body.username or "").strip(): - raise HTTPException(status_code=400, detail="cli_username_required") - device_type = _normalize_webcrt_device_type(body.device_type) - if device_type not in WEBCRT_DEVICE_TYPES: - raise HTTPException(status_code=400, detail="unsupported_device_type") - - existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() - now = _now() - - if existing is not None: - src = str(existing.source or "").strip() - if src != WEBCRT_NE_SOURCE: - # Do not overwrite inventory / UME-synced assets; just open them. - return row_to_out(existing), "reused" - - existing.name = str(body.name or "").strip() or existing.name or ip - existing.vendor = _normalize_vendor(body.vendor) if str(body.vendor or "").strip() else ( - "Other" if device_type == "linux" else existing.vendor - ) - existing.device_type = device_type - existing.port = int(body.port or existing.port or 22) - existing.protocol = _normalize_protocol(body.protocol) - existing.username = str(body.username or "").strip() - if str(body.password or "").strip(): - existing.password_enc = encrypt_secret(body.password) - elif not str(existing.password_enc or "").strip() and not ( - body.hop_enabled - and _normalize_hop_vendor(body.hop_vendor) == "bastion" - and _normalize_hop_target_auth_mode(body.hop_target_auth_mode) == "bastion_managed" - ): - raise HTTPException(status_code=400, detail="password_required") - existing.source = WEBCRT_NE_SOURCE - _apply_hop_create(existing, body) - existing.updated_at = now - db.commit() - db.refresh(existing) - return row_to_out(existing), "updated" - - if not str(body.password or "").strip() and not ( - body.hop_enabled - and _normalize_hop_vendor(body.hop_vendor) == "bastion" - and _normalize_hop_target_auth_mode(body.hop_target_auth_mode) == "bastion_managed" - ): - raise HTTPException(status_code=400, detail="password_required") - - vendor = _normalize_vendor(body.vendor) - if device_type == "linux" and not str(body.vendor or "").strip(): - vendor = "Other" - - row = ManagedNE( - name=str(body.name or "").strip() or ip, - vendor=vendor, - device_type=device_type, - ip_address=ip, - port=int(body.port or 22), - protocol=_normalize_protocol(body.protocol), - username=str(body.username or "").strip(), - password_enc=encrypt_secret(body.password) if str(body.password or "").strip() else "", - enable_secret_enc="", - connect_status="unknown", - tags=str(body.tags or "").strip(), - remark=str(body.remark or "").strip(), - source=WEBCRT_NE_SOURCE, - source_ref="", - created_at=now, - updated_at=now, - ) - _apply_hop_create(row, body) - db.add(row) - db.commit() - db.refresh(row) - return row_to_out(row), "created" - - -def _next_webcrt_session_name(db: Session, base: str) -> str: - """Return base, or ``base (1)``, ``base (2)``, … among WebCRT session names.""" - root = str(base or "").strip() or "session" - rows = ( - db.query(ManagedNE.name) - .filter(ManagedNE.source == WEBCRT_NE_SOURCE) - .all() - ) - taken = {str(r[0] or "").strip() for r in rows if str(r[0] or "").strip()} - if root not in taken: - return root - n = 1 - while f"{root} ({n})" in taken: - n += 1 - return f"{root} ({n})" - - -def upsert_webcrt_session_host( - db: Session, - *, - name: str = "", - ip_address: str, - port: int = 22, - protocol: str = "ssh", - username: str = "", - password: str = "", - save_password: bool = False, -) -> tuple[ManagedNeOut, str]: - """Create a WebCRT session host (linux, no hop). Always inserts a new row. - - Same IP is allowed; session name auto-suffixes ``(1)``, ``(2)``, … on collision. - Telnet never persists a password. SSH persists password only when ``save_password``. - Returns ``(ne_out, \"created\")``. - """ - _require_crypto() - ip = _normalize_ip(ip_address) - if not ip: - raise HTTPException(status_code=400, detail="ip_address_required") - proto = _normalize_protocol(protocol) - user = str(username or "").strip() - pwd = str(password or "") - if proto == "ssh" and not user: - raise HTTPException(status_code=400, detail="cli_username_required") - if proto == "ssh" and save_password and not pwd.strip(): - raise HTTPException(status_code=400, detail="password_required") - - now = _now() - display_name = _next_webcrt_session_name(db, str(name or "").strip() or ip) - - password_enc = "" - if proto == "ssh" and save_password and pwd.strip(): - password_enc = encrypt_secret(pwd) - - row = ManagedNE( - name=display_name, - vendor="Other", - # generic → Netmiko terminal_server: SSH auth then raw PTY (no linux session prep). - device_type="generic", - ip_address=ip, - port=int(port or (23 if proto == "telnet" else 22)), - protocol=proto, - username=user, - password_enc=password_enc, - enable_secret_enc="", - connect_status="unknown", - tags="", - remark="", - source=WEBCRT_NE_SOURCE, - source_ref="", - created_at=now, - updated_at=now, - ) - db.add(row) - try: - db.commit() - except Exception as exc: - db.rollback() - # Stale unique index on ip_address → restart API after migration, or drop constraint manually. - from sqlalchemy.exc import IntegrityError - - if isinstance(exc, IntegrityError): - raise HTTPException( - status_code=409, - detail="ip_address_conflict_restart_required", - ) from exc - raise - db.refresh(row) - return row_to_out(row), "created" - - -def update_managed_ne(db: Session, ne_id: str, body: ManagedNeUpdate) -> ManagedNeOut: - row = db.get(ManagedNE, ne_id) - if not row: - raise HTTPException(status_code=404, detail="managed_ne_not_found") - data = body.model_dump(exclude_unset=True) - if "ip_address" in data: - ip = _normalize_ip(data["ip_address"]) - if not ip: - raise HTTPException(status_code=400, detail="ip_address_required") - other = db.query(ManagedNE).filter(ManagedNE.ip_address == ip, ManagedNE.id != ne_id).first() - if other: - raise HTTPException(status_code=400, detail="ip_address_exists") - row.ip_address = ip - if "device_type" in data: - if data["device_type"] not in SUPPORTED_DEVICE_TYPES: - raise HTTPException(status_code=400, detail="unsupported_device_type") - row.device_type = data["device_type"] - if "vendor" in data: - v = str(data["vendor"] or "").strip() - row.vendor = v if v in SUPPORTED_VENDORS else "Other" - if "name" in data: - row.name = str(data["name"] or "").strip() - if "port" in data and data["port"] is not None: - row.port = int(data["port"]) - if "protocol" in data and data["protocol"] is not None: - row.protocol = _normalize_protocol(data["protocol"]) - if "username" in data and data["username"] is not None: - row.username = str(data["username"]).strip() - if "tags" in data and data["tags"] is not None: - row.tags = str(data["tags"]).strip() - if "remark" in data and data["remark"] is not None: - row.remark = str(data["remark"]).strip() - if "password" in data and data["password"]: - _require_crypto() - row.password_enc = encrypt_secret(str(data["password"])) - hop_keys = ( - "hop_enabled", - "hop_vendor", - "hop_host", - "hop_port", - "hop_protocol", - "hop_username", - "hop_password", - "hop_command_template", - "hop_vrf", - "hop_target_auth_mode", - ) - if any(k in data for k in hop_keys): - _apply_hop_update(row, data) - row.updated_at = _now() - db.commit() - db.refresh(row) - return row_to_out(row) - - -def batch_apply_hop_proxy(db: Session, ids: list[str], hop: HopProxyConfig) -> dict[str, Any]: - """Apply the same jump-host (proxy) settings to multiple managed NEs.""" - hop_host = str(hop.hop_host or "").strip() - hop_user = str(hop.hop_username or "").strip() - hop_pass = str(hop.hop_password or "").strip() - if hop_pass: - _require_crypto() - if not hop_host: - raise HTTPException(status_code=400, detail="hop_host_required") - if not hop_user: - raise HTTPException(status_code=400, detail="hop_username_required") - hop_auth_mode = _normalize_hop_target_auth_mode(hop.hop_target_auth_mode) - if not hop_pass and hop_auth_mode != "bastion_managed": - raise HTTPException(status_code=400, detail="hop_password_required") - - hop_vendor = _normalize_hop_vendor(hop.hop_vendor) - template = str(hop.hop_command_template or "").strip() - if hop_vendor == "bastion" and not template: - template = default_bastion_username_template() - elif hop_vendor not in ("linux", "bastion") and not template: - template = default_hop_command_template(hop_vendor, hop.hop_protocol, hop.hop_vrf) - - ne_ids = [str(x).strip() for x in ids if str(x).strip()] - if not ne_ids: - raise HTTPException(status_code=400, detail="ids_required") - - rows = db.query(ManagedNE).filter(ManagedNE.id.in_(ne_ids)).all() - found_ids = {str(r.id) for r in rows} - missing = [x for x in ne_ids if x not in found_ids] - if missing: - raise HTTPException(status_code=404, detail=f"managed_ne_not_found: {','.join(missing[:5])}") - - now = _now() - for row in rows: - row.hop_enabled = True - row.hop_vendor = hop_vendor - row.hop_host = hop_host - row.hop_port = int(hop.hop_port or 22) - row.hop_protocol = _normalize_protocol(hop.hop_protocol) - row.hop_username = hop_user - if hop_pass: - row.hop_password_enc = encrypt_secret(hop_pass) - row.hop_command_template = template - row.hop_vrf = str(hop.hop_vrf or "").strip() - row.hop_target_auth_mode = hop_auth_mode - row.updated_at = now - db.commit() - return {"ok": True, "updated": len(rows)} - - -def batch_apply_account(db: Session, ids: list[str], account: BatchAccountConfig) -> dict[str, Any]: - user = str(account.username or "").strip() - pwd = str(account.password or "") - if not user and not pwd: - raise HTTPException(status_code=400, detail="username_or_password_required") - if pwd: - _require_crypto() - pwd_enc = encrypt_secret(pwd) - else: - pwd_enc = "" - ne_ids = [str(x).strip() for x in ids if str(x).strip()] - if not ne_ids: - raise HTTPException(status_code=400, detail="ids_required") - rows = db.query(ManagedNE).filter(ManagedNE.id.in_(ne_ids)).all() - found_ids = {str(r.id) for r in rows} - missing = [x for x in ne_ids if x not in found_ids] - if missing: - raise HTTPException(status_code=404, detail=f"managed_ne_not_found: {','.join(missing[:5])}") - now = _now() - for row in rows: - if user: - row.username = user - if pwd: - row.password_enc = pwd_enc - row.updated_at = now - db.commit() - return {"ok": True, "updated": len(rows)} - - -def delete_managed_ne(db: Session, ne_id: str) -> dict[str, bool]: - from .topology_inventory_lifecycle import detach_fabric_from_managed - - row = db.get(ManagedNE, ne_id) - if not row: - raise HTTPException(status_code=404, detail="managed_ne_not_found") - detach_fabric_from_managed(db, [str(row.id)]) - db.delete(row) - db.commit() - return {"ok": True} - - -def get_managed_ne_stats(db: Session) -> dict[str, Any]: - """Return total counts by connect_status, and tag statistics.""" - from sqlalchemy import func - - rows = db.query(ManagedNE.connect_status, func.count(ManagedNE.id)).group_by(ManagedNE.connect_status).all() - by_status: dict[str, int] = {} - total = 0 - for status, cnt in rows: - by_status[str(status or "unknown")] = int(cnt) - total += int(cnt) - - # Tag statistics & per-tag connect_status aggregation (space-separated) - def _bump(bucket: dict[str, int], status: str) -> None: - s = str(status or "unknown") - bucket[s] = int(bucket.get(s, 0)) + 1 - - tag_counts: dict[str, int] = {} - no_tag_count = 0 - per_tag_by_status: dict[str, dict[str, int]] = {} - per_tag_total: dict[str, int] = {} - - for connect_status, tags_str in db.query(ManagedNE.connect_status, ManagedNE.tags).all(): - status = str(connect_status or "unknown") - tags_val = str(tags_str or "").strip() - if not tags_val: - no_tag_count += 1 - per_tag_total["__no_tag__"] = int(per_tag_total.get("__no_tag__", 0)) + 1 - per_tag_by_status.setdefault("__no_tag__", {}) - _bump(per_tag_by_status["__no_tag__"], status) - continue - for t in tags_val.split(): - if not t: - continue - tag_counts[t] = int(tag_counts.get(t, 0)) + 1 - per_tag_total[t] = int(per_tag_total.get(t, 0)) + 1 - per_tag_by_status.setdefault(t, {}) - _bump(per_tag_by_status[t], status) - - return { - "total": total, - "by_status": by_status, - "no_tag_count": int(no_tag_count), - "tag_counts": {k: int(tag_counts[k]) for k in sorted(tag_counts.keys())}, - "tags": sorted(tag_counts.keys()), - "per_tag": { - k: {"total": int(per_tag_total.get(k, 0)), "by_status": per_tag_by_status.get(k, {})} - for k in sorted(per_tag_total.keys(), key=lambda x: ("0" if x == "__no_tag__" else "1") + x) - }, - } - - -def get_ids_by_tag(db: Session, tag: str | None) -> list[str]: - """ - Return NE ids by tag. - - - tag is None: all ids - - tag == "__no_tag__": ids where tags is empty/blank - - otherwise: ids where tag exists in space-separated tags list - """ - result: list[str] = [] - norm = str(tag).strip() if tag is not None else None - for ne_id, tags_str in db.query(ManagedNE.id, ManagedNE.tags).all(): - tags_val = str(tags_str or "").strip() - if norm is None: - result.append(str(ne_id)) - elif norm == "__no_tag__": - if not tags_val: - result.append(str(ne_id)) - else: - if norm in tags_val.split(): - result.append(str(ne_id)) - return result - - -def batch_delete_managed_ne(db: Session, ids: list[str]) -> dict[str, Any]: - from .topology_inventory_lifecycle import detach_fabric_from_managed - - ne_ids = [str(x).strip() for x in ids if str(x).strip()] - if not ne_ids: - raise HTTPException(status_code=400, detail="ids_required") - rows = db.query(ManagedNE).filter(ManagedNE.id.in_(ne_ids)).all() - found_ids = {str(r.id) for r in rows} - missing = [x for x in ne_ids if x not in found_ids] - if missing: - raise HTTPException(status_code=404, detail=f"managed_ne_not_found: {','.join(missing[:5])}") - detach_fabric_from_managed(db, [str(r.id) for r in rows]) - for row in rows: - db.delete(row) - db.commit() - return {"ok": True, "deleted": len(rows)} - - -def sync_ume_inventory_to_managed_ne(db: Session) -> UmeManagedSyncResult: - rows = db.query(UmeInventoryNE).all() - by_source_ref = { - str(x.source_ref or ""): x - for x in db.query(ManagedNE).filter(ManagedNE.source == UME_SYNC_SOURCE).all() - } - inventory_ids = {str(x.ne_id or "").strip() for x in rows if str(x.ne_id or "").strip()} - inserted = 0 - updated = 0 - now = _now() - for inv in rows: - source_ref = str(inv.ne_id or "").strip() - ip = _normalize_ip(str(inv.ip_address or "")) - if not source_ref or not ip: - continue - existing = by_source_ref.get(source_ref) - if existing is None: - existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() - device_type, vendor = _infer_managed_ne_type_vendor(str(inv.ne_type or ""), str(inv.vendor or "")) - display_name = str(inv.host_name or "").strip() or str(inv.ne_name or "").strip() or ip - existing_tags = str(existing.tags or "").strip() if existing is not None else "" - if existing is None: - existing = ManagedNE( - ip_address=ip, - created_at=now, - source=UME_SYNC_SOURCE, - source_ref=source_ref, - ) - db.add(existing) - inserted += 1 - else: - updated += 1 - existing.name = display_name - existing.vendor = vendor - existing.device_type = device_type - existing.port = int(existing.port or 22 or 22) - existing.protocol = _normalize_protocol(str(existing.protocol or "ssh")) - existing.tags = _merge_tags(existing_tags, UME_SYNC_TAG) - existing.source = UME_SYNC_SOURCE - existing.source_ref = source_ref - existing.updated_at = now - from .topology_inventory_lifecycle import detach_fabric_from_managed - - stale = [ - row - for row in db.query(ManagedNE).filter(ManagedNE.source == UME_SYNC_SOURCE).all() - if (not str(row.source_ref or "").strip()) - or str(row.source_ref or "").strip() not in inventory_ids - ] - if stale: - detach_fabric_from_managed(db, [str(r.id) for r in stale]) - for row in stale: - db.delete(row) - deleted = len(stale) - db.commit() - return UmeManagedSyncResult( - inserted=inserted, - updated=updated, - deleted=deleted, - total_inventory=len(inventory_ids), - ) - - -def delete_ume_synced_managed_ne(db: Session) -> UmeManagedDeleteResult: - from .topology_inventory_lifecycle import detach_fabric_from_managed - - rows = db.query(ManagedNE).filter(ManagedNE.source == UME_SYNC_SOURCE).all() - deleted = len(rows) - if rows: - detach_fabric_from_managed(db, [str(r.id) for r in rows]) - for row in rows: - db.delete(row) - db.commit() - return UmeManagedDeleteResult(deleted=deleted) - - -def build_managed_ne_import_template(fmt: str = "xlsx") -> tuple[str, bytes, str]: - """Return (filename, content, media_type) for bulk-import template.""" - rows = [ - { - "device_type": "cisco_ios", - "ip": "192.168.0.1", - "username": "admin", - "password": "your_password", - "port": 22, - "protocol": "ssh", - "name": "Core-SW1", - "vendor": "Cisco", - "tags": "core", - "remark": "", - }, - { - "device_type": "zte_zxros", - "ip": "2.2.2.2", - "username": "target-user", - "password": "", - "port": 22, - "protocol": "ssh", - "name": "PE-01", - "vendor": "ZTE", - "tags": "edge bastion", - "remark": "no direct password, use batch proxy", - }, - ] - df = pd.DataFrame(rows, columns=list(IMPORT_COLUMNS)) - buf = BytesIO() - kind = str(fmt or "xlsx").strip().lower() - if kind == "csv": - df.to_csv(buf, index=False, encoding="utf-8-sig") - return ( - "managed_ne_import_template.csv", - buf.getvalue(), - "text/csv; charset=utf-8", - ) - device_types_df = pd.DataFrame({"device_type": list(SUPPORTED_DEVICE_TYPES)}) - with pd.ExcelWriter(buf, engine="openpyxl") as writer: - df.to_excel(writer, sheet_name="import", index=False) - device_types_df.to_excel(writer, sheet_name="device_type", index=False) - return ( - "managed_ne_import_template.xlsx", - buf.getvalue(), - "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", - ) - - -def import_managed_ne(db: Session, content: bytes, filename: str) -> ImportResult: - _require_crypto() - name = str(filename or "").lower() - try: - if name.endswith(".csv"): - df = pd.read_csv(BytesIO(content)) - else: - df = pd.read_excel(BytesIO(content)) - except Exception as exc: - raise HTTPException(status_code=400, detail=f"import_parse_failed: {exc}") from exc - df.columns = [str(c).strip().lower() for c in df.columns] - missing = [c for c in IMPORT_COLUMNS if c not in df.columns] - if missing: - raise HTTPException(status_code=400, detail=f"import_missing_columns: {','.join(missing)}") - inserted = 0 - updated = 0 - failed: list[ImportFailure] = [] - for idx, row in df.iterrows(): - row_no = int(idx) + 2 - try: - ip = _normalize_ip(_import_cell_str(row.get("ip", ""))) - if not ip: - failed.append(ImportFailure(row=row_no, reason="ip_required")) - continue - device_type = _import_cell_str(row.get("device_type", "")) - if device_type not in SUPPORTED_DEVICE_TYPES: - failed.append(ImportFailure(row=row_no, reason="unsupported_device_type")) - continue - username = _import_cell_str(row.get("username", "")) - password = _import_cell_str(row.get("password", "")) - if not username: - failed.append(ImportFailure(row=row_no, reason="username_required")) - continue - port_raw = row.get("port", 22) - try: - port = int(port_raw) - except (TypeError, ValueError): - port = 22 - protocol = _normalize_protocol(str(row.get("protocol", "ssh"))) - display_name = _import_cell_str(row.get("name", "")) or ip - vendor_raw = _import_cell_str(row.get("vendor", "")) or "Other" - vendor = "Other" - for v in SUPPORTED_VENDORS: - if v.lower() == vendor_raw.lower(): - vendor = v - break - existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() - now = _now() - if existing is None: - existing = ManagedNE(ip_address=ip, created_at=now) - db.add(existing) - inserted += 1 - else: - updated += 1 - existing.name = display_name - existing.vendor = vendor - existing.device_type = device_type - existing.port = port - existing.protocol = protocol - existing.username = username - existing.password_enc = encrypt_secret(password) if password else "" - tags_val = _import_cell_str(row.get("tags", "")) - remark_val = _import_cell_str(row.get("remark", "")) - if tags_val: - existing.tags = tags_val - if remark_val: - existing.remark = remark_val - existing.updated_at = now - except CredentialCryptoError as exc: - failed.append(ImportFailure(row=row_no, reason=str(exc))) - except Exception as exc: - failed.append(ImportFailure(row=row_no, reason=str(exc)[:200])) - db.commit() - return ImportResult(inserted=inserted, updated=updated, failed=failed) - - -def get_device_credentials(row: ManagedNE) -> dict[str, Any]: - hop_enabled = bool(row.hop_enabled) - hop_password = "" - if hop_enabled and str(row.hop_password_enc or "").strip(): - hop_password = decrypt_secret(row.hop_password_enc) - return { - "id": str(row.id), - "vendor": str(row.vendor or ""), - "device_type": str(row.device_type or ""), - "ip_address": str(row.ip_address or ""), - "port": int(row.port or 22), - "protocol": str(row.protocol or "ssh"), - "username": str(row.username or ""), - "password": decrypt_secret(row.password_enc), - "enable_secret": decrypt_secret(row.enable_secret_enc), - "name": str(row.name or ""), - "hop_enabled": hop_enabled, - "hop_vendor": str(row.hop_vendor or "zte"), - "hop_host": str(row.hop_host or ""), - "hop_port": int(row.hop_port or 22), - "hop_protocol": str(row.hop_protocol or "ssh"), - "hop_username": str(row.hop_username or ""), - "hop_password": hop_password, - "hop_command_template": str(row.hop_command_template or ""), - "hop_vrf": str(row.hop_vrf or ""), - "hop_target_auth_mode": str(row.hop_target_auth_mode or "bastion_managed"), - } diff --git a/netx_api/ne_service_common.py b/netx_api/ne_service_common.py new file mode 100644 index 0000000..dfc7ebc --- /dev/null +++ b/netx_api/ne_service_common.py @@ -0,0 +1,252 @@ +"""Managed NE shared helpers, constants, and credential extraction.""" +from __future__ import annotations + +from datetime import datetime +import re +from typing import Any + +import pandas as pd +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .device_types import ( + SUPPORTED_DEVICE_TYPES, + SUPPORTED_VENDORS, + WEBCRT_DEVICE_TYPES, + WEBCRT_NE_SOURCE, +) +from .models import ManagedNE +from .ne_crypto import CredentialCryptoError, credentials_configured, decrypt_secret, encrypt_secret +from .ne_schemas import ManagedNeCreate, ManagedNeOut, ManagedNeUpdate +from .ne_session_factory import default_bastion_username_template, default_hop_command_template + +IMPORT_COLUMNS = ( + "device_type", + "ip", + "username", + "password", + "port", + "protocol", + "name", + "vendor", + "tags", + "remark", +) + +UME_SYNC_SOURCE = "ume_sync" +UME_SYNC_TAG = "UME" +# Re-export for callers (WebCRT quick-connect). +WEBCRT_SOURCE = WEBCRT_NE_SOURCE +_BUILTIN_NE_TYPE_RULES: list[tuple[re.Pattern[str], str, str]] = [ + (re.compile(r"ZXR|ZXCTN|M6000|\bBN\b", re.I), "zte_zxros", "ZTE"), + (re.compile(r"NE40|CE\b|ATN|MA5800|OptiX", re.I), "huawei", "Huawei"), + (re.compile(r"ASR|NCS|IOS.?XR|XR\b", re.I), "cisco_xr", "Cisco"), + (re.compile(r"Catalyst|Nexus|C9[0-9]{3}|ISR", re.I), "cisco_ios", "Cisco"), +] + +def _now() -> datetime: + return datetime.utcnow() + + +def _require_crypto() -> None: + if not credentials_configured(): + raise HTTPException(status_code=503, detail="credential_secret_key_not_configured") + + +def _normalize_ip(ip: str) -> str: + return str(ip or "").strip() + + +def _normalize_protocol(protocol: str) -> str: + p = str(protocol or "ssh").strip().lower() + return p if p in ("ssh", "telnet") else "ssh" + + +def _normalize_hop_vendor(vendor: str) -> str: + v = str(vendor or "zte").strip().lower() + return v if v in ("zte", "linux", "huawei", "cisco", "bastion") else "zte" + + +def _normalize_hop_target_auth_mode(mode: str) -> str: + m = str(mode or "bastion_managed").strip().lower() + return m if m in ("bastion_managed", "manual") else "bastion_managed" + + +def _normalize_vendor(vendor: str) -> str: + raw = str(vendor or "").strip() + if not raw: + return "Other" + for item in SUPPORTED_VENDORS: + if item.lower() == raw.lower(): + return item + return "Other" + + +def _merge_tags(tags: str, *extras: str) -> str: + seen: set[str] = set() + out: list[str] = [] + for token in str(tags or "").split(): + t = token.strip() + if t and t not in seen: + seen.add(t) + out.append(t) + for extra in extras: + t = str(extra or "").strip() + if t and t not in seen: + seen.add(t) + out.append(t) + return " ".join(out) + + +def _infer_managed_ne_type_vendor(ne_type: str, vendor: str) -> tuple[str, str]: + raw_vendor = _normalize_vendor(vendor) + text = str(ne_type or "").strip() + for pattern, device_type, inferred_vendor in _BUILTIN_NE_TYPE_RULES: + if pattern.search(text): + dt = device_type if device_type in SUPPORTED_DEVICE_TYPES else "zte_zxros" + return dt, _normalize_vendor(inferred_vendor or raw_vendor) + if raw_vendor == "Huawei": + return "huawei", "Huawei" + if raw_vendor == "Cisco": + return "cisco_ios", "Cisco" + if raw_vendor == "ZTE": + return "zte_zxros", "ZTE" + return "zte_zxros", raw_vendor + + +def _validate_hop_on_create(body: ManagedNeCreate) -> None: + if not body.hop_enabled: + return + if not str(body.hop_host or "").strip(): + raise HTTPException(status_code=400, detail="hop_host_required") + if not str(body.hop_username or "").strip(): + raise HTTPException(status_code=400, detail="hop_username_required") + if not str(body.hop_password or "").strip(): + raise HTTPException(status_code=400, detail="hop_password_required") + hop_vendor = _normalize_hop_vendor(body.hop_vendor) + if hop_vendor == "bastion" and _normalize_hop_target_auth_mode(body.hop_target_auth_mode) == "manual": + if not str(body.password or "").strip(): + raise HTTPException(status_code=400, detail="password_required") + + +def _parse_import_bool(value: Any) -> bool: + return str(value or "").strip().lower() in ("1", "true", "yes", "y", "on") + + +def _import_cell_str(value: Any) -> str: + if value is None or (isinstance(value, float) and pd.isna(value)): + return "" + text = str(value).strip() + return "" if text.lower() == "nan" else text + + +def _apply_hop_create(row: ManagedNE, body: ManagedNeCreate) -> None: + row.hop_enabled = bool(body.hop_enabled) + row.hop_vendor = _normalize_hop_vendor(body.hop_vendor) + row.hop_host = str(body.hop_host or "").strip() + row.hop_port = int(body.hop_port or 22) + row.hop_protocol = _normalize_protocol(body.hop_protocol) + row.hop_username = str(body.hop_username or "").strip() + row.hop_password_enc = encrypt_secret(body.hop_password) if body.hop_enabled else "" + row.hop_command_template = str(body.hop_command_template or "").strip() + row.hop_vrf = str(body.hop_vrf or "").strip() + row.hop_target_auth_mode = _normalize_hop_target_auth_mode(body.hop_target_auth_mode) + + +def _apply_hop_update(row: ManagedNE, data: dict[str, Any]) -> None: + if "hop_enabled" in data and data["hop_enabled"] is not None: + row.hop_enabled = bool(data["hop_enabled"]) + if "hop_vendor" in data and data["hop_vendor"] is not None: + row.hop_vendor = _normalize_hop_vendor(data["hop_vendor"]) + if "hop_host" in data and data["hop_host"] is not None: + row.hop_host = str(data["hop_host"]).strip() + if "hop_port" in data and data["hop_port"] is not None: + row.hop_port = int(data["hop_port"]) + if "hop_protocol" in data and data["hop_protocol"] is not None: + row.hop_protocol = _normalize_protocol(data["hop_protocol"]) + if "hop_username" in data and data["hop_username"] is not None: + row.hop_username = str(data["hop_username"]).strip() + if "hop_password" in data and data["hop_password"]: + _require_crypto() + row.hop_password_enc = encrypt_secret(str(data["hop_password"])) + if "hop_command_template" in data and data["hop_command_template"] is not None: + row.hop_command_template = str(data["hop_command_template"]).strip() + if "hop_vrf" in data and data["hop_vrf"] is not None: + row.hop_vrf = str(data["hop_vrf"]).strip() + if "hop_target_auth_mode" in data and data["hop_target_auth_mode"] is not None: + row.hop_target_auth_mode = _normalize_hop_target_auth_mode(data["hop_target_auth_mode"]) + if row.hop_enabled: + if not str(row.hop_host or "").strip(): + raise HTTPException(status_code=400, detail="hop_host_required") + if not str(row.hop_username or "").strip(): + raise HTTPException(status_code=400, detail="hop_username_required") + if ( + not str(row.hop_password_enc or "").strip() + and _normalize_hop_target_auth_mode(row.hop_target_auth_mode) != "bastion_managed" + ): + raise HTTPException(status_code=400, detail="hop_password_required") + + +def row_to_out(row: ManagedNE) -> ManagedNeOut: + status = str(row.connect_status or "unknown") + if status not in ("unknown", "testing", "pass", "fail"): + status = "unknown" + return ManagedNeOut( + id=str(row.id), + name=str(row.name or ""), + vendor=str(row.vendor or "Other"), + device_type=str(row.device_type or ""), + ip_address=str(row.ip_address or ""), + port=int(row.port or 22), + protocol=str(row.protocol or "ssh"), + username=str(row.username or ""), + connect_status=status, # type: ignore[arg-type] + connect_message=str(row.connect_message or "")[:500], + connect_detail=str(row.connect_detail or "")[:8000], + connect_tested_at=row.connect_tested_at, + tags=str(row.tags or ""), + remark=str(row.remark or ""), + source=str(row.source or ""), + source_ref=str(row.source_ref or ""), + hop_enabled=bool(row.hop_enabled), + hop_vendor=str(row.hop_vendor or "zte"), + hop_host=str(row.hop_host or ""), + hop_port=int(row.hop_port or 22), + hop_protocol=str(row.hop_protocol or "ssh"), + hop_username=str(row.hop_username or ""), + hop_command_template=str(row.hop_command_template or ""), + hop_vrf=str(row.hop_vrf or ""), + hop_target_auth_mode=str(row.hop_target_auth_mode or "bastion_managed"), + created_at=row.created_at, + updated_at=row.updated_at, + ) + + + +def get_device_credentials(row: ManagedNE) -> dict[str, Any]: + hop_enabled = bool(row.hop_enabled) + hop_password = "" + if hop_enabled and str(row.hop_password_enc or "").strip(): + hop_password = decrypt_secret(row.hop_password_enc) + return { + "id": str(row.id), + "vendor": str(row.vendor or ""), + "device_type": str(row.device_type or ""), + "ip_address": str(row.ip_address or ""), + "port": int(row.port or 22), + "protocol": str(row.protocol or "ssh"), + "username": str(row.username or ""), + "password": decrypt_secret(row.password_enc), + "enable_secret": decrypt_secret(row.enable_secret_enc), + "name": str(row.name or ""), + "hop_enabled": hop_enabled, + "hop_vendor": str(row.hop_vendor or "zte"), + "hop_host": str(row.hop_host or ""), + "hop_port": int(row.hop_port or 22), + "hop_protocol": str(row.hop_protocol or "ssh"), + "hop_username": str(row.hop_username or ""), + "hop_password": hop_password, + "hop_command_template": str(row.hop_command_template or ""), + "hop_vrf": str(row.hop_vrf or ""), + "hop_target_auth_mode": str(row.hop_target_auth_mode or "bastion_managed"), + } diff --git a/netx_api/ne_service_crud.py b/netx_api/ne_service_crud.py new file mode 100644 index 0000000..fdeb0dc --- /dev/null +++ b/netx_api/ne_service_crud.py @@ -0,0 +1,358 @@ +"""Managed NE CRUD, batch hop/account, and stats.""" +from __future__ import annotations + +from typing import Any +from uuid import uuid4 + +from fastapi import HTTPException +from sqlalchemy import or_ +from sqlalchemy.orm import Session + +from .device_types import SUPPORTED_DEVICE_TYPES +from .models import ManagedNE +from .ne_crypto import encrypt_secret +from .ne_schemas import ( + BatchAccountConfig, + HopProxyConfig, + ManagedNeCreate, + ManagedNeOut, + ManagedNeUpdate, +) +from .ne_service_common import ( + _apply_hop_create, + _apply_hop_update, + _normalize_ip, + _normalize_protocol, + _normalize_vendor, + _now, + _require_crypto, + _validate_hop_on_create, + row_to_out, +) + +def list_managed_ne( + db: Session, + *, + keyword: str | None = None, + vendor: str | None = None, + connect_status: str | None = None, + page: int = 1, + page_size: int = 50, +) -> dict[str, Any]: + stmt = db.query(ManagedNE) + kw = str(keyword or "").strip() + v = str(vendor or "").strip() + cs = str(connect_status or "").strip() + if kw: + like = f"%{kw}%" + stmt = stmt.filter( + or_( + ManagedNE.name.ilike(like), + ManagedNE.ip_address.ilike(like), + ManagedNE.username.ilike(like), + ManagedNE.tags.ilike(like), + ManagedNE.vendor.ilike(like), + ManagedNE.device_type.ilike(like), + ) + ) + if v: + stmt = stmt.filter(ManagedNE.vendor == v) + if cs: + stmt = stmt.filter(ManagedNE.connect_status == cs) + total = int(stmt.count()) + rows = ( + stmt.order_by(ManagedNE.updated_at.desc()) + .offset((page - 1) * page_size) + .limit(page_size) + .all() + ) + return { + "total": total, + "page": page, + "page_size": page_size, + "items": [row_to_out(x).model_dump() for x in rows], + } + + +def get_managed_ne(db: Session, ne_id: str) -> ManagedNeOut: + row = db.get(ManagedNE, ne_id) + if not row: + raise HTTPException(status_code=404, detail="managed_ne_not_found") + return row_to_out(row) + + +def create_managed_ne(db: Session, body: ManagedNeCreate) -> ManagedNeOut: + _require_crypto() + _validate_hop_on_create(body) + ip = _normalize_ip(body.ip_address) + if not ip: + raise HTTPException(status_code=400, detail="ip_address_required") + if body.device_type not in SUPPORTED_DEVICE_TYPES: + raise HTTPException(status_code=400, detail="unsupported_device_type") + existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() + if existing: + raise HTTPException(status_code=400, detail="ip_address_exists") + now = _now() + row = ManagedNE( + name=str(body.name or "").strip() or ip, + vendor=body.vendor, + device_type=body.device_type, + ip_address=ip, + port=int(body.port or 22), + protocol=_normalize_protocol(body.protocol), + username=str(body.username or "").strip(), + password_enc=encrypt_secret(body.password) if str(body.password or "").strip() else "", + enable_secret_enc="", + connect_status="unknown", + tags=str(body.tags or "").strip(), + remark=str(body.remark or "").strip(), + source="", + source_ref="", + created_at=now, + updated_at=now, + ) + _apply_hop_create(row, body) + db.add(row) + db.commit() + db.refresh(row) + return row_to_out(row) + + +def update_managed_ne(db: Session, ne_id: str, body: ManagedNeUpdate) -> ManagedNeOut: + row = db.get(ManagedNE, ne_id) + if not row: + raise HTTPException(status_code=404, detail="managed_ne_not_found") + data = body.model_dump(exclude_unset=True) + if "ip_address" in data: + ip = _normalize_ip(data["ip_address"]) + if not ip: + raise HTTPException(status_code=400, detail="ip_address_required") + other = db.query(ManagedNE).filter(ManagedNE.ip_address == ip, ManagedNE.id != ne_id).first() + if other: + raise HTTPException(status_code=400, detail="ip_address_exists") + row.ip_address = ip + if "device_type" in data: + if data["device_type"] not in SUPPORTED_DEVICE_TYPES: + raise HTTPException(status_code=400, detail="unsupported_device_type") + row.device_type = data["device_type"] + if "vendor" in data: + v = str(data["vendor"] or "").strip() + row.vendor = v if v in SUPPORTED_VENDORS else "Other" + if "name" in data: + row.name = str(data["name"] or "").strip() + if "port" in data and data["port"] is not None: + row.port = int(data["port"]) + if "protocol" in data and data["protocol"] is not None: + row.protocol = _normalize_protocol(data["protocol"]) + if "username" in data and data["username"] is not None: + row.username = str(data["username"]).strip() + if "tags" in data and data["tags"] is not None: + row.tags = str(data["tags"]).strip() + if "remark" in data and data["remark"] is not None: + row.remark = str(data["remark"]).strip() + if "password" in data and data["password"]: + _require_crypto() + row.password_enc = encrypt_secret(str(data["password"])) + hop_keys = ( + "hop_enabled", + "hop_vendor", + "hop_host", + "hop_port", + "hop_protocol", + "hop_username", + "hop_password", + "hop_command_template", + "hop_vrf", + "hop_target_auth_mode", + ) + if any(k in data for k in hop_keys): + _apply_hop_update(row, data) + row.updated_at = _now() + db.commit() + db.refresh(row) + return row_to_out(row) + + +def batch_apply_hop_proxy(db: Session, ids: list[str], hop: HopProxyConfig) -> dict[str, Any]: + """Apply the same jump-host (proxy) settings to multiple managed NEs.""" + hop_host = str(hop.hop_host or "").strip() + hop_user = str(hop.hop_username or "").strip() + hop_pass = str(hop.hop_password or "").strip() + if hop_pass: + _require_crypto() + if not hop_host: + raise HTTPException(status_code=400, detail="hop_host_required") + if not hop_user: + raise HTTPException(status_code=400, detail="hop_username_required") + hop_auth_mode = _normalize_hop_target_auth_mode(hop.hop_target_auth_mode) + if not hop_pass and hop_auth_mode != "bastion_managed": + raise HTTPException(status_code=400, detail="hop_password_required") + + hop_vendor = _normalize_hop_vendor(hop.hop_vendor) + template = str(hop.hop_command_template or "").strip() + if hop_vendor == "bastion" and not template: + template = default_bastion_username_template() + elif hop_vendor not in ("linux", "bastion") and not template: + template = default_hop_command_template(hop_vendor, hop.hop_protocol, hop.hop_vrf) + + ne_ids = [str(x).strip() for x in ids if str(x).strip()] + if not ne_ids: + raise HTTPException(status_code=400, detail="ids_required") + + rows = db.query(ManagedNE).filter(ManagedNE.id.in_(ne_ids)).all() + found_ids = {str(r.id) for r in rows} + missing = [x for x in ne_ids if x not in found_ids] + if missing: + raise HTTPException(status_code=404, detail=f"managed_ne_not_found: {','.join(missing[:5])}") + + now = _now() + for row in rows: + row.hop_enabled = True + row.hop_vendor = hop_vendor + row.hop_host = hop_host + row.hop_port = int(hop.hop_port or 22) + row.hop_protocol = _normalize_protocol(hop.hop_protocol) + row.hop_username = hop_user + if hop_pass: + row.hop_password_enc = encrypt_secret(hop_pass) + row.hop_command_template = template + row.hop_vrf = str(hop.hop_vrf or "").strip() + row.hop_target_auth_mode = hop_auth_mode + row.updated_at = now + db.commit() + return {"ok": True, "updated": len(rows)} + + +def batch_apply_account(db: Session, ids: list[str], account: BatchAccountConfig) -> dict[str, Any]: + user = str(account.username or "").strip() + pwd = str(account.password or "") + if not user and not pwd: + raise HTTPException(status_code=400, detail="username_or_password_required") + if pwd: + _require_crypto() + pwd_enc = encrypt_secret(pwd) + else: + pwd_enc = "" + ne_ids = [str(x).strip() for x in ids if str(x).strip()] + if not ne_ids: + raise HTTPException(status_code=400, detail="ids_required") + rows = db.query(ManagedNE).filter(ManagedNE.id.in_(ne_ids)).all() + found_ids = {str(r.id) for r in rows} + missing = [x for x in ne_ids if x not in found_ids] + if missing: + raise HTTPException(status_code=404, detail=f"managed_ne_not_found: {','.join(missing[:5])}") + now = _now() + for row in rows: + if user: + row.username = user + if pwd: + row.password_enc = pwd_enc + row.updated_at = now + db.commit() + return {"ok": True, "updated": len(rows)} + + +def delete_managed_ne(db: Session, ne_id: str) -> dict[str, bool]: + from .topology_inventory_lifecycle import detach_fabric_from_managed + + row = db.get(ManagedNE, ne_id) + if not row: + raise HTTPException(status_code=404, detail="managed_ne_not_found") + detach_fabric_from_managed(db, [str(row.id)]) + db.delete(row) + db.commit() + return {"ok": True} + + +def get_managed_ne_stats(db: Session) -> dict[str, Any]: + """Return total counts by connect_status, and tag statistics.""" + from sqlalchemy import func + + rows = db.query(ManagedNE.connect_status, func.count(ManagedNE.id)).group_by(ManagedNE.connect_status).all() + by_status: dict[str, int] = {} + total = 0 + for status, cnt in rows: + by_status[str(status or "unknown")] = int(cnt) + total += int(cnt) + + # Tag statistics & per-tag connect_status aggregation (space-separated) + def _bump(bucket: dict[str, int], status: str) -> None: + s = str(status or "unknown") + bucket[s] = int(bucket.get(s, 0)) + 1 + + tag_counts: dict[str, int] = {} + no_tag_count = 0 + per_tag_by_status: dict[str, dict[str, int]] = {} + per_tag_total: dict[str, int] = {} + + for connect_status, tags_str in db.query(ManagedNE.connect_status, ManagedNE.tags).all(): + status = str(connect_status or "unknown") + tags_val = str(tags_str or "").strip() + if not tags_val: + no_tag_count += 1 + per_tag_total["__no_tag__"] = int(per_tag_total.get("__no_tag__", 0)) + 1 + per_tag_by_status.setdefault("__no_tag__", {}) + _bump(per_tag_by_status["__no_tag__"], status) + continue + for t in tags_val.split(): + if not t: + continue + tag_counts[t] = int(tag_counts.get(t, 0)) + 1 + per_tag_total[t] = int(per_tag_total.get(t, 0)) + 1 + per_tag_by_status.setdefault(t, {}) + _bump(per_tag_by_status[t], status) + + return { + "total": total, + "by_status": by_status, + "no_tag_count": int(no_tag_count), + "tag_counts": {k: int(tag_counts[k]) for k in sorted(tag_counts.keys())}, + "tags": sorted(tag_counts.keys()), + "per_tag": { + k: {"total": int(per_tag_total.get(k, 0)), "by_status": per_tag_by_status.get(k, {})} + for k in sorted(per_tag_total.keys(), key=lambda x: ("0" if x == "__no_tag__" else "1") + x) + }, + } + + +def get_ids_by_tag(db: Session, tag: str | None) -> list[str]: + """ + Return NE ids by tag. + + - tag is None: all ids + - tag == "__no_tag__": ids where tags is empty/blank + - otherwise: ids where tag exists in space-separated tags list + """ + result: list[str] = [] + norm = str(tag).strip() if tag is not None else None + for ne_id, tags_str in db.query(ManagedNE.id, ManagedNE.tags).all(): + tags_val = str(tags_str or "").strip() + if norm is None: + result.append(str(ne_id)) + elif norm == "__no_tag__": + if not tags_val: + result.append(str(ne_id)) + else: + if norm in tags_val.split(): + result.append(str(ne_id)) + return result + + +def batch_delete_managed_ne(db: Session, ids: list[str]) -> dict[str, Any]: + from .topology_inventory_lifecycle import detach_fabric_from_managed + + ne_ids = [str(x).strip() for x in ids if str(x).strip()] + if not ne_ids: + raise HTTPException(status_code=400, detail="ids_required") + rows = db.query(ManagedNE).filter(ManagedNE.id.in_(ne_ids)).all() + found_ids = {str(r.id) for r in rows} + missing = [x for x in ne_ids if x not in found_ids] + if missing: + raise HTTPException(status_code=404, detail=f"managed_ne_not_found: {','.join(missing[:5])}") + detach_fabric_from_managed(db, [str(r.id) for r in rows]) + for row in rows: + db.delete(row) + db.commit() + return {"ok": True, "deleted": len(rows)} + + diff --git a/netx_api/ne_service_import.py b/netx_api/ne_service_import.py new file mode 100644 index 0000000..22415c8 --- /dev/null +++ b/netx_api/ne_service_import.py @@ -0,0 +1,247 @@ +"""Managed NE Excel import and UME inventory sync into managed_ne.""" +from __future__ import annotations + +from io import BytesIO +from typing import Any + +import pandas as pd +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .models import ManagedNE, UmeInventoryNE +from .ne_crypto import CredentialCryptoError, encrypt_secret +from .device_types import SUPPORTED_DEVICE_TYPES, SUPPORTED_VENDORS +from .ne_schemas import ( + ImportFailure, + ImportResult, + UmeManagedDeleteResult, + UmeManagedSyncResult, +) +from .ne_service_common import ( + IMPORT_COLUMNS, + UME_SYNC_SOURCE, + UME_SYNC_TAG, + _import_cell_str, + _infer_managed_ne_type_vendor, + _merge_tags, + _normalize_ip, + _normalize_protocol, + _normalize_vendor, + _now, + _parse_import_bool, + _require_crypto, +) + +# Template lists all columns; CSV/XLS import only requires the core set. +_REQUIRED_IMPORT_COLUMNS = ( + "device_type", + "ip", + "username", + "password", + "port", + "protocol", + "name", + "vendor", +) + +def sync_ume_inventory_to_managed_ne(db: Session) -> UmeManagedSyncResult: + rows = db.query(UmeInventoryNE).all() + by_source_ref = { + str(x.source_ref or ""): x + for x in db.query(ManagedNE).filter(ManagedNE.source == UME_SYNC_SOURCE).all() + } + inventory_ids = {str(x.ne_id or "").strip() for x in rows if str(x.ne_id or "").strip()} + inserted = 0 + updated = 0 + now = _now() + for inv in rows: + source_ref = str(inv.ne_id or "").strip() + ip = _normalize_ip(str(inv.ip_address or "")) + if not source_ref or not ip: + continue + existing = by_source_ref.get(source_ref) + if existing is None: + existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() + device_type, vendor = _infer_managed_ne_type_vendor(str(inv.ne_type or ""), str(inv.vendor or "")) + display_name = str(inv.host_name or "").strip() or str(inv.ne_name or "").strip() or ip + existing_tags = str(existing.tags or "").strip() if existing is not None else "" + if existing is None: + existing = ManagedNE( + ip_address=ip, + created_at=now, + source=UME_SYNC_SOURCE, + source_ref=source_ref, + ) + db.add(existing) + inserted += 1 + else: + updated += 1 + existing.name = display_name + existing.vendor = vendor + existing.device_type = device_type + existing.port = int(existing.port or 22 or 22) + existing.protocol = _normalize_protocol(str(existing.protocol or "ssh")) + existing.tags = _merge_tags(existing_tags, UME_SYNC_TAG) + existing.source = UME_SYNC_SOURCE + existing.source_ref = source_ref + existing.updated_at = now + from .topology_inventory_lifecycle import detach_fabric_from_managed + + stale = [ + row + for row in db.query(ManagedNE).filter(ManagedNE.source == UME_SYNC_SOURCE).all() + if (not str(row.source_ref or "").strip()) + or str(row.source_ref or "").strip() not in inventory_ids + ] + if stale: + detach_fabric_from_managed(db, [str(r.id) for r in stale]) + for row in stale: + db.delete(row) + deleted = len(stale) + db.commit() + return UmeManagedSyncResult( + inserted=inserted, + updated=updated, + deleted=deleted, + total_inventory=len(inventory_ids), + ) + + +def delete_ume_synced_managed_ne(db: Session) -> UmeManagedDeleteResult: + from .topology_inventory_lifecycle import detach_fabric_from_managed + + rows = db.query(ManagedNE).filter(ManagedNE.source == UME_SYNC_SOURCE).all() + deleted = len(rows) + if rows: + detach_fabric_from_managed(db, [str(r.id) for r in rows]) + for row in rows: + db.delete(row) + db.commit() + return UmeManagedDeleteResult(deleted=deleted) + + +def build_managed_ne_import_template(fmt: str = "xlsx") -> tuple[str, bytes, str]: + """Return (filename, content, media_type) for bulk-import template.""" + rows = [ + { + "device_type": "cisco_ios", + "ip": "192.168.0.1", + "username": "admin", + "password": "your_password", + "port": 22, + "protocol": "ssh", + "name": "Core-SW1", + "vendor": "Cisco", + "tags": "core", + "remark": "", + }, + { + "device_type": "zte_zxros", + "ip": "2.2.2.2", + "username": "target-user", + "password": "", + "port": 22, + "protocol": "ssh", + "name": "PE-01", + "vendor": "ZTE", + "tags": "edge bastion", + "remark": "no direct password, use batch proxy", + }, + ] + df = pd.DataFrame(rows, columns=list(IMPORT_COLUMNS)) + buf = BytesIO() + kind = str(fmt or "xlsx").strip().lower() + if kind == "csv": + df.to_csv(buf, index=False, encoding="utf-8-sig") + return ( + "managed_ne_import_template.csv", + buf.getvalue(), + "text/csv; charset=utf-8", + ) + device_types_df = pd.DataFrame({"device_type": list(SUPPORTED_DEVICE_TYPES)}) + with pd.ExcelWriter(buf, engine="openpyxl") as writer: + df.to_excel(writer, sheet_name="import", index=False) + device_types_df.to_excel(writer, sheet_name="device_type", index=False) + return ( + "managed_ne_import_template.xlsx", + buf.getvalue(), + "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", + ) + + +def import_managed_ne(db: Session, content: bytes, filename: str) -> ImportResult: + _require_crypto() + name = str(filename or "").lower() + try: + if name.endswith(".csv"): + df = pd.read_csv(BytesIO(content)) + else: + df = pd.read_excel(BytesIO(content)) + except Exception as exc: + raise HTTPException(status_code=400, detail=f"import_parse_failed: {exc}") from exc + df.columns = [str(c).strip().lower() for c in df.columns] + missing = [c for c in _REQUIRED_IMPORT_COLUMNS if c not in df.columns] + if missing: + raise HTTPException(status_code=400, detail=f"import_missing_columns: {','.join(missing)}") + inserted = 0 + updated = 0 + failed: list[ImportFailure] = [] + for idx, row in df.iterrows(): + row_no = int(idx) + 2 + try: + ip = _normalize_ip(_import_cell_str(row.get("ip", ""))) + if not ip: + failed.append(ImportFailure(row=row_no, reason="ip_required")) + continue + device_type = _import_cell_str(row.get("device_type", "")) + if device_type not in SUPPORTED_DEVICE_TYPES: + failed.append(ImportFailure(row=row_no, reason="unsupported_device_type")) + continue + username = _import_cell_str(row.get("username", "")) + password = _import_cell_str(row.get("password", "")) + if not username: + failed.append(ImportFailure(row=row_no, reason="username_required")) + continue + port_raw = row.get("port", 22) + try: + port = int(port_raw) + except (TypeError, ValueError): + port = 22 + protocol = _normalize_protocol(str(row.get("protocol", "ssh"))) + display_name = _import_cell_str(row.get("name", "")) or ip + vendor_raw = _import_cell_str(row.get("vendor", "")) or "Other" + vendor = "Other" + for v in SUPPORTED_VENDORS: + if v.lower() == vendor_raw.lower(): + vendor = v + break + existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() + now = _now() + if existing is None: + existing = ManagedNE(ip_address=ip, created_at=now) + db.add(existing) + inserted += 1 + else: + updated += 1 + existing.name = display_name + existing.vendor = vendor + existing.device_type = device_type + existing.port = port + existing.protocol = protocol + existing.username = username + existing.password_enc = encrypt_secret(password) if password else "" + tags_val = _import_cell_str(row.get("tags", "")) + remark_val = _import_cell_str(row.get("remark", "")) + if tags_val: + existing.tags = tags_val + if remark_val: + existing.remark = remark_val + existing.updated_at = now + except CredentialCryptoError as exc: + failed.append(ImportFailure(row=row_no, reason=str(exc))) + except Exception as exc: + failed.append(ImportFailure(row=row_no, reason=str(exc)[:200])) + db.commit() + return ImportResult(inserted=inserted, updated=updated, failed=failed) + + diff --git a/netx_api/ne_service_webcrt.py b/netx_api/ne_service_webcrt.py new file mode 100644 index 0000000..443408e --- /dev/null +++ b/netx_api/ne_service_webcrt.py @@ -0,0 +1,210 @@ +"""WebCRT managed-NE host upsert helpers.""" +from __future__ import annotations + +from typing import Any + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .device_types import WEBCRT_DEVICE_TYPES, WEBCRT_NE_SOURCE +from .models import ManagedNE +from .ne_crypto import encrypt_secret +from .ne_schemas import ManagedNeCreate, ManagedNeOut +from .ne_service_common import ( + WEBCRT_SOURCE, + _apply_hop_create, + _normalize_hop_target_auth_mode, + _normalize_hop_vendor, + _normalize_ip, + _normalize_protocol, + _normalize_vendor, + _now, + _require_crypto, + _validate_hop_on_create, + row_to_out, +) + +def _normalize_webcrt_device_type(device_type: str) -> str: + dt = str(device_type or "").strip() + low = dt.lower() + if low in ("linux", "linux_ssh", "linux_telnet"): + return "linux" + if low in ("generic", "generic_ssh", "generic_telnet", "terminal_server", "generic_termserver"): + return "generic" + return dt + + +def upsert_webcrt_managed_ne(db: Session, body: ManagedNeCreate) -> tuple[ManagedNeOut, str]: + """Create/update a WebCRT-origin NE, or reuse an existing inventory NE by IP. + + Returns ``(ne_out, action)`` where action is ``created`` | ``updated`` | ``reused``. + """ + _require_crypto() + _validate_hop_on_create(body) + ip = _normalize_ip(body.ip_address) + if not ip: + raise HTTPException(status_code=400, detail="ip_address_required") + if not str(body.username or "").strip(): + raise HTTPException(status_code=400, detail="cli_username_required") + device_type = _normalize_webcrt_device_type(body.device_type) + if device_type not in WEBCRT_DEVICE_TYPES: + raise HTTPException(status_code=400, detail="unsupported_device_type") + + existing = db.query(ManagedNE).filter(ManagedNE.ip_address == ip).first() + now = _now() + + if existing is not None: + src = str(existing.source or "").strip() + if src != WEBCRT_NE_SOURCE: + # Do not overwrite inventory / UME-synced assets; just open them. + return row_to_out(existing), "reused" + + existing.name = str(body.name or "").strip() or existing.name or ip + existing.vendor = _normalize_vendor(body.vendor) if str(body.vendor or "").strip() else ( + "Other" if device_type == "linux" else existing.vendor + ) + existing.device_type = device_type + existing.port = int(body.port or existing.port or 22) + existing.protocol = _normalize_protocol(body.protocol) + existing.username = str(body.username or "").strip() + if str(body.password or "").strip(): + existing.password_enc = encrypt_secret(body.password) + elif not str(existing.password_enc or "").strip() and not ( + body.hop_enabled + and _normalize_hop_vendor(body.hop_vendor) == "bastion" + and _normalize_hop_target_auth_mode(body.hop_target_auth_mode) == "bastion_managed" + ): + raise HTTPException(status_code=400, detail="password_required") + existing.source = WEBCRT_NE_SOURCE + _apply_hop_create(existing, body) + existing.updated_at = now + db.commit() + db.refresh(existing) + return row_to_out(existing), "updated" + + if not str(body.password or "").strip() and not ( + body.hop_enabled + and _normalize_hop_vendor(body.hop_vendor) == "bastion" + and _normalize_hop_target_auth_mode(body.hop_target_auth_mode) == "bastion_managed" + ): + raise HTTPException(status_code=400, detail="password_required") + + vendor = _normalize_vendor(body.vendor) + if device_type == "linux" and not str(body.vendor or "").strip(): + vendor = "Other" + + row = ManagedNE( + name=str(body.name or "").strip() or ip, + vendor=vendor, + device_type=device_type, + ip_address=ip, + port=int(body.port or 22), + protocol=_normalize_protocol(body.protocol), + username=str(body.username or "").strip(), + password_enc=encrypt_secret(body.password) if str(body.password or "").strip() else "", + enable_secret_enc="", + connect_status="unknown", + tags=str(body.tags or "").strip(), + remark=str(body.remark or "").strip(), + source=WEBCRT_NE_SOURCE, + source_ref="", + created_at=now, + updated_at=now, + ) + _apply_hop_create(row, body) + db.add(row) + db.commit() + db.refresh(row) + return row_to_out(row), "created" + + +def _next_webcrt_session_name(db: Session, base: str) -> str: + """Return base, or ``base (1)``, ``base (2)``, … among WebCRT session names.""" + root = str(base or "").strip() or "session" + rows = ( + db.query(ManagedNE.name) + .filter(ManagedNE.source == WEBCRT_NE_SOURCE) + .all() + ) + taken = {str(r[0] or "").strip() for r in rows if str(r[0] or "").strip()} + if root not in taken: + return root + n = 1 + while f"{root} ({n})" in taken: + n += 1 + return f"{root} ({n})" + + +def upsert_webcrt_session_host( + db: Session, + *, + name: str = "", + ip_address: str, + port: int = 22, + protocol: str = "ssh", + username: str = "", + password: str = "", + save_password: bool = False, +) -> tuple[ManagedNeOut, str]: + """Create a WebCRT session host (linux, no hop). Always inserts a new row. + + Same IP is allowed; session name auto-suffixes ``(1)``, ``(2)``, … on collision. + Telnet never persists a password. SSH persists password only when ``save_password``. + Returns ``(ne_out, \"created\")``. + """ + _require_crypto() + ip = _normalize_ip(ip_address) + if not ip: + raise HTTPException(status_code=400, detail="ip_address_required") + proto = _normalize_protocol(protocol) + user = str(username or "").strip() + pwd = str(password or "") + if proto == "ssh" and not user: + raise HTTPException(status_code=400, detail="cli_username_required") + if proto == "ssh" and save_password and not pwd.strip(): + raise HTTPException(status_code=400, detail="password_required") + + now = _now() + display_name = _next_webcrt_session_name(db, str(name or "").strip() or ip) + + password_enc = "" + if proto == "ssh" and save_password and pwd.strip(): + password_enc = encrypt_secret(pwd) + + row = ManagedNE( + name=display_name, + vendor="Other", + # generic → Netmiko terminal_server: SSH auth then raw PTY (no linux session prep). + device_type="generic", + ip_address=ip, + port=int(port or (23 if proto == "telnet" else 22)), + protocol=proto, + username=user, + password_enc=password_enc, + enable_secret_enc="", + connect_status="unknown", + tags="", + remark="", + source=WEBCRT_NE_SOURCE, + source_ref="", + created_at=now, + updated_at=now, + ) + db.add(row) + try: + db.commit() + except Exception as exc: + db.rollback() + # Stale unique index on ip_address → restart API after migration, or drop constraint manually. + from sqlalchemy.exc import IntegrityError + + if isinstance(exc, IntegrityError): + raise HTTPException( + status_code=409, + detail="ip_address_conflict_restart_required", + ) from exc + raise + db.refresh(row) + return row_to_out(row), "created" + + diff --git a/netx_api/ume_alarm_ws.py b/netx_api/ume_alarm_ws.py index ef04895..04f0773 100644 --- a/netx_api/ume_alarm_ws.py +++ b/netx_api/ume_alarm_ws.py @@ -249,7 +249,7 @@ def get_ws_connection_status() -> dict[str, Any]: detail = str(_ws_connection_detail or "") return { "state": state, - "label": f"ws:{state}", + "label": _WS_CONNECTION_LABELS.get(state, state), "detail": detail, } diff --git a/netx_api/ume_alarms_router.py b/netx_api/ume_alarms_router.py new file mode 100644 index 0000000..f94e93c --- /dev/null +++ b/netx_api/ume_alarms_router.py @@ -0,0 +1,509 @@ +"""UME current/history alarms, aggregates, diagnostics.""" +from __future__ import annotations + +import logging +from datetime import datetime, timezone +from typing import Any + +from fastapi import APIRouter, Depends, HTTPException, Query +from sqlalchemy import or_ +from sqlalchemy.orm import Session + +from .config import settings +from .db import get_db +from .key_alert_config import ( + get_key_alert_monitor_config, + invalidate_key_alert_config_cache, + set_key_alert_monitor_config, +) +from .key_alert_matcher import ( + invalidate_key_alert_rule_cache, + normalize_match_type, + parse_rule_ne_types_payload, + rule_match_type, + rule_match_value, + rule_ne_types, + rule_storage_key, + serialize_rule_ne_types, +) +from .models import ( + UmeAlarmCurrent, + UmeAlarmHistory, + UmeInventoryNE, + UmeKeyAlertForwardLog, + UmeKeyAlertRule, + UmeSyncJob, +) +from .oclaw_alarm_forwarder import ( + forwarder_status, + request_forwarder_reconnect, +) +from .ume_alarm_ws import ( + cancel_alarm_subscription_manual, + clear_local_alarm_subscription_manual, + establish_alarm_subscription_manual, + get_alarms_coordination_status, + get_subscription_status, + get_ws_connection_status, + get_ws_logs, + request_ws_reconnect, +) +from .ume_support import ( + UME_KNOWN_RUNTIME_TASKS, + _aggregate_rows, + _ensure_utc, + _list_runtime_tasks, + _request_force_sync_after_resume, + _runtime_pause_task, + _runtime_resume_task, + _ume_alarm_host_name, + _ume_alarm_ne_group_key, + _ume_client, + _ume_error_kind, + _clear_force_resume_hints, +) +from .ume_sync_service import sync_alarms_current, sync_alarms_history_full, sync_inventory_full +from .ume_token_store import clear_shared_token + +_log = logging.getLogger("netx.ume.router") +router = APIRouter(tags=["ume"]) + +@router.get("/v1/ume/alarms") +def ume_list_alarms( + severity: str | None = Query(default=None), + is_cleared: str | None = Query(default=None), + ne_id: str | None = Query(default=None), + host_name: str | None = Query(default=None), + keyword: str | None = Query(default=None), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=50, ge=1, le=500), + db: Session = Depends(get_db), +) -> dict[str, Any]: + stmt = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( + UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id + ) + if severity and str(severity).strip(): + stmt = stmt.filter(UmeAlarmCurrent.perceived_severity == str(severity).strip()) + if is_cleared and str(is_cleared).strip(): + stmt = stmt.filter(UmeAlarmCurrent.is_cleared == str(is_cleared).strip()) + if ne_id and str(ne_id).strip(): + stmt = stmt.filter(UmeAlarmCurrent.ne_id == str(ne_id).strip()) + hn = str(host_name or "").strip() + if hn: + stmt = stmt.filter( + UmeAlarmCurrent.host_name.contains(hn) | UmeInventoryNE.host_name.contains(hn) + ) + kw = str(keyword or "").strip() + if kw: + stmt = stmt.filter( + UmeAlarmCurrent.alarm_key.contains(kw) + | UmeAlarmCurrent.object_name.contains(kw) + | UmeAlarmCurrent.native_probable_cause.contains(kw) + | UmeAlarmCurrent.notification_id.contains(kw) + | UmeAlarmCurrent.host_name.contains(kw) + | UmeInventoryNE.ne_name.contains(kw) + | UmeInventoryNE.user_label.contains(kw) + | UmeInventoryNE.ip_address.contains(kw) + | UmeInventoryNE.host_name.contains(kw) + ) + total = int(stmt.count()) + rows = ( + stmt.order_by( + UmeAlarmCurrent.time_created.desc(), + UmeAlarmCurrent.last_seen_at.desc(), + UmeAlarmCurrent.alarm_key.desc(), + ) + .offset((page - 1) * page_size) + .limit(page_size) + .all() + ) + items = [ + { + "alarm_key": str(alarm.alarm_key or ""), + "ne_id": str(alarm.ne_id or ""), + "ne_name": str((ne.ne_name if ne else "") or ""), + "user_label": str((ne.user_label if ne else "") or ""), + "host_name": _ume_alarm_host_name(alarm, ne), + "ne_type": str((ne.ne_type if ne else "") or ""), + "object_name": str(alarm.object_name or ""), + "event_type": str(alarm.event_type or ""), + "native_probable_cause": str(alarm.native_probable_cause or ""), + "notification_id": str(alarm.notification_id or ""), + "perceived_severity": str(alarm.perceived_severity or ""), + "is_cleared": str(alarm.is_cleared or ""), + "time_created": str(alarm.time_created or ""), + "last_seen_at": (_ensure_utc(alarm.last_seen_at) or datetime.now(timezone.utc)).isoformat(), + } + for alarm, ne in rows + ] + return {"total": total, "page": page, "page_size": page_size, "items": items} + + +@router.get("/v1/ume/alarms/fields") +def ume_alarms_fields() -> dict[str, Any]: + """List all queryable field names for UME raw alarm query.""" + alarm_cols = [str(c.name) for c in UmeAlarmCurrent.__table__.columns] # type: ignore[attr-defined] + ne_cols = [str(c.name) for c in UmeInventoryNE.__table__.columns] # type: ignore[attr-defined] + selectable_fields = [f"alarm_{x}" for x in alarm_cols] + [f"ne_{x}" for x in ne_cols] + ["ne_exists"] + order_by_allowed = ["last_seen_at", "time_created", "perceived_severity", "event_type", "ne_id"] + return { + "alarm_fields": alarm_cols, + "ne_fields": ne_cols, + "selectable_fields": selectable_fields, + "order_by_allowed": order_by_allowed, + } + + +def _serialize_ume_alarm_raw_row( + alarm: UmeAlarmCurrent, ne: UmeInventoryNE | None, selected_fields: set[str] | None = None +) -> dict[str, Any]: + selected = selected_fields or set() + use_all = len(selected) == 0 + out: dict[str, Any] = {} + for c in UmeAlarmCurrent.__table__.columns: # type: ignore[attr-defined] + name = str(c.name) + v = getattr(alarm, name, None) + key = f"alarm_{name}" + if not use_all and key not in selected: + continue + if hasattr(v, "isoformat"): + try: + if isinstance(v, datetime): + out[key] = (_ensure_utc(v) or v).isoformat() + else: + out[key] = v.isoformat() + continue + except Exception: + pass + out[key] = v + if ne is None: + if use_all or "ne_exists" in selected: + out["ne_exists"] = False + return out + if use_all or "ne_exists" in selected: + out["ne_exists"] = True + for c in UmeInventoryNE.__table__.columns: # type: ignore[attr-defined] + name = str(c.name) + v = getattr(ne, name, None) + key = f"ne_{name}" + if not use_all and key not in selected: + continue + if hasattr(v, "isoformat"): + try: + if isinstance(v, datetime): + out[key] = (_ensure_utc(v) or v).isoformat() + else: + out[key] = v.isoformat() + continue + except Exception: + pass + out[key] = v + return out + + +def _extract_ume_raw_group_field(alarm: UmeAlarmCurrent, ne: UmeInventoryNE | None, field: str) -> str: + key = str(field or "").strip() + if not key: + return "" + if key.startswith("alarm_"): + attr = key[len("alarm_") :] + return str(getattr(alarm, attr, "") or "") + if key.startswith("ne_"): + attr = key[len("ne_") :] + if key == "ne_exists": + return "1" if ne is not None else "0" + if key == "ne_host_name": + hn = str(getattr(alarm, "host_name", "") or "").strip() + if hn: + return hn + if ne is None: + return "" + return str(getattr(ne, attr, "") or "") + return "" + + +@router.get("/v1/ume/alarms/raw") +def ume_alarms_raw( + severity: str | None = Query(default=None), + is_cleared: str | None = Query(default=None), + ne_id: str | None = Query(default=None), + event_type: str | None = Query(default=None), + keyword: str | None = Query(default=None), + time_from: str | None = Query(default=None), + time_to: str | None = Query(default=None), + order_by: str = Query(default="last_seen_at"), + order: str = Query(default="desc"), + select_fields: str | None = Query(default=None, description="comma-separated alarm_*/ne_* fields"), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=50, ge=1, le=500), + db: Session = Depends(get_db), +) -> dict[str, Any]: + stmt = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( + UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id + ) + if severity and str(severity).strip(): + stmt = stmt.filter(UmeAlarmCurrent.perceived_severity == str(severity).strip()) + if is_cleared and str(is_cleared).strip(): + stmt = stmt.filter(UmeAlarmCurrent.is_cleared == str(is_cleared).strip()) + if ne_id and str(ne_id).strip(): + stmt = stmt.filter(UmeAlarmCurrent.ne_id == str(ne_id).strip()) + if event_type and str(event_type).strip(): + stmt = stmt.filter(UmeAlarmCurrent.event_type.contains(str(event_type).strip())) + kw = str(keyword or "").strip() + if kw: + stmt = stmt.filter( + UmeAlarmCurrent.alarm_key.contains(kw) + | UmeAlarmCurrent.object_name.contains(kw) + | UmeAlarmCurrent.native_probable_cause.contains(kw) + | UmeAlarmCurrent.event_type.contains(kw) + | UmeInventoryNE.ne_name.contains(kw) + | UmeInventoryNE.user_label.contains(kw) + | UmeInventoryNE.ip_address.contains(kw) + ) + dt_from = _parse_time(time_from) + dt_to = _parse_time(time_to) + if dt_from: + stmt = stmt.filter(UmeAlarmCurrent.last_seen_at >= dt_from.replace(tzinfo=None)) + if dt_to: + stmt = stmt.filter(UmeAlarmCurrent.last_seen_at <= dt_to.replace(tzinfo=None)) + + allowed_order_by = { + "last_seen_at": UmeAlarmCurrent.last_seen_at, + "time_created": UmeAlarmCurrent.time_created, + "perceived_severity": UmeAlarmCurrent.perceived_severity, + "event_type": UmeAlarmCurrent.event_type, + "ne_id": UmeAlarmCurrent.ne_id, + } + col = allowed_order_by.get(str(order_by or "").strip(), UmeAlarmCurrent.last_seen_at) + if str(order or "").strip().lower() == "asc": + stmt = stmt.order_by(col.asc()) + else: + stmt = stmt.order_by(col.desc()) + + selected_fields: set[str] = set() + fields_meta = ume_alarms_fields() + selectable_fields = set(str(x) for x in (fields_meta.get("selectable_fields") or [])) + order_by_allowed = [str(x) for x in (fields_meta.get("order_by_allowed") or [])] + if select_fields and str(select_fields).strip(): + selected_fields = {x.strip() for x in str(select_fields).split(",") if x.strip()} + invalid = [x for x in selected_fields if x not in selectable_fields] + if invalid: + raise HTTPException(status_code=400, detail=f"invalid_select_fields:{','.join(sorted(invalid)[:20])}") + + total = int(stmt.count()) + rows = stmt.offset((int(page) - 1) * int(page_size)).limit(int(page_size)).all() + return { + "total": total, + "page": int(page), + "page_size": int(page_size), + "select_fields": sorted(selected_fields) if selected_fields else [], + "meta": { + "available_fields": sorted(selectable_fields), + "order_by_allowed": order_by_allowed, + "time_filter_field": "last_seen_at", + }, + "items": [_serialize_ume_alarm_raw_row(alarm, ne, selected_fields) for alarm, ne in rows], + } + + +@router.get("/v1/ume/alarms/aggregate/raw") +def ume_alarms_aggregate_raw( + group_by: str = Query(default="alarm_perceived_severity"), + group_by2: str | None = Query(default=None), + severity: str | None = Query(default=None), + is_cleared: str | None = Query(default=None), + ne_id: str | None = Query(default=None), + event_type: str | None = Query(default=None), + keyword: str | None = Query(default=None), + time_from: str | None = Query(default=None), + time_to: str | None = Query(default=None), + limit: int = Query(default=200, ge=1, le=2000), + db: Session = Depends(get_db), +) -> dict[str, Any]: + fields_meta = ume_alarms_fields() + selectable_fields = set(str(x) for x in (fields_meta.get("selectable_fields") or [])) + g1 = str(group_by or "").strip() + g2 = str(group_by2 or "").strip() + if g1 not in selectable_fields: + raise HTTPException(status_code=400, detail=f"invalid_group_by:{g1}") + if g2 and g2 not in selectable_fields: + raise HTTPException(status_code=400, detail=f"invalid_group_by2:{g2}") + + stmt = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( + UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id + ) + if severity and str(severity).strip(): + stmt = stmt.filter(UmeAlarmCurrent.perceived_severity == str(severity).strip()) + if is_cleared and str(is_cleared).strip(): + stmt = stmt.filter(UmeAlarmCurrent.is_cleared == str(is_cleared).strip()) + if ne_id and str(ne_id).strip(): + stmt = stmt.filter(UmeAlarmCurrent.ne_id == str(ne_id).strip()) + if event_type and str(event_type).strip(): + stmt = stmt.filter(UmeAlarmCurrent.event_type.contains(str(event_type).strip())) + kw = str(keyword or "").strip() + if kw: + stmt = stmt.filter( + UmeAlarmCurrent.alarm_key.contains(kw) + | UmeAlarmCurrent.object_name.contains(kw) + | UmeAlarmCurrent.native_probable_cause.contains(kw) + | UmeAlarmCurrent.event_type.contains(kw) + | UmeInventoryNE.ne_name.contains(kw) + | UmeInventoryNE.user_label.contains(kw) + | UmeInventoryNE.ip_address.contains(kw) + ) + dt_from = _parse_time(time_from) + dt_to = _parse_time(time_to) + if dt_from: + stmt = stmt.filter(UmeAlarmCurrent.last_seen_at >= dt_from.replace(tzinfo=None)) + if dt_to: + stmt = stmt.filter(UmeAlarmCurrent.last_seen_at <= dt_to.replace(tzinfo=None)) + + rows = stmt.order_by(UmeAlarmCurrent.last_seen_at.desc()).all() + counts: dict[tuple[str, str], int] = {} + for alarm, ne in rows: + k1 = _extract_ume_raw_group_field(alarm, ne, g1) + k2 = _extract_ume_raw_group_field(alarm, ne, g2) if g2 else "" + kk = (k1, k2) + counts[kk] = int(counts.get(kk, 0)) + 1 + buckets = sorted(counts.items(), key=lambda x: x[1], reverse=True)[: int(limit)] + return { + "total": len(rows), + "group_by": g1, + "group_by2": g2 or None, + "meta": { + "available_fields": sorted(selectable_fields), + "group_by_allowed": sorted(selectable_fields), + "applied_filters": { + "severity": str(severity or "").strip() or None, + "is_cleared": str(is_cleared or "").strip() or None, + "ne_id": str(ne_id or "").strip() or None, + "event_type": str(event_type or "").strip() or None, + "keyword": str(keyword or "").strip() or None, + "time_from": str(time_from or "").strip() or None, + "time_to": str(time_to or "").strip() or None, + }, + "time_filter_field": "last_seen_at", + "limit": int(limit), + }, + "buckets": [ + {"key": k1, "key2": (k2 if g2 else None), "count": int(v)} + for (k1, k2), v in buckets + ], + } + + +@router.get("/v1/ume/alarms/aggregate") +def ume_alarms_aggregate(db: Session = Depends(get_db)) -> dict[str, Any]: + rows = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( + UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id + ).all() + by_severity = _aggregate_rows(rows, lambda x: x[0].perceived_severity) + by_ne = _aggregate_rows(rows, lambda x: _ume_alarm_ne_group_key(x[0], x[1])) + return {"total": len(rows), "by_severity": by_severity, "by_ne": by_ne} + + +@router.get("/v1/ume/diagnostics") +def ume_diagnostics( + lang: str | None = Query(default=None), + db: Session = Depends(get_db), +) -> dict[str, Any]: + rows = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( + UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id + ).all() + by_severity = _aggregate_rows(rows, lambda x: x[0].perceived_severity) + by_alarm_code = _aggregate_rows(rows, lambda x: x[0].event_type)[:10] + by_ne = _aggregate_rows(rows, lambda x: _ume_alarm_ne_group_key(x[0], x[1]))[:10] + + lang_norm = _normalize_netx_lang(lang) + proto_counts: dict[str, int] = {} + for alarm, ne in rows: + blob = " | ".join( + [ + str(alarm.event_type or ""), + str(alarm.native_probable_cause or ""), + str(alarm.object_name or ""), + str(ne.ne_name if ne else ""), + str(ne.user_label if ne else ""), + str(ne.ip_address if ne else ""), + ] + ) + bucket = _protocol_bucket_label(blob, lang=lang_norm) + proto_counts[bucket] = int(proto_counts.get(bucket, 0)) + 1 + protocol_summary = sorted(proto_counts.items(), key=lambda x: x[1], reverse=True)[:10] + + return { + "source": "ume_alarms_current", + "total_alarms": len(rows), + "severity_summary": [{"key": k, "count": v} for k, v in by_severity], + "top_alarm_codes": [{"key": k, "count": v} for k, v in by_alarm_code], + "top_ne": [{"key": k, "count": v} for k, v in by_ne], + "protocol_summary": [{"key": k, "count": v} for k, v in protocol_summary], + } + + +@router.get("/v1/ume/alarms/history") +def ume_list_alarms_history( + severity: str | None = Query(default=None), + ne_id: str | None = Query(default=None), + keyword: str | None = Query(default=None), + time_from: str | None = Query(default=None), + time_to: str | None = Query(default=None), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=50, ge=1, le=500), + db: Session = Depends(get_db), +) -> dict[str, Any]: + stmt = db.query(UmeAlarmHistory, UmeInventoryNE).outerjoin( + UmeInventoryNE, UmeAlarmHistory.ne_id == UmeInventoryNE.ne_id + ) + if severity and str(severity).strip(): + stmt = stmt.filter(UmeAlarmHistory.perceived_severity == str(severity).strip()) + if ne_id and str(ne_id).strip(): + stmt = stmt.filter(UmeAlarmHistory.ne_id == str(ne_id).strip()) + kw = str(keyword or "").strip() + if kw: + stmt = stmt.filter( + UmeAlarmHistory.alarm_key.contains(kw) + | UmeAlarmHistory.object_name.contains(kw) + | UmeAlarmHistory.native_probable_cause.contains(kw) + | UmeInventoryNE.ne_name.contains(kw) + | UmeInventoryNE.user_label.contains(kw) + | UmeInventoryNE.ip_address.contains(kw) + ) + dt_from = _parse_time(time_from) + dt_to = _parse_time(time_to) + if dt_from: + stmt = stmt.filter(UmeAlarmHistory.last_seen_at >= dt_from.replace(tzinfo=None)) + if dt_to: + stmt = stmt.filter(UmeAlarmHistory.last_seen_at <= dt_to.replace(tzinfo=None)) + total = int(stmt.count()) + rows = stmt.order_by(UmeAlarmHistory.last_seen_at.desc()).offset((page - 1) * page_size).limit(page_size).all() + items = [ + { + "alarm_key": str(alarm.alarm_key or ""), + "ne_id": str(alarm.ne_id or ""), + "ne_name": str((ne.ne_name if ne else "") or ""), + "user_label": str((ne.user_label if ne else "") or ""), + "object_name": str(alarm.object_name or ""), + "event_type": str(alarm.event_type or ""), + "native_probable_cause": str(alarm.native_probable_cause or ""), + "perceived_severity": str(alarm.perceived_severity or ""), + "is_cleared": str(alarm.is_cleared or ""), + "time_created": str(alarm.time_created or ""), + "last_seen_at": (_ensure_utc(alarm.last_seen_at) or datetime.now(timezone.utc)).isoformat(), + } + for alarm, ne in rows + ] + return {"total": total, "page": page, "page_size": page_size, "items": items} + + +@router.get("/v1/ume/alarms/history/aggregate") +def ume_alarms_history_aggregate(db: Session = Depends(get_db)) -> dict[str, Any]: + rows = db.query(UmeAlarmHistory, UmeInventoryNE).outerjoin( + UmeInventoryNE, UmeAlarmHistory.ne_id == UmeInventoryNE.ne_id + ).all() + by_severity = _aggregate_rows(rows, lambda x: x[0].perceived_severity) + by_ne = _aggregate_rows(rows, lambda x: (x[1].user_label if x[1] else "") or (x[1].ne_name if x[1] else "") or x[0].ne_id) + by_date = _aggregate_rows(rows, lambda x: str(x[0].time_created or "")[:10]) + return {"total": len(rows), "by_severity": by_severity, "by_ne": by_ne, "by_date": by_date} + + diff --git a/netx_api/ume_inventory_router.py b/netx_api/ume_inventory_router.py new file mode 100644 index 0000000..bd67b81 --- /dev/null +++ b/netx_api/ume_inventory_router.py @@ -0,0 +1,177 @@ +"""UME inventory NE list/detail.""" +from __future__ import annotations + +import logging +from datetime import datetime, timezone +from typing import Any + +from fastapi import APIRouter, Depends, HTTPException, Query +from sqlalchemy import or_ +from sqlalchemy.orm import Session + +from .config import settings +from .db import get_db +from .key_alert_config import ( + get_key_alert_monitor_config, + invalidate_key_alert_config_cache, + set_key_alert_monitor_config, +) +from .key_alert_matcher import ( + invalidate_key_alert_rule_cache, + normalize_match_type, + parse_rule_ne_types_payload, + rule_match_type, + rule_match_value, + rule_ne_types, + rule_storage_key, + serialize_rule_ne_types, +) +from .models import ( + UmeAlarmCurrent, + UmeAlarmHistory, + UmeInventoryNE, + UmeKeyAlertForwardLog, + UmeKeyAlertRule, + UmeSyncJob, +) +from .oclaw_alarm_forwarder import ( + forwarder_status, + request_forwarder_reconnect, +) +from .ume_alarm_ws import ( + cancel_alarm_subscription_manual, + clear_local_alarm_subscription_manual, + establish_alarm_subscription_manual, + get_alarms_coordination_status, + get_subscription_status, + get_ws_connection_status, + get_ws_logs, + request_ws_reconnect, +) +from .ume_support import ( + UME_KNOWN_RUNTIME_TASKS, + _aggregate_rows, + _ensure_utc, + _list_runtime_tasks, + _request_force_sync_after_resume, + _runtime_pause_task, + _runtime_resume_task, + _ume_alarm_host_name, + _ume_alarm_ne_group_key, + _ume_client, + _ume_error_kind, + _clear_force_resume_hints, +) +from .ume_sync_service import sync_alarms_current, sync_alarms_history_full, sync_inventory_full +from .ume_token_store import clear_shared_token + +_log = logging.getLogger("netx.ume.router") +router = APIRouter(tags=["ume"]) + +@router.get("/v1/ume/inventory/ne-types") +def ume_list_inventory_ne_types( + limit: int = Query(default=500, ge=1, le=2000), + db: Session = Depends(get_db), +) -> dict[str, Any]: + from sqlalchemy import func + + rows = ( + db.query( + UmeInventoryNE.ne_type, + func.count(UmeInventoryNE.ne_id).label("ne_count"), + ) + .filter(UmeInventoryNE.ne_type != "") + .group_by(UmeInventoryNE.ne_type) + .order_by(func.count(UmeInventoryNE.ne_id).desc(), UmeInventoryNE.ne_type.asc()) + .limit(limit) + .all() + ) + items = [{"ne_type": str(ne_type or ""), "ne_count": int(ne_count or 0)} for ne_type, ne_count in rows if str(ne_type or "").strip()] + return {"items": items, "total": len(items)} + + +@router.get("/v1/ume/inventory/ne") +def ume_list_ne( + keyword: str | None = Query(default=None), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=50, ge=1, le=500), + db: Session = Depends(get_db), +) -> dict[str, Any]: + stmt = db.query(UmeInventoryNE) + kw = str(keyword or "").strip() + if kw: + like = f"%{kw}%" + stmt = stmt.filter( + UmeInventoryNE.ne_id.ilike(like) + | UmeInventoryNE.ne_name.ilike(like) + | UmeInventoryNE.user_label.ilike(like) + | UmeInventoryNE.ip_address.ilike(like) + | UmeInventoryNE.host_name.ilike(like) + ) + total = int(stmt.count()) + rows = stmt.order_by(UmeInventoryNE.ne_id.asc()).offset((page - 1) * page_size).limit(page_size).all() + items = [ + { + "ne_id": str(x.ne_id or ""), + "ne_name": str(x.ne_name or ""), + "user_label": str(x.user_label or ""), + "ip_address": str(x.ip_address or ""), + "ipv6_address": str(x.ipv6_address or ""), + "ne_type": str(x.ne_type or ""), + "device_level": str(x.device_level or ""), + "host_name": str(x.host_name or ""), + "location": str(x.location or ""), + "hardware_version": str(x.hardware_version or ""), + "loopback": str(x.loopback or ""), + "consistent_state": str(x.consistent_state or ""), + "interface_version": str(x.interface_version or ""), + "mac": str(x.mac or ""), + "admin_status": str(x.admin_status or ""), + "address_type": str(x.address_type or ""), + "connection_status": str(x.connection_status or ""), + "maintain_status": str(x.maintain_status or ""), + "net_mask": str(x.net_mask or ""), + "create_time": str(x.create_time or ""), + "creator": str(x.creator or ""), + "last_seen_at": (_ensure_utc(x.last_seen_at) or datetime.now(timezone.utc)).isoformat(), + } + for x in rows + ] + return {"total": total, "page": page, "page_size": page_size, "items": items} + + +@router.get("/v1/ume/inventory/ne/{ne_id}") +def ume_get_ne(ne_id: str, db: Session = Depends(get_db)) -> dict[str, Any]: + row = db.get(UmeInventoryNE, ne_id) + if not row: + raise HTTPException(status_code=404, detail="ume_ne_not_found") + return { + "ne_id": str(row.ne_id or ""), + "ne_name": str(row.ne_name or ""), + "user_label": str(row.user_label or ""), + "ip_address": str(row.ip_address or ""), + "ipv6_address": str(row.ipv6_address or ""), + "ne_type": str(row.ne_type or ""), + "device_level": str(row.device_level or ""), + "host_name": str(row.host_name or ""), + "location": str(row.location or ""), + "hardware_version": str(row.hardware_version or ""), + "loopback": str(row.loopback or ""), + "consistent_state": str(row.consistent_state or ""), + "interface_version": str(row.interface_version or ""), + "mac": str(row.mac or ""), + "admin_status": str(row.admin_status or ""), + "address_type": str(row.address_type or ""), + "connection_status": str(row.connection_status or ""), + "maintain_status": str(row.maintain_status or ""), + "net_mask": str(row.net_mask or ""), + "create_time": str(row.create_time or ""), + "creator": str(row.creator or ""), + "vendor": str(row.vendor or ""), + "source_type": str(row.source_type or ""), + "first_seen_at": (_ensure_utc(row.first_seen_at) or datetime.now(timezone.utc)).isoformat(), + "last_seen_at": (_ensure_utc(row.last_seen_at) or datetime.now(timezone.utc)).isoformat(), + "raw_json": str(row.raw_json or "{}"), + } + + diff --git a/netx_api/ume_key_alert_router.py b/netx_api/ume_key_alert_router.py new file mode 100644 index 0000000..20d33c7 --- /dev/null +++ b/netx_api/ume_key_alert_router.py @@ -0,0 +1,339 @@ +"""UME key-alert rules / monitor / keyword helpers.""" +from __future__ import annotations + +import logging +from datetime import datetime, timezone +from typing import Any + +from fastapi import APIRouter, Depends, HTTPException, Query +from sqlalchemy import or_ +from sqlalchemy.orm import Session + +from .config import settings +from .db import get_db +from .key_alert_config import ( + get_key_alert_monitor_config, + invalidate_key_alert_config_cache, + set_key_alert_monitor_config, +) +from .key_alert_matcher import ( + invalidate_key_alert_rule_cache, + normalize_match_type, + parse_rule_ne_types_payload, + rule_match_type, + rule_match_value, + rule_ne_types, + rule_storage_key, + serialize_rule_ne_types, +) +from .models import ( + UmeAlarmCurrent, + UmeAlarmHistory, + UmeInventoryNE, + UmeKeyAlertForwardLog, + UmeKeyAlertRule, + UmeSyncJob, +) +from .oclaw_alarm_forwarder import ( + forwarder_status, + request_forwarder_reconnect, +) +from .ume_alarm_ws import ( + cancel_alarm_subscription_manual, + clear_local_alarm_subscription_manual, + establish_alarm_subscription_manual, + get_alarms_coordination_status, + get_subscription_status, + get_ws_connection_status, + get_ws_logs, + request_ws_reconnect, +) +from .ume_support import ( + UME_KNOWN_RUNTIME_TASKS, + _aggregate_rows, + _ensure_utc, + _list_runtime_tasks, + _request_force_sync_after_resume, + _runtime_pause_task, + _runtime_resume_task, + _ume_alarm_host_name, + _ume_alarm_ne_group_key, + _ume_client, + _ume_error_kind, + _clear_force_resume_hints, +) +from .ume_sync_service import sync_alarms_current, sync_alarms_history_full, sync_inventory_full +from .ume_token_store import clear_shared_token + +_log = logging.getLogger("netx.ume.router") +router = APIRouter(tags=["ume"]) + +@router.get("/v1/ume/key-alert-rules") +def ume_list_key_alert_rules( + db: Session = Depends(get_db), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=50, ge=1, le=200), + keyword: str = Query(default=""), + enabled: str | None = Query(default=None), + match_type: str | None = Query(default=None), +) -> dict[str, Any]: + from sqlalchemy import func, or_ + + q = db.query(UmeKeyAlertRule) + kw = str(keyword or "").strip() + if kw: + like = f"%{kw}%" + q = q.filter( + or_( + UmeKeyAlertRule.notification_id.ilike(like), + UmeKeyAlertRule.match_value.ilike(like), + UmeKeyAlertRule.label.ilike(like), + ) + ) + if enabled is not None: + en = str(enabled).strip().lower() + if en in {"1", "true", "yes", "on"}: + q = q.filter(UmeKeyAlertRule.enabled == 1) + elif en in {"0", "false", "no", "off"}: + q = q.filter(UmeKeyAlertRule.enabled == 0) + if match_type: + mt = normalize_match_type(str(match_type)) + q = q.filter(UmeKeyAlertRule.match_type == mt) + + total = int(q.count()) + rows = ( + q.order_by(UmeKeyAlertRule.notification_id.asc()) + .offset((page - 1) * page_size) + .limit(page_size) + .all() + ) + stat_rows = ( + db.query( + UmeKeyAlertForwardLog.rule_key, + func.count(UmeKeyAlertForwardLog.id).label("attempts"), + func.sum(UmeKeyAlertForwardLog.oclaw_ok).label("published_ok"), + func.max(UmeKeyAlertForwardLog.forwarded_at).label("last_forwarded_at"), + ) + .filter(UmeKeyAlertForwardLog.rule_key != "") + .group_by(UmeKeyAlertForwardLog.rule_key) + .all() + ) + stat_map = { + str(rk or ""): { + "attempts": int(attempts or 0), + "published_ok": int(published_ok or 0), + "last_forwarded_at": (_ensure_utc(last_at) or datetime.now(timezone.utc)).isoformat() if last_at else "", + } + for rk, attempts, published_ok, last_at in stat_rows + if str(rk or "").strip() + } + items = [ + { + "notification_id": str(row.notification_id or ""), + "match_type": rule_match_type(row), + "match_value": rule_match_value(row), + "enabled": bool(int(row.enabled or 0)), + "label": str(row.label or ""), + "ne_types": rule_ne_types(row), + "created_at": (_ensure_utc(row.created_at) or datetime.now(timezone.utc)).isoformat(), + "updated_at": (_ensure_utc(row.updated_at) or datetime.now(timezone.utc)).isoformat(), + "forward_stats": stat_map.get(str(row.notification_id or ""), { + "attempts": 0, + "published_ok": 0, + "last_forwarded_at": "", + }), + } + for row in rows + ] + fwd = forwarder_status() + return {"items": items, "total": total, "page": page, "page_size": page_size, "forwarder": fwd} + + +@router.get("/v1/ume/key-alert-monitor") +def ume_key_alert_monitor( + db: Session = Depends(get_db), + page: int = Query(default=1, ge=1), + page_size: int = Query(default=50, ge=1, le=200), + keyword: str = Query(default=""), + enabled: str | None = Query(default=None), + match_type: str | None = Query(default=None), +) -> dict[str, Any]: + base = ume_list_key_alert_rules( + db=db, + page=page, + page_size=page_size, + keyword=keyword, + enabled=enabled, + match_type=match_type, + ) + return { + "ok": True, + "rules": base.get("items") or [], + "total": int(base.get("total") or 0), + "page": int(base.get("page") or page), + "page_size": int(base.get("page_size") or page_size), + "config": get_key_alert_monitor_config(db), + "forwarder": base.get("forwarder") or forwarder_status(), + } + + +@router.patch("/v1/ume/key-alert-monitor/config") +def ume_update_key_alert_monitor_config(payload: dict[str, Any], db: Session = Depends(get_db)) -> dict[str, Any]: + if "forward_on_clear" not in payload: + raise HTTPException(status_code=400, detail="forward_on_clear_required") + config = set_key_alert_monitor_config(db, forward_on_clear=bool(payload.get("forward_on_clear"))) + return {"ok": True, "config": config} + + +@router.post("/v1/ume/key-alert-rules") +def ume_upsert_key_alert_rule(payload: dict[str, Any], db: Session = Depends(get_db)) -> dict[str, Any]: + match_type = normalize_match_type(str(payload.get("match_type") or "notification_id")) + match_value = str(payload.get("match_value") or payload.get("notification_id") or "").strip() + if not match_value: + raise HTTPException(status_code=400, detail="match_value_required") + label = str(payload.get("label") or "").strip() + if not label: + raise HTTPException(status_code=400, detail="label_required") + enabled = 1 if bool(payload.get("enabled", True)) else 0 + ne_types_list = parse_rule_ne_types_payload(payload.get("ne_types")) + now = datetime.now(timezone.utc).replace(tzinfo=None) + try: + storage_key = rule_storage_key(match_type=match_type, value=match_value) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + row = db.get(UmeKeyAlertRule, storage_key) + if row is None: + row = UmeKeyAlertRule(notification_id=storage_key, created_at=now, updated_at=now) + db.add(row) + row.match_type = match_type + row.match_value = match_value + row.enabled = enabled + row.label = label + row.ne_types = serialize_rule_ne_types(ne_types_list) + row.updated_at = now + saved = { + "notification_id": storage_key, + "match_type": match_type, + "match_value": match_value, + "enabled": bool(enabled), + "label": label, + "ne_types": ne_types_list, + } + try: + db.commit() + except Exception as exc: + db.rollback() + msg = str(exc).lower() + if "match_type" in msg or "match_value" in msg or "ne_types" in msg or "undefinedcolumn" in msg: + raise HTTPException( + status_code=503, + detail="key_alert_schema_outdated: restart netx API to apply database migration", + ) from exc + raise + invalidate_key_alert_rule_cache() + return {"ok": True, "item": saved} + + +@router.patch("/v1/ume/key-alert-rules/{rule_key:path}") +def ume_patch_key_alert_rule(rule_key: str, payload: dict[str, Any], db: Session = Depends(get_db)) -> dict[str, Any]: + key = str(rule_key or "").strip() + if not key: + raise HTTPException(status_code=400, detail="rule_key_required") + row = db.get(UmeKeyAlertRule, key) + if row is None: + raise HTTPException(status_code=404, detail="rule_not_found") + has_enabled = "enabled" in payload + has_ne_types = "ne_types" in payload + if not has_enabled and not has_ne_types: + raise HTTPException(status_code=400, detail="patch_fields_required") + now = datetime.now(timezone.utc).replace(tzinfo=None) + if has_enabled: + row.enabled = 1 if bool(payload.get("enabled")) else 0 + if has_ne_types: + row.ne_types = serialize_rule_ne_types(parse_rule_ne_types_payload(payload.get("ne_types"))) + row.updated_at = now + db.commit() + invalidate_key_alert_rule_cache() + return { + "ok": True, + "item": { + "notification_id": key, + "match_type": rule_match_type(row), + "match_value": rule_match_value(row), + "enabled": bool(int(row.enabled or 0)), + "label": str(row.label or ""), + "ne_types": rule_ne_types(row), + }, + } + + +@router.delete("/v1/ume/key-alert-rules/{rule_key:path}") +def ume_delete_key_alert_rule(rule_key: str, db: Session = Depends(get_db)) -> dict[str, Any]: + key = str(rule_key or "").strip() + row = db.get(UmeKeyAlertRule, key) + if row is None: + raise HTTPException(status_code=404, detail="rule_not_found") + db.delete(row) + db.commit() + invalidate_key_alert_rule_cache() + return {"ok": True, "deleted": key} + + +@router.get("/v1/ume/alarm-keywords") +def ume_list_alarm_keywords( + limit: int = Query(default=200, ge=1, le=2000), + db: Session = Depends(get_db), +) -> dict[str, Any]: + from sqlalchemy import func + + rows = ( + db.query( + UmeAlarmCurrent.native_probable_cause, + func.count(UmeAlarmCurrent.alarm_key).label("cnt"), + ) + .filter(UmeAlarmCurrent.native_probable_cause != "") + .group_by(UmeAlarmCurrent.native_probable_cause) + .order_by(func.count(UmeAlarmCurrent.alarm_key).desc(), UmeAlarmCurrent.native_probable_cause.asc()) + .limit(limit) + .all() + ) + items = [ + { + "keyword": str(cause or ""), + "alarm_count": int(cnt or 0), + } + for cause, cnt in rows + if str(cause or "").strip() + ] + return {"items": items, "total": len(items)} + + +@router.get("/v1/ume/notification-ids") +def ume_list_notification_ids( + limit: int = Query(default=200, ge=1, le=2000), + db: Session = Depends(get_db), +) -> dict[str, Any]: + from sqlalchemy import func + + rows = ( + db.query( + UmeAlarmCurrent.notification_id, + func.max(UmeAlarmCurrent.native_probable_cause).label("cause_sample"), + ) + .filter(UmeAlarmCurrent.notification_id != "") + .group_by(UmeAlarmCurrent.notification_id) + .order_by(UmeAlarmCurrent.notification_id.asc()) + .limit(limit) + .all() + ) + items = [ + { + "notification_id": str(nid or ""), + "native_probable_cause_sample": str(cause or ""), + } + for nid, cause in rows + if str(nid or "").strip() + ] + return {"items": items, "total": len(items), "forwarder": forwarder_status()} + + diff --git a/netx_api/ume_router.py b/netx_api/ume_router.py index 0eb8c64..f3da5ca 100644 --- a/netx_api/ume_router.py +++ b/netx_api/ume_router.py @@ -1,1156 +1,29 @@ -"""UME REST routes (token, sync, inventory, alarms, key-alert, runtime).""" +"""UME REST routes aggregator (token / key-alert / sync / inventory / alarms).""" from __future__ import annotations -import logging -from datetime import datetime, timezone -from typing import Any +from fastapi import APIRouter -from fastapi import APIRouter, Depends, HTTPException, Query -from sqlalchemy import or_ -from sqlalchemy.orm import Session +from .ume_alarms_router import ( + _extract_ume_raw_group_field, + _serialize_ume_alarm_raw_row, + router as alarms_router, + ume_alarms_fields, +) +from .ume_inventory_router import router as inventory_router +from .ume_key_alert_router import router as key_alert_router +from .ume_sync_router import router as sync_router +from .ume_token_router import router as token_router -from .config import settings -from .db import get_db -from .key_alert_config import ( - get_key_alert_monitor_config, - invalidate_key_alert_config_cache, - set_key_alert_monitor_config, -) -from .key_alert_matcher import ( - invalidate_key_alert_rule_cache, - normalize_match_type, - parse_rule_ne_types_payload, - rule_match_type, - rule_match_value, - rule_ne_types, - rule_storage_key, - serialize_rule_ne_types, -) -from .models import ( - UmeAlarmCurrent, - UmeAlarmHistory, - UmeInventoryNE, - UmeKeyAlertForwardLog, - UmeKeyAlertRule, - UmeSyncJob, -) -from .oclaw_alarm_forwarder import ( - forwarder_status, - request_forwarder_reconnect, -) -from .ume_alarm_ws import ( - cancel_alarm_subscription_manual, - clear_local_alarm_subscription_manual, - establish_alarm_subscription_manual, - get_alarms_coordination_status, - get_subscription_status, - get_ws_connection_status, - get_ws_logs, - request_ws_reconnect, -) -from .ume_support import ( - UME_KNOWN_RUNTIME_TASKS, - _aggregate_rows, - _ensure_utc, - _list_runtime_tasks, - _request_force_sync_after_resume, - _runtime_pause_task, - _runtime_resume_task, - _ume_alarm_host_name, - _ume_alarm_ne_group_key, - _ume_client, - _ume_error_kind, - _clear_force_resume_hints, -) -from .ume_sync_service import sync_alarms_current, sync_alarms_history_full, sync_inventory_full -from .ume_token_store import clear_shared_token - -_log = logging.getLogger("netx.ume.router") router = APIRouter(tags=["ume"]) - -@router.get("/v1/ume/token/status") -def ume_token_status() -> dict[str, Any]: - client = _ume_client() - st = client.token_status() - return {"ok": True, **st} - - -@router.post("/v1/ume/token/refresh") -def ume_token_refresh() -> dict[str, Any]: - client = _ume_client() - try: - before = client.token_status() - token = client.refresh_if_needed() - after = client.token_status() - return { - "ok": True, - "token": token, - "changed": bool(before.get("token_preview") != after.get("token_preview")), - **after, - } - except Exception as exc: - msg = str(exc)[:240] - return {"ok": False, "error_kind": _ume_error_kind(msg), "error": msg} - - -@router.post("/v1/ume/token/disconnect") -def ume_token_disconnect() -> dict[str, Any]: - client = _ume_client() - ok = bool(client.logout_token()) - st = client.token_status() - return {"ok": ok, **st} - - -@router.get("/v1/ume/alarm-subscription/status") -def ume_alarm_subscription_status(limit: int = 80) -> dict[str, Any]: - st = get_subscription_status() - ws_task = _UME_RUNTIME_TASKS.get("alarms_current_ws_consumer") or {} - log_limit = max(10, min(int(limit or 80), 100)) - return { - "ok": True, - **st, - **get_alarms_coordination_status(), - "ws_connection": get_ws_connection_status(), - "ws_consumer_status": str(ws_task.get("status") or ""), - "ws_consumer_last_error": str(ws_task.get("last_error") or ""), - "ws_consumer_last_run_at": ws_task.get("last_run_at"), - "ws_logs": get_ws_logs(limit=log_limit), - } - - -@router.post("/v1/ume/alarm-subscription/establish") -def ume_alarm_subscription_establish( - payload: dict[str, Any] | None = None, - db: Session = Depends(get_db), -) -> dict[str, Any]: - client = _ume_client() - body = payload or {} - force_reestablish = bool(body.get("force_reestablish")) - try: - st = establish_alarm_subscription_manual(client, db, force_reestablish=force_reestablish) - return {"ok": True, "created": not bool(st.get("already_exists")), **st} - except Exception as exc: - msg = str(exc)[:240] - raise HTTPException(status_code=502, detail=msg) from exc - - -@router.post("/v1/ume/alarm-subscription/cancel") -def ume_alarm_subscription_cancel( - payload: dict[str, Any] | None = None, - db: Session = Depends(get_db), -) -> dict[str, Any]: - client = _ume_client() - body = payload or {} - force_clear_local = bool(body.get("force_clear_local")) - try: - st = cancel_alarm_subscription_manual(client, db, force_clear_local=force_clear_local) - if st.get("needs_local_cleanup"): - return st - return {"ok": True, **st} - except Exception as exc: - msg = str(exc)[:240] - raise HTTPException(status_code=502, detail=msg) from exc - - -@router.post("/v1/ume/alarm-subscription/clear-local") -def ume_alarm_subscription_clear_local(db: Session = Depends(get_db)) -> dict[str, Any]: - try: - st = clear_local_alarm_subscription_manual(db) - return {"ok": True, "cleared_local": True, **st} - except Exception as exc: - msg = str(exc)[:240] - raise HTTPException(status_code=502, detail=msg) from exc - - -@router.get("/v1/ume/key-alert-rules") -def ume_list_key_alert_rules( - db: Session = Depends(get_db), - page: int = Query(default=1, ge=1), - page_size: int = Query(default=50, ge=1, le=200), - keyword: str = Query(default=""), - enabled: str | None = Query(default=None), - match_type: str | None = Query(default=None), -) -> dict[str, Any]: - from sqlalchemy import func, or_ - - q = db.query(UmeKeyAlertRule) - kw = str(keyword or "").strip() - if kw: - like = f"%{kw}%" - q = q.filter( - or_( - UmeKeyAlertRule.notification_id.ilike(like), - UmeKeyAlertRule.match_value.ilike(like), - UmeKeyAlertRule.label.ilike(like), - ) - ) - if enabled is not None: - en = str(enabled).strip().lower() - if en in {"1", "true", "yes", "on"}: - q = q.filter(UmeKeyAlertRule.enabled == 1) - elif en in {"0", "false", "no", "off"}: - q = q.filter(UmeKeyAlertRule.enabled == 0) - if match_type: - mt = normalize_match_type(str(match_type)) - q = q.filter(UmeKeyAlertRule.match_type == mt) - - total = int(q.count()) - rows = ( - q.order_by(UmeKeyAlertRule.notification_id.asc()) - .offset((page - 1) * page_size) - .limit(page_size) - .all() - ) - stat_rows = ( - db.query( - UmeKeyAlertForwardLog.rule_key, - func.count(UmeKeyAlertForwardLog.id).label("attempts"), - func.sum(UmeKeyAlertForwardLog.oclaw_ok).label("published_ok"), - func.max(UmeKeyAlertForwardLog.forwarded_at).label("last_forwarded_at"), - ) - .filter(UmeKeyAlertForwardLog.rule_key != "") - .group_by(UmeKeyAlertForwardLog.rule_key) - .all() - ) - stat_map = { - str(rk or ""): { - "attempts": int(attempts or 0), - "published_ok": int(published_ok or 0), - "last_forwarded_at": (_ensure_utc(last_at) or datetime.now(timezone.utc)).isoformat() if last_at else "", - } - for rk, attempts, published_ok, last_at in stat_rows - if str(rk or "").strip() - } - items = [ - { - "notification_id": str(row.notification_id or ""), - "match_type": rule_match_type(row), - "match_value": rule_match_value(row), - "enabled": bool(int(row.enabled or 0)), - "label": str(row.label or ""), - "ne_types": rule_ne_types(row), - "created_at": (_ensure_utc(row.created_at) or datetime.now(timezone.utc)).isoformat(), - "updated_at": (_ensure_utc(row.updated_at) or datetime.now(timezone.utc)).isoformat(), - "forward_stats": stat_map.get(str(row.notification_id or ""), { - "attempts": 0, - "published_ok": 0, - "last_forwarded_at": "", - }), - } - for row in rows - ] - fwd = forwarder_status() - return {"items": items, "total": total, "page": page, "page_size": page_size, "forwarder": fwd} - - -@router.get("/v1/ume/key-alert-monitor") -def ume_key_alert_monitor( - db: Session = Depends(get_db), - page: int = Query(default=1, ge=1), - page_size: int = Query(default=50, ge=1, le=200), - keyword: str = Query(default=""), - enabled: str | None = Query(default=None), - match_type: str | None = Query(default=None), -) -> dict[str, Any]: - base = ume_list_key_alert_rules( - db=db, - page=page, - page_size=page_size, - keyword=keyword, - enabled=enabled, - match_type=match_type, - ) - return { - "ok": True, - "rules": base.get("items") or [], - "total": int(base.get("total") or 0), - "page": int(base.get("page") or page), - "page_size": int(base.get("page_size") or page_size), - "config": get_key_alert_monitor_config(db), - "forwarder": base.get("forwarder") or forwarder_status(), - } - - -@router.patch("/v1/ume/key-alert-monitor/config") -def ume_update_key_alert_monitor_config(payload: dict[str, Any], db: Session = Depends(get_db)) -> dict[str, Any]: - if "forward_on_clear" not in payload: - raise HTTPException(status_code=400, detail="forward_on_clear_required") - config = set_key_alert_monitor_config(db, forward_on_clear=bool(payload.get("forward_on_clear"))) - return {"ok": True, "config": config} - - -@router.post("/v1/ume/key-alert-rules") -def ume_upsert_key_alert_rule(payload: dict[str, Any], db: Session = Depends(get_db)) -> dict[str, Any]: - match_type = normalize_match_type(str(payload.get("match_type") or "notification_id")) - match_value = str(payload.get("match_value") or payload.get("notification_id") or "").strip() - if not match_value: - raise HTTPException(status_code=400, detail="match_value_required") - label = str(payload.get("label") or "").strip() - if not label: - raise HTTPException(status_code=400, detail="label_required") - enabled = 1 if bool(payload.get("enabled", True)) else 0 - ne_types_list = parse_rule_ne_types_payload(payload.get("ne_types")) - now = datetime.now(timezone.utc).replace(tzinfo=None) - try: - storage_key = rule_storage_key(match_type=match_type, value=match_value) - except ValueError as exc: - raise HTTPException(status_code=400, detail=str(exc)) from exc - row = db.get(UmeKeyAlertRule, storage_key) - if row is None: - row = UmeKeyAlertRule(notification_id=storage_key, created_at=now, updated_at=now) - db.add(row) - row.match_type = match_type - row.match_value = match_value - row.enabled = enabled - row.label = label - row.ne_types = serialize_rule_ne_types(ne_types_list) - row.updated_at = now - saved = { - "notification_id": storage_key, - "match_type": match_type, - "match_value": match_value, - "enabled": bool(enabled), - "label": label, - "ne_types": ne_types_list, - } - try: - db.commit() - except Exception as exc: - db.rollback() - msg = str(exc).lower() - if "match_type" in msg or "match_value" in msg or "ne_types" in msg or "undefinedcolumn" in msg: - raise HTTPException( - status_code=503, - detail="key_alert_schema_outdated: restart netx API to apply database migration", - ) from exc - raise - invalidate_key_alert_rule_cache() - return {"ok": True, "item": saved} - - -@router.patch("/v1/ume/key-alert-rules/{rule_key:path}") -def ume_patch_key_alert_rule(rule_key: str, payload: dict[str, Any], db: Session = Depends(get_db)) -> dict[str, Any]: - key = str(rule_key or "").strip() - if not key: - raise HTTPException(status_code=400, detail="rule_key_required") - row = db.get(UmeKeyAlertRule, key) - if row is None: - raise HTTPException(status_code=404, detail="rule_not_found") - has_enabled = "enabled" in payload - has_ne_types = "ne_types" in payload - if not has_enabled and not has_ne_types: - raise HTTPException(status_code=400, detail="patch_fields_required") - now = datetime.now(timezone.utc).replace(tzinfo=None) - if has_enabled: - row.enabled = 1 if bool(payload.get("enabled")) else 0 - if has_ne_types: - row.ne_types = serialize_rule_ne_types(parse_rule_ne_types_payload(payload.get("ne_types"))) - row.updated_at = now - db.commit() - invalidate_key_alert_rule_cache() - return { - "ok": True, - "item": { - "notification_id": key, - "match_type": rule_match_type(row), - "match_value": rule_match_value(row), - "enabled": bool(int(row.enabled or 0)), - "label": str(row.label or ""), - "ne_types": rule_ne_types(row), - }, - } - - -@router.delete("/v1/ume/key-alert-rules/{rule_key:path}") -def ume_delete_key_alert_rule(rule_key: str, db: Session = Depends(get_db)) -> dict[str, Any]: - key = str(rule_key or "").strip() - row = db.get(UmeKeyAlertRule, key) - if row is None: - raise HTTPException(status_code=404, detail="rule_not_found") - db.delete(row) - db.commit() - invalidate_key_alert_rule_cache() - return {"ok": True, "deleted": key} - - -@router.get("/v1/ume/alarm-keywords") -def ume_list_alarm_keywords( - limit: int = Query(default=200, ge=1, le=2000), - db: Session = Depends(get_db), -) -> dict[str, Any]: - from sqlalchemy import func - - rows = ( - db.query( - UmeAlarmCurrent.native_probable_cause, - func.count(UmeAlarmCurrent.alarm_key).label("cnt"), - ) - .filter(UmeAlarmCurrent.native_probable_cause != "") - .group_by(UmeAlarmCurrent.native_probable_cause) - .order_by(func.count(UmeAlarmCurrent.alarm_key).desc(), UmeAlarmCurrent.native_probable_cause.asc()) - .limit(limit) - .all() - ) - items = [ - { - "keyword": str(cause or ""), - "alarm_count": int(cnt or 0), - } - for cause, cnt in rows - if str(cause or "").strip() - ] - return {"items": items, "total": len(items)} - - -@router.get("/v1/ume/notification-ids") -def ume_list_notification_ids( - limit: int = Query(default=200, ge=1, le=2000), - db: Session = Depends(get_db), -) -> dict[str, Any]: - from sqlalchemy import func - - rows = ( - db.query( - UmeAlarmCurrent.notification_id, - func.max(UmeAlarmCurrent.native_probable_cause).label("cause_sample"), - ) - .filter(UmeAlarmCurrent.notification_id != "") - .group_by(UmeAlarmCurrent.notification_id) - .order_by(UmeAlarmCurrent.notification_id.asc()) - .limit(limit) - .all() - ) - items = [ - { - "notification_id": str(nid or ""), - "native_probable_cause_sample": str(cause or ""), - } - for nid, cause in rows - if str(nid or "").strip() - ] - return {"items": items, "total": len(items), "forwarder": forwarder_status()} - - -@router.post("/v1/ume/sync") -def ume_sync(payload: dict[str, Any] | None = None, db: Session = Depends(get_db)) -> dict[str, Any]: - body = payload or {} - domains = body.get("domains") - if not isinstance(domains, list) or not domains: - domains = ["inventory", "alarms_current", "alarms_history"] - domain_set = {str(x).strip().lower() for x in domains if str(x).strip()} - trigger_mode = str(body.get("trigger_mode") or "manual").strip().lower() - if trigger_mode not in {"manual", "schedule"}: - trigger_mode = "manual" - - client = _ume_client() - out: dict[str, Any] = {"ok": True, "jobs": []} - try: - if "inventory" in domain_set: - job = sync_inventory_full(db, client, trigger_mode=trigger_mode) - out["jobs"].append( - { - "domain": "inventory", - "status": job.status, - "pulled_count": int(job.pulled_count or 0), - "inserted_count": int(job.inserted_count or 0), - "updated_count": int(job.updated_count or 0), - "error_message": str(job.error_message or ""), - } - ) - if "alarms" in domain_set or "alarms_current" in domain_set: - paused_ws_for_sync = False - if is_wss_active_for_current_alarms() and trigger_mode == "manual": - _runtime_pause_task("alarms_current_ws_consumer") - request_ws_reconnect() - paused_ws_for_sync = True - try: - job, batch = sync_alarms_current( - db, - client, - trigger_mode=trigger_mode, - wss_active=is_wss_active_for_current_alarms(), - ) - finally: - if paused_ws_for_sync: - _runtime_resume_task("alarms_current_ws_consumer") - request_ws_reconnect() - out["jobs"].append( - { - "domain": "alarms_current", - "status": job.status, - "batch_id": str(batch.batch_id), - "pulled_count": int(job.pulled_count or 0), - "inserted_count": int(job.inserted_count or 0), - "updated_count": int(job.updated_count or 0), - "error_message": str(job.error_message or ""), - } - ) - if "alarms_history" in domain_set: - job, batch = sync_alarms_history_full(db, client, trigger_mode=trigger_mode) - out["jobs"].append( - { - "domain": "alarms_history", - "status": job.status, - "batch_id": str(batch.batch_id), - "pulled_count": int(job.pulled_count or 0), - "inserted_count": int(job.inserted_count or 0), - "updated_count": int(job.updated_count or 0), - "error_message": str(job.error_message or ""), - } - ) - except Exception as exc: - out["ok"] = False - out["error"] = str(exc)[:240] - return out - - -def _ume_sync_job_deleted_count(row: UmeSyncJob) -> int: - """Single reconcile delete count: inventory uses deleted_inventory_ne; current alarms uses deleted_stale_current_alarms.""" - raw = str(getattr(row, "details_json", "") or "").strip() - if not raw: - return 0 - try: - obj = json.loads(raw) - except Exception: - return 0 - if not isinstance(obj, dict): - return 0 - inv = cur = 0 - try: - inv = max(0, int(obj.get("deleted_inventory_ne") or 0)) - except Exception: - pass - try: - cur = max(0, int(obj.get("deleted_stale_current_alarms") or 0)) - except Exception: - pass - return int(inv + cur) - - -@router.get("/v1/ume/sync/status") -def ume_sync_status( - page: int = Query(default=1, ge=1), - page_size: int = Query(default=20, ge=1, le=200), - db: Session = Depends(get_db), -) -> dict[str, Any]: - q = db.query(UmeSyncJob) - total = int(q.count()) - rows = ( - q.order_by(UmeSyncJob.id.desc()) - .offset((int(page) - 1) * int(page_size)) - .limit(int(page_size)) - .all() - ) - items = [] - latest_by_domain: dict[str, dict[str, Any]] = {} - for r in rows: - item = { - "id": int(r.id), - "domain": str(r.domain or ""), - "status": str(r.status or ""), - "trigger_mode": str(r.trigger_mode or ""), - "pulled_count": int(r.pulled_count or 0), - "inserted_count": int(r.inserted_count or 0), - "updated_count": int(r.updated_count or 0), - "deleted": int(_ume_sync_job_deleted_count(r)), - "error_message": str(r.error_message or ""), - "started_at": (_ensure_utc(r.started_at) or datetime.now(timezone.utc)).isoformat(), - "ended_at": (_ensure_utc(r.ended_at).isoformat() if r.ended_at else None), - } - items.append(item) - if item["domain"] and item["domain"] not in latest_by_domain: - latest_by_domain[item["domain"]] = item - return { - "total": total, - "page": page, - "page_size": page_size, - "items": items, - "latest_by_domain": latest_by_domain, - "runtime_tasks": _list_runtime_tasks(), - "alarm_subscription": get_subscription_status(), - } - - -@router.post("/v1/ume/runtime/tasks/{task}/pause") -def ume_runtime_task_pause(task: str) -> dict[str, Any]: - tid = str(task or "").strip() - if tid not in UME_KNOWN_RUNTIME_TASKS: - raise HTTPException(status_code=404, detail="unknown_runtime_task") - _runtime_pause_task(tid) - if tid in ("alarms_current_auto_sync", "inventory_auto_sync"): - _clear_force_resume_hints(tid) - if tid == "alarms_current_ws_consumer": - request_ws_reconnect() - if tid == "oclaw_alarm_forwarder": - request_forwarder_reconnect() - _set_runtime_task(tid, status="paused", last_error="") - return {"ok": True, "task": tid, "runtime_tasks": _list_runtime_tasks()} - - -@router.post("/v1/ume/runtime/tasks/{task}/resume") -def ume_runtime_task_resume(task: str) -> dict[str, Any]: - tid = str(task or "").strip() - if tid not in UME_KNOWN_RUNTIME_TASKS: - raise HTTPException(status_code=404, detail="unknown_runtime_task") - _runtime_resume_task(tid) - if tid in ("alarms_current_auto_sync", "inventory_auto_sync"): - _request_force_sync_after_resume(tid) - resume_hint = RT_RESUMED_SYNC_SOON - elif tid == "alarms_current_ws_consumer": - request_ws_reconnect() - resume_hint = RT_RESUMED_WSS_RECONNECT - elif tid == "oclaw_alarm_forwarder": - request_forwarder_reconnect() - resume_hint = RT_RESUMED_OCLAW_WSS_RECONNECT - else: - resume_hint = RT_RESUMED - _set_runtime_task(tid, status="running", last_error=resume_hint) - return {"ok": True, "task": tid, "runtime_tasks": _list_runtime_tasks()} - - -@router.get("/v1/ume/inventory/ne-types") -def ume_list_inventory_ne_types( - limit: int = Query(default=500, ge=1, le=2000), - db: Session = Depends(get_db), -) -> dict[str, Any]: - from sqlalchemy import func - - rows = ( - db.query( - UmeInventoryNE.ne_type, - func.count(UmeInventoryNE.ne_id).label("ne_count"), - ) - .filter(UmeInventoryNE.ne_type != "") - .group_by(UmeInventoryNE.ne_type) - .order_by(func.count(UmeInventoryNE.ne_id).desc(), UmeInventoryNE.ne_type.asc()) - .limit(limit) - .all() - ) - items = [{"ne_type": str(ne_type or ""), "ne_count": int(ne_count or 0)} for ne_type, ne_count in rows if str(ne_type or "").strip()] - return {"items": items, "total": len(items)} - - -@router.get("/v1/ume/inventory/ne") -def ume_list_ne( - keyword: str | None = Query(default=None), - page: int = Query(default=1, ge=1), - page_size: int = Query(default=50, ge=1, le=500), - db: Session = Depends(get_db), -) -> dict[str, Any]: - stmt = db.query(UmeInventoryNE) - kw = str(keyword or "").strip() - if kw: - like = f"%{kw}%" - stmt = stmt.filter( - UmeInventoryNE.ne_id.ilike(like) - | UmeInventoryNE.ne_name.ilike(like) - | UmeInventoryNE.user_label.ilike(like) - | UmeInventoryNE.ip_address.ilike(like) - | UmeInventoryNE.host_name.ilike(like) - ) - total = int(stmt.count()) - rows = stmt.order_by(UmeInventoryNE.ne_id.asc()).offset((page - 1) * page_size).limit(page_size).all() - items = [ - { - "ne_id": str(x.ne_id or ""), - "ne_name": str(x.ne_name or ""), - "user_label": str(x.user_label or ""), - "ip_address": str(x.ip_address or ""), - "ipv6_address": str(x.ipv6_address or ""), - "ne_type": str(x.ne_type or ""), - "device_level": str(x.device_level or ""), - "host_name": str(x.host_name or ""), - "location": str(x.location or ""), - "hardware_version": str(x.hardware_version or ""), - "loopback": str(x.loopback or ""), - "consistent_state": str(x.consistent_state or ""), - "interface_version": str(x.interface_version or ""), - "mac": str(x.mac or ""), - "admin_status": str(x.admin_status or ""), - "address_type": str(x.address_type or ""), - "connection_status": str(x.connection_status or ""), - "maintain_status": str(x.maintain_status or ""), - "net_mask": str(x.net_mask or ""), - "create_time": str(x.create_time or ""), - "creator": str(x.creator or ""), - "last_seen_at": (_ensure_utc(x.last_seen_at) or datetime.now(timezone.utc)).isoformat(), - } - for x in rows - ] - return {"total": total, "page": page, "page_size": page_size, "items": items} - - -@router.get("/v1/ume/inventory/ne/{ne_id}") -def ume_get_ne(ne_id: str, db: Session = Depends(get_db)) -> dict[str, Any]: - row = db.get(UmeInventoryNE, ne_id) - if not row: - raise HTTPException(status_code=404, detail="ume_ne_not_found") - return { - "ne_id": str(row.ne_id or ""), - "ne_name": str(row.ne_name or ""), - "user_label": str(row.user_label or ""), - "ip_address": str(row.ip_address or ""), - "ipv6_address": str(row.ipv6_address or ""), - "ne_type": str(row.ne_type or ""), - "device_level": str(row.device_level or ""), - "host_name": str(row.host_name or ""), - "location": str(row.location or ""), - "hardware_version": str(row.hardware_version or ""), - "loopback": str(row.loopback or ""), - "consistent_state": str(row.consistent_state or ""), - "interface_version": str(row.interface_version or ""), - "mac": str(row.mac or ""), - "admin_status": str(row.admin_status or ""), - "address_type": str(row.address_type or ""), - "connection_status": str(row.connection_status or ""), - "maintain_status": str(row.maintain_status or ""), - "net_mask": str(row.net_mask or ""), - "create_time": str(row.create_time or ""), - "creator": str(row.creator or ""), - "vendor": str(row.vendor or ""), - "source_type": str(row.source_type or ""), - "first_seen_at": (_ensure_utc(row.first_seen_at) or datetime.now(timezone.utc)).isoformat(), - "last_seen_at": (_ensure_utc(row.last_seen_at) or datetime.now(timezone.utc)).isoformat(), - "raw_json": str(row.raw_json or "{}"), - } - - -@router.get("/v1/ume/alarms") -def ume_list_alarms( - severity: str | None = Query(default=None), - is_cleared: str | None = Query(default=None), - ne_id: str | None = Query(default=None), - host_name: str | None = Query(default=None), - keyword: str | None = Query(default=None), - page: int = Query(default=1, ge=1), - page_size: int = Query(default=50, ge=1, le=500), - db: Session = Depends(get_db), -) -> dict[str, Any]: - stmt = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( - UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id - ) - if severity and str(severity).strip(): - stmt = stmt.filter(UmeAlarmCurrent.perceived_severity == str(severity).strip()) - if is_cleared and str(is_cleared).strip(): - stmt = stmt.filter(UmeAlarmCurrent.is_cleared == str(is_cleared).strip()) - if ne_id and str(ne_id).strip(): - stmt = stmt.filter(UmeAlarmCurrent.ne_id == str(ne_id).strip()) - hn = str(host_name or "").strip() - if hn: - stmt = stmt.filter( - UmeAlarmCurrent.host_name.contains(hn) | UmeInventoryNE.host_name.contains(hn) - ) - kw = str(keyword or "").strip() - if kw: - stmt = stmt.filter( - UmeAlarmCurrent.alarm_key.contains(kw) - | UmeAlarmCurrent.object_name.contains(kw) - | UmeAlarmCurrent.native_probable_cause.contains(kw) - | UmeAlarmCurrent.notification_id.contains(kw) - | UmeAlarmCurrent.host_name.contains(kw) - | UmeInventoryNE.ne_name.contains(kw) - | UmeInventoryNE.user_label.contains(kw) - | UmeInventoryNE.ip_address.contains(kw) - | UmeInventoryNE.host_name.contains(kw) - ) - total = int(stmt.count()) - rows = ( - stmt.order_by( - UmeAlarmCurrent.time_created.desc(), - UmeAlarmCurrent.last_seen_at.desc(), - UmeAlarmCurrent.alarm_key.desc(), - ) - .offset((page - 1) * page_size) - .limit(page_size) - .all() - ) - items = [ - { - "alarm_key": str(alarm.alarm_key or ""), - "ne_id": str(alarm.ne_id or ""), - "ne_name": str((ne.ne_name if ne else "") or ""), - "user_label": str((ne.user_label if ne else "") or ""), - "host_name": _ume_alarm_host_name(alarm, ne), - "ne_type": str((ne.ne_type if ne else "") or ""), - "object_name": str(alarm.object_name or ""), - "event_type": str(alarm.event_type or ""), - "native_probable_cause": str(alarm.native_probable_cause or ""), - "notification_id": str(alarm.notification_id or ""), - "perceived_severity": str(alarm.perceived_severity or ""), - "is_cleared": str(alarm.is_cleared or ""), - "time_created": str(alarm.time_created or ""), - "last_seen_at": (_ensure_utc(alarm.last_seen_at) or datetime.now(timezone.utc)).isoformat(), - } - for alarm, ne in rows - ] - return {"total": total, "page": page, "page_size": page_size, "items": items} - - -@router.get("/v1/ume/alarms/fields") -def ume_alarms_fields() -> dict[str, Any]: - """List all queryable field names for UME raw alarm query.""" - alarm_cols = [str(c.name) for c in UmeAlarmCurrent.__table__.columns] # type: ignore[attr-defined] - ne_cols = [str(c.name) for c in UmeInventoryNE.__table__.columns] # type: ignore[attr-defined] - selectable_fields = [f"alarm_{x}" for x in alarm_cols] + [f"ne_{x}" for x in ne_cols] + ["ne_exists"] - order_by_allowed = ["last_seen_at", "time_created", "perceived_severity", "event_type", "ne_id"] - return { - "alarm_fields": alarm_cols, - "ne_fields": ne_cols, - "selectable_fields": selectable_fields, - "order_by_allowed": order_by_allowed, - } - - -def _serialize_ume_alarm_raw_row( - alarm: UmeAlarmCurrent, ne: UmeInventoryNE | None, selected_fields: set[str] | None = None -) -> dict[str, Any]: - selected = selected_fields or set() - use_all = len(selected) == 0 - out: dict[str, Any] = {} - for c in UmeAlarmCurrent.__table__.columns: # type: ignore[attr-defined] - name = str(c.name) - v = getattr(alarm, name, None) - key = f"alarm_{name}" - if not use_all and key not in selected: - continue - if hasattr(v, "isoformat"): - try: - if isinstance(v, datetime): - out[key] = (_ensure_utc(v) or v).isoformat() - else: - out[key] = v.isoformat() - continue - except Exception: - pass - out[key] = v - if ne is None: - if use_all or "ne_exists" in selected: - out["ne_exists"] = False - return out - if use_all or "ne_exists" in selected: - out["ne_exists"] = True - for c in UmeInventoryNE.__table__.columns: # type: ignore[attr-defined] - name = str(c.name) - v = getattr(ne, name, None) - key = f"ne_{name}" - if not use_all and key not in selected: - continue - if hasattr(v, "isoformat"): - try: - if isinstance(v, datetime): - out[key] = (_ensure_utc(v) or v).isoformat() - else: - out[key] = v.isoformat() - continue - except Exception: - pass - out[key] = v - return out - - -def _extract_ume_raw_group_field(alarm: UmeAlarmCurrent, ne: UmeInventoryNE | None, field: str) -> str: - key = str(field or "").strip() - if not key: - return "" - if key.startswith("alarm_"): - attr = key[len("alarm_") :] - return str(getattr(alarm, attr, "") or "") - if key.startswith("ne_"): - attr = key[len("ne_") :] - if key == "ne_exists": - return "1" if ne is not None else "0" - if key == "ne_host_name": - hn = str(getattr(alarm, "host_name", "") or "").strip() - if hn: - return hn - if ne is None: - return "" - return str(getattr(ne, attr, "") or "") - return "" - - -@router.get("/v1/ume/alarms/raw") -def ume_alarms_raw( - severity: str | None = Query(default=None), - is_cleared: str | None = Query(default=None), - ne_id: str | None = Query(default=None), - event_type: str | None = Query(default=None), - keyword: str | None = Query(default=None), - time_from: str | None = Query(default=None), - time_to: str | None = Query(default=None), - order_by: str = Query(default="last_seen_at"), - order: str = Query(default="desc"), - select_fields: str | None = Query(default=None, description="comma-separated alarm_*/ne_* fields"), - page: int = Query(default=1, ge=1), - page_size: int = Query(default=50, ge=1, le=500), - db: Session = Depends(get_db), -) -> dict[str, Any]: - stmt = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( - UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id - ) - if severity and str(severity).strip(): - stmt = stmt.filter(UmeAlarmCurrent.perceived_severity == str(severity).strip()) - if is_cleared and str(is_cleared).strip(): - stmt = stmt.filter(UmeAlarmCurrent.is_cleared == str(is_cleared).strip()) - if ne_id and str(ne_id).strip(): - stmt = stmt.filter(UmeAlarmCurrent.ne_id == str(ne_id).strip()) - if event_type and str(event_type).strip(): - stmt = stmt.filter(UmeAlarmCurrent.event_type.contains(str(event_type).strip())) - kw = str(keyword or "").strip() - if kw: - stmt = stmt.filter( - UmeAlarmCurrent.alarm_key.contains(kw) - | UmeAlarmCurrent.object_name.contains(kw) - | UmeAlarmCurrent.native_probable_cause.contains(kw) - | UmeAlarmCurrent.event_type.contains(kw) - | UmeInventoryNE.ne_name.contains(kw) - | UmeInventoryNE.user_label.contains(kw) - | UmeInventoryNE.ip_address.contains(kw) - ) - dt_from = _parse_time(time_from) - dt_to = _parse_time(time_to) - if dt_from: - stmt = stmt.filter(UmeAlarmCurrent.last_seen_at >= dt_from.replace(tzinfo=None)) - if dt_to: - stmt = stmt.filter(UmeAlarmCurrent.last_seen_at <= dt_to.replace(tzinfo=None)) - - allowed_order_by = { - "last_seen_at": UmeAlarmCurrent.last_seen_at, - "time_created": UmeAlarmCurrent.time_created, - "perceived_severity": UmeAlarmCurrent.perceived_severity, - "event_type": UmeAlarmCurrent.event_type, - "ne_id": UmeAlarmCurrent.ne_id, - } - col = allowed_order_by.get(str(order_by or "").strip(), UmeAlarmCurrent.last_seen_at) - if str(order or "").strip().lower() == "asc": - stmt = stmt.order_by(col.asc()) - else: - stmt = stmt.order_by(col.desc()) - - selected_fields: set[str] = set() - fields_meta = ume_alarms_fields() - selectable_fields = set(str(x) for x in (fields_meta.get("selectable_fields") or [])) - order_by_allowed = [str(x) for x in (fields_meta.get("order_by_allowed") or [])] - if select_fields and str(select_fields).strip(): - selected_fields = {x.strip() for x in str(select_fields).split(",") if x.strip()} - invalid = [x for x in selected_fields if x not in selectable_fields] - if invalid: - raise HTTPException(status_code=400, detail=f"invalid_select_fields:{','.join(sorted(invalid)[:20])}") - - total = int(stmt.count()) - rows = stmt.offset((int(page) - 1) * int(page_size)).limit(int(page_size)).all() - return { - "total": total, - "page": int(page), - "page_size": int(page_size), - "select_fields": sorted(selected_fields) if selected_fields else [], - "meta": { - "available_fields": sorted(selectable_fields), - "order_by_allowed": order_by_allowed, - "time_filter_field": "last_seen_at", - }, - "items": [_serialize_ume_alarm_raw_row(alarm, ne, selected_fields) for alarm, ne in rows], - } - - -@router.get("/v1/ume/alarms/aggregate/raw") -def ume_alarms_aggregate_raw( - group_by: str = Query(default="alarm_perceived_severity"), - group_by2: str | None = Query(default=None), - severity: str | None = Query(default=None), - is_cleared: str | None = Query(default=None), - ne_id: str | None = Query(default=None), - event_type: str | None = Query(default=None), - keyword: str | None = Query(default=None), - time_from: str | None = Query(default=None), - time_to: str | None = Query(default=None), - limit: int = Query(default=200, ge=1, le=2000), - db: Session = Depends(get_db), -) -> dict[str, Any]: - fields_meta = ume_alarms_fields() - selectable_fields = set(str(x) for x in (fields_meta.get("selectable_fields") or [])) - g1 = str(group_by or "").strip() - g2 = str(group_by2 or "").strip() - if g1 not in selectable_fields: - raise HTTPException(status_code=400, detail=f"invalid_group_by:{g1}") - if g2 and g2 not in selectable_fields: - raise HTTPException(status_code=400, detail=f"invalid_group_by2:{g2}") - - stmt = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( - UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id - ) - if severity and str(severity).strip(): - stmt = stmt.filter(UmeAlarmCurrent.perceived_severity == str(severity).strip()) - if is_cleared and str(is_cleared).strip(): - stmt = stmt.filter(UmeAlarmCurrent.is_cleared == str(is_cleared).strip()) - if ne_id and str(ne_id).strip(): - stmt = stmt.filter(UmeAlarmCurrent.ne_id == str(ne_id).strip()) - if event_type and str(event_type).strip(): - stmt = stmt.filter(UmeAlarmCurrent.event_type.contains(str(event_type).strip())) - kw = str(keyword or "").strip() - if kw: - stmt = stmt.filter( - UmeAlarmCurrent.alarm_key.contains(kw) - | UmeAlarmCurrent.object_name.contains(kw) - | UmeAlarmCurrent.native_probable_cause.contains(kw) - | UmeAlarmCurrent.event_type.contains(kw) - | UmeInventoryNE.ne_name.contains(kw) - | UmeInventoryNE.user_label.contains(kw) - | UmeInventoryNE.ip_address.contains(kw) - ) - dt_from = _parse_time(time_from) - dt_to = _parse_time(time_to) - if dt_from: - stmt = stmt.filter(UmeAlarmCurrent.last_seen_at >= dt_from.replace(tzinfo=None)) - if dt_to: - stmt = stmt.filter(UmeAlarmCurrent.last_seen_at <= dt_to.replace(tzinfo=None)) - - rows = stmt.order_by(UmeAlarmCurrent.last_seen_at.desc()).all() - counts: dict[tuple[str, str], int] = {} - for alarm, ne in rows: - k1 = _extract_ume_raw_group_field(alarm, ne, g1) - k2 = _extract_ume_raw_group_field(alarm, ne, g2) if g2 else "" - kk = (k1, k2) - counts[kk] = int(counts.get(kk, 0)) + 1 - buckets = sorted(counts.items(), key=lambda x: x[1], reverse=True)[: int(limit)] - return { - "total": len(rows), - "group_by": g1, - "group_by2": g2 or None, - "meta": { - "available_fields": sorted(selectable_fields), - "group_by_allowed": sorted(selectable_fields), - "applied_filters": { - "severity": str(severity or "").strip() or None, - "is_cleared": str(is_cleared or "").strip() or None, - "ne_id": str(ne_id or "").strip() or None, - "event_type": str(event_type or "").strip() or None, - "keyword": str(keyword or "").strip() or None, - "time_from": str(time_from or "").strip() or None, - "time_to": str(time_to or "").strip() or None, - }, - "time_filter_field": "last_seen_at", - "limit": int(limit), - }, - "buckets": [ - {"key": k1, "key2": (k2 if g2 else None), "count": int(v)} - for (k1, k2), v in buckets - ], - } - - -@router.get("/v1/ume/alarms/aggregate") -def ume_alarms_aggregate(db: Session = Depends(get_db)) -> dict[str, Any]: - rows = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( - UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id - ).all() - by_severity = _aggregate_rows(rows, lambda x: x[0].perceived_severity) - by_ne = _aggregate_rows(rows, lambda x: _ume_alarm_ne_group_key(x[0], x[1])) - return {"total": len(rows), "by_severity": by_severity, "by_ne": by_ne} - - -@router.get("/v1/ume/diagnostics") -def ume_diagnostics( - lang: str | None = Query(default=None), - db: Session = Depends(get_db), -) -> dict[str, Any]: - rows = db.query(UmeAlarmCurrent, UmeInventoryNE).outerjoin( - UmeInventoryNE, UmeAlarmCurrent.ne_id == UmeInventoryNE.ne_id - ).all() - by_severity = _aggregate_rows(rows, lambda x: x[0].perceived_severity) - by_alarm_code = _aggregate_rows(rows, lambda x: x[0].event_type)[:10] - by_ne = _aggregate_rows(rows, lambda x: _ume_alarm_ne_group_key(x[0], x[1]))[:10] - - lang_norm = _normalize_netx_lang(lang) - proto_counts: dict[str, int] = {} - for alarm, ne in rows: - blob = " | ".join( - [ - str(alarm.event_type or ""), - str(alarm.native_probable_cause or ""), - str(alarm.object_name or ""), - str(ne.ne_name if ne else ""), - str(ne.user_label if ne else ""), - str(ne.ip_address if ne else ""), - ] - ) - bucket = _protocol_bucket_label(blob, lang=lang_norm) - proto_counts[bucket] = int(proto_counts.get(bucket, 0)) + 1 - protocol_summary = sorted(proto_counts.items(), key=lambda x: x[1], reverse=True)[:10] - - return { - "source": "ume_alarms_current", - "total_alarms": len(rows), - "severity_summary": [{"key": k, "count": v} for k, v in by_severity], - "top_alarm_codes": [{"key": k, "count": v} for k, v in by_alarm_code], - "top_ne": [{"key": k, "count": v} for k, v in by_ne], - "protocol_summary": [{"key": k, "count": v} for k, v in protocol_summary], - } - - -@router.get("/v1/ume/alarms/history") -def ume_list_alarms_history( - severity: str | None = Query(default=None), - ne_id: str | None = Query(default=None), - keyword: str | None = Query(default=None), - time_from: str | None = Query(default=None), - time_to: str | None = Query(default=None), - page: int = Query(default=1, ge=1), - page_size: int = Query(default=50, ge=1, le=500), - db: Session = Depends(get_db), -) -> dict[str, Any]: - stmt = db.query(UmeAlarmHistory, UmeInventoryNE).outerjoin( - UmeInventoryNE, UmeAlarmHistory.ne_id == UmeInventoryNE.ne_id - ) - if severity and str(severity).strip(): - stmt = stmt.filter(UmeAlarmHistory.perceived_severity == str(severity).strip()) - if ne_id and str(ne_id).strip(): - stmt = stmt.filter(UmeAlarmHistory.ne_id == str(ne_id).strip()) - kw = str(keyword or "").strip() - if kw: - stmt = stmt.filter( - UmeAlarmHistory.alarm_key.contains(kw) - | UmeAlarmHistory.object_name.contains(kw) - | UmeAlarmHistory.native_probable_cause.contains(kw) - | UmeInventoryNE.ne_name.contains(kw) - | UmeInventoryNE.user_label.contains(kw) - | UmeInventoryNE.ip_address.contains(kw) - ) - dt_from = _parse_time(time_from) - dt_to = _parse_time(time_to) - if dt_from: - stmt = stmt.filter(UmeAlarmHistory.last_seen_at >= dt_from.replace(tzinfo=None)) - if dt_to: - stmt = stmt.filter(UmeAlarmHistory.last_seen_at <= dt_to.replace(tzinfo=None)) - total = int(stmt.count()) - rows = stmt.order_by(UmeAlarmHistory.last_seen_at.desc()).offset((page - 1) * page_size).limit(page_size).all() - items = [ - { - "alarm_key": str(alarm.alarm_key or ""), - "ne_id": str(alarm.ne_id or ""), - "ne_name": str((ne.ne_name if ne else "") or ""), - "user_label": str((ne.user_label if ne else "") or ""), - "object_name": str(alarm.object_name or ""), - "event_type": str(alarm.event_type or ""), - "native_probable_cause": str(alarm.native_probable_cause or ""), - "perceived_severity": str(alarm.perceived_severity or ""), - "is_cleared": str(alarm.is_cleared or ""), - "time_created": str(alarm.time_created or ""), - "last_seen_at": (_ensure_utc(alarm.last_seen_at) or datetime.now(timezone.utc)).isoformat(), - } - for alarm, ne in rows - ] - return {"total": total, "page": page, "page_size": page_size, "items": items} - - -@router.get("/v1/ume/alarms/history/aggregate") -def ume_alarms_history_aggregate(db: Session = Depends(get_db)) -> dict[str, Any]: - rows = db.query(UmeAlarmHistory, UmeInventoryNE).outerjoin( - UmeInventoryNE, UmeAlarmHistory.ne_id == UmeInventoryNE.ne_id - ).all() - by_severity = _aggregate_rows(rows, lambda x: x[0].perceived_severity) - by_ne = _aggregate_rows(rows, lambda x: (x[1].user_label if x[1] else "") or (x[1].ne_name if x[1] else "") or x[0].ne_id) - by_date = _aggregate_rows(rows, lambda x: str(x[0].time_created or "")[:10]) - return {"total": len(rows), "by_severity": by_severity, "by_ne": by_ne, "by_date": by_date} - - +router.include_router(token_router) +router.include_router(key_alert_router) +router.include_router(sync_router) +router.include_router(inventory_router) +router.include_router(alarms_router) + +__all__ = [ + "_extract_ume_raw_group_field", + "_serialize_ume_alarm_raw_row", + "router", + "ume_alarms_fields", +] diff --git a/netx_api/ume_sync_router.py b/netx_api/ume_sync_router.py new file mode 100644 index 0000000..3c19641 --- /dev/null +++ b/netx_api/ume_sync_router.py @@ -0,0 +1,247 @@ +"""UME sync jobs and runtime pause/resume.""" +from __future__ import annotations + +import logging +from datetime import datetime, timezone +from typing import Any + +from fastapi import APIRouter, Depends, HTTPException, Query +from sqlalchemy import or_ +from sqlalchemy.orm import Session + +from .config import settings +from .db import get_db +from .key_alert_config import ( + get_key_alert_monitor_config, + invalidate_key_alert_config_cache, + set_key_alert_monitor_config, +) +from .key_alert_matcher import ( + invalidate_key_alert_rule_cache, + normalize_match_type, + parse_rule_ne_types_payload, + rule_match_type, + rule_match_value, + rule_ne_types, + rule_storage_key, + serialize_rule_ne_types, +) +from .models import ( + UmeAlarmCurrent, + UmeAlarmHistory, + UmeInventoryNE, + UmeKeyAlertForwardLog, + UmeKeyAlertRule, + UmeSyncJob, +) +from .oclaw_alarm_forwarder import ( + forwarder_status, + request_forwarder_reconnect, +) +from .ume_alarm_ws import ( + cancel_alarm_subscription_manual, + clear_local_alarm_subscription_manual, + establish_alarm_subscription_manual, + get_alarms_coordination_status, + get_subscription_status, + get_ws_connection_status, + get_ws_logs, + request_ws_reconnect, +) +from .ume_support import ( + UME_KNOWN_RUNTIME_TASKS, + _aggregate_rows, + _ensure_utc, + _list_runtime_tasks, + _request_force_sync_after_resume, + _runtime_pause_task, + _runtime_resume_task, + _ume_alarm_host_name, + _ume_alarm_ne_group_key, + _ume_client, + _ume_error_kind, + _clear_force_resume_hints, +) +from .ume_sync_service import sync_alarms_current, sync_alarms_history_full, sync_inventory_full +from .ume_token_store import clear_shared_token + +_log = logging.getLogger("netx.ume.router") +router = APIRouter(tags=["ume"]) + +@router.post("/v1/ume/sync") +def ume_sync(payload: dict[str, Any] | None = None, db: Session = Depends(get_db)) -> dict[str, Any]: + body = payload or {} + domains = body.get("domains") + if not isinstance(domains, list) or not domains: + domains = ["inventory", "alarms_current", "alarms_history"] + domain_set = {str(x).strip().lower() for x in domains if str(x).strip()} + trigger_mode = str(body.get("trigger_mode") or "manual").strip().lower() + if trigger_mode not in {"manual", "schedule"}: + trigger_mode = "manual" + + client = _ume_client() + out: dict[str, Any] = {"ok": True, "jobs": []} + try: + if "inventory" in domain_set: + job = sync_inventory_full(db, client, trigger_mode=trigger_mode) + out["jobs"].append( + { + "domain": "inventory", + "status": job.status, + "pulled_count": int(job.pulled_count or 0), + "inserted_count": int(job.inserted_count or 0), + "updated_count": int(job.updated_count or 0), + "error_message": str(job.error_message or ""), + } + ) + if "alarms" in domain_set or "alarms_current" in domain_set: + paused_ws_for_sync = False + if is_wss_active_for_current_alarms() and trigger_mode == "manual": + _runtime_pause_task("alarms_current_ws_consumer") + request_ws_reconnect() + paused_ws_for_sync = True + try: + job, batch = sync_alarms_current( + db, + client, + trigger_mode=trigger_mode, + wss_active=is_wss_active_for_current_alarms(), + ) + finally: + if paused_ws_for_sync: + _runtime_resume_task("alarms_current_ws_consumer") + request_ws_reconnect() + out["jobs"].append( + { + "domain": "alarms_current", + "status": job.status, + "batch_id": str(batch.batch_id), + "pulled_count": int(job.pulled_count or 0), + "inserted_count": int(job.inserted_count or 0), + "updated_count": int(job.updated_count or 0), + "error_message": str(job.error_message or ""), + } + ) + if "alarms_history" in domain_set: + job, batch = sync_alarms_history_full(db, client, trigger_mode=trigger_mode) + out["jobs"].append( + { + "domain": "alarms_history", + "status": job.status, + "batch_id": str(batch.batch_id), + "pulled_count": int(job.pulled_count or 0), + "inserted_count": int(job.inserted_count or 0), + "updated_count": int(job.updated_count or 0), + "error_message": str(job.error_message or ""), + } + ) + except Exception as exc: + out["ok"] = False + out["error"] = str(exc)[:240] + return out + + +def _ume_sync_job_deleted_count(row: UmeSyncJob) -> int: + """Single reconcile delete count: inventory uses deleted_inventory_ne; current alarms uses deleted_stale_current_alarms.""" + raw = str(getattr(row, "details_json", "") or "").strip() + if not raw: + return 0 + try: + obj = json.loads(raw) + except Exception: + return 0 + if not isinstance(obj, dict): + return 0 + inv = cur = 0 + try: + inv = max(0, int(obj.get("deleted_inventory_ne") or 0)) + except Exception: + pass + try: + cur = max(0, int(obj.get("deleted_stale_current_alarms") or 0)) + except Exception: + pass + return int(inv + cur) + + +@router.get("/v1/ume/sync/status") +def ume_sync_status( + page: int = Query(default=1, ge=1), + page_size: int = Query(default=20, ge=1, le=200), + db: Session = Depends(get_db), +) -> dict[str, Any]: + q = db.query(UmeSyncJob) + total = int(q.count()) + rows = ( + q.order_by(UmeSyncJob.id.desc()) + .offset((int(page) - 1) * int(page_size)) + .limit(int(page_size)) + .all() + ) + items = [] + latest_by_domain: dict[str, dict[str, Any]] = {} + for r in rows: + item = { + "id": int(r.id), + "domain": str(r.domain or ""), + "status": str(r.status or ""), + "trigger_mode": str(r.trigger_mode or ""), + "pulled_count": int(r.pulled_count or 0), + "inserted_count": int(r.inserted_count or 0), + "updated_count": int(r.updated_count or 0), + "deleted": int(_ume_sync_job_deleted_count(r)), + "error_message": str(r.error_message or ""), + "started_at": (_ensure_utc(r.started_at) or datetime.now(timezone.utc)).isoformat(), + "ended_at": (_ensure_utc(r.ended_at).isoformat() if r.ended_at else None), + } + items.append(item) + if item["domain"] and item["domain"] not in latest_by_domain: + latest_by_domain[item["domain"]] = item + return { + "total": total, + "page": page, + "page_size": page_size, + "items": items, + "latest_by_domain": latest_by_domain, + "runtime_tasks": _list_runtime_tasks(), + "alarm_subscription": get_subscription_status(), + } + + +@router.post("/v1/ume/runtime/tasks/{task}/pause") +def ume_runtime_task_pause(task: str) -> dict[str, Any]: + tid = str(task or "").strip() + if tid not in UME_KNOWN_RUNTIME_TASKS: + raise HTTPException(status_code=404, detail="unknown_runtime_task") + _runtime_pause_task(tid) + if tid in ("alarms_current_auto_sync", "inventory_auto_sync"): + _clear_force_resume_hints(tid) + if tid == "alarms_current_ws_consumer": + request_ws_reconnect() + if tid == "oclaw_alarm_forwarder": + request_forwarder_reconnect() + _set_runtime_task(tid, status="paused", last_error="") + return {"ok": True, "task": tid, "runtime_tasks": _list_runtime_tasks()} + + +@router.post("/v1/ume/runtime/tasks/{task}/resume") +def ume_runtime_task_resume(task: str) -> dict[str, Any]: + tid = str(task or "").strip() + if tid not in UME_KNOWN_RUNTIME_TASKS: + raise HTTPException(status_code=404, detail="unknown_runtime_task") + _runtime_resume_task(tid) + if tid in ("alarms_current_auto_sync", "inventory_auto_sync"): + _request_force_sync_after_resume(tid) + resume_hint = RT_RESUMED_SYNC_SOON + elif tid == "alarms_current_ws_consumer": + request_ws_reconnect() + resume_hint = RT_RESUMED_WSS_RECONNECT + elif tid == "oclaw_alarm_forwarder": + request_forwarder_reconnect() + resume_hint = RT_RESUMED_OCLAW_WSS_RECONNECT + else: + resume_hint = RT_RESUMED + _set_runtime_task(tid, status="running", last_error=resume_hint) + return {"ok": True, "task": tid, "runtime_tasks": _list_runtime_tasks()} + + diff --git a/netx_api/ume_token_router.py b/netx_api/ume_token_router.py new file mode 100644 index 0000000..e3f7821 --- /dev/null +++ b/netx_api/ume_token_router.py @@ -0,0 +1,164 @@ +"""UME token + alarm subscription routes.""" +from __future__ import annotations + +import logging +from datetime import datetime, timezone +from typing import Any + +from fastapi import APIRouter, Depends, HTTPException, Query +from sqlalchemy import or_ +from sqlalchemy.orm import Session + +from .config import settings +from .db import get_db +from .key_alert_config import ( + get_key_alert_monitor_config, + invalidate_key_alert_config_cache, + set_key_alert_monitor_config, +) +from .key_alert_matcher import ( + invalidate_key_alert_rule_cache, + normalize_match_type, + parse_rule_ne_types_payload, + rule_match_type, + rule_match_value, + rule_ne_types, + rule_storage_key, + serialize_rule_ne_types, +) +from .models import ( + UmeAlarmCurrent, + UmeAlarmHistory, + UmeInventoryNE, + UmeKeyAlertForwardLog, + UmeKeyAlertRule, + UmeSyncJob, +) +from .oclaw_alarm_forwarder import ( + forwarder_status, + request_forwarder_reconnect, +) +from .ume_alarm_ws import ( + cancel_alarm_subscription_manual, + clear_local_alarm_subscription_manual, + establish_alarm_subscription_manual, + get_alarms_coordination_status, + get_subscription_status, + get_ws_connection_status, + get_ws_logs, + request_ws_reconnect, +) +from .ume_support import ( + UME_KNOWN_RUNTIME_TASKS, + _aggregate_rows, + _ensure_utc, + _list_runtime_tasks, + _request_force_sync_after_resume, + _runtime_pause_task, + _runtime_resume_task, + _ume_alarm_host_name, + _ume_alarm_ne_group_key, + _ume_client, + _ume_error_kind, + _clear_force_resume_hints, +) +from .ume_sync_service import sync_alarms_current, sync_alarms_history_full, sync_inventory_full +from .ume_token_store import clear_shared_token + +_log = logging.getLogger("netx.ume.router") +router = APIRouter(tags=["ume"]) + +@router.get("/v1/ume/token/status") +def ume_token_status() -> dict[str, Any]: + client = _ume_client() + st = client.token_status() + return {"ok": True, **st} + + +@router.post("/v1/ume/token/refresh") +def ume_token_refresh() -> dict[str, Any]: + client = _ume_client() + try: + before = client.token_status() + token = client.refresh_if_needed() + after = client.token_status() + return { + "ok": True, + "token": token, + "changed": bool(before.get("token_preview") != after.get("token_preview")), + **after, + } + except Exception as exc: + msg = str(exc)[:240] + return {"ok": False, "error_kind": _ume_error_kind(msg), "error": msg} + + +@router.post("/v1/ume/token/disconnect") +def ume_token_disconnect() -> dict[str, Any]: + client = _ume_client() + ok = bool(client.logout_token()) + st = client.token_status() + return {"ok": ok, **st} + + +@router.get("/v1/ume/alarm-subscription/status") +def ume_alarm_subscription_status(limit: int = 80) -> dict[str, Any]: + st = get_subscription_status() + ws_task = _UME_RUNTIME_TASKS.get("alarms_current_ws_consumer") or {} + log_limit = max(10, min(int(limit or 80), 100)) + return { + "ok": True, + **st, + **get_alarms_coordination_status(), + "ws_connection": get_ws_connection_status(), + "ws_consumer_status": str(ws_task.get("status") or ""), + "ws_consumer_last_error": str(ws_task.get("last_error") or ""), + "ws_consumer_last_run_at": ws_task.get("last_run_at"), + "ws_logs": get_ws_logs(limit=log_limit), + } + + +@router.post("/v1/ume/alarm-subscription/establish") +def ume_alarm_subscription_establish( + payload: dict[str, Any] | None = None, + db: Session = Depends(get_db), +) -> dict[str, Any]: + client = _ume_client() + body = payload or {} + force_reestablish = bool(body.get("force_reestablish")) + try: + st = establish_alarm_subscription_manual(client, db, force_reestablish=force_reestablish) + return {"ok": True, "created": not bool(st.get("already_exists")), **st} + except Exception as exc: + msg = str(exc)[:240] + raise HTTPException(status_code=502, detail=msg) from exc + + +@router.post("/v1/ume/alarm-subscription/cancel") +def ume_alarm_subscription_cancel( + payload: dict[str, Any] | None = None, + db: Session = Depends(get_db), +) -> dict[str, Any]: + client = _ume_client() + body = payload or {} + force_clear_local = bool(body.get("force_clear_local")) + try: + st = cancel_alarm_subscription_manual(client, db, force_clear_local=force_clear_local) + if st.get("needs_local_cleanup"): + return st + return {"ok": True, **st} + except Exception as exc: + msg = str(exc)[:240] + raise HTTPException(status_code=502, detail=msg) from exc + + +@router.post("/v1/ume/alarm-subscription/clear-local") +def ume_alarm_subscription_clear_local(db: Session = Depends(get_db)) -> dict[str, Any]: + try: + st = clear_local_alarm_subscription_manual(db) + return {"ok": True, "cleared_local": True, **st} + except Exception as exc: + msg = str(exc)[:240] + raise HTTPException(status_code=502, detail=msg) from exc + + diff --git a/tests/test_managed_ne.py b/tests/test_managed_ne.py index 8a219c9..568d06e 100644 --- a/tests/test_managed_ne.py +++ b/tests/test_managed_ne.py @@ -77,6 +77,7 @@ class ManagedNeApiTests(unittest.TestCase): ) ManagedNE.__table__.create(bind=self.engine, checkfirst=True) UmeInventoryNE.__table__.create(bind=self.engine, checkfirst=True) + Base.metadata.create_all(bind=self.engine) self.Session = sessionmaker(bind=self.engine, autoflush=False, autocommit=False) def override_get_db(): @@ -89,9 +90,17 @@ class ManagedNeApiTests(unittest.TestCase): app.dependency_overrides[get_db] = override_get_db self._session_patch = patch("netx_api.ne_connect.SessionLocal", self.Session) self._session_patch.start() + self._auth_patches = [ + patch("netx_api.auth_middleware.settings.auth_enabled", False), + patch("netx_api.auth_deps.settings.auth_enabled", False), + ] + for p in self._auth_patches: + p.start() self.client = TestClient(app) def tearDown(self): + for p in getattr(self, "_auth_patches", []): + p.stop() app.dependency_overrides.clear() self._session_patch.stop() settings.credential_secret_key = self._orig_key