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 <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-08-02 17:12:55 +08:00
parent a2f91f6ee2
commit 96b76294b4
10 changed files with 2129 additions and 1980 deletions

309
netx_api/ume_alarm_apply.py Normal file
View file

@ -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) "<net_id>#<suffix>"
# 2) "<net_id>, <x>, <y>"
# 3) "<net_id> <x> <y>"
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)

View file

@ -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

502
netx_api/ume_sync_pull.py Normal file
View file

@ -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)

View file

@ -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) "<net_id>#<suffix>"
# 2) "<net_id>, <x>, <y>"
# 3) "<net_id> <x> <y>"
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",
]

View file

@ -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

View file

@ -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",

File diff suppressed because it is too large Load diff

View file

@ -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()

View file

@ -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 ``<r1>`` 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,
}

View file

@ -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, "<HOP>")
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(