From 96b76294b454973fcf9596699c25dfdc96951c5c Mon Sep 17 00:00:00 2001 From: oliver Date: Sun, 2 Aug 2026 17:12:55 +0800 Subject: [PATCH] Split UME sync and WebCRT session into focused modules. Keep public facades stable, re-export symbols tests patch, and fix missing threading import plus reaper fixture state. Co-authored-by: Cursor --- netx_api/ume_alarm_apply.py | 309 ++++++++ netx_api/ume_sync_common.py | 85 ++ netx_api/ume_sync_pull.py | 502 ++++++++++++ netx_api/ume_sync_service.py | 879 +-------------------- netx_api/webcrt_channel.py | 1 + netx_api/webcrt_service.py | 29 + netx_api/webcrt_session.py | 1133 +-------------------------- netx_api/webcrt_session_model.py | 487 ++++++++++++ netx_api/webcrt_session_registry.py | 637 +++++++++++++++ tests/test_webcrt.py | 47 +- 10 files changed, 2129 insertions(+), 1980 deletions(-) create mode 100644 netx_api/ume_alarm_apply.py create mode 100644 netx_api/ume_sync_common.py create mode 100644 netx_api/ume_sync_pull.py create mode 100644 netx_api/webcrt_session_model.py create mode 100644 netx_api/webcrt_session_registry.py diff --git a/netx_api/ume_alarm_apply.py b/netx_api/ume_alarm_apply.py new file mode 100644 index 0000000..55e15fa --- /dev/null +++ b/netx_api/ume_alarm_apply.py @@ -0,0 +1,309 @@ +"""UME alarm normalize / upsert / notification apply.""" +from __future__ import annotations + +import hashlib +import json +import re +import threading +import time +from datetime import datetime +from typing import Any + +from sqlalchemy.dialects.postgresql import insert as pg_insert +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session + +from .config import settings +from .models import UmeAlarmCurrent +from .ume_sync_common import _lookup_host_name, _pick, _s, _utc_now_naive + +def _alarm_key(alarm: dict[str, Any]) -> str: + key = _s(_pick(alarm, "alarmKey", "alarm-key","alarmkey","id")) + if key: + # Keep full upstream key; use a stable digest only for pathological ultra-long keys. + if len(key) > 512: + return "sha256:" + hashlib.sha256(key.encode("utf-8", errors="ignore")).hexdigest() + return key + parts = [ + _s(_pick(alarm, "objectName", "object-name")), + _s(_pick(alarm, "eventType", "event-type")), + _s(_pick(alarm, "timeCreated", "time-created")), + _s(_pick(alarm, "nativeProbableCause", "native-probable-cause")), + ] + merged = "|".join(x for x in parts if x) + raw = merged or f"fallback-{_utc_now_naive().timestamp()}" + if len(raw) > 512: + return "sha256:" + hashlib.sha256(raw.encode("utf-8", errors="ignore")).hexdigest() + return raw + + +def _derive_ne_id_from_alarm(alarm: dict[str, Any]) -> str: + ne_id = _s(_pick(alarm, "ne-id", "neId", "ne_id")) + if ne_id: + return ne_id + + def _uuid_like(s: str) -> bool: + return bool( + re.fullmatch( + r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}", + str(s or "").strip(), + ) + ) + + # Prefer UUID net_id in objectName, e.g. "... ME{33a3e8f4-a76e-40fd-a0ba-045371a5f234} ..." + object_name = _s(_pick(alarm, "objectName", "object-name")) + if object_name: + m_obj = re.search(r"ME\{([^}]+)\}", object_name, flags=re.IGNORECASE) + if m_obj: + candidate = _s(m_obj.group(1)) + if _uuid_like(candidate): + return candidate + + alarm_key = _s(_pick(alarm, "alarmKey", "alarm-key", "alarmkey")) + if not alarm_key: + return "" + + # If alarmkey contains ME{...}, only accept it when it looks like a UUID. + m0 = re.search(r"ME\{([^}]+)\}", alarm_key, flags=re.IGNORECASE) + if m0: + candidate = _s(m0.group(1)) + if _uuid_like(candidate): + return candidate + + # Common UME formats observed: + # 1) "#" + # 2) ", , " + # 3) " " + if "#" in alarm_key: + candidate = _s(alarm_key.split("#", 1)[0]) + return candidate if _uuid_like(candidate) else "" + if "," in alarm_key: + candidate = _s(alarm_key.split(",", 1)[0]) + return candidate if _uuid_like(candidate) else "" + parts = [p for p in alarm_key.split() if p] + if len(parts) >= 2 and ":" not in parts[0]: + candidate = _s(parts[0]) + return candidate if _uuid_like(candidate) else "" + return "" + + +def _normalize_yang_key(key: str) -> str: + raw = str(key or "").strip() + if not raw: + return "" + if ":" in raw: + return raw.rsplit(":", 1)[-1] + return raw + + +def normalize_yang_alarm(raw: dict[str, Any]) -> dict[str, Any]: + """Flatten YANG namespace-prefixed keys (e.g. zte-alarms:alarmkey) for _pick().""" + out: dict[str, Any] = {} + for k, v in raw.items(): + nk = _normalize_yang_key(str(k)) + if not nk: + continue + if nk in out and out[nk] not in (None, ""): + continue + out[nk] = v + return out + + +def _is_alarm_cleared(alarm: dict[str, Any]) -> bool: + val = _pick(alarm, "isCleared", "is-cleared") + if isinstance(val, bool): + return val + text = _s(val).lower() + return text in {"true", "1", "yes"} + + +_cleared_tombstone_lock = threading.Lock() +_cleared_tombstones: dict[str, float] = {} + + +def _mark_alarm_cleared_tombstone(alarm_key: str) -> None: + key = str(alarm_key or "").strip() + if not key: + return + ttl_s = max(60, int(getattr(settings, "ume_alarm_cleared_tombstone_s", 300) or 300)) + expires = time.time() + ttl_s + with _cleared_tombstone_lock: + _cleared_tombstones[key] = expires + if len(_cleared_tombstones) > 50000: + now = time.time() + stale = [k for k, exp in _cleared_tombstones.items() if exp <= now] + for k in stale[:10000]: + _cleared_tombstones.pop(k, None) + + +def _is_alarm_cleared_tombstone(alarm_key: str) -> bool: + key = str(alarm_key or "").strip() + if not key: + return False + now = time.time() + with _cleared_tombstone_lock: + exp = _cleared_tombstones.get(key) + if exp is None: + return False + if exp <= now: + _cleared_tombstones.pop(key, None) + return False + return True + + +def notification_id_from_norm(norm: dict[str, Any]) -> str: + return _s(_pick(norm, "notificationId", "notification-id")) + + +def _alarm_row_from_norm(key: str, norm: dict[str, Any], *, touch_ts: datetime, first_seen_at: datetime) -> dict[str, Any]: + return { + "alarm_key": key, + "ne_id": _s(_derive_ne_id_from_alarm(norm)), + "host_name": "", + "object_name": _s(_pick(norm, "objectName", "object-name")), + "event_type": _s(_pick(norm, "eventType", "event-type")), + "native_probable_cause": _s(_pick(norm, "nativeProbableCause", "native-probable-cause")), + "perceived_severity": _s(_pick(norm, "perceivedSeverity", "perceived-severity")), + "is_cleared": _s(_pick(norm, "isCleared", "is-cleared")), + "time_created": _s(_pick(norm, "timeCreated", "time-created")), + "root_cause_alarm_indication": _s( + _pick(norm, "rootCauseAlarmIndication", "root-cause-alarm-indication") + ), + "notification_id": notification_id_from_norm(norm), + "first_seen_at": first_seen_at, + "last_seen_at": touch_ts, + "raw_json": json.dumps(norm, ensure_ascii=False, default=str), + } + + +def _apply_row_to_model(db: Session, existing: UmeAlarmCurrent, norm: dict[str, Any], *, touch_ts: datetime) -> None: + existing.ne_id = _s(_derive_ne_id_from_alarm(norm)) + existing.object_name = _s(_pick(norm, "objectName", "object-name")) + existing.event_type = _s(_pick(norm, "eventType", "event-type")) + existing.native_probable_cause = _s(_pick(norm, "nativeProbableCause", "native-probable-cause")) + existing.perceived_severity = _s(_pick(norm, "perceivedSeverity", "perceived-severity")) + existing.is_cleared = _s(_pick(norm, "isCleared", "is-cleared")) + existing.time_created = _s(_pick(norm, "timeCreated", "time-created")) + existing.root_cause_alarm_indication = _s( + _pick(norm, "rootCauseAlarmIndication", "root-cause-alarm-indication") + ) + existing.notification_id = notification_id_from_norm(norm) + prev_seen = existing.last_seen_at + if prev_seen is None or touch_ts >= prev_seen: + existing.last_seen_at = touch_ts + existing.raw_json = json.dumps(norm, ensure_ascii=False, default=str) + + +def _upsert_alarm_current(db: Session, key: str, norm: dict[str, Any], *, touch_ts: datetime) -> tuple[str, bool]: + bind = db.get_bind() + dialect = str(getattr(getattr(bind, "dialect", None), "name", "") or "").lower() + existing = db.get(UmeAlarmCurrent, key) + if existing is not None: + _apply_row_to_model(db, existing, norm, touch_ts=touch_ts) + existing.host_name = _lookup_host_name(db, existing.ne_id) + return "updated", True + + if dialect == "postgresql": + row = _alarm_row_from_norm(key, norm, touch_ts=touch_ts, first_seen_at=touch_ts) + row["host_name"] = _lookup_host_name(db, row["ne_id"]) + ins = pg_insert(UmeAlarmCurrent).values(**row) + excluded = ins.excluded + stmt = ins.on_conflict_do_update( + index_elements=[UmeAlarmCurrent.alarm_key], + set_={ + "ne_id": excluded.ne_id, + "host_name": excluded.host_name, + "object_name": excluded.object_name, + "event_type": excluded.event_type, + "native_probable_cause": excluded.native_probable_cause, + "perceived_severity": excluded.perceived_severity, + "is_cleared": excluded.is_cleared, + "time_created": excluded.time_created, + "root_cause_alarm_indication": excluded.root_cause_alarm_indication, + "notification_id": excluded.notification_id, + "last_seen_at": func.greatest(UmeAlarmCurrent.last_seen_at, excluded.last_seen_at), + "raw_json": excluded.raw_json, + }, + ) + db.execute(stmt) + return "inserted", True + + try: + with db.begin_nested(): + model = UmeAlarmCurrent(alarm_key=key, first_seen_at=touch_ts) + db.add(model) + db.flush() + _apply_row_to_model(db, model, norm, touch_ts=touch_ts) + model.host_name = _lookup_host_name(db, model.ne_id) + return "inserted", True + except IntegrityError: + existing = db.get(UmeAlarmCurrent, key) + if existing is None: + return "skipped", False + _apply_row_to_model(db, existing, norm, touch_ts=touch_ts) + existing.host_name = _lookup_host_name(db, existing.ne_id) + return "updated", True + + +def extract_alarm_from_notification(payload: dict[str, Any]) -> dict[str, Any] | None: + """Parse alarm-notification from a WS/REST notification envelope.""" + if not isinstance(payload, dict): + return None + + def _find_alarm_notification(node: Any) -> dict[str, Any] | None: + if isinstance(node, dict): + for k, v in node.items(): + key = str(k).lower() + if key in {"alarm-notification", "alarm_notification"} and isinstance(v, dict): + return normalize_yang_alarm(v) + found = _find_alarm_notification(v) + if found is not None: + return found + elif isinstance(node, list): + for item in node: + found = _find_alarm_notification(item) + if found is not None: + return found + return None + + direct = _find_alarm_notification(payload) + if direct is not None: + return direct + return normalize_yang_alarm(payload) if payload else None + + +def apply_alarm_to_current( + db: Session, + alarm: dict[str, Any], + *, + touch_ts: datetime, + source: str = "", +) -> tuple[str, bool]: + """ + Apply one alarm to ume_alarms_current. + Returns (action, changed) where action is inserted|updated|deleted|skipped. + """ + norm = normalize_yang_alarm(alarm) if alarm else {} + if not norm: + return "skipped", False + + if _is_alarm_cleared(norm): + key = _alarm_key(norm) + if not key: + return "skipped", False + existing = db.get(UmeAlarmCurrent, key) + if existing is None: + _mark_alarm_cleared_tombstone(key) + return "deleted", False + db.delete(existing) + _mark_alarm_cleared_tombstone(key) + return "deleted", True + + key = _alarm_key(norm) + if not key: + return "skipped", False + if str(source or "").strip().lower() == "rest" and _is_alarm_cleared_tombstone(key): + return "skipped", False + + return _upsert_alarm_current(db, key, norm, touch_ts=touch_ts) + diff --git a/netx_api/ume_sync_common.py b/netx_api/ume_sync_common.py new file mode 100644 index 0000000..5a523b2 --- /dev/null +++ b/netx_api/ume_sync_common.py @@ -0,0 +1,85 @@ +"""UME sync shared helpers (string/pick/host-name).""" +from __future__ import annotations + +from datetime import datetime +from typing import Any + +from sqlalchemy.orm import Session + +from .models import UmeAlarmCurrent, UmeAlarmHistory, UmeInventoryNE +from .timeutil import utcnow_naive + + +def _utc_now_naive() -> datetime: + return utcnow_naive() + + +def _s(v: Any) -> str: + if v is None: + return "" + text = str(v).strip() + if text.lower() == "nan": + return "" + return text + + +def _pick(d: dict[str, Any], *keys: str) -> Any: + for key in keys: + if key in d: + return d.get(key) + return None + + +def _lookup_host_name(db: Session, ne_id: str) -> str: + nid = _s(ne_id) + if not nid: + return "" + row = db.get(UmeInventoryNE, nid) + if row is None: + return "" + return _s(row.host_name) + + +def _propagate_host_name_to_alarms(db: Session, ne_id: str, host_name: str) -> None: + nid = _s(ne_id) + if not nid: + return + hn = _s(host_name) + db.query(UmeAlarmCurrent).filter(UmeAlarmCurrent.ne_id == nid).update( + {UmeAlarmCurrent.host_name: hn}, + synchronize_session=False, + ) + db.query(UmeAlarmHistory).filter(UmeAlarmHistory.ne_id == nid).update( + {UmeAlarmHistory.host_name: hn}, + synchronize_session=False, + ) + + +def _backfill_alarm_host_names(db: Session, model: type[UmeAlarmCurrent] | type[UmeAlarmHistory]) -> int: + """Set alarm.host_name from ume_inventory_ne for all rows with matching ne_id.""" + table = str(getattr(model, "__tablename__", "") or "") + if not table: + return 0 + bind = db.get_bind() + if bind is not None and str(bind.dialect.name).lower() == "postgresql": + res = db.execute( + sql_text( + f""" + UPDATE {table} AS a + SET host_name = COALESCE(NULLIF(TRIM(ne.host_name), ''), '') + FROM ume_inventory_ne AS ne + WHERE a.ne_id <> '' AND a.ne_id = ne.ne_id + """ + ) + ) + return int(res.rowcount or 0) + ne_map = {str(r.ne_id or ""): _s(r.host_name) for r in db.query(UmeInventoryNE).all() if str(r.ne_id or "")} + updated = 0 + for alarm in db.query(model).filter(model.ne_id != "").all(): # type: ignore[arg-type] + hn = ne_map.get(str(alarm.ne_id or ""), "") + if str(alarm.host_name or "") != hn: + alarm.host_name = hn + updated += 1 + return updated + + diff --git a/netx_api/ume_sync_pull.py b/netx_api/ume_sync_pull.py new file mode 100644 index 0000000..236d38f --- /dev/null +++ b/netx_api/ume_sync_pull.py @@ -0,0 +1,502 @@ +"""UME inventory and alarm pull/sync jobs.""" +from __future__ import annotations + +import json +import logging +import threading +from datetime import datetime +from typing import Any + +from sqlalchemy import func +from sqlalchemy import text as sql_text +from sqlalchemy.orm import Session + +from .config import settings +from .models import ( + UmeAlarmBatch, + UmeAlarmCurrent, + UmeAlarmHistory, + UmeInventoryNE, + UmeSyncJob, +) +from .ume_alarm_apply import ( + _alarm_key, + _derive_ne_id_from_alarm, + _is_alarm_cleared, + _upsert_alarm_current, + apply_alarm_to_current, + normalize_yang_alarm, + notification_id_from_norm, +) +from .ume_client import UMEClient +from .ume_sync_common import ( + _backfill_alarm_host_names, + _lookup_host_name, + _pick, + _propagate_host_name_to_alarms, + _s, + _utc_now_naive, +) + +_sync_log = logging.getLogger("netx.ume.sync") +_ALARMS_CURRENT_SYNC_LOCK = threading.Lock() + +def _build_sync_job(domain: str, trigger_mode: str) -> UmeSyncJob: + return UmeSyncJob( + domain=domain, + status="running", + trigger_mode=trigger_mode, + started_at=_utc_now_naive(), + ) + + +def _snapshot_reconcile_ok(meta: dict[str, Any]) -> bool: + """True when paging finished normally (full snapshot); avoid deleting local rows on partial pulls.""" + if not bool(meta.get("is_end_of_reply")): + return False + if bool(meta.get("graceful_end_by_iterator_error")): + return False + warnings = meta.get("warnings") or [] + if not isinstance(warnings, list): + return False + if "duplicate_page_detected" in [str(w) for w in warnings]: + return False + if str(meta.get("paging_note") or "").strip() == "duplicate_page_detected": + return False + return True + + +def _collect_marker_pages( + fetch_page: Any, + *, + max_pages: int, + iterator_500_as_end: bool = False, +) -> tuple[list[list[dict[str, Any]]], dict[str, Any]]: + page_no = 0 + next_marker = "" + is_end_of_reply = False + graceful_end_by_iterator_error = False + paging_note = "" + warnings: list[str] = [] + last_page_signature = "" + pages: list[list[dict[str, Any]]] = [] + + while True: + page_no += 1 + if page_no > max_pages: + raise RuntimeError(f"ume_alarms_pagination_exceeded:max_pages={max_pages}") + + try: + rows, diag = fetch_page(next_marker or None) + except Exception as exc: + msg = str(exc or "") + low = msg.lower() + if iterator_500_as_end and pages and "ume_request_failed:500" in low and "iterator" in low and "null" in low: + graceful_end_by_iterator_error = True + paging_note = msg[:240] + break + raise + + rows = [x for x in rows if isinstance(x, dict)] + pages.append(rows) + + if page_no == 1 or page_no % 25 == 0: + _sync_log.info( + "ume marker page=%s rows=%s is_end=%s marker_len=%s", + page_no, + len(rows), + getattr(diag, "is_end_of_reply", None), + len(str(getattr(diag, "marker", "") or "")), + ) + + # Protection against repeated pages causing infinite loops. + cur_sig = "|".join(sorted(_alarm_key(x) for x in rows)) + if cur_sig and cur_sig == last_page_signature: + warnings.append("duplicate_page_detected") + paging_note = "duplicate_page_detected" + break + last_page_signature = cur_sig + + has_is_end = diag.is_end_of_reply is not None + is_end_of_reply = bool(diag.is_end_of_reply) if has_is_end else False + next_marker = str(diag.marker or "").strip() + + if has_is_end and is_end_of_reply: + break + if has_is_end and (not is_end_of_reply) and (not next_marker): + warnings.append("marker_missing_when_not_end") + paging_note = "marker_missing_when_not_end" + break + if (not has_is_end) and (not next_marker): + if rows: + warnings.append("marker_missing_stop") + paging_note = "marker_missing_stop" + break + + meta = { + "page_count": page_no, + "last_marker": next_marker, + "is_end_of_reply": is_end_of_reply, + "graceful_end_by_iterator_error": graceful_end_by_iterator_error, + "paging_note": paging_note, + "warnings": warnings, + } + return pages, meta + + +def sync_inventory_full(db: Session, client: UMEClient, *, trigger_mode: str = "manual") -> UmeSyncJob: + job = _build_sync_job("inventory", trigger_mode) + db.add(job) + db.flush() + db.commit() + _sync_log.info("inventory sync job %s committed as running (trigger=%s)", getattr(job, "id", "?"), trigger_mode) + pulled = inserted = updated = 0 + try: + limit_max = int(getattr(settings, "ume_limit_max", 5000) or 5000) + limit_max = max(1, limit_max) + page_size = int(getattr(settings, "ume_marker_page_limit", getattr(settings, "ume_page_size", 1000)) or 1000) + page_size = max(1, min(page_size, limit_max)) + max_pages = int(getattr(settings, "ume_marker_max_pages", getattr(settings, "ume_max_pages", 2000)) or 2000) + max_pages = max(1, min(max_pages, 20000)) + + pages, inv_meta = _collect_marker_pages( + lambda marker: client.get_network_elements(limit=page_size, marker=marker), + max_pages=max_pages, + iterator_500_as_end=False, + ) + ne_rows = [row for page in pages for row in page] + now = _utc_now_naive() + pulled = len(ne_rows) + seen_ne_ids: set[str] = set() + for row in ne_rows: + ne_id = _s(_pick(row, "ne-id", "ne_id", "id")) + if not ne_id: + continue + seen_ne_ids.add(ne_id) + existing = db.get(UmeInventoryNE, ne_id) + if existing is None: + existing = UmeInventoryNE( + ne_id=ne_id, + first_seen_at=now, + ) + db.add(existing) + inserted += 1 + else: + updated += 1 + existing.ne_name = _s(_pick(row, "name", "ne-name")) + existing.user_label = _s(_pick(row, "user-label", "user_label")) + existing.ip_address = _s(_pick(row, "ip-Address", "ip-address", "ip")) + existing.ipv6_address = _s(_pick(row, "ipv6-address", "ipv6_address")) + existing.ne_type = _s(_pick(row, "type", "ne-type")) + existing.device_level = _s(_pick(row, "device-level")) + existing.host_name = _s(_pick(row, "host-name")) + existing.location = _s(_pick(row, "location")) + existing.hardware_version = _s(_pick(row, "hardware-version")) + existing.loopback = _s(_pick(row, "loopback")) + existing.consistent_state = _s(_pick(row, "consistent-state")) + existing.interface_version = _s(_pick(row, "interface-version")) + existing.mac = _s(_pick(row, "mac")) + existing.admin_status = _s(_pick(row, "admin-status")) + existing.address_type = _s(_pick(row, "address-type")) + existing.connection_status = _s(_pick(row, "connection-status")) + existing.maintain_status = _s(_pick(row, "maintain-status")) + existing.net_mask = _s(_pick(row, "net-mask")) + existing.create_time = _s(_pick(row, "create-time")) + existing.creator = _s(_pick(row, "creator")) + existing.vendor = _s(_pick(row, "vendor-name")) or "ZTE" + existing.last_seen_at = now + existing.raw_json = json.dumps(row, ensure_ascii=False, default=str) + _propagate_host_name_to_alarms(db, ne_id, existing.host_name) + + db.flush() + deleted_ne = 0 + if _snapshot_reconcile_ok(inv_meta): + from .topology_inventory_lifecycle import detach_fabric_from_ume + + if seen_ne_ids: + stale_ids = [ + str(x[0]) + for x in db.query(UmeInventoryNE.ne_id) + .filter(~UmeInventoryNE.ne_id.in_(list(seen_ne_ids))) + .all() + if str(x[0] or "").strip() + ] + else: + stale_ids = [ + str(x[0]) + for x in db.query(UmeInventoryNE.ne_id).all() + if str(x[0] or "").strip() + ] + if stale_ids: + detach_fabric_from_ume(db, stale_ids) + deleted_ne = int( + db.query(UmeInventoryNE) + .filter(UmeInventoryNE.ne_id.in_(stale_ids)) + .delete(synchronize_session=False) + ) + + job.details_json = json.dumps( + { + "inventory_reconcile": _snapshot_reconcile_ok(inv_meta), + "deleted_inventory_ne": deleted_ne, + "paging": { + "is_end_of_reply": bool(inv_meta.get("is_end_of_reply")), + "warnings": list(inv_meta.get("warnings") or []), + }, + }, + ensure_ascii=False, + ) + + job.status = "done" + except Exception as exc: + job.status = "failed" + job.error_message = str(exc)[:1024] + finally: + job.pulled_count = int(pulled) + job.inserted_count = int(inserted) + job.updated_count = int(updated) + job.ended_at = _utc_now_naive() + db.commit() + db.refresh(job) + return job + + +def _reconcile_stale_current_alarms( + db: Session, + *, + sync_batch_ts: datetime, + seen_keys: set[str], + wss_active: bool, +) -> int: + """Remove local current alarms missing from REST snapshot. + + When WSS is active, skip deletes (WSS may have keys not yet in REST); manual sync only upserts. + """ + del seen_keys + if wss_active: + return 0 + return int( + db.query(UmeAlarmCurrent) + .filter(UmeAlarmCurrent.last_seen_at < sync_batch_ts) + .delete(synchronize_session=False) + ) + + +def _sync_alarms_common( + db: Session, + client: UMEClient, + *, + is_uncleared: bool, + trigger_mode: str, + wss_active: bool = False, +) -> tuple[UmeSyncJob, UmeAlarmBatch]: + domain = "alarms_history" if is_uncleared else "alarms_current" + job = _build_sync_job(domain, trigger_mode) + db.add(job) + batch = UmeAlarmBatch( + kind="history" if is_uncleared else "current", + status="running", + started_at=_utc_now_naive(), + ) + db.add(batch) + db.flush() + db.commit() + _sync_log.info( + "alarms sync job domain=%s id=%s committed as running (trigger=%s)", + domain, + getattr(job, "id", "?"), + trigger_mode, + ) + pulled = inserted = updated = 0 + deleted_stale_current = 0 + host_names_backfilled = 0 + reconcile_mode = "" + seen_keys: set[str] = set() + paging_mode = "marker" + paging_note = "" + page_no = 0 + next_marker = "" + is_end_of_reply = False + graceful_end_by_iterator_error = False + warnings: list[str] = [] + meta: dict[str, Any] = {} + sync_batch_ts = _utc_now_naive() + try: + limit_max = int(getattr(settings, "ume_limit_max", 5000) or 5000) + limit_max = max(1, limit_max) + page_size = int(getattr(settings, "ume_marker_page_limit", getattr(settings, "ume_page_size", 1000)) or 1000) + page_size = max(1, min(page_size, limit_max)) + max_pages = int(getattr(settings, "ume_marker_max_pages", getattr(settings, "ume_max_pages", 2000)) or 2000) + max_pages = max(1, min(max_pages, 20000)) + + def upsert_alarm_history(alarm: dict[str, Any], *, touch_ts: datetime) -> None: + nonlocal inserted, updated + key = _alarm_key(alarm) + existing = db.get(UmeAlarmHistory, key) + if existing is None: + existing = UmeAlarmHistory(alarm_key=key, first_seen_at=touch_ts) + db.add(existing) + inserted += 1 + else: + updated += 1 + existing.ne_id = _s(_derive_ne_id_from_alarm(alarm)) + existing.host_name = _lookup_host_name(db, existing.ne_id) + existing.object_name = _s(_pick(alarm, "objectName", "object-name")) + existing.event_type = _s(_pick(alarm, "eventType", "event-type")) + existing.native_probable_cause = _s(_pick(alarm, "nativeProbableCause", "native-probable-cause")) + existing.perceived_severity = _s(_pick(alarm, "perceivedSeverity", "perceived-severity")) + existing.is_cleared = _s(_pick(alarm, "isCleared", "is-cleared")) + existing.time_created = _s(_pick(alarm, "timeCreated", "time-created")) + existing.root_cause_alarm_indication = _s( + _pick(alarm, "rootCauseAlarmIndication", "root-cause-alarm-indication") + ) + existing.notification_id = notification_id_from_norm(alarm) + existing.last_seen_at = touch_ts + existing.raw_json = json.dumps(alarm, ensure_ascii=False, default=str) + + iterator_500_as_end = bool(getattr(settings, "ume_iterator_500_as_end", True)) + pages, meta = _collect_marker_pages( + lambda marker: client.get_alarms(is_uncleared=is_uncleared, limit=page_size, marker=marker), + max_pages=max_pages, + iterator_500_as_end=iterator_500_as_end, + ) + sync_batch_ts = _utc_now_naive() + seen_keys = set() + for rows in pages: + pulled += len(rows) + for alarm in rows: + if is_uncleared: + upsert_alarm_history(alarm, touch_ts=sync_batch_ts) + else: + key = _alarm_key(alarm) + if key: + seen_keys.add(key) + action, _changed = apply_alarm_to_current( + db, + alarm, + touch_ts=sync_batch_ts, + source="rest", + ) + if action == "inserted": + inserted += 1 + elif action == "updated": + updated += 1 + db.flush() + page_no = int(meta.get("page_count") or 0) + next_marker = str(meta.get("last_marker") or "") + is_end_of_reply = bool(meta.get("is_end_of_reply")) + graceful_end_by_iterator_error = bool(meta.get("graceful_end_by_iterator_error")) + paging_note = str(meta.get("paging_note") or "") + warnings = [str(x) for x in (meta.get("warnings") or []) if str(x)] + + reconcile_mode = "full" + if not is_uncleared and _snapshot_reconcile_ok(meta): + if wss_active: + reconcile_mode = "upsert_only" + deleted_stale_current = _reconcile_stale_current_alarms( + db, + sync_batch_ts=sync_batch_ts, + seen_keys=seen_keys, + wss_active=wss_active, + ) + + alarm_model = UmeAlarmHistory if is_uncleared else UmeAlarmCurrent + host_names_backfilled = _backfill_alarm_host_names(db, alarm_model) + + batch.total_rows = int(pulled) + batch.success_rows = int(inserted + updated) + batch.failed_rows = max(0, int(pulled) - int(inserted + updated)) + batch.status = "done" + batch.ended_at = _utc_now_naive() + batch.raw_json = json.dumps( + { + "pulled": pulled, + "inserted": inserted, + "updated": updated, + "paging_mode": paging_mode, + "page_count": page_no, + "last_marker": next_marker, + "is_end_of_reply": is_end_of_reply, + "graceful_end_by_iterator_error": graceful_end_by_iterator_error, + "warnings": warnings, + "deleted_stale_current_alarms": int(deleted_stale_current), + "host_names_backfilled": int(host_names_backfilled), + "current_snapshot_reconcile": (not is_uncleared) and _snapshot_reconcile_ok(meta), + "reconcile_mode": reconcile_mode if not is_uncleared else "", + "wss_active_during_sync": bool(wss_active) if not is_uncleared else False, + "seen_keys_count": len(seen_keys) if not is_uncleared else 0, + }, + ensure_ascii=False, + ) + + job.status = "done" + except Exception as exc: + msg = str(exc)[:1024] + reconcile_mode = "failed" + batch.status = "failed" + batch.error_message = msg + batch.ended_at = _utc_now_naive() + job.status = "failed" + job.error_message = msg + finally: + job.pulled_count = int(pulled) + job.inserted_count = int(inserted) + job.updated_count = int(updated) + job.ended_at = _utc_now_naive() + job.details_json = json.dumps( + { + "batch_id": batch.batch_id, + "kind": batch.kind, + "status": batch.status, + "paging_mode": paging_mode, + "paging_note": paging_note, + "page_count": page_no, + "last_marker": next_marker, + "is_end_of_reply": is_end_of_reply, + "graceful_end_by_iterator_error": graceful_end_by_iterator_error, + "warnings": warnings, + "deleted_stale_current_alarms": int(deleted_stale_current), + "host_names_backfilled": int(host_names_backfilled), + "current_snapshot_reconcile": (not is_uncleared) and _snapshot_reconcile_ok(meta), + "reconcile_mode": reconcile_mode if not is_uncleared else "", + "wss_active_during_sync": bool(wss_active) if not is_uncleared else False, + "seen_keys_count": len(seen_keys) if not is_uncleared else 0, + }, + ensure_ascii=False, + ) + db.commit() + db.refresh(job) + db.refresh(batch) + return job, batch + + +def sync_alarms_current( + db: Session, + client: UMEClient, + *, + trigger_mode: str = "manual", + wss_active: bool | None = None, +) -> tuple[UmeSyncJob, UmeAlarmBatch]: + if not _ALARMS_CURRENT_SYNC_LOCK.acquire(blocking=False): + _sync_log.warning("alarms_current sync skipped: another sync is in progress") + raise RuntimeError("alarms_current_sync_busy") + try: + if wss_active is None: + from .ume_alarm_ws import is_wss_active_for_current_alarms + + wss_active = is_wss_active_for_current_alarms() + return _sync_alarms_common( + db, + client, + is_uncleared=False, + trigger_mode=trigger_mode, + wss_active=bool(wss_active), + ) + finally: + _ALARMS_CURRENT_SYNC_LOCK.release() + + +def sync_alarms_history_full( + db: Session, client: UMEClient, *, trigger_mode: str = "manual" +) -> tuple[UmeSyncJob, UmeAlarmBatch]: + return _sync_alarms_common(db, client, is_uncleared=True, trigger_mode=trigger_mode) diff --git a/netx_api/ume_sync_service.py b/netx_api/ume_sync_service.py index a3a8222..326c9e1 100644 --- a/netx_api/ume_sync_service.py +++ b/netx_api/ume_sync_service.py @@ -1,855 +1,32 @@ +"""UME sync service facade (inventory + alarms).""" from __future__ import annotations -import json -import logging -import threading -import hashlib -import re -import threading -import time -from datetime import datetime, timezone -from typing import Any - -_sync_log = logging.getLogger("netx.ume.sync") -_ALARMS_CURRENT_SYNC_LOCK = threading.Lock() - -from sqlalchemy import func -from sqlalchemy import text as sql_text -from sqlalchemy.dialects.postgresql import insert as pg_insert -from sqlalchemy.exc import IntegrityError -from sqlalchemy.orm import Session - from .config import settings -from .models import ( - UmeAlarmBatch, - UmeAlarmCurrent, - UmeAlarmHistory, - UmeInventoryNE, - UmeSyncJob, +from .ume_alarm_apply import ( + _alarm_key, + _derive_ne_id_from_alarm, + _is_alarm_cleared, + apply_alarm_to_current, + extract_alarm_from_notification, + normalize_yang_alarm, + notification_id_from_norm, ) -from .ume_client import UMEClient - - -def _s(v: Any) -> str: - if v is None: - return "" - text = str(v).strip() - if text.lower() == "nan": - return "" - return text - - -def _utc_now_naive() -> datetime: - return datetime.now(timezone.utc).replace(tzinfo=None) - - -def _pick(d: dict[str, Any], *keys: str) -> Any: - for key in keys: - if key in d: - return d.get(key) - return None - - -def _lookup_host_name(db: Session, ne_id: str) -> str: - nid = _s(ne_id) - if not nid: - return "" - row = db.get(UmeInventoryNE, nid) - if row is None: - return "" - return _s(row.host_name) - - -def _propagate_host_name_to_alarms(db: Session, ne_id: str, host_name: str) -> None: - nid = _s(ne_id) - if not nid: - return - hn = _s(host_name) - db.query(UmeAlarmCurrent).filter(UmeAlarmCurrent.ne_id == nid).update( - {UmeAlarmCurrent.host_name: hn}, - synchronize_session=False, - ) - db.query(UmeAlarmHistory).filter(UmeAlarmHistory.ne_id == nid).update( - {UmeAlarmHistory.host_name: hn}, - synchronize_session=False, - ) - - -def _backfill_alarm_host_names(db: Session, model: type[UmeAlarmCurrent] | type[UmeAlarmHistory]) -> int: - """Set alarm.host_name from ume_inventory_ne for all rows with matching ne_id.""" - table = str(getattr(model, "__tablename__", "") or "") - if not table: - return 0 - bind = db.get_bind() - if bind is not None and str(bind.dialect.name).lower() == "postgresql": - res = db.execute( - sql_text( - f""" - UPDATE {table} AS a - SET host_name = COALESCE(NULLIF(TRIM(ne.host_name), ''), '') - FROM ume_inventory_ne AS ne - WHERE a.ne_id <> '' AND a.ne_id = ne.ne_id - """ - ) - ) - return int(res.rowcount or 0) - ne_map = {str(r.ne_id or ""): _s(r.host_name) for r in db.query(UmeInventoryNE).all() if str(r.ne_id or "")} - updated = 0 - for alarm in db.query(model).filter(model.ne_id != "").all(): # type: ignore[arg-type] - hn = ne_map.get(str(alarm.ne_id or ""), "") - if str(alarm.host_name or "") != hn: - alarm.host_name = hn - updated += 1 - return updated - - -def _alarm_key(alarm: dict[str, Any]) -> str: - key = _s(_pick(alarm, "alarmKey", "alarm-key","alarmkey","id")) - if key: - # Keep full upstream key; use a stable digest only for pathological ultra-long keys. - if len(key) > 512: - return "sha256:" + hashlib.sha256(key.encode("utf-8", errors="ignore")).hexdigest() - return key - parts = [ - _s(_pick(alarm, "objectName", "object-name")), - _s(_pick(alarm, "eventType", "event-type")), - _s(_pick(alarm, "timeCreated", "time-created")), - _s(_pick(alarm, "nativeProbableCause", "native-probable-cause")), - ] - merged = "|".join(x for x in parts if x) - raw = merged or f"fallback-{datetime.utcnow().timestamp()}" - if len(raw) > 512: - return "sha256:" + hashlib.sha256(raw.encode("utf-8", errors="ignore")).hexdigest() - return raw - - -def _derive_ne_id_from_alarm(alarm: dict[str, Any]) -> str: - ne_id = _s(_pick(alarm, "ne-id", "neId", "ne_id")) - if ne_id: - return ne_id - - def _uuid_like(s: str) -> bool: - return bool( - re.fullmatch( - r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}", - str(s or "").strip(), - ) - ) - - # Prefer UUID net_id in objectName, e.g. "... ME{33a3e8f4-a76e-40fd-a0ba-045371a5f234} ..." - object_name = _s(_pick(alarm, "objectName", "object-name")) - if object_name: - m_obj = re.search(r"ME\{([^}]+)\}", object_name, flags=re.IGNORECASE) - if m_obj: - candidate = _s(m_obj.group(1)) - if _uuid_like(candidate): - return candidate - - alarm_key = _s(_pick(alarm, "alarmKey", "alarm-key", "alarmkey")) - if not alarm_key: - return "" - - # If alarmkey contains ME{...}, only accept it when it looks like a UUID. - m0 = re.search(r"ME\{([^}]+)\}", alarm_key, flags=re.IGNORECASE) - if m0: - candidate = _s(m0.group(1)) - if _uuid_like(candidate): - return candidate - - # Common UME formats observed: - # 1) "#" - # 2) ", , " - # 3) " " - if "#" in alarm_key: - candidate = _s(alarm_key.split("#", 1)[0]) - return candidate if _uuid_like(candidate) else "" - if "," in alarm_key: - candidate = _s(alarm_key.split(",", 1)[0]) - return candidate if _uuid_like(candidate) else "" - parts = [p for p in alarm_key.split() if p] - if len(parts) >= 2 and ":" not in parts[0]: - candidate = _s(parts[0]) - return candidate if _uuid_like(candidate) else "" - return "" - - -def _normalize_yang_key(key: str) -> str: - raw = str(key or "").strip() - if not raw: - return "" - if ":" in raw: - return raw.rsplit(":", 1)[-1] - return raw - - -def normalize_yang_alarm(raw: dict[str, Any]) -> dict[str, Any]: - """Flatten YANG namespace-prefixed keys (e.g. zte-alarms:alarmkey) for _pick().""" - out: dict[str, Any] = {} - for k, v in raw.items(): - nk = _normalize_yang_key(str(k)) - if not nk: - continue - if nk in out and out[nk] not in (None, ""): - continue - out[nk] = v - return out - - -def _is_alarm_cleared(alarm: dict[str, Any]) -> bool: - val = _pick(alarm, "isCleared", "is-cleared") - if isinstance(val, bool): - return val - text = _s(val).lower() - return text in {"true", "1", "yes"} - - -_cleared_tombstone_lock = threading.Lock() -_cleared_tombstones: dict[str, float] = {} - - -def _mark_alarm_cleared_tombstone(alarm_key: str) -> None: - key = str(alarm_key or "").strip() - if not key: - return - ttl_s = max(60, int(getattr(settings, "ume_alarm_cleared_tombstone_s", 300) or 300)) - expires = time.time() + ttl_s - with _cleared_tombstone_lock: - _cleared_tombstones[key] = expires - if len(_cleared_tombstones) > 50000: - now = time.time() - stale = [k for k, exp in _cleared_tombstones.items() if exp <= now] - for k in stale[:10000]: - _cleared_tombstones.pop(k, None) - - -def _is_alarm_cleared_tombstone(alarm_key: str) -> bool: - key = str(alarm_key or "").strip() - if not key: - return False - now = time.time() - with _cleared_tombstone_lock: - exp = _cleared_tombstones.get(key) - if exp is None: - return False - if exp <= now: - _cleared_tombstones.pop(key, None) - return False - return True - - -def notification_id_from_norm(norm: dict[str, Any]) -> str: - return _s(_pick(norm, "notificationId", "notification-id")) - - -def _alarm_row_from_norm(key: str, norm: dict[str, Any], *, touch_ts: datetime, first_seen_at: datetime) -> dict[str, Any]: - return { - "alarm_key": key, - "ne_id": _s(_derive_ne_id_from_alarm(norm)), - "host_name": "", - "object_name": _s(_pick(norm, "objectName", "object-name")), - "event_type": _s(_pick(norm, "eventType", "event-type")), - "native_probable_cause": _s(_pick(norm, "nativeProbableCause", "native-probable-cause")), - "perceived_severity": _s(_pick(norm, "perceivedSeverity", "perceived-severity")), - "is_cleared": _s(_pick(norm, "isCleared", "is-cleared")), - "time_created": _s(_pick(norm, "timeCreated", "time-created")), - "root_cause_alarm_indication": _s( - _pick(norm, "rootCauseAlarmIndication", "root-cause-alarm-indication") - ), - "notification_id": notification_id_from_norm(norm), - "first_seen_at": first_seen_at, - "last_seen_at": touch_ts, - "raw_json": json.dumps(norm, ensure_ascii=False, default=str), - } - - -def _apply_row_to_model(db: Session, existing: UmeAlarmCurrent, norm: dict[str, Any], *, touch_ts: datetime) -> None: - existing.ne_id = _s(_derive_ne_id_from_alarm(norm)) - existing.object_name = _s(_pick(norm, "objectName", "object-name")) - existing.event_type = _s(_pick(norm, "eventType", "event-type")) - existing.native_probable_cause = _s(_pick(norm, "nativeProbableCause", "native-probable-cause")) - existing.perceived_severity = _s(_pick(norm, "perceivedSeverity", "perceived-severity")) - existing.is_cleared = _s(_pick(norm, "isCleared", "is-cleared")) - existing.time_created = _s(_pick(norm, "timeCreated", "time-created")) - existing.root_cause_alarm_indication = _s( - _pick(norm, "rootCauseAlarmIndication", "root-cause-alarm-indication") - ) - existing.notification_id = notification_id_from_norm(norm) - prev_seen = existing.last_seen_at - if prev_seen is None or touch_ts >= prev_seen: - existing.last_seen_at = touch_ts - existing.raw_json = json.dumps(norm, ensure_ascii=False, default=str) - - -def _upsert_alarm_current(db: Session, key: str, norm: dict[str, Any], *, touch_ts: datetime) -> tuple[str, bool]: - bind = db.get_bind() - dialect = str(getattr(getattr(bind, "dialect", None), "name", "") or "").lower() - existing = db.get(UmeAlarmCurrent, key) - if existing is not None: - _apply_row_to_model(db, existing, norm, touch_ts=touch_ts) - existing.host_name = _lookup_host_name(db, existing.ne_id) - return "updated", True - - if dialect == "postgresql": - row = _alarm_row_from_norm(key, norm, touch_ts=touch_ts, first_seen_at=touch_ts) - row["host_name"] = _lookup_host_name(db, row["ne_id"]) - ins = pg_insert(UmeAlarmCurrent).values(**row) - excluded = ins.excluded - stmt = ins.on_conflict_do_update( - index_elements=[UmeAlarmCurrent.alarm_key], - set_={ - "ne_id": excluded.ne_id, - "host_name": excluded.host_name, - "object_name": excluded.object_name, - "event_type": excluded.event_type, - "native_probable_cause": excluded.native_probable_cause, - "perceived_severity": excluded.perceived_severity, - "is_cleared": excluded.is_cleared, - "time_created": excluded.time_created, - "root_cause_alarm_indication": excluded.root_cause_alarm_indication, - "notification_id": excluded.notification_id, - "last_seen_at": func.greatest(UmeAlarmCurrent.last_seen_at, excluded.last_seen_at), - "raw_json": excluded.raw_json, - }, - ) - db.execute(stmt) - return "inserted", True - - try: - with db.begin_nested(): - model = UmeAlarmCurrent(alarm_key=key, first_seen_at=touch_ts) - db.add(model) - db.flush() - _apply_row_to_model(db, model, norm, touch_ts=touch_ts) - model.host_name = _lookup_host_name(db, model.ne_id) - return "inserted", True - except IntegrityError: - existing = db.get(UmeAlarmCurrent, key) - if existing is None: - return "skipped", False - _apply_row_to_model(db, existing, norm, touch_ts=touch_ts) - existing.host_name = _lookup_host_name(db, existing.ne_id) - return "updated", True - - -def extract_alarm_from_notification(payload: dict[str, Any]) -> dict[str, Any] | None: - """Parse alarm-notification from a WS/REST notification envelope.""" - if not isinstance(payload, dict): - return None - - def _find_alarm_notification(node: Any) -> dict[str, Any] | None: - if isinstance(node, dict): - for k, v in node.items(): - key = str(k).lower() - if key in {"alarm-notification", "alarm_notification"} and isinstance(v, dict): - return normalize_yang_alarm(v) - found = _find_alarm_notification(v) - if found is not None: - return found - elif isinstance(node, list): - for item in node: - found = _find_alarm_notification(item) - if found is not None: - return found - return None - - direct = _find_alarm_notification(payload) - if direct is not None: - return direct - return normalize_yang_alarm(payload) if payload else None - - -def apply_alarm_to_current( - db: Session, - alarm: dict[str, Any], - *, - touch_ts: datetime, - source: str = "", -) -> tuple[str, bool]: - """ - Apply one alarm to ume_alarms_current. - Returns (action, changed) where action is inserted|updated|deleted|skipped. - """ - norm = normalize_yang_alarm(alarm) if alarm else {} - if not norm: - return "skipped", False - - if _is_alarm_cleared(norm): - key = _alarm_key(norm) - if not key: - return "skipped", False - existing = db.get(UmeAlarmCurrent, key) - if existing is None: - _mark_alarm_cleared_tombstone(key) - return "deleted", False - db.delete(existing) - _mark_alarm_cleared_tombstone(key) - return "deleted", True - - key = _alarm_key(norm) - if not key: - return "skipped", False - if str(source or "").strip().lower() == "rest" and _is_alarm_cleared_tombstone(key): - return "skipped", False - - return _upsert_alarm_current(db, key, norm, touch_ts=touch_ts) - - -def _build_sync_job(domain: str, trigger_mode: str) -> UmeSyncJob: - return UmeSyncJob( - domain=domain, - status="running", - trigger_mode=trigger_mode, - started_at=_utc_now_naive(), - ) - - -def _snapshot_reconcile_ok(meta: dict[str, Any]) -> bool: - """True when paging finished normally (full snapshot); avoid deleting local rows on partial pulls.""" - if not bool(meta.get("is_end_of_reply")): - return False - if bool(meta.get("graceful_end_by_iterator_error")): - return False - warnings = meta.get("warnings") or [] - if not isinstance(warnings, list): - return False - if "duplicate_page_detected" in [str(w) for w in warnings]: - return False - if str(meta.get("paging_note") or "").strip() == "duplicate_page_detected": - return False - return True - - -def _collect_marker_pages( - fetch_page: Any, - *, - max_pages: int, - iterator_500_as_end: bool = False, -) -> tuple[list[list[dict[str, Any]]], dict[str, Any]]: - page_no = 0 - next_marker = "" - is_end_of_reply = False - graceful_end_by_iterator_error = False - paging_note = "" - warnings: list[str] = [] - last_page_signature = "" - pages: list[list[dict[str, Any]]] = [] - - while True: - page_no += 1 - if page_no > max_pages: - raise RuntimeError(f"ume_alarms_pagination_exceeded:max_pages={max_pages}") - - try: - rows, diag = fetch_page(next_marker or None) - except Exception as exc: - msg = str(exc or "") - low = msg.lower() - if iterator_500_as_end and pages and "ume_request_failed:500" in low and "iterator" in low and "null" in low: - graceful_end_by_iterator_error = True - paging_note = msg[:240] - break - raise - - rows = [x for x in rows if isinstance(x, dict)] - pages.append(rows) - - if page_no == 1 or page_no % 25 == 0: - _sync_log.info( - "ume marker page=%s rows=%s is_end=%s marker_len=%s", - page_no, - len(rows), - getattr(diag, "is_end_of_reply", None), - len(str(getattr(diag, "marker", "") or "")), - ) - - # Protection against repeated pages causing infinite loops. - cur_sig = "|".join(sorted(_alarm_key(x) for x in rows)) - if cur_sig and cur_sig == last_page_signature: - warnings.append("duplicate_page_detected") - paging_note = "duplicate_page_detected" - break - last_page_signature = cur_sig - - has_is_end = diag.is_end_of_reply is not None - is_end_of_reply = bool(diag.is_end_of_reply) if has_is_end else False - next_marker = str(diag.marker or "").strip() - - if has_is_end and is_end_of_reply: - break - if has_is_end and (not is_end_of_reply) and (not next_marker): - warnings.append("marker_missing_when_not_end") - paging_note = "marker_missing_when_not_end" - break - if (not has_is_end) and (not next_marker): - if rows: - warnings.append("marker_missing_stop") - paging_note = "marker_missing_stop" - break - - meta = { - "page_count": page_no, - "last_marker": next_marker, - "is_end_of_reply": is_end_of_reply, - "graceful_end_by_iterator_error": graceful_end_by_iterator_error, - "paging_note": paging_note, - "warnings": warnings, - } - return pages, meta - - -def sync_inventory_full(db: Session, client: UMEClient, *, trigger_mode: str = "manual") -> UmeSyncJob: - job = _build_sync_job("inventory", trigger_mode) - db.add(job) - db.flush() - db.commit() - _sync_log.info("inventory sync job %s committed as running (trigger=%s)", getattr(job, "id", "?"), trigger_mode) - pulled = inserted = updated = 0 - try: - limit_max = int(getattr(settings, "ume_limit_max", 5000) or 5000) - limit_max = max(1, limit_max) - page_size = int(getattr(settings, "ume_marker_page_limit", getattr(settings, "ume_page_size", 1000)) or 1000) - page_size = max(1, min(page_size, limit_max)) - max_pages = int(getattr(settings, "ume_marker_max_pages", getattr(settings, "ume_max_pages", 2000)) or 2000) - max_pages = max(1, min(max_pages, 20000)) - - pages, inv_meta = _collect_marker_pages( - lambda marker: client.get_network_elements(limit=page_size, marker=marker), - max_pages=max_pages, - iterator_500_as_end=False, - ) - ne_rows = [row for page in pages for row in page] - now = _utc_now_naive() - pulled = len(ne_rows) - seen_ne_ids: set[str] = set() - for row in ne_rows: - ne_id = _s(_pick(row, "ne-id", "ne_id", "id")) - if not ne_id: - continue - seen_ne_ids.add(ne_id) - existing = db.get(UmeInventoryNE, ne_id) - if existing is None: - existing = UmeInventoryNE( - ne_id=ne_id, - first_seen_at=now, - ) - db.add(existing) - inserted += 1 - else: - updated += 1 - existing.ne_name = _s(_pick(row, "name", "ne-name")) - existing.user_label = _s(_pick(row, "user-label", "user_label")) - existing.ip_address = _s(_pick(row, "ip-Address", "ip-address", "ip")) - existing.ipv6_address = _s(_pick(row, "ipv6-address", "ipv6_address")) - existing.ne_type = _s(_pick(row, "type", "ne-type")) - existing.device_level = _s(_pick(row, "device-level")) - existing.host_name = _s(_pick(row, "host-name")) - existing.location = _s(_pick(row, "location")) - existing.hardware_version = _s(_pick(row, "hardware-version")) - existing.loopback = _s(_pick(row, "loopback")) - existing.consistent_state = _s(_pick(row, "consistent-state")) - existing.interface_version = _s(_pick(row, "interface-version")) - existing.mac = _s(_pick(row, "mac")) - existing.admin_status = _s(_pick(row, "admin-status")) - existing.address_type = _s(_pick(row, "address-type")) - existing.connection_status = _s(_pick(row, "connection-status")) - existing.maintain_status = _s(_pick(row, "maintain-status")) - existing.net_mask = _s(_pick(row, "net-mask")) - existing.create_time = _s(_pick(row, "create-time")) - existing.creator = _s(_pick(row, "creator")) - existing.vendor = _s(_pick(row, "vendor-name")) or "ZTE" - existing.last_seen_at = now - existing.raw_json = json.dumps(row, ensure_ascii=False, default=str) - _propagate_host_name_to_alarms(db, ne_id, existing.host_name) - - db.flush() - deleted_ne = 0 - if _snapshot_reconcile_ok(inv_meta): - from .topology_inventory_lifecycle import detach_fabric_from_ume - - if seen_ne_ids: - stale_ids = [ - str(x[0]) - for x in db.query(UmeInventoryNE.ne_id) - .filter(~UmeInventoryNE.ne_id.in_(list(seen_ne_ids))) - .all() - if str(x[0] or "").strip() - ] - else: - stale_ids = [ - str(x[0]) - for x in db.query(UmeInventoryNE.ne_id).all() - if str(x[0] or "").strip() - ] - if stale_ids: - detach_fabric_from_ume(db, stale_ids) - deleted_ne = int( - db.query(UmeInventoryNE) - .filter(UmeInventoryNE.ne_id.in_(stale_ids)) - .delete(synchronize_session=False) - ) - - job.details_json = json.dumps( - { - "inventory_reconcile": _snapshot_reconcile_ok(inv_meta), - "deleted_inventory_ne": deleted_ne, - "paging": { - "is_end_of_reply": bool(inv_meta.get("is_end_of_reply")), - "warnings": list(inv_meta.get("warnings") or []), - }, - }, - ensure_ascii=False, - ) - - job.status = "done" - except Exception as exc: - job.status = "failed" - job.error_message = str(exc)[:1024] - finally: - job.pulled_count = int(pulled) - job.inserted_count = int(inserted) - job.updated_count = int(updated) - job.ended_at = _utc_now_naive() - db.commit() - db.refresh(job) - return job - - -def _reconcile_stale_current_alarms( - db: Session, - *, - sync_batch_ts: datetime, - seen_keys: set[str], - wss_active: bool, -) -> int: - """Remove local current alarms missing from REST snapshot. - - When WSS is active, skip deletes (WSS may have keys not yet in REST); manual sync only upserts. - """ - del seen_keys - if wss_active: - return 0 - return int( - db.query(UmeAlarmCurrent) - .filter(UmeAlarmCurrent.last_seen_at < sync_batch_ts) - .delete(synchronize_session=False) - ) - - -def _sync_alarms_common( - db: Session, - client: UMEClient, - *, - is_uncleared: bool, - trigger_mode: str, - wss_active: bool = False, -) -> tuple[UmeSyncJob, UmeAlarmBatch]: - domain = "alarms_history" if is_uncleared else "alarms_current" - job = _build_sync_job(domain, trigger_mode) - db.add(job) - batch = UmeAlarmBatch( - kind="history" if is_uncleared else "current", - status="running", - started_at=_utc_now_naive(), - ) - db.add(batch) - db.flush() - db.commit() - _sync_log.info( - "alarms sync job domain=%s id=%s committed as running (trigger=%s)", - domain, - getattr(job, "id", "?"), - trigger_mode, - ) - pulled = inserted = updated = 0 - deleted_stale_current = 0 - host_names_backfilled = 0 - reconcile_mode = "" - seen_keys: set[str] = set() - paging_mode = "marker" - paging_note = "" - page_no = 0 - next_marker = "" - is_end_of_reply = False - graceful_end_by_iterator_error = False - warnings: list[str] = [] - meta: dict[str, Any] = {} - sync_batch_ts = _utc_now_naive() - try: - limit_max = int(getattr(settings, "ume_limit_max", 5000) or 5000) - limit_max = max(1, limit_max) - page_size = int(getattr(settings, "ume_marker_page_limit", getattr(settings, "ume_page_size", 1000)) or 1000) - page_size = max(1, min(page_size, limit_max)) - max_pages = int(getattr(settings, "ume_marker_max_pages", getattr(settings, "ume_max_pages", 2000)) or 2000) - max_pages = max(1, min(max_pages, 20000)) - - def upsert_alarm_history(alarm: dict[str, Any], *, touch_ts: datetime) -> None: - nonlocal inserted, updated - key = _alarm_key(alarm) - existing = db.get(UmeAlarmHistory, key) - if existing is None: - existing = UmeAlarmHistory(alarm_key=key, first_seen_at=touch_ts) - db.add(existing) - inserted += 1 - else: - updated += 1 - existing.ne_id = _s(_derive_ne_id_from_alarm(alarm)) - existing.host_name = _lookup_host_name(db, existing.ne_id) - existing.object_name = _s(_pick(alarm, "objectName", "object-name")) - existing.event_type = _s(_pick(alarm, "eventType", "event-type")) - existing.native_probable_cause = _s(_pick(alarm, "nativeProbableCause", "native-probable-cause")) - existing.perceived_severity = _s(_pick(alarm, "perceivedSeverity", "perceived-severity")) - existing.is_cleared = _s(_pick(alarm, "isCleared", "is-cleared")) - existing.time_created = _s(_pick(alarm, "timeCreated", "time-created")) - existing.root_cause_alarm_indication = _s( - _pick(alarm, "rootCauseAlarmIndication", "root-cause-alarm-indication") - ) - existing.notification_id = notification_id_from_norm(alarm) - existing.last_seen_at = touch_ts - existing.raw_json = json.dumps(alarm, ensure_ascii=False, default=str) - - iterator_500_as_end = bool(getattr(settings, "ume_iterator_500_as_end", True)) - pages, meta = _collect_marker_pages( - lambda marker: client.get_alarms(is_uncleared=is_uncleared, limit=page_size, marker=marker), - max_pages=max_pages, - iterator_500_as_end=iterator_500_as_end, - ) - sync_batch_ts = _utc_now_naive() - seen_keys = set() - for rows in pages: - pulled += len(rows) - for alarm in rows: - if is_uncleared: - upsert_alarm_history(alarm, touch_ts=sync_batch_ts) - else: - key = _alarm_key(alarm) - if key: - seen_keys.add(key) - action, _changed = apply_alarm_to_current( - db, - alarm, - touch_ts=sync_batch_ts, - source="rest", - ) - if action == "inserted": - inserted += 1 - elif action == "updated": - updated += 1 - db.flush() - page_no = int(meta.get("page_count") or 0) - next_marker = str(meta.get("last_marker") or "") - is_end_of_reply = bool(meta.get("is_end_of_reply")) - graceful_end_by_iterator_error = bool(meta.get("graceful_end_by_iterator_error")) - paging_note = str(meta.get("paging_note") or "") - warnings = [str(x) for x in (meta.get("warnings") or []) if str(x)] - - reconcile_mode = "full" - if not is_uncleared and _snapshot_reconcile_ok(meta): - if wss_active: - reconcile_mode = "upsert_only" - deleted_stale_current = _reconcile_stale_current_alarms( - db, - sync_batch_ts=sync_batch_ts, - seen_keys=seen_keys, - wss_active=wss_active, - ) - - alarm_model = UmeAlarmHistory if is_uncleared else UmeAlarmCurrent - host_names_backfilled = _backfill_alarm_host_names(db, alarm_model) - - batch.total_rows = int(pulled) - batch.success_rows = int(inserted + updated) - batch.failed_rows = max(0, int(pulled) - int(inserted + updated)) - batch.status = "done" - batch.ended_at = _utc_now_naive() - batch.raw_json = json.dumps( - { - "pulled": pulled, - "inserted": inserted, - "updated": updated, - "paging_mode": paging_mode, - "page_count": page_no, - "last_marker": next_marker, - "is_end_of_reply": is_end_of_reply, - "graceful_end_by_iterator_error": graceful_end_by_iterator_error, - "warnings": warnings, - "deleted_stale_current_alarms": int(deleted_stale_current), - "host_names_backfilled": int(host_names_backfilled), - "current_snapshot_reconcile": (not is_uncleared) and _snapshot_reconcile_ok(meta), - "reconcile_mode": reconcile_mode if not is_uncleared else "", - "wss_active_during_sync": bool(wss_active) if not is_uncleared else False, - "seen_keys_count": len(seen_keys) if not is_uncleared else 0, - }, - ensure_ascii=False, - ) - - job.status = "done" - except Exception as exc: - msg = str(exc)[:1024] - reconcile_mode = "failed" - batch.status = "failed" - batch.error_message = msg - batch.ended_at = _utc_now_naive() - job.status = "failed" - job.error_message = msg - finally: - job.pulled_count = int(pulled) - job.inserted_count = int(inserted) - job.updated_count = int(updated) - job.ended_at = _utc_now_naive() - job.details_json = json.dumps( - { - "batch_id": batch.batch_id, - "kind": batch.kind, - "status": batch.status, - "paging_mode": paging_mode, - "paging_note": paging_note, - "page_count": page_no, - "last_marker": next_marker, - "is_end_of_reply": is_end_of_reply, - "graceful_end_by_iterator_error": graceful_end_by_iterator_error, - "warnings": warnings, - "deleted_stale_current_alarms": int(deleted_stale_current), - "host_names_backfilled": int(host_names_backfilled), - "current_snapshot_reconcile": (not is_uncleared) and _snapshot_reconcile_ok(meta), - "reconcile_mode": reconcile_mode if not is_uncleared else "", - "wss_active_during_sync": bool(wss_active) if not is_uncleared else False, - "seen_keys_count": len(seen_keys) if not is_uncleared else 0, - }, - ensure_ascii=False, - ) - db.commit() - db.refresh(job) - db.refresh(batch) - return job, batch - - -def sync_alarms_current( - db: Session, - client: UMEClient, - *, - trigger_mode: str = "manual", - wss_active: bool | None = None, -) -> tuple[UmeSyncJob, UmeAlarmBatch]: - if not _ALARMS_CURRENT_SYNC_LOCK.acquire(blocking=False): - _sync_log.warning("alarms_current sync skipped: another sync is in progress") - raise RuntimeError("alarms_current_sync_busy") - try: - if wss_active is None: - from .ume_alarm_ws import is_wss_active_for_current_alarms - - wss_active = is_wss_active_for_current_alarms() - return _sync_alarms_common( - db, - client, - is_uncleared=False, - trigger_mode=trigger_mode, - wss_active=bool(wss_active), - ) - finally: - _ALARMS_CURRENT_SYNC_LOCK.release() - - -def sync_alarms_history_full( - db: Session, client: UMEClient, *, trigger_mode: str = "manual" -) -> tuple[UmeSyncJob, UmeAlarmBatch]: - return _sync_alarms_common(db, client, is_uncleared=True, trigger_mode=trigger_mode) +from .ume_sync_common import _pick, _s, _utc_now_naive +from .ume_sync_pull import sync_alarms_current, sync_alarms_history_full, sync_inventory_full + +__all__ = [ + "_alarm_key", + "_derive_ne_id_from_alarm", + "_is_alarm_cleared", + "_pick", + "_s", + "_utc_now_naive", + "apply_alarm_to_current", + "extract_alarm_from_notification", + "normalize_yang_alarm", + "notification_id_from_norm", + "settings", + "sync_alarms_current", + "sync_alarms_history_full", + "sync_inventory_full", +] diff --git a/netx_api/webcrt_channel.py b/netx_api/webcrt_channel.py index 0b582f0..b205e5b 100644 --- a/netx_api/webcrt_channel.py +++ b/netx_api/webcrt_channel.py @@ -6,6 +6,7 @@ import json import logging import queue import re +import threading import time from datetime import datetime, timezone from pathlib import Path diff --git a/netx_api/webcrt_service.py b/netx_api/webcrt_service.py index 34cae9c..5321cfc 100644 --- a/netx_api/webcrt_service.py +++ b/netx_api/webcrt_service.py @@ -1,10 +1,20 @@ """WebCRT facade — re-exports channel helpers and session registry.""" from __future__ import annotations +from .config import settings +from .ne_session_factory import get_cli_hop_guard, open_netmiko_connection from .webcrt_channel import ( + _BoundedByteQueue, + _audit, + _capture_raw_channel, _decode_bytes, _encode_text, + _is_prompt_only_echo, + _looks_like_cli_prompt, + _looks_like_login_prompt, + _looks_like_password_change_prompt, _normalize_encoding, + _session_log_path, channel_return, map_network_cli_enter, map_network_cli_keys, @@ -27,12 +37,28 @@ from .webcrt_session import ( mark_attached, wait_session_ready, ) +from .webcrt_session_registry import ( + _reap_sessions, + _sessions, + _sessions_lock, +) __all__ = [ "WebcrtSession", + "_BoundedByteQueue", + "_audit", + "_capture_raw_channel", "_decode_bytes", "_encode_text", + "_is_prompt_only_echo", + "_looks_like_cli_prompt", + "_looks_like_login_prompt", + "_looks_like_password_change_prompt", "_normalize_encoding", + "_reap_sessions", + "_session_log_path", + "_sessions", + "_sessions_lock", "_webcrt_creds_ready", "active_session_count", "channel_return", @@ -40,14 +66,17 @@ __all__ = [ "create_session", "detach_session", "find_ssh_session_for_ne", + "get_cli_hop_guard", "get_session", "list_sessions", "map_network_cli_enter", "map_network_cli_keys", "mark_attached", "normalize_cli_transcript", + "open_netmiko_connection", "prepare_bootstrap_output", "read_session_log_tail", + "settings", "uses_network_cli_keymap", "wait_session_ready", "webcrt_data_root", diff --git a/netx_api/webcrt_session.py b/netx_api/webcrt_session.py index 916099e..b094ee6 100644 --- a/netx_api/webcrt_session.py +++ b/netx_api/webcrt_session.py @@ -1,1111 +1,30 @@ -"""WebCRT interactive sessions and process-local registry.""" +"""WebCRT interactive sessions and process-local registry (facade).""" from __future__ import annotations -import io -import json -import logging -import threading -import time -import uuid -from dataclasses import dataclass, field -from datetime import datetime, timezone -from pathlib import Path -from typing import Any - -from fastapi import HTTPException -from netmiko import ConnectHandler -from sqlalchemy.orm import Session - -from .config import settings -from .ne_crypto import CredentialCryptoError -from .ne_session_factory import ( - close_netmiko_connection, - extract_cli_prompt_marker, - get_cli_hop_guard, - open_netmiko_connection, - should_close_cli_hop_session, -) -from .webcrt_channel import ( - _BoundedByteQueue, - _audit, - _capture_raw_channel, - _decode_bytes, - _drain_channel, - _drain_raw_channel, - _encode_text, - _is_prompt_only_echo, - _looks_like_cli_prompt, - _looks_like_login_prompt, - _looks_like_password_change_prompt, - _normalize_encoding, - _prime_interactive_channel, - _session_log_path, - _session_log_text, - _utc_iso, - _utc_now, - channel_return, - map_network_cli_enter, - map_network_cli_keys, - normalize_cli_transcript, - prepare_bootstrap_output, - read_session_log_tail, - uses_network_cli_keymap, - webcrt_data_root, +from .webcrt_session_model import WebcrtSession +from .webcrt_session_registry import ( + _webcrt_creds_ready, + active_session_count, + close_session, + create_session, + detach_session, + find_ssh_session_for_ne, + get_session, + list_sessions, + mark_attached, + wait_session_ready, ) -_log = logging.getLogger("netx.webcrt") - -_sessions_lock = threading.Lock() -_sessions: dict[str, "WebcrtSession"] = {} -_reaper_started = False - - -@dataclass -class WebcrtSession: - session_id: str - ne_id: str - ne_name: str - ne_ip: str - protocol: str - cols: int - rows: int - device_type: str = "" - vendor: str = "" - cli_keymap: bool = True - encoding: str = "utf-8" - keepalive_sec: int = 0 - conn: ConnectHandler | None = None - created_at: float = field(default_factory=time.time) - last_activity: float = field(default_factory=time.time) - attached: bool = False - detach_deadline: float | None = None - closed: bool = False - close_reason: str = "" - state: str = "connecting" - connect_error: str = "" - connect_started_at: float = field(default_factory=time.time) - connect_finished_at: float | None = None - bootstrap_output: bytes = b"" - # First WS attach gets login bootstrap; later attaches prefer session-log tail. - bootstrap_replayed: bool = False - needs_live_prompt: bool = True - # React StrictMode remounts open a second WS before the first fully tears down. - # Only the newest attach_gen may consume out_queue / mark detach. - attach_gen: int = 0 - out_queue: _BoundedByteQueue = field( - default_factory=lambda: _BoundedByteQueue(int(getattr(settings, "webcrt_out_queue_max", 2000) or 2000)) - ) - # Vendor CLI hop (Huawei/ZTE/Cisco): close when nested target session returns to hop. - cli_hop_guard: bool = False - cli_hop_prompt: str = "" - post_login_commands: list[str] = field(default_factory=list) - bytes_in: int = 0 - bytes_out: int = 0 - _reader: threading.Thread | None = field(default=None, repr=False) - _write_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) - _stdout_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) - _hop_scan_buf: str = field(default="", repr=False) - _cli_hop_seen_other_prompt: bool = field(default=False, repr=False) - _log_fh: Any = field(default=None, repr=False) - _ready_event: threading.Event = field(default_factory=threading.Event, repr=False) - # SFTP channel on the same SSH transport as the interactive shell (direct SSH only). - sftp_ready: bool = False - _sftp: Any = field(default=None, repr=False) - _sftp_lock: threading.RLock = field(default_factory=threading.RLock, repr=False) - - def touch(self) -> None: - self.last_activity = time.time() - - def close_sftp(self) -> None: - with self._sftp_lock: - sftp = self._sftp - self._sftp = None - self.sftp_ready = False - if sftp is None: - return - try: - sftp.close() - except Exception: - pass - - def _ssh_transport_unlocked(self) -> Any: - """Caller must hold ``_sftp_lock``. Returns an active Paramiko Transport.""" - if self.closed or self.conn is None: - raise RuntimeError("session_closed") - if str(self.protocol or "ssh").lower() != "ssh": - raise RuntimeError("sftp_requires_ssh") - if self.cli_hop_guard: - raise RuntimeError("sftp_hop_not_supported") - channel = getattr(self.conn, "remote_conn", None) - transport = None - if channel is not None and hasattr(channel, "get_transport"): - try: - transport = channel.get_transport() - except Exception: - transport = None - if transport is None or not bool(getattr(transport, "is_active", lambda: False)()): - raise RuntimeError("ssh_transport_unavailable") - return transport - - def _ensure_sftp_unlocked(self) -> Any: - """Caller must hold ``_sftp_lock``. Shared probe client (sftp_ready).""" - import paramiko - - if self._sftp is not None: - sock = getattr(self._sftp, "sock", None) - if sock is not None and not bool(getattr(sock, "closed", False)): - return self._sftp - try: - self._sftp.close() - except Exception: - pass - self._sftp = None - transport = self._ssh_transport_unlocked() - self._sftp = paramiko.SFTPClient.from_transport(transport) - if self._sftp is None: - raise RuntimeError("sftp_open_failed") - self.sftp_ready = True - return self._sftp - - def open_sftp(self) -> Any: - """Open/reuse an SFTP client on this session's SSH transport.""" - with self._sftp_lock: - return self._ensure_sftp_unlocked() - - def open_ephemeral_sftp(self) -> Any: - """Open a dedicated SFTP channel for one operation; caller must ``close()`` it. - - Only holds ``_sftp_lock`` briefly while resolving the SSH transport, so long - list/upload/download work does not block other SFTP ops on the same session. - """ - import paramiko - - with self._sftp_lock: - transport = self._ssh_transport_unlocked() - # Keep probe client warm for UI sftp_ready without sharing it for I/O. - try: - self._ensure_sftp_unlocked() - except Exception: - pass - sftp = paramiko.SFTPClient.from_transport(transport) - if sftp is None: - raise RuntimeError("sftp_open_failed") - return sftp - - def run_sftp(self, fn: Any) -> Any: - """Run ``fn(sftp)`` on an ephemeral channel (does not hold the lock during ``fn``).""" - sftp = self.open_ephemeral_sftp() - try: - return fn(sftp) - finally: - try: - sftp.close() - except Exception: - pass - - def try_attach_sftp(self) -> bool: - """Best-effort SFTP channel open after SSH login (does not fail the shell).""" - if str(self.protocol or "ssh").lower() != "ssh" or self.cli_hop_guard: - self.sftp_ready = False - return False - try: - self.open_sftp() - self.sftp_ready = True - return True - except Exception: - self.sftp_ready = False - _log.debug("webcrt sftp attach skipped session=%s", self.session_id, exc_info=True) - return False - - def open_session_log(self) -> None: - if not bool(getattr(settings, "webcrt_session_log_enabled", True)): - return - if self._log_fh is not None: - return - try: - self._log_fh = _session_log_path(self.session_id).open("a", encoding="utf-8", errors="replace") - self._log_fh.write(f"# session={self.session_id} ne={self.ne_id} ip={self.ne_ip} ts={_utc_iso()}\n") - self._log_fh.flush() - except Exception: - _log.debug("webcrt session log open failed", exc_info=True) - self._log_fh = None - - def append_session_log(self, text: str) -> None: - if not text or self._log_fh is None: - return - try: - self._log_fh.write(text) - self._log_fh.flush() - except Exception: - pass - - def close_session_log(self) -> None: - fh = self._log_fh - self._log_fh = None - if fh is None: - return - try: - fh.write(f"\n# closed reason={self.close_reason} ts={_utc_iso()}\n") - fh.close() - except Exception: - pass - - def take_stdout(self, attach_gen: int, *, timeout: float = 0.25) -> bytes | None | str: - """Exclusive stdout take for one WS attach generation. - - Returns: - bytes — device output chunk - None — device reader closed (session end) - \"stale\" — a newer WebSocket owns this session; caller must stop - \"empty\" — no data within timeout (keep polling) - """ - deadline = time.time() + max(0.05, float(timeout)) - while True: - with self._stdout_lock: - if attach_gen != self.attach_gen: - return "stale" - remaining = deadline - time.time() - if remaining <= 0: - return "empty" - # Slice waits so we can notice attach_gen bumps without busy-spinning. - try: - chunk = self.out_queue.get(timeout=min(0.05, remaining)) - except queue.Empty: - continue - with self._stdout_lock: - if attach_gen != self.attach_gen: - # Put back including EOF sentinel so the new owner still sees close. - self.out_queue.put(chunk) - return "stale" - return chunk # bytes | None - - def write_stdin(self, data: str) -> None: - if self.closed or self.conn is None: - raise RuntimeError("session_closed") - text = str(data or "") - if not text: - return - if self.cli_keymap: - text = map_network_cli_keys( - text, - device_type=self.device_type, - vendor=self.vendor, - protocol=self.protocol, - ) - text = map_network_cli_enter(text, self.conn) - if not text: - return - with self._write_lock: - # Prefer raw channel I/O for interactive typing (char echo / backspace). - channel = getattr(self.conn, "remote_conn", None) - try: - if channel is not None and hasattr(channel, "send") and callable(channel.send): - payload = _encode_text(text, self.encoding) - # Paramiko may write partially when the window is full. - view = memoryview(payload) - while len(view): - n = int(channel.send(view) or 0) - if n <= 0: - time.sleep(0.01) - continue - view = view[n:] - self.bytes_in += len(payload) - elif channel is not None and hasattr(channel, "write") and callable(channel.write): - payload = _encode_text(text, self.encoding) - channel.write(payload) - self.bytes_in += len(payload) - else: - self.conn.write_channel(text) - self.bytes_in += len(text) - except Exception: - self.conn.write_channel(text) - self.bytes_in += len(text) - self.touch() - - def send_break(self) -> None: - """Send SSH break / Telnet IAC BREAK to interrupt paging or hung commands.""" - if self.closed or self.conn is None: - raise RuntimeError("session_closed") - channel = getattr(self.conn, "remote_conn", None) - with self._write_lock: - sent = False - if channel is not None and hasattr(channel, "send_break") and callable(channel.send_break): - try: - channel.send_break(0) - sent = True - except Exception: - _log.debug("send_break failed session=%s", self.session_id, exc_info=True) - if not sent and channel is not None and hasattr(channel, "send") and callable(channel.send): - # Telnet IAC BREAK = 255 243 - try: - channel.send(b"\xff\xf3") - sent = True - except Exception: - pass - if not sent: - # Fallback: Ctrl-C often interrupts device CLI more-pages. - try: - self.conn.write_channel("\x03") - except Exception: - raise RuntimeError("break_failed") - self.touch() - - def resize(self, cols: int, rows: int) -> None: - if self.closed or self.conn is None: - return - c = max(20, min(500, int(cols or 80))) - r = max(5, min(200, int(rows or 24))) - self.cols = c - self.rows = r - channel = getattr(self.conn, "remote_conn", None) - if channel is not None and hasattr(channel, "resize_pty"): - try: - channel.resize_pty(width=c, height=r) - except Exception: - _log.debug("resize_pty failed session=%s", self.session_id, exc_info=True) - self.touch() - - def start_reader(self) -> None: - if self._reader and self._reader.is_alive(): - return - self._reader = threading.Thread( - target=self._reader_loop, - name=f"webcrt-reader-{self.session_id[:8]}", - daemon=True, - ) - self._reader.start() - - def _reader_loop(self) -> None: - conn = self.conn - if conn is None: - self.out_queue.put(None) - return - channel = getattr(conn, "remote_conn", None) - hop_return = False - poll = max(0.002, float(getattr(settings, "webcrt_reader_poll_sec", 0.01) or 0.01)) - try: - while not self.closed: - chunk = b"" - try: - if channel is not None and hasattr(channel, "recv_ready") and hasattr(channel, "recv"): - # Paramiko SSH: prefer short blocking recv over fixed spin-sleep. - ready = False - try: - ready = bool(channel.recv_ready()) - except Exception: - ready = False - if ready: - chunk = channel.recv(16384) - if not chunk: - break - elif hasattr(channel, "exit_status_ready") and channel.exit_status_ready(): - break - else: - # Brief block: settimeout + recv wakes sooner than sleep(0.04). - prev_timeout = None - try: - prev_timeout = channel.gettimeout() - except Exception: - prev_timeout = None - try: - channel.settimeout(poll) - chunk = channel.recv(16384) - except Exception: - chunk = b"" - finally: - try: - channel.settimeout(prev_timeout) - except Exception: - pass - if not chunk: - continue - elif channel is not None and hasattr(channel, "read_very_eager"): - # Telnet: do NOT use conn.read_channel() — Netmiko strips ANSI. - data = channel.read_very_eager() - if data: - chunk = ( - data - if isinstance(data, (bytes, bytearray)) - else _encode_text(str(data), self.encoding) - ) - else: - time.sleep(poll) - continue - else: - text = conn.read_channel() - if text: - chunk = _encode_text(str(text), self.encoding) - else: - time.sleep(poll) - continue - except Exception as exc: - if self.closed: - break - _log.debug("webcrt reader error session=%s: %s", self.session_id, exc) - time.sleep(0.05) - continue - if chunk: - self.touch() - self.bytes_out += len(chunk) - self.out_queue.put(chunk) - try: - self.append_session_log(_decode_bytes(chunk, self.encoding)) - except Exception: - pass - if self.cli_hop_guard and self._note_cli_hop_output(chunk): - hop_return = True - notice = ( - "\r\n*** WebCRT: 目标会话已结束,已断开代理连接 " - "(target session ended; closing hop proxy) ***\r\n" - ) - self.out_queue.put(notice.encode("utf-8", errors="replace")) - break - finally: - if hop_return and not self.closed: - try: - close_session(self.session_id, reason="cli_hop_return") - except Exception: - self.close("cli_hop_return") - self.out_queue.put(None) - - def _note_cli_hop_output(self, chunk: bytes) -> bool: - """Accumulate stdout and return True when nested CLI hop has returned to proxy.""" - try: - text = _decode_bytes(chunk, self.encoding) - except Exception: - text = str(chunk) - self._hop_scan_buf = (self._hop_scan_buf + text)[-12000:] - marker = str(self.cli_hop_prompt or "").strip() - last = extract_cli_prompt_marker(self._hop_scan_buf) - if last and (not marker or last != marker): - self._cli_hop_seen_other_prompt = True - return should_close_cli_hop_session( - self._hop_scan_buf, - self.cli_hop_prompt, - seen_other_prompt=self._cli_hop_seen_other_prompt, - ) - - def run_post_login_commands(self) -> None: - cmds = [str(c).rstrip("\r\n") for c in (self.post_login_commands or []) if str(c).strip()] - if not cmds or self.closed or self.conn is None: - return - for cmd in cmds[:20]: - try: - self.write_stdin(cmd + "\r") - time.sleep(0.15) - except Exception: - _log.debug("post_login command failed session=%s", self.session_id, exc_info=True) - break - - def close(self, reason: str = "closed") -> None: - if self.closed: - return - self.closed = True - self.state = "closed" - self.close_reason = reason or "closed" - self._ready_event.set() - self.close_sftp() - try: - close_netmiko_connection(self.conn) - except Exception: - pass - self.conn = None - try: - self.out_queue.put_nowait(None) - except Exception: - pass - self.close_session_log() - - -def _ensure_reaper() -> None: - global _reaper_started - with _sessions_lock: - if _reaper_started: - return - _reaper_started = True - t = threading.Thread(target=_reaper_loop, name="webcrt-reaper", daemon=True) - t.start() - - -def _reaper_loop() -> None: - while True: - try: - _reap_sessions() - except Exception: - _log.exception("webcrt reaper failed") - time.sleep(2) - - -def _reap_sessions() -> None: - idle = max(60, int(settings.webcrt_idle_timeout_sec or 1800)) - attach = max(10, int(settings.webcrt_attach_timeout_sec or 60)) - anti_idle = max(0, int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0)) - anti_payload = str(getattr(settings, "webcrt_anti_idle_payload", " ") or " ") - now = time.time() - to_close: list[tuple[WebcrtSession, str]] = [] - to_nudge: list[WebcrtSession] = [] - with _sessions_lock: - for sess in list(_sessions.values()): - if sess.closed: - _sessions.pop(sess.session_id, None) - continue - if sess.state == "connecting": - # Connecting sessions use connect timeout, not attach timeout alone. - connect_budget = max(30, int(settings.webcrt_connect_timeout_sec or 90)) + 30 - if (now - sess.connect_started_at) > connect_budget: - to_close.append((sess, "connect_timeout")) - continue - if sess.attached: - if (now - sess.last_activity) > idle: - to_close.append((sess, "idle_timeout")) - elif ( - anti_idle > 0 - and sess.state == "ready" - and sess.conn is not None - and (now - sess.last_activity) >= anti_idle - ): - to_nudge.append(sess) - continue - # Not attached: either never attached, or briefly detached for reconnect. - if sess.detach_deadline is not None: - if now >= sess.detach_deadline: - to_close.append((sess, "detach_timeout")) - else: - # Start attach clock after connect finishes (not HTTP create time), - # so slow auth + UI mount does not race attach_timeout. - anchor = float(sess.connect_finished_at or sess.created_at or now) - if (now - anchor) > attach: - to_close.append((sess, "attach_timeout")) - elif (now - sess.last_activity) > idle: - to_close.append((sess, "idle_timeout")) - for sess in to_nudge: - try: - # Touch without changing visible prompt when payload is empty/null-ish. - payload = anti_payload - if payload == "\\0": - payload = "\x00" - if payload: - sess.write_stdin(payload) - else: - sess.touch() - except Exception: - _log.debug("webcrt anti-idle failed session=%s", sess.session_id, exc_info=True) - for sess, reason in to_close: - close_session(sess.session_id, reason=reason) - - -def active_session_count() -> int: - with _sessions_lock: - return sum(1 for s in _sessions.values() if not s.closed) - - -def get_session(session_id: str) -> WebcrtSession | None: - with _sessions_lock: - sess = _sessions.get(session_id) - if sess is None or sess.closed: - return None - return sess - - -def find_ssh_session_for_ne(ne_id: str) -> WebcrtSession | None: - """Prefer a ready/attached interactive SSH session for SFTP channel reuse.""" - nid = str(ne_id or "").strip() - if not nid: - return None - with _sessions_lock: - candidates = [ - s - for s in _sessions.values() - if (not s.closed) - and str(s.ne_id) == nid - and str(s.protocol or "ssh").lower() == "ssh" - and s.conn is not None - and s.state in ("ready", "connecting") - and not s.cli_hop_guard - ] - if not candidates: - return None - # Prefer attached + ready sessions. - candidates.sort(key=lambda s: (0 if s.attached and s.state == "ready" else 1, -s.last_activity)) - return candidates[0] - - -def wait_session_ready(session_id: str, *, timeout: float = 120.0) -> WebcrtSession: - """Block until async connect finishes (ready or error). Used by tests and WS.""" - deadline = time.time() + max(1.0, float(timeout)) - while time.time() < deadline: - sess = get_session(session_id) - if sess is None: - raise HTTPException(status_code=404, detail="webcrt_session_not_found") - if sess.state == "ready": - return sess - if sess.state == "error": - raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") - sess._ready_event.wait(timeout=0.25) - raise HTTPException(status_code=504, detail="connect_timeout") - - -def _webcrt_creds_ready(creds: dict[str, Any]) -> bool: - """True when WebCRT can open a session with the resolved credentials. - - Bastion-managed hops store the target password on the bastion side, so an empty - NE password is valid (same as connectivity test). Direct / manual / Linux hops - still require a target password for SSH. - - Telnet (no hop) allows empty username/password so the user can authenticate - interactively in the terminal (SecureCRT-style). - """ - hop_enabled = bool(creds.get("hop_enabled")) - hop_vendor = str(creds.get("hop_vendor") or "").strip().lower() - auth_mode = str(creds.get("hop_target_auth_mode") or "bastion_managed").strip().lower() - protocol = str(creds.get("protocol") or "ssh").strip().lower() - if hop_enabled and hop_vendor == "bastion" and auth_mode == "bastion_managed": - return bool( - str(creds.get("hop_host") or "").strip() - and str(creds.get("hop_username") or "").strip() - and str(creds.get("hop_password") or "") - ) - if protocol == "telnet" and not hop_enabled: - return True - if not str(creds.get("username") or "").strip(): - return False - return bool(str(creds.get("password") or "")) - - -def _finish_connect( - sess: WebcrtSession, - *, - creds: dict[str, Any], - device: dict[str, Any], - connect_timeout: int, - client: str, -) -> None: - log_buf = io.BytesIO() - try: - conn = open_netmiko_connection( - creds, - session_timeout=connect_timeout, - session_log=log_buf, - cols=sess.cols, - rows=sess.rows, - interactive=True, - keepalive=int(sess.keepalive_sec or 0), - ) - except Exception as exc: - partial = _session_log_text(log_buf).strip() - from .ne_cli_errors import format_cli_failure - - classified = format_cli_failure(exc, partial) - detail = f"connect_failed:{classified}" - if partial: - detail = f"{detail}\n--- device transcript ---\n{partial[-4000:]}" - sess.state = "error" - sess.connect_error = detail - sess.connect_finished_at = time.time() - sess._ready_event.set() - _audit( - "session_open_failed", - session_id=sess.session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - source=str(device.get("source") or ""), - client=client or "", - error=str(exc)[:500], - transcript_len=len(partial), - ) - return - - channel = getattr(conn, "remote_conn", None) - if channel is not None and hasattr(channel, "resize_pty"): - try: - channel.resize_pty(width=sess.cols, height=sess.rows) - except Exception: - pass - - pre_log = _session_log_text(log_buf) - # Pull post-auth banner/MOTD from the PTY. With interactive no-op session_preparation - # (generic_termserver), Netmiko session_log is often empty — do not discard these bytes. - try: - early = _capture_raw_channel(conn, duration=0.35) - except Exception: - early = "" - seed = f"{pre_log}{early}" - already_prompted = _looks_like_cli_prompt(seed) - primed = "" - # Do not send Enter at Username:/Password: or Huawei password-change [Y/N]: - # (Netmiko telnet_login already answers password-change with "N"). - if _looks_like_login_prompt(seed) or _looks_like_password_change_prompt(seed): - try: - primed = _capture_raw_channel(conn, duration=0.9) - except Exception: - primed = "" - else: - try: - primed = _prime_interactive_channel(conn, already_prompted=already_prompted) - except Exception: - primed = "" - combined = f"{seed}{primed}" - # Final settle: keep stragglers in bootstrap (normalize collapses duplicate prompts). - try: - combined += _capture_raw_channel(conn, duration=0.35) - except Exception: - pass - if not str(combined).strip(): - try: - combined = _drain_channel(conn, rounds=6, wait=0.08) - except Exception: - combined = "" - bootstrap = prepare_bootstrap_output(combined) - # Discard lone punctuation left on the wire (would glue onto ```` in xterm). - try: - leftover = _capture_raw_channel(conn, duration=0.12) - except Exception: - leftover = "" - if leftover and leftover.strip() not in {":", ">", "#", "]", "$"}: - bootstrap = prepare_bootstrap_output(f"{bootstrap}{leftover}") - - hop_guard = get_cli_hop_guard(conn) - sess.conn = conn - sess.cli_hop_guard = bool(hop_guard) - sess.cli_hop_prompt = str((hop_guard or {}).get("hop_prompt") or "") - sess.bootstrap_output = _encode_text(str(bootstrap or ""), sess.encoding) - # Nudge Enter on WS attach only when we still need a shell prompt. - # Never when already at CLI prompt or Username:/Password: (would empty-submit login). - sess.needs_live_prompt = ( - not _looks_like_cli_prompt(bootstrap) and not _looks_like_login_prompt(bootstrap) - ) - sess.open_session_log() - if bootstrap: - sess.append_session_log(bootstrap if bootstrap.endswith("\n") else bootstrap + "\n") - sess.start_reader() - # Drop late prompt echoes that race into the queue right after reader start. - prompt_hint = "" - if bootstrap: - prompt_hint = str(bootstrap).replace("\r\n", "\n").replace("\r", "\n").strip().split("\n")[-1].strip() - settle_deadline = time.time() + 0.45 - while time.time() < settle_deadline: - try: - chunk = sess.out_queue.get_nowait() - except queue.Empty: - time.sleep(0.02) - continue - if chunk is None: - sess.out_queue.put(None) - break - try: - text = _decode_bytes(chunk, sess.encoding) - except Exception: - text = "" - if _is_prompt_only_echo(text, prompt_hint): - continue - # Non-prompt data: put back and stop settling. - sess.out_queue.put(chunk) - break - try: - sess.run_post_login_commands() - except Exception: - _log.debug("post_login failed session=%s", sess.session_id, exc_info=True) - # Same SSH transport: open SFTP channel when the device supports it. - sftp_ok = sess.try_attach_sftp() - sess.state = "ready" - sess.connect_finished_at = time.time() - sess._ready_event.set() - elapsed_ms = int((sess.connect_finished_at - sess.connect_started_at) * 1000) - _audit( - "session_created", - session_id=sess.session_id, - ne_id=sess.ne_id, - ne_name=sess.ne_name, - ne_ip=sess.ne_ip, - protocol=sess.protocol, - encoding=sess.encoding, - source=str(device.get("source") or ""), - hop_enabled=bool(creds.get("hop_enabled")), - hop_vendor=str(creds.get("hop_vendor") or "") if creds.get("hop_enabled") else "", - cli_hop_guard=bool(hop_guard), - cli_hop_prompt=str((hop_guard or {}).get("hop_prompt") or ""), - sftp_ready=bool(sftp_ok), - client=client or "", - connect_ms=elapsed_ms, - active=active_session_count(), - ) - - -def create_session( - db: Session, - *, - ne_id: str | None = None, - ume_ne_id: str | None = None, - cols: int = 80, - rows: int = 24, - client: str = "", - encoding: str = "utf-8", - keepalive_sec: int | None = None, - post_login_commands: list[str] | None = None, - async_connect: bool = True, - username_override: str | None = None, - password_override: str | None = None, -) -> dict[str, Any]: - from .cli_resolve import resolve_cli_target - - _ensure_reaper() - max_sessions = max(1, int(settings.webcrt_max_sessions or 20)) - if active_session_count() >= max_sessions: - raise HTTPException(status_code=429, detail="webcrt_session_limit") - - mid = str(ne_id or "").strip() - uid = str(ume_ne_id or "").strip() - try: - creds, device = resolve_cli_target(db, managed_ne_id=mid or None, ume_ne_id=uid or None) - except HTTPException: - raise - except CredentialCryptoError as exc: - raise HTTPException(status_code=400, detail=str(exc) or "credential_crypto_error") from exc - except Exception as exc: - raise HTTPException(status_code=400, detail=f"credential_error:{exc}") from exc - - # One-shot credentials for SecureCRT-style "do not save password" / retry. - if username_override is not None and str(username_override).strip(): - creds["username"] = str(username_override).strip() - if password_override is not None: - creds["password"] = str(password_override) - - protocol = str(device.get("protocol") or creds.get("protocol") or "ssh").strip().lower() - creds["protocol"] = protocol - # Netmiko telnet drivers dislike a completely missing username; use a placeholder - # for the wire only (interactive login still happens in the terminal). - if protocol == "telnet" and not bool(creds.get("hop_enabled")) and not str(creds.get("username") or "").strip(): - creds["username"] = "telnet" - - if not _webcrt_creds_ready(creds): - raise HTTPException(status_code=400, detail="credentials_incomplete") - - session_id = str(uuid.uuid4()) - c = max(20, min(500, int(cols or 80))) - r = max(5, min(200, int(rows or 24))) - connect_timeout = max(30, int(settings.webcrt_connect_timeout_sec or 90)) - target_id = str(device.get("id") or mid or uid) - target_ip = str(device.get("ip_address") or "") - target_name = str(device.get("name") or target_ip) - device_type = str(device.get("device_type") or creds.get("device_type") or "") - vendor = str(device.get("vendor") or creds.get("vendor") or "") - cli_keymap = uses_network_cli_keymap(device_type, vendor) - enc = _normalize_encoding(encoding) - if keepalive_sec is None: - ka = max(0, int(getattr(settings, "webcrt_keepalive_sec", 0) or 0)) - else: - ka = max(0, min(600, int(keepalive_sec))) - - sess = WebcrtSession( - session_id=session_id, - ne_id=target_id, - ne_name=target_name, - ne_ip=target_ip, - protocol=protocol, - cols=c, - rows=r, - device_type=device_type, - vendor=vendor, - cli_keymap=cli_keymap, - encoding=enc, - keepalive_sec=ka, - state="connecting", - post_login_commands=list(post_login_commands or [])[:20], - ) - with _sessions_lock: - _sessions[session_id] = sess - - _audit( - "session_connecting", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - protocol=sess.protocol, - encoding=enc, - client=client or "", - async_connect=bool(async_connect), - ) - - if async_connect: - t = threading.Thread( - target=_finish_connect, - kwargs={ - "sess": sess, - "creds": creds, - "device": device, - "connect_timeout": connect_timeout, - "client": client or "", - }, - name=f"webcrt-connect-{session_id[:8]}", - daemon=True, - ) - t.start() - else: - _finish_connect( - sess, - creds=creds, - device=device, - connect_timeout=connect_timeout, - client=client or "", - ) - if sess.state == "error": - with _sessions_lock: - _sessions.pop(session_id, None) - raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") - - return { - "session_id": session_id, - "ne_id": sess.ne_id, - "ne_name": sess.ne_name, - "ne_ip": sess.ne_ip, - "source": str(device.get("source") or ""), - "protocol": sess.protocol, - "cols": sess.cols, - "rows": sess.rows, - "encoding": enc, - "keepalive_sec": ka, - "state": sess.state, - "ws_path": f"/v1/webcrt/sessions/{session_id}/ws", - "cli_hop": bool(sess.cli_hop_guard), - "sftp_ready": bool(sess.sftp_ready), - } - - -def mark_attached(session_id: str) -> tuple[WebcrtSession, int]: - sess = get_session(session_id) - if sess is None: - raise HTTPException(status_code=404, detail="webcrt_session_not_found") - # Allow re-attach after brief WS drop (React StrictMode remount / network blip). - # Bump generation so the previous WS pump stops and does not steal echo bytes. - sess.attach_gen += 1 - attach_gen = sess.attach_gen - sess.attached = True - sess.detach_deadline = None - sess.touch() - _audit( - "session_attached", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - attach_gen=attach_gen, - state=sess.state, - ) - return sess, attach_gen - - -def detach_session( - session_id: str, - *, - grace_sec: float = 8.0, - client: str = "", - attach_gen: int | None = None, -) -> dict[str, Any]: - """Mark session unattached but keep device channel open briefly for reconnect.""" - sess = get_session(session_id) - if sess is None: - return {"ok": True, "session_id": session_id, "detached": False} - if sess.closed: - return {"ok": True, "session_id": session_id, "detached": False} - # Ignore detach from an older StrictMode WS once a newer attach owns the session. - if attach_gen is not None and attach_gen != sess.attach_gen: - return { - "ok": True, - "session_id": session_id, - "detached": False, - "ignored_stale_attach": True, - } - sess.attached = False - sess.detach_deadline = time.time() + max(1.0, float(grace_sec)) - sess.touch() - _audit( - "session_detached", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - grace_sec=grace_sec, - client=client or "", - attach_gen=attach_gen, - ) - return {"ok": True, "session_id": session_id, "detached": True} - - -def close_session(session_id: str, *, reason: str = "closed", client: str = "") -> dict[str, Any]: - with _sessions_lock: - sess = _sessions.pop(session_id, None) - if sess is None: - return {"ok": True, "session_id": session_id, "closed": False} - if not sess.closed: - sess.close(reason) - _audit( - "session_closed", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - reason=reason, - client=client or "", - bytes_in=sess.bytes_in, - bytes_out=sess.bytes_out, - queue_dropped=getattr(sess.out_queue, "dropped", 0), - active=active_session_count(), - ) - return {"ok": True, "session_id": session_id, "closed": True, "reason": reason} - - -def list_sessions() -> dict[str, Any]: - with _sessions_lock: - items = [] - for s in _sessions.values(): - if s.closed: - continue - state = str(s.state or "unknown") - attached = bool(s.attached) - # Lifecycle for ops UI: distinguish login vs live vs grace-period detach. - if state == "connecting": - lifecycle = "connecting" - elif state == "error": - lifecycle = "error" - elif state == "ready" and attached: - lifecycle = "ready" - elif state == "ready" and not attached: - lifecycle = "detached" - else: - lifecycle = state - elapsed_ms = None - if state == "connecting": - elapsed_ms = int(max(0.0, time.time() - float(s.connect_started_at or time.time())) * 1000) - items.append( - { - "session_id": s.session_id, - "ne_id": s.ne_id, - "ne_name": s.ne_name, - "ne_ip": s.ne_ip, - "protocol": s.protocol, - "encoding": s.encoding, - "keepalive_sec": int(s.keepalive_sec or 0), - "state": state, - "lifecycle": lifecycle, - "attached": attached, - "detach_deadline": s.detach_deadline, - "connect_error": str(s.connect_error or "")[:500], - "elapsed_ms": elapsed_ms, - "created_at": datetime.fromtimestamp(s.created_at, tz=timezone.utc).isoformat(), - "last_activity": datetime.fromtimestamp(s.last_activity, tz=timezone.utc).isoformat(), - "bytes_in": s.bytes_in, - "bytes_out": s.bytes_out, - "queue_depth": s.out_queue.qsize(), - "queue_dropped": getattr(s.out_queue, "dropped", 0), - "connect_ms": ( - int((s.connect_finished_at - s.connect_started_at) * 1000) - if s.connect_finished_at - else None - ), - } - ) - return { - "total": len(items), - "max_sessions": max(1, int(settings.webcrt_max_sessions or 20)), - "idle_timeout_sec": max(60, int(settings.webcrt_idle_timeout_sec or 1800)), - "keepalive_sec": int(getattr(settings, "webcrt_keepalive_sec", 0) or 0), - "anti_idle_sec": int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0), - "items": items, - } +__all__ = [ + "WebcrtSession", + "_webcrt_creds_ready", + "active_session_count", + "close_session", + "create_session", + "detach_session", + "find_ssh_session_for_ne", + "get_session", + "list_sessions", + "mark_attached", + "wait_session_ready", +] diff --git a/netx_api/webcrt_session_model.py b/netx_api/webcrt_session_model.py new file mode 100644 index 0000000..cfb0e93 --- /dev/null +++ b/netx_api/webcrt_session_model.py @@ -0,0 +1,487 @@ +"""WebCRT interactive session object (I/O, SFTP, close).""" +from __future__ import annotations + +import logging +import queue +import threading +import time +from dataclasses import dataclass, field +from typing import Any + +from netmiko import ConnectHandler + +from .config import settings +from .ne_session_factory import ( + close_netmiko_connection, + extract_cli_prompt_marker, + should_close_cli_hop_session, +) +from .webcrt_channel import ( + _BoundedByteQueue, + _decode_bytes, + _encode_text, + _session_log_path, + _utc_iso, + map_network_cli_enter, + map_network_cli_keys, +) + +_log = logging.getLogger("netx.webcrt") + +@dataclass +class WebcrtSession: + session_id: str + ne_id: str + ne_name: str + ne_ip: str + protocol: str + cols: int + rows: int + device_type: str = "" + vendor: str = "" + cli_keymap: bool = True + encoding: str = "utf-8" + keepalive_sec: int = 0 + conn: ConnectHandler | None = None + created_at: float = field(default_factory=time.time) + last_activity: float = field(default_factory=time.time) + attached: bool = False + detach_deadline: float | None = None + closed: bool = False + close_reason: str = "" + state: str = "connecting" + connect_error: str = "" + connect_started_at: float = field(default_factory=time.time) + connect_finished_at: float | None = None + bootstrap_output: bytes = b"" + # First WS attach gets login bootstrap; later attaches prefer session-log tail. + bootstrap_replayed: bool = False + needs_live_prompt: bool = True + # React StrictMode remounts open a second WS before the first fully tears down. + # Only the newest attach_gen may consume out_queue / mark detach. + attach_gen: int = 0 + out_queue: _BoundedByteQueue = field( + default_factory=lambda: _BoundedByteQueue(int(getattr(settings, "webcrt_out_queue_max", 2000) or 2000)) + ) + # Vendor CLI hop (Huawei/ZTE/Cisco): close when nested target session returns to hop. + cli_hop_guard: bool = False + cli_hop_prompt: str = "" + post_login_commands: list[str] = field(default_factory=list) + bytes_in: int = 0 + bytes_out: int = 0 + _reader: threading.Thread | None = field(default=None, repr=False) + _write_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) + _stdout_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) + _hop_scan_buf: str = field(default="", repr=False) + _cli_hop_seen_other_prompt: bool = field(default=False, repr=False) + _log_fh: Any = field(default=None, repr=False) + _ready_event: threading.Event = field(default_factory=threading.Event, repr=False) + # SFTP channel on the same SSH transport as the interactive shell (direct SSH only). + sftp_ready: bool = False + _sftp: Any = field(default=None, repr=False) + _sftp_lock: threading.RLock = field(default_factory=threading.RLock, repr=False) + + def touch(self) -> None: + self.last_activity = time.time() + + def close_sftp(self) -> None: + with self._sftp_lock: + sftp = self._sftp + self._sftp = None + self.sftp_ready = False + if sftp is None: + return + try: + sftp.close() + except Exception: + pass + + def _ssh_transport_unlocked(self) -> Any: + """Caller must hold ``_sftp_lock``. Returns an active Paramiko Transport.""" + if self.closed or self.conn is None: + raise RuntimeError("session_closed") + if str(self.protocol or "ssh").lower() != "ssh": + raise RuntimeError("sftp_requires_ssh") + if self.cli_hop_guard: + raise RuntimeError("sftp_hop_not_supported") + channel = getattr(self.conn, "remote_conn", None) + transport = None + if channel is not None and hasattr(channel, "get_transport"): + try: + transport = channel.get_transport() + except Exception: + transport = None + if transport is None or not bool(getattr(transport, "is_active", lambda: False)()): + raise RuntimeError("ssh_transport_unavailable") + return transport + + def _ensure_sftp_unlocked(self) -> Any: + """Caller must hold ``_sftp_lock``. Shared probe client (sftp_ready).""" + import paramiko + + if self._sftp is not None: + sock = getattr(self._sftp, "sock", None) + if sock is not None and not bool(getattr(sock, "closed", False)): + return self._sftp + try: + self._sftp.close() + except Exception: + pass + self._sftp = None + transport = self._ssh_transport_unlocked() + self._sftp = paramiko.SFTPClient.from_transport(transport) + if self._sftp is None: + raise RuntimeError("sftp_open_failed") + self.sftp_ready = True + return self._sftp + + def open_sftp(self) -> Any: + """Open/reuse an SFTP client on this session's SSH transport.""" + with self._sftp_lock: + return self._ensure_sftp_unlocked() + + def open_ephemeral_sftp(self) -> Any: + """Open a dedicated SFTP channel for one operation; caller must ``close()`` it. + + Only holds ``_sftp_lock`` briefly while resolving the SSH transport, so long + list/upload/download work does not block other SFTP ops on the same session. + """ + import paramiko + + with self._sftp_lock: + transport = self._ssh_transport_unlocked() + # Keep probe client warm for UI sftp_ready without sharing it for I/O. + try: + self._ensure_sftp_unlocked() + except Exception: + pass + sftp = paramiko.SFTPClient.from_transport(transport) + if sftp is None: + raise RuntimeError("sftp_open_failed") + return sftp + + def run_sftp(self, fn: Any) -> Any: + """Run ``fn(sftp)`` on an ephemeral channel (does not hold the lock during ``fn``).""" + sftp = self.open_ephemeral_sftp() + try: + return fn(sftp) + finally: + try: + sftp.close() + except Exception: + pass + + def try_attach_sftp(self) -> bool: + """Best-effort SFTP channel open after SSH login (does not fail the shell).""" + if str(self.protocol or "ssh").lower() != "ssh" or self.cli_hop_guard: + self.sftp_ready = False + return False + try: + self.open_sftp() + self.sftp_ready = True + return True + except Exception: + self.sftp_ready = False + _log.debug("webcrt sftp attach skipped session=%s", self.session_id, exc_info=True) + return False + + def open_session_log(self) -> None: + if not bool(getattr(settings, "webcrt_session_log_enabled", True)): + return + if self._log_fh is not None: + return + try: + self._log_fh = _session_log_path(self.session_id).open("a", encoding="utf-8", errors="replace") + self._log_fh.write(f"# session={self.session_id} ne={self.ne_id} ip={self.ne_ip} ts={_utc_iso()}\n") + self._log_fh.flush() + except Exception: + _log.debug("webcrt session log open failed", exc_info=True) + self._log_fh = None + + def append_session_log(self, text: str) -> None: + if not text or self._log_fh is None: + return + try: + self._log_fh.write(text) + self._log_fh.flush() + except Exception: + pass + + def close_session_log(self) -> None: + fh = self._log_fh + self._log_fh = None + if fh is None: + return + try: + fh.write(f"\n# closed reason={self.close_reason} ts={_utc_iso()}\n") + fh.close() + except Exception: + pass + + def take_stdout(self, attach_gen: int, *, timeout: float = 0.25) -> bytes | None | str: + """Exclusive stdout take for one WS attach generation. + + Returns: + bytes — device output chunk + None — device reader closed (session end) + \"stale\" — a newer WebSocket owns this session; caller must stop + \"empty\" — no data within timeout (keep polling) + """ + deadline = time.time() + max(0.05, float(timeout)) + while True: + with self._stdout_lock: + if attach_gen != self.attach_gen: + return "stale" + remaining = deadline - time.time() + if remaining <= 0: + return "empty" + # Slice waits so we can notice attach_gen bumps without busy-spinning. + try: + chunk = self.out_queue.get(timeout=min(0.05, remaining)) + except queue.Empty: + continue + with self._stdout_lock: + if attach_gen != self.attach_gen: + # Put back including EOF sentinel so the new owner still sees close. + self.out_queue.put(chunk) + return "stale" + return chunk # bytes | None + + def write_stdin(self, data: str) -> None: + if self.closed or self.conn is None: + raise RuntimeError("session_closed") + text = str(data or "") + if not text: + return + if self.cli_keymap: + text = map_network_cli_keys( + text, + device_type=self.device_type, + vendor=self.vendor, + protocol=self.protocol, + ) + text = map_network_cli_enter(text, self.conn) + if not text: + return + with self._write_lock: + # Prefer raw channel I/O for interactive typing (char echo / backspace). + channel = getattr(self.conn, "remote_conn", None) + try: + if channel is not None and hasattr(channel, "send") and callable(channel.send): + payload = _encode_text(text, self.encoding) + # Paramiko may write partially when the window is full. + view = memoryview(payload) + while len(view): + n = int(channel.send(view) or 0) + if n <= 0: + time.sleep(0.01) + continue + view = view[n:] + self.bytes_in += len(payload) + elif channel is not None and hasattr(channel, "write") and callable(channel.write): + payload = _encode_text(text, self.encoding) + channel.write(payload) + self.bytes_in += len(payload) + else: + self.conn.write_channel(text) + self.bytes_in += len(text) + except Exception: + self.conn.write_channel(text) + self.bytes_in += len(text) + self.touch() + + def send_break(self) -> None: + """Send SSH break / Telnet IAC BREAK to interrupt paging or hung commands.""" + if self.closed or self.conn is None: + raise RuntimeError("session_closed") + channel = getattr(self.conn, "remote_conn", None) + with self._write_lock: + sent = False + if channel is not None and hasattr(channel, "send_break") and callable(channel.send_break): + try: + channel.send_break(0) + sent = True + except Exception: + _log.debug("send_break failed session=%s", self.session_id, exc_info=True) + if not sent and channel is not None and hasattr(channel, "send") and callable(channel.send): + # Telnet IAC BREAK = 255 243 + try: + channel.send(b"\xff\xf3") + sent = True + except Exception: + pass + if not sent: + # Fallback: Ctrl-C often interrupts device CLI more-pages. + try: + self.conn.write_channel("\x03") + except Exception: + raise RuntimeError("break_failed") + self.touch() + + def resize(self, cols: int, rows: int) -> None: + if self.closed or self.conn is None: + return + c = max(20, min(500, int(cols or 80))) + r = max(5, min(200, int(rows or 24))) + self.cols = c + self.rows = r + channel = getattr(self.conn, "remote_conn", None) + if channel is not None and hasattr(channel, "resize_pty"): + try: + channel.resize_pty(width=c, height=r) + except Exception: + _log.debug("resize_pty failed session=%s", self.session_id, exc_info=True) + self.touch() + + def start_reader(self) -> None: + if self._reader and self._reader.is_alive(): + return + self._reader = threading.Thread( + target=self._reader_loop, + name=f"webcrt-reader-{self.session_id[:8]}", + daemon=True, + ) + self._reader.start() + + def _reader_loop(self) -> None: + conn = self.conn + if conn is None: + self.out_queue.put(None) + return + channel = getattr(conn, "remote_conn", None) + hop_return = False + poll = max(0.002, float(getattr(settings, "webcrt_reader_poll_sec", 0.01) or 0.01)) + try: + while not self.closed: + chunk = b"" + try: + if channel is not None and hasattr(channel, "recv_ready") and hasattr(channel, "recv"): + # Paramiko SSH: prefer short blocking recv over fixed spin-sleep. + ready = False + try: + ready = bool(channel.recv_ready()) + except Exception: + ready = False + if ready: + chunk = channel.recv(16384) + if not chunk: + break + elif hasattr(channel, "exit_status_ready") and channel.exit_status_ready(): + break + else: + # Brief block: settimeout + recv wakes sooner than sleep(0.04). + prev_timeout = None + try: + prev_timeout = channel.gettimeout() + except Exception: + prev_timeout = None + try: + channel.settimeout(poll) + chunk = channel.recv(16384) + except Exception: + chunk = b"" + finally: + try: + channel.settimeout(prev_timeout) + except Exception: + pass + if not chunk: + continue + elif channel is not None and hasattr(channel, "read_very_eager"): + # Telnet: do NOT use conn.read_channel() — Netmiko strips ANSI. + data = channel.read_very_eager() + if data: + chunk = ( + data + if isinstance(data, (bytes, bytearray)) + else _encode_text(str(data), self.encoding) + ) + else: + time.sleep(poll) + continue + else: + text = conn.read_channel() + if text: + chunk = _encode_text(str(text), self.encoding) + else: + time.sleep(poll) + continue + except Exception as exc: + if self.closed: + break + _log.debug("webcrt reader error session=%s: %s", self.session_id, exc) + time.sleep(0.05) + continue + if chunk: + self.touch() + self.bytes_out += len(chunk) + self.out_queue.put(chunk) + try: + self.append_session_log(_decode_bytes(chunk, self.encoding)) + except Exception: + pass + if self.cli_hop_guard and self._note_cli_hop_output(chunk): + hop_return = True + notice = ( + "\r\n*** WebCRT: 目标会话已结束,已断开代理连接 " + "(target session ended; closing hop proxy) ***\r\n" + ) + self.out_queue.put(notice.encode("utf-8", errors="replace")) + break + finally: + if hop_return and not self.closed: + try: + from .webcrt_session_registry import close_session + + close_session(self.session_id, reason="cli_hop_return") + except Exception: + self.close("cli_hop_return") + self.out_queue.put(None) + + def _note_cli_hop_output(self, chunk: bytes) -> bool: + """Accumulate stdout and return True when nested CLI hop has returned to proxy.""" + try: + text = _decode_bytes(chunk, self.encoding) + except Exception: + text = str(chunk) + self._hop_scan_buf = (self._hop_scan_buf + text)[-12000:] + marker = str(self.cli_hop_prompt or "").strip() + last = extract_cli_prompt_marker(self._hop_scan_buf) + if last and (not marker or last != marker): + self._cli_hop_seen_other_prompt = True + return should_close_cli_hop_session( + self._hop_scan_buf, + self.cli_hop_prompt, + seen_other_prompt=self._cli_hop_seen_other_prompt, + ) + + def run_post_login_commands(self) -> None: + cmds = [str(c).rstrip("\r\n") for c in (self.post_login_commands or []) if str(c).strip()] + if not cmds or self.closed or self.conn is None: + return + for cmd in cmds[:20]: + try: + self.write_stdin(cmd + "\r") + time.sleep(0.15) + except Exception: + _log.debug("post_login command failed session=%s", self.session_id, exc_info=True) + break + + def close(self, reason: str = "closed") -> None: + if self.closed: + return + self.closed = True + self.state = "closed" + self.close_reason = reason or "closed" + self._ready_event.set() + self.close_sftp() + try: + close_netmiko_connection(self.conn) + except Exception: + pass + self.conn = None + try: + self.out_queue.put_nowait(None) + except Exception: + pass + self.close_session_log() diff --git a/netx_api/webcrt_session_registry.py b/netx_api/webcrt_session_registry.py new file mode 100644 index 0000000..4c4918d --- /dev/null +++ b/netx_api/webcrt_session_registry.py @@ -0,0 +1,637 @@ +"""WebCRT process-local session registry, reaper, and connect lifecycle.""" +from __future__ import annotations + +import io +import logging +import queue +import threading +import time +import uuid +from datetime import datetime, timezone +from typing import Any + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .config import settings +from .ne_crypto import CredentialCryptoError +from .ne_session_factory import ( + get_cli_hop_guard, + open_netmiko_connection, +) +from .webcrt_channel import ( + _audit, + _capture_raw_channel, + _decode_bytes, + _drain_channel, + _encode_text, + _is_prompt_only_echo, + _looks_like_cli_prompt, + _looks_like_login_prompt, + _looks_like_password_change_prompt, + _normalize_encoding, + _prime_interactive_channel, + _session_log_text, + prepare_bootstrap_output, + uses_network_cli_keymap, +) +from .webcrt_session_model import WebcrtSession + +_log = logging.getLogger("netx.webcrt") + +_sessions_lock = threading.Lock() +_sessions: dict[str, WebcrtSession] = {} +_reaper_started = False + +def _ensure_reaper() -> None: + global _reaper_started + with _sessions_lock: + if _reaper_started: + return + _reaper_started = True + t = threading.Thread(target=_reaper_loop, name="webcrt-reaper", daemon=True) + t.start() + + +def _reaper_loop() -> None: + while True: + try: + _reap_sessions() + except Exception: + _log.exception("webcrt reaper failed") + time.sleep(2) + + +def _reap_sessions() -> None: + idle = max(60, int(settings.webcrt_idle_timeout_sec or 1800)) + attach = max(10, int(settings.webcrt_attach_timeout_sec or 60)) + anti_idle = max(0, int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0)) + anti_payload = str(getattr(settings, "webcrt_anti_idle_payload", " ") or " ") + now = time.time() + to_close: list[tuple[WebcrtSession, str]] = [] + to_nudge: list[WebcrtSession] = [] + with _sessions_lock: + for sess in list(_sessions.values()): + if sess.closed: + _sessions.pop(sess.session_id, None) + continue + if sess.state == "connecting": + # Connecting sessions use connect timeout, not attach timeout alone. + connect_budget = max(30, int(settings.webcrt_connect_timeout_sec or 90)) + 30 + if (now - sess.connect_started_at) > connect_budget: + to_close.append((sess, "connect_timeout")) + continue + if sess.attached: + if (now - sess.last_activity) > idle: + to_close.append((sess, "idle_timeout")) + elif ( + anti_idle > 0 + and sess.state == "ready" + and sess.conn is not None + and (now - sess.last_activity) >= anti_idle + ): + to_nudge.append(sess) + continue + # Not attached: either never attached, or briefly detached for reconnect. + if sess.detach_deadline is not None: + if now >= sess.detach_deadline: + to_close.append((sess, "detach_timeout")) + else: + # Start attach clock after connect finishes (not HTTP create time), + # so slow auth + UI mount does not race attach_timeout. + anchor = float(sess.connect_finished_at or sess.created_at or now) + if (now - anchor) > attach: + to_close.append((sess, "attach_timeout")) + elif (now - sess.last_activity) > idle: + to_close.append((sess, "idle_timeout")) + for sess in to_nudge: + try: + # Touch without changing visible prompt when payload is empty/null-ish. + payload = anti_payload + if payload == "\\0": + payload = "\x00" + if payload: + sess.write_stdin(payload) + else: + sess.touch() + except Exception: + _log.debug("webcrt anti-idle failed session=%s", sess.session_id, exc_info=True) + for sess, reason in to_close: + close_session(sess.session_id, reason=reason) + + +def active_session_count() -> int: + with _sessions_lock: + return sum(1 for s in _sessions.values() if not s.closed) + + +def get_session(session_id: str) -> WebcrtSession | None: + with _sessions_lock: + sess = _sessions.get(session_id) + if sess is None or sess.closed: + return None + return sess + + +def find_ssh_session_for_ne(ne_id: str) -> WebcrtSession | None: + """Prefer a ready/attached interactive SSH session for SFTP channel reuse.""" + nid = str(ne_id or "").strip() + if not nid: + return None + with _sessions_lock: + candidates = [ + s + for s in _sessions.values() + if (not s.closed) + and str(s.ne_id) == nid + and str(s.protocol or "ssh").lower() == "ssh" + and s.conn is not None + and s.state in ("ready", "connecting") + and not s.cli_hop_guard + ] + if not candidates: + return None + # Prefer attached + ready sessions. + candidates.sort(key=lambda s: (0 if s.attached and s.state == "ready" else 1, -s.last_activity)) + return candidates[0] + + +def wait_session_ready(session_id: str, *, timeout: float = 120.0) -> WebcrtSession: + """Block until async connect finishes (ready or error). Used by tests and WS.""" + deadline = time.time() + max(1.0, float(timeout)) + while time.time() < deadline: + sess = get_session(session_id) + if sess is None: + raise HTTPException(status_code=404, detail="webcrt_session_not_found") + if sess.state == "ready": + return sess + if sess.state == "error": + raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") + sess._ready_event.wait(timeout=0.25) + raise HTTPException(status_code=504, detail="connect_timeout") + + +def _webcrt_creds_ready(creds: dict[str, Any]) -> bool: + """True when WebCRT can open a session with the resolved credentials. + + Bastion-managed hops store the target password on the bastion side, so an empty + NE password is valid (same as connectivity test). Direct / manual / Linux hops + still require a target password for SSH. + + Telnet (no hop) allows empty username/password so the user can authenticate + interactively in the terminal (SecureCRT-style). + """ + hop_enabled = bool(creds.get("hop_enabled")) + hop_vendor = str(creds.get("hop_vendor") or "").strip().lower() + auth_mode = str(creds.get("hop_target_auth_mode") or "bastion_managed").strip().lower() + protocol = str(creds.get("protocol") or "ssh").strip().lower() + if hop_enabled and hop_vendor == "bastion" and auth_mode == "bastion_managed": + return bool( + str(creds.get("hop_host") or "").strip() + and str(creds.get("hop_username") or "").strip() + and str(creds.get("hop_password") or "") + ) + if protocol == "telnet" and not hop_enabled: + return True + if not str(creds.get("username") or "").strip(): + return False + return bool(str(creds.get("password") or "")) + + +def _finish_connect( + sess: WebcrtSession, + *, + creds: dict[str, Any], + device: dict[str, Any], + connect_timeout: int, + client: str, +) -> None: + log_buf = io.BytesIO() + try: + conn = open_netmiko_connection( + creds, + session_timeout=connect_timeout, + session_log=log_buf, + cols=sess.cols, + rows=sess.rows, + interactive=True, + keepalive=int(sess.keepalive_sec or 0), + ) + except Exception as exc: + partial = _session_log_text(log_buf).strip() + from .ne_cli_errors import format_cli_failure + + classified = format_cli_failure(exc, partial) + detail = f"connect_failed:{classified}" + if partial: + detail = f"{detail}\n--- device transcript ---\n{partial[-4000:]}" + sess.state = "error" + sess.connect_error = detail + sess.connect_finished_at = time.time() + sess._ready_event.set() + _audit( + "session_open_failed", + session_id=sess.session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + source=str(device.get("source") or ""), + client=client or "", + error=str(exc)[:500], + transcript_len=len(partial), + ) + return + + channel = getattr(conn, "remote_conn", None) + if channel is not None and hasattr(channel, "resize_pty"): + try: + channel.resize_pty(width=sess.cols, height=sess.rows) + except Exception: + pass + + pre_log = _session_log_text(log_buf) + # Pull post-auth banner/MOTD from the PTY. With interactive no-op session_preparation + # (generic_termserver), Netmiko session_log is often empty — do not discard these bytes. + try: + early = _capture_raw_channel(conn, duration=0.35) + except Exception: + early = "" + seed = f"{pre_log}{early}" + already_prompted = _looks_like_cli_prompt(seed) + primed = "" + # Do not send Enter at Username:/Password: or Huawei password-change [Y/N]: + # (Netmiko telnet_login already answers password-change with "N"). + if _looks_like_login_prompt(seed) or _looks_like_password_change_prompt(seed): + try: + primed = _capture_raw_channel(conn, duration=0.9) + except Exception: + primed = "" + else: + try: + primed = _prime_interactive_channel(conn, already_prompted=already_prompted) + except Exception: + primed = "" + combined = f"{seed}{primed}" + # Final settle: keep stragglers in bootstrap (normalize collapses duplicate prompts). + try: + combined += _capture_raw_channel(conn, duration=0.35) + except Exception: + pass + if not str(combined).strip(): + try: + combined = _drain_channel(conn, rounds=6, wait=0.08) + except Exception: + combined = "" + bootstrap = prepare_bootstrap_output(combined) + # Discard lone punctuation left on the wire (would glue onto ```` in xterm). + try: + leftover = _capture_raw_channel(conn, duration=0.12) + except Exception: + leftover = "" + if leftover and leftover.strip() not in {":", ">", "#", "]", "$"}: + bootstrap = prepare_bootstrap_output(f"{bootstrap}{leftover}") + + hop_guard = get_cli_hop_guard(conn) + sess.conn = conn + sess.cli_hop_guard = bool(hop_guard) + sess.cli_hop_prompt = str((hop_guard or {}).get("hop_prompt") or "") + sess.bootstrap_output = _encode_text(str(bootstrap or ""), sess.encoding) + # Nudge Enter on WS attach only when we still need a shell prompt. + # Never when already at CLI prompt or Username:/Password: (would empty-submit login). + sess.needs_live_prompt = ( + not _looks_like_cli_prompt(bootstrap) and not _looks_like_login_prompt(bootstrap) + ) + sess.open_session_log() + if bootstrap: + sess.append_session_log(bootstrap if bootstrap.endswith("\n") else bootstrap + "\n") + sess.start_reader() + # Drop late prompt echoes that race into the queue right after reader start. + prompt_hint = "" + if bootstrap: + prompt_hint = str(bootstrap).replace("\r\n", "\n").replace("\r", "\n").strip().split("\n")[-1].strip() + settle_deadline = time.time() + 0.45 + while time.time() < settle_deadline: + try: + chunk = sess.out_queue.get_nowait() + except queue.Empty: + time.sleep(0.02) + continue + if chunk is None: + sess.out_queue.put(None) + break + try: + text = _decode_bytes(chunk, sess.encoding) + except Exception: + text = "" + if _is_prompt_only_echo(text, prompt_hint): + continue + # Non-prompt data: put back and stop settling. + sess.out_queue.put(chunk) + break + try: + sess.run_post_login_commands() + except Exception: + _log.debug("post_login failed session=%s", sess.session_id, exc_info=True) + # Same SSH transport: open SFTP channel when the device supports it. + sftp_ok = sess.try_attach_sftp() + sess.state = "ready" + sess.connect_finished_at = time.time() + sess._ready_event.set() + elapsed_ms = int((sess.connect_finished_at - sess.connect_started_at) * 1000) + _audit( + "session_created", + session_id=sess.session_id, + ne_id=sess.ne_id, + ne_name=sess.ne_name, + ne_ip=sess.ne_ip, + protocol=sess.protocol, + encoding=sess.encoding, + source=str(device.get("source") or ""), + hop_enabled=bool(creds.get("hop_enabled")), + hop_vendor=str(creds.get("hop_vendor") or "") if creds.get("hop_enabled") else "", + cli_hop_guard=bool(hop_guard), + cli_hop_prompt=str((hop_guard or {}).get("hop_prompt") or ""), + sftp_ready=bool(sftp_ok), + client=client or "", + connect_ms=elapsed_ms, + active=active_session_count(), + ) + + +def create_session( + db: Session, + *, + ne_id: str | None = None, + ume_ne_id: str | None = None, + cols: int = 80, + rows: int = 24, + client: str = "", + encoding: str = "utf-8", + keepalive_sec: int | None = None, + post_login_commands: list[str] | None = None, + async_connect: bool = True, + username_override: str | None = None, + password_override: str | None = None, +) -> dict[str, Any]: + from .cli_resolve import resolve_cli_target + + _ensure_reaper() + max_sessions = max(1, int(settings.webcrt_max_sessions or 20)) + if active_session_count() >= max_sessions: + raise HTTPException(status_code=429, detail="webcrt_session_limit") + + mid = str(ne_id or "").strip() + uid = str(ume_ne_id or "").strip() + try: + creds, device = resolve_cli_target(db, managed_ne_id=mid or None, ume_ne_id=uid or None) + except HTTPException: + raise + except CredentialCryptoError as exc: + raise HTTPException(status_code=400, detail=str(exc) or "credential_crypto_error") from exc + except Exception as exc: + raise HTTPException(status_code=400, detail=f"credential_error:{exc}") from exc + + # One-shot credentials for SecureCRT-style "do not save password" / retry. + if username_override is not None and str(username_override).strip(): + creds["username"] = str(username_override).strip() + if password_override is not None: + creds["password"] = str(password_override) + + protocol = str(device.get("protocol") or creds.get("protocol") or "ssh").strip().lower() + creds["protocol"] = protocol + # Netmiko telnet drivers dislike a completely missing username; use a placeholder + # for the wire only (interactive login still happens in the terminal). + if protocol == "telnet" and not bool(creds.get("hop_enabled")) and not str(creds.get("username") or "").strip(): + creds["username"] = "telnet" + + if not _webcrt_creds_ready(creds): + raise HTTPException(status_code=400, detail="credentials_incomplete") + + session_id = str(uuid.uuid4()) + c = max(20, min(500, int(cols or 80))) + r = max(5, min(200, int(rows or 24))) + connect_timeout = max(30, int(settings.webcrt_connect_timeout_sec or 90)) + target_id = str(device.get("id") or mid or uid) + target_ip = str(device.get("ip_address") or "") + target_name = str(device.get("name") or target_ip) + device_type = str(device.get("device_type") or creds.get("device_type") or "") + vendor = str(device.get("vendor") or creds.get("vendor") or "") + cli_keymap = uses_network_cli_keymap(device_type, vendor) + enc = _normalize_encoding(encoding) + if keepalive_sec is None: + ka = max(0, int(getattr(settings, "webcrt_keepalive_sec", 0) or 0)) + else: + ka = max(0, min(600, int(keepalive_sec))) + + sess = WebcrtSession( + session_id=session_id, + ne_id=target_id, + ne_name=target_name, + ne_ip=target_ip, + protocol=protocol, + cols=c, + rows=r, + device_type=device_type, + vendor=vendor, + cli_keymap=cli_keymap, + encoding=enc, + keepalive_sec=ka, + state="connecting", + post_login_commands=list(post_login_commands or [])[:20], + ) + with _sessions_lock: + _sessions[session_id] = sess + + _audit( + "session_connecting", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + protocol=sess.protocol, + encoding=enc, + client=client or "", + async_connect=bool(async_connect), + ) + + if async_connect: + t = threading.Thread( + target=_finish_connect, + kwargs={ + "sess": sess, + "creds": creds, + "device": device, + "connect_timeout": connect_timeout, + "client": client or "", + }, + name=f"webcrt-connect-{session_id[:8]}", + daemon=True, + ) + t.start() + else: + _finish_connect( + sess, + creds=creds, + device=device, + connect_timeout=connect_timeout, + client=client or "", + ) + if sess.state == "error": + with _sessions_lock: + _sessions.pop(session_id, None) + raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") + + return { + "session_id": session_id, + "ne_id": sess.ne_id, + "ne_name": sess.ne_name, + "ne_ip": sess.ne_ip, + "source": str(device.get("source") or ""), + "protocol": sess.protocol, + "cols": sess.cols, + "rows": sess.rows, + "encoding": enc, + "keepalive_sec": ka, + "state": sess.state, + "ws_path": f"/v1/webcrt/sessions/{session_id}/ws", + "cli_hop": bool(sess.cli_hop_guard), + "sftp_ready": bool(sess.sftp_ready), + } + + +def mark_attached(session_id: str) -> tuple[WebcrtSession, int]: + sess = get_session(session_id) + if sess is None: + raise HTTPException(status_code=404, detail="webcrt_session_not_found") + # Allow re-attach after brief WS drop (React StrictMode remount / network blip). + # Bump generation so the previous WS pump stops and does not steal echo bytes. + sess.attach_gen += 1 + attach_gen = sess.attach_gen + sess.attached = True + sess.detach_deadline = None + sess.touch() + _audit( + "session_attached", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + attach_gen=attach_gen, + state=sess.state, + ) + return sess, attach_gen + + +def detach_session( + session_id: str, + *, + grace_sec: float = 8.0, + client: str = "", + attach_gen: int | None = None, +) -> dict[str, Any]: + """Mark session unattached but keep device channel open briefly for reconnect.""" + sess = get_session(session_id) + if sess is None: + return {"ok": True, "session_id": session_id, "detached": False} + if sess.closed: + return {"ok": True, "session_id": session_id, "detached": False} + # Ignore detach from an older StrictMode WS once a newer attach owns the session. + if attach_gen is not None and attach_gen != sess.attach_gen: + return { + "ok": True, + "session_id": session_id, + "detached": False, + "ignored_stale_attach": True, + } + sess.attached = False + sess.detach_deadline = time.time() + max(1.0, float(grace_sec)) + sess.touch() + _audit( + "session_detached", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + grace_sec=grace_sec, + client=client or "", + attach_gen=attach_gen, + ) + return {"ok": True, "session_id": session_id, "detached": True} + + +def close_session(session_id: str, *, reason: str = "closed", client: str = "") -> dict[str, Any]: + with _sessions_lock: + sess = _sessions.pop(session_id, None) + if sess is None: + return {"ok": True, "session_id": session_id, "closed": False} + if not sess.closed: + sess.close(reason) + _audit( + "session_closed", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + reason=reason, + client=client or "", + bytes_in=sess.bytes_in, + bytes_out=sess.bytes_out, + queue_dropped=getattr(sess.out_queue, "dropped", 0), + active=active_session_count(), + ) + return {"ok": True, "session_id": session_id, "closed": True, "reason": reason} + + +def list_sessions() -> dict[str, Any]: + with _sessions_lock: + items = [] + for s in _sessions.values(): + if s.closed: + continue + state = str(s.state or "unknown") + attached = bool(s.attached) + # Lifecycle for ops UI: distinguish login vs live vs grace-period detach. + if state == "connecting": + lifecycle = "connecting" + elif state == "error": + lifecycle = "error" + elif state == "ready" and attached: + lifecycle = "ready" + elif state == "ready" and not attached: + lifecycle = "detached" + else: + lifecycle = state + elapsed_ms = None + if state == "connecting": + elapsed_ms = int(max(0.0, time.time() - float(s.connect_started_at or time.time())) * 1000) + items.append( + { + "session_id": s.session_id, + "ne_id": s.ne_id, + "ne_name": s.ne_name, + "ne_ip": s.ne_ip, + "protocol": s.protocol, + "encoding": s.encoding, + "keepalive_sec": int(s.keepalive_sec or 0), + "state": state, + "lifecycle": lifecycle, + "attached": attached, + "detach_deadline": s.detach_deadline, + "connect_error": str(s.connect_error or "")[:500], + "elapsed_ms": elapsed_ms, + "created_at": datetime.fromtimestamp(s.created_at, tz=timezone.utc).isoformat(), + "last_activity": datetime.fromtimestamp(s.last_activity, tz=timezone.utc).isoformat(), + "bytes_in": s.bytes_in, + "bytes_out": s.bytes_out, + "queue_depth": s.out_queue.qsize(), + "queue_dropped": getattr(s.out_queue, "dropped", 0), + "connect_ms": ( + int((s.connect_finished_at - s.connect_started_at) * 1000) + if s.connect_finished_at + else None + ), + } + ) + return { + "total": len(items), + "max_sessions": max(1, int(settings.webcrt_max_sessions or 20)), + "idle_timeout_sec": max(60, int(settings.webcrt_idle_timeout_sec or 1800)), + "keepalive_sec": int(getattr(settings, "webcrt_keepalive_sec", 0) or 0), + "anti_idle_sec": int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0), + "items": items, + } diff --git a/tests/test_webcrt.py b/tests/test_webcrt.py index d6c02d9..cb70437 100644 --- a/tests/test_webcrt.py +++ b/tests/test_webcrt.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock, patch from fastapi import HTTPException from netx_api import webcrt_service as svc +from netx_api import webcrt_session_registry as reg class _FakeConn: @@ -127,8 +128,8 @@ class WebcrtServiceTests(unittest.TestCase): self.assertIn("R2#", text) self.assertNotIn("MagicMock", text) - @patch.object(svc, "_audit") - @patch.object(svc, "open_netmiko_connection") + @patch.object(reg, "_audit") + @patch.object(reg, "open_netmiko_connection") @patch("netx_api.cli_resolve.resolve_cli_target") def test_bootstrap_from_channel_when_session_log_empty( self, @@ -170,8 +171,8 @@ class WebcrtServiceTests(unittest.TestCase): self.assertIn("R2#", boot) svc.close_session(out["session_id"], reason="test") - @patch.object(svc, "_audit") - @patch.object(svc, "open_netmiko_connection") + @patch.object(reg, "_audit") + @patch.object(reg, "open_netmiko_connection") @patch("netx_api.cli_resolve.resolve_cli_target") def test_create_session_password_override( self, @@ -215,8 +216,8 @@ class WebcrtServiceTests(unittest.TestCase): self.assertEqual(called_creds.get("password"), "once") svc.close_session(out["session_id"], reason="test") - @patch.object(svc, "_audit") - @patch.object(svc, "open_netmiko_connection") + @patch.object(reg, "_audit") + @patch.object(reg, "open_netmiko_connection") @patch("netx_api.cli_resolve.resolve_cli_target") def test_create_session_limit( self, @@ -252,8 +253,8 @@ class WebcrtServiceTests(unittest.TestCase): svc.create_session(db, ne_id="ne-a", cols=80, rows=24, client="test", async_connect=False) self.assertEqual(ctx.exception.status_code, 429) - @patch.object(svc, "_audit") - @patch.object(svc, "open_netmiko_connection") + @patch.object(reg, "_audit") + @patch.object(reg, "open_netmiko_connection") @patch("netx_api.cli_resolve.resolve_cli_target") def test_create_session_passes_hop_creds( self, @@ -319,9 +320,9 @@ class WebcrtServiceTests(unittest.TestCase): self.assertEqual(fake.written[len(before) :], ["\n"]) svc.close_session(out["session_id"], reason="test") - @patch.object(svc, "_audit") - @patch.object(svc, "get_cli_hop_guard") - @patch.object(svc, "open_netmiko_connection") + @patch.object(reg, "_audit") + @patch.object(reg, "get_cli_hop_guard") + @patch.object(reg, "open_netmiko_connection") @patch("netx_api.cli_resolve.resolve_cli_target") def test_create_session_reports_cli_hop( self, @@ -361,8 +362,8 @@ class WebcrtServiceTests(unittest.TestCase): self.assertEqual(sess.cli_hop_prompt, "") svc.close_session(out["session_id"], reason="test") - @patch.object(svc, "_audit") - @patch.object(svc, "open_netmiko_connection") + @patch.object(reg, "_audit") + @patch.object(reg, "open_netmiko_connection") @patch("netx_api.cli_resolve.resolve_cli_target") def test_create_session_bastion_managed_without_target_password( self, @@ -451,7 +452,7 @@ class WebcrtServiceTests(unittest.TestCase): ) ) - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_attach_gen_exclusive_stdout_and_stale_detach(self, _mock_audit: MagicMock) -> None: conn = _FakeConn() sess = svc.WebcrtSession( @@ -490,7 +491,7 @@ class WebcrtServiceTests(unittest.TestCase): self.assertFalse(sess.attached) svc.close_session("race", reason="test") - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_detach_grace_keeps_session_until_deadline(self, _mock_audit: MagicMock) -> None: conn = _FakeConn() sess = svc.WebcrtSession( @@ -524,7 +525,7 @@ class WebcrtServiceTests(unittest.TestCase): svc._reap_sessions() self.assertIsNone(svc.get_session("grace")) - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_list_sessions_lifecycle_ready_vs_detached(self, _mock_audit: MagicMock) -> None: conn = _FakeConn() sess = svc.WebcrtSession( @@ -560,7 +561,7 @@ class WebcrtServiceTests(unittest.TestCase): svc.close_session("life", reason="test") self.assertEqual(svc.list_sessions()["total"], 0) - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_attach_timeout_reaper(self, _mock_audit: MagicMock) -> None: conn = _FakeConn() sess = svc.WebcrtSession( @@ -574,6 +575,7 @@ class WebcrtServiceTests(unittest.TestCase): conn=conn, # type: ignore[arg-type] ) sess.created_at = time.time() - 120 + sess.state = "ready" with svc._sessions_lock: svc._sessions["stale"] = sess with patch.object(svc.settings, "webcrt_attach_timeout_sec", 30): @@ -581,7 +583,7 @@ class WebcrtServiceTests(unittest.TestCase): svc._reap_sessions() self.assertIsNone(svc.get_session("stale")) - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_attach_timeout_uses_connect_finished_at(self, _mock_audit: MagicMock) -> None: """Slow connect should not burn the attach window from HTTP create time.""" conn = _FakeConn() @@ -622,7 +624,7 @@ class WebcrtServiceTests(unittest.TestCase): except OSError: pass - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_idle_timeout_reaper(self, _mock_audit: MagicMock) -> None: conn = _FakeConn() sess = svc.WebcrtSession( @@ -636,6 +638,7 @@ class WebcrtServiceTests(unittest.TestCase): conn=conn, # type: ignore[arg-type] ) sess.attached = True + sess.state = "ready" sess.last_activity = time.time() - 9999 with svc._sessions_lock: svc._sessions["idle"] = sess @@ -644,7 +647,7 @@ class WebcrtServiceTests(unittest.TestCase): svc._reap_sessions() self.assertIsNone(svc.get_session("idle")) - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_cli_hop_return_closes_session(self, _mock_audit: MagicMock) -> None: """Vendor CLI hop: nested target exit must tear down WebCRT (no hop shell).""" conn = _FakeConn() @@ -889,7 +892,7 @@ class WebcrtServiceTests(unittest.TestCase): _parse_chmod_mode("bad") self.assertEqual(cm.exception.detail, "sftp_chmod_invalid_mode") - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_find_ssh_session_for_ne_prefers_attached(self, _mock_audit: MagicMock) -> None: conn = _FakeConn() sess = svc.WebcrtSession( @@ -911,7 +914,7 @@ class WebcrtServiceTests(unittest.TestCase): self.assertIsNone(svc.find_ssh_session_for_ne("other")) svc.close_session("sftpne", reason="test") - @patch.object(svc, "_audit") + @patch.object(reg, "_audit") def test_reattach_clears_detach_deadline(self, _mock_audit: MagicMock) -> None: conn = _FakeConn() sess = svc.WebcrtSession(