From 136f40cdaeb013c52a787f4908888047527ea665 Mon Sep 17 00:00:00 2001 From: oliver Date: Sun, 2 Aug 2026 16:50:52 +0800 Subject: [PATCH] Split WebCRT and API startup; default collectors to worker process. Extract channel/session modules and CLI guard, move UME sidebands out of main, and default NETX_RUN_INLINE_SCHEDULERS off so ops run python -m netx_api.worker. Co-authored-by: Cursor --- .env.example | 2 + PROD_MIN_CHECKLIST.md | 2 +- netx_api/app_startup.py | 144 +++ netx_api/config.py | 5 +- netx_api/integrations_router.py | 87 +- netx_api/main.py | 603 +----------- netx_api/ne_exec.py | 65 +- netx_api/ne_exec_guard.py | 74 ++ netx_api/ume_runtime.py | 339 ++++++- netx_api/ume_support.py | 1 + netx_api/webcrt_channel.py | 425 +++++++++ netx_api/webcrt_service.py | 1547 +------------------------------ netx_api/webcrt_session.py | 1111 ++++++++++++++++++++++ netx_api/worker.py | 3 +- tests/test_ne_exec.py | 6 + tests/test_schema_patches.py | 5 +- 16 files changed, 2269 insertions(+), 2150 deletions(-) create mode 100644 netx_api/app_startup.py create mode 100644 netx_api/ne_exec_guard.py create mode 100644 netx_api/webcrt_channel.py create mode 100644 netx_api/webcrt_session.py diff --git a/.env.example b/.env.example index d0519a7..780b88b 100644 --- a/.env.example +++ b/.env.example @@ -54,6 +54,8 @@ NETX_UME_NOTIFICATION_TOPIC=ALARM # NETX_ALEMBIC_UPGRADE_ON_START=false # NETX_SKIP_LEGACY_STARTUP_DDL=false # NETX_SQL_READONLY_DATABASE_URL=postgresql+psycopg://netx_ro:xxx@127.0.0.1:5432/netx +# Device collectors: default off in API — run `python -m netx_api.worker`. +# Lab single-process: NETX_RUN_INLINE_SCHEDULERS=true # NETX_RUN_INLINE_SCHEDULERS=true # NETX_AUDIT_ASYNC=true # NETX_AUDIT_SAMPLE_N=1 diff --git a/PROD_MIN_CHECKLIST.md b/PROD_MIN_CHECKLIST.md index b26754e..906ef7c 100644 --- a/PROD_MIN_CHECKLIST.md +++ b/PROD_MIN_CHECKLIST.md @@ -12,7 +12,7 @@ ## Runtime - Ensure PostgreSQL backup policy exists (daily logical backup + retention). - Schema: API auto-runs `alembic upgrade head` on start (see [docs/ALEMBIC.md](docs/ALEMBIC.md)). No manual migrate flag required for normal deploys. -- Optional: `NETX_RUN_INLINE_SCHEDULERS=false` and run `python -m netx_api.worker` for collectors. +- Collectors: default is external worker (`python -m netx_api.worker`). Only set `NETX_RUN_INLINE_SCHEDULERS=true` for single-process lab. Check `/health/ready` → `schedulers.mode`. - Run `oclaw` and `netx` under process managers (systemd/Windows service/pm2 equivalent). - Enable auto-restart and startup-at-boot for both services. diff --git a/netx_api/app_startup.py b/netx_api/app_startup.py new file mode 100644 index 0000000..4469811 --- /dev/null +++ b/netx_api/app_startup.py @@ -0,0 +1,144 @@ +"""API process startup orchestration (schema, recovery, schedulers, UME sidebands).""" + +from __future__ import annotations + +import logging + +from .auth_service import bootstrap_admin_if_needed +from .config import settings +from .db import Base, SessionLocal, engine +from .schema_patches import ( + apply_all_legacy_startup_ddl, + apply_auth_schema_patches, + run_alembic_upgrade_to_head, +) +from .security_bootstrap import assert_secure_defaults_or_exit +from .ume_runtime import start_api_sideband_threads, start_device_schedulers +import netx_api.ume_support as ume_support +from .ume_alarm_ws import ( + begin_startup_alarm_sync_gate, + complete_startup_alarm_sync_gate, +) + +_log = logging.getLogger("netx.ume.schedule") + + +def _configure_ume_diag_logging() -> None: + fmt = logging.Formatter("%(asctime)s %(levelname)s %(name)s: %(message)s") + for name in ("netx.ume.schedule", "netx.ume.sync"): + lg = logging.getLogger(name) + if lg.handlers: + continue + h = logging.StreamHandler() + h.setFormatter(fmt) + lg.addHandler(h) + lg.setLevel(logging.INFO) + lg.propagate = False + + +def run_api_startup() -> None: + """Full API boot sequence previously inlined in ``main.on_startup``.""" + assert_secure_defaults_or_exit() + _configure_ume_diag_logging() + Base.metadata.create_all(bind=engine) + + alembic_ok = True + if bool(getattr(settings, "alembic_upgrade_on_start", True)): + try: + run_alembic_upgrade_to_head() + except Exception: + alembic_ok = False + _log.exception("startup: alembic upgrade head failed") + + skip_ddl = bool(getattr(settings, "skip_legacy_startup_ddl", True)) + try: + with engine.begin() as conn: + apply_auth_schema_patches(conn) + except Exception: + _log.exception("startup: auth schema patches failed") + if skip_ddl and alembic_ok: + _log.info("startup: schema via Alembic (legacy inline DDL skipped)") + else: + if skip_ddl and not alembic_ok: + _log.warning("startup: Alembic failed — falling back to legacy schema patches") + try: + apply_all_legacy_startup_ddl(engine) + except Exception: + _log.exception("startup: legacy schema patches failed") + + ume_support._reset_runtime_pause_flags() + ume_support._fail_stale_running_sync_jobs_on_startup() + try: + from .topology_service import bootstrap_topology_tree, reclaim_stale_discover_jobs + + db_topo = SessionLocal() + try: + bootstrap_topology_tree(db_topo) + closed = reclaim_stale_discover_jobs(db_topo, force_all_open=True) + if closed: + _log.warning("startup: closed %s orphaned topology discover jobs", closed) + finally: + db_topo.close() + except Exception: + _log.exception("startup: topology discover job cleanup failed") + + if ume_support._needs_startup_alarm_sync_before_ws(): + begin_startup_alarm_sync_gate() + _log.info( + "startup: WSS blocked until initial REST current-alarm sync completes (delay=%ss)", + ume_support._startup_alarm_pull_delay_s(), + ) + else: + complete_startup_alarm_sync_gate() + + db = SessionLocal() + try: + try: + bootstrap_admin_if_needed(db) + except Exception: + _log.exception("startup: auth bootstrap admin failed") + from .collection_recovery import recover_collection_jobs_on_startup + + resumed = recover_collection_jobs_on_startup(db) + if resumed: + _log.info("startup: resumed %s pending ne collection runs", resumed) + from .config_sync_recovery import recover_config_sync_on_startup + from .config_sync_service import ensure_policy + from .port_traffic_recovery import recover_port_traffic_on_startup + + ensure_policy(db) + cfg_resumed = recover_config_sync_on_startup(db) + if cfg_resumed: + _log.info("startup: resumed %s config_sync task(s) from interrupted cycle", cfg_resumed) + try: + from .lldp_collect_service import ensure_policy as ensure_lldp_collect_policy + + ensure_lldp_collect_policy(db) + except Exception: + _log.exception("startup: lldp_collect policy ensure failed") + pt_cleared = recover_port_traffic_on_startup(db) + if pt_cleared: + _log.info("startup: cleared %s port_traffic stuck collect_running flag(s)", pt_cleared) + try: + from .port_traffic_migrate import backfill_port_traffic_series + + backfill_port_traffic_series(db) + except Exception: + _log.exception("startup: port_traffic series backfill failed") + except Exception: + _log.exception("startup: ne collection / config_sync recovery failed") + finally: + db.close() + + if bool(getattr(settings, "run_inline_schedulers", False)): + try: + start_device_schedulers() + except Exception: + _log.exception("startup: device schedulers init failed") + else: + _log.info( + "startup: inline schedulers disabled — run `python -m netx_api.worker` for " + "config_sync / lldp_collect / port_traffic" + ) + + start_api_sideband_threads() diff --git a/netx_api/config.py b/netx_api/config.py index c467120..147d536 100644 --- a/netx_api/config.py +++ b/netx_api/config.py @@ -140,8 +140,9 @@ class Settings(BaseSettings): alembic_upgrade_on_start: bool = True # Optional dedicated SQLAlchemy URL for /v1/sql/* (read-only DB role recommended). sql_readonly_database_url: str = "" - # When false, API skips config_sync / lldp / port_traffic schedulers (run `python -m netx_api.worker`). - run_inline_schedulers: bool = True + # When false (default), API skips config_sync / lldp / port_traffic schedulers — + # run `python -m netx_api.worker` alongside the API. Set true only for single-process lab. + run_inline_schedulers: bool = False settings = Settings() diff --git a/netx_api/integrations_router.py b/netx_api/integrations_router.py index fa1430e..9eb27db 100644 --- a/netx_api/integrations_router.py +++ b/netx_api/integrations_router.py @@ -2,13 +2,16 @@ from __future__ import annotations +import time from typing import Any from fastapi import APIRouter, Depends from sqlalchemy import text as sql_text from sqlalchemy.orm import Session +from .config import settings from .db import get_db +from .oclaw_alarm_forwarder import forwarder_status router = APIRouter(tags=["health"]) @@ -21,9 +24,87 @@ def health_live() -> dict[str, str]: @router.get("/health/ready", status_code=200) def health_ready(db: Session = Depends(get_db)) -> dict[str, Any]: - """Readiness — verifies database connectivity.""" + """Readiness — DB plus scheduler deployment hint.""" + out: dict[str, Any] = {"status": "ok", "probe": "ready"} try: db.execute(sql_text("select 1")) - return {"status": "ok", "probe": "ready", "db": "up"} + out["db"] = "up" except Exception as exc: - return {"status": "down", "probe": "ready", "db": "down", "error": str(exc)[:240]} + return { + "status": "down", + "probe": "ready", + "db": "down", + "error": str(exc)[:240], + } + inline = bool(getattr(settings, "run_inline_schedulers", False)) + out["schedulers"] = { + "inline": inline, + "mode": "inline" if inline else "external_worker", + "hint": None + if inline + else "run `python -m netx_api.worker` for config_sync / lldp_collect / port_traffic", + } + return out + + +@router.get("/v1/integrations/status") +def integrations_status(db: Session = Depends(get_db)) -> dict: + """netx API + DB + oclaw bridge status.""" + netx_api = {"status": "up"} + + db_status: dict = {"status": "unknown"} + try: + t0 = time.monotonic() + db.execute(sql_text("select 1")) + db_status = {"status": "up", "latency_ms": int((time.monotonic() - t0) * 1000)} + except Exception as exc: + db_status = {"status": "down", "error": str(exc)[:240]} + + oclaw_status: dict = {"status": "unknown"} + fwd = forwarder_status() + if not bool(fwd.get("enabled")): + oclaw_status = { + "status": "unknown", + "mode": "ws", + "enabled": False, + "connected": False, + "error_kind": "disabled", + "error": "NETX_OCLAW_ALARM_WS_ENABLED=false or missing token/url", + "forwarder": fwd, + } + elif bool(fwd.get("paused")): + oclaw_status = { + "status": "unknown", + "mode": "ws", + "enabled": True, + "connected": False, + "error_kind": "paused", + "error": "oclaw_alarm_forwarder runtime task paused", + "forwarder": fwd, + } + elif bool(fwd.get("connected")): + oclaw_status = { + "status": "up", + "mode": "ws", + "enabled": True, + "connected": True, + "queue_size": int(fwd.get("queue_size") or 0), + "published_ok": int(fwd.get("published_ok") or 0), + "published_fail": int(fwd.get("published_fail") or 0), + "url": str(fwd.get("url") or ""), + "forwarder": fwd, + } + else: + oclaw_status = { + "status": "down", + "mode": "ws", + "enabled": True, + "connected": False, + "error_kind": "ws_disconnected", + "error": "oclaw netx-bridge WebSocket not connected", + "queue_size": int(fwd.get("queue_size") or 0), + "url": str(fwd.get("url") or ""), + "forwarder": fwd, + } + + return {"netx_api": netx_api, "db": db_status, "oclaw_bridge": oclaw_status} diff --git a/netx_api/main.py b/netx_api/main.py index 236d3ed..b647233 100644 --- a/netx_api/main.py +++ b/netx_api/main.py @@ -1,44 +1,30 @@ +"""FastAPI application entry — routers + thin lifecycle hooks.""" + from __future__ import annotations -import csv -import json import logging -from datetime import datetime, timezone -from io import StringIO import time -import re -import threading -_schedule_log = logging.getLogger("netx.ume.schedule") -_BOOT_MONO = time.monotonic() -from fastapi import Depends, FastAPI, File, HTTPException, Query, UploadFile -from fastapi.responses import Response -from sqlalchemy import text as sql_text -from sqlalchemy.orm import Session -from typing import Any import uvicorn +from fastapi import FastAPI -from .ap_client import analyze_with_oclaw, health_with_oclaw from .auth_middleware import AuthAuditMiddleware from .auth_router import router as auth_router -from .auth_service import bootstrap_admin_if_needed -from .config import settings -from .db import Base, SessionLocal, engine, get_db -from .collection_router import router as collection_router +from .alarms_router import router as alarms_router from .cli_router import router as cli_router +from .collection_router import router as collection_router +from .config import settings from .config_sync_router import router as config_sync_router -from .port_traffic_router import router as port_traffic_router -from .managed_ne_router import router as managed_ne_router -from .webcrt_router import router as webcrt_router -from .topology_router import router as topology_router +from .integrations_router import router as integrations_router from .lldp_collect_router import router as lldp_collect_router +from .managed_ne_router import router as managed_ne_router from .ops_router import router as ops_router +from .parser_config import load_parser_config +from .port_traffic_router import router as port_traffic_router from .sql_router import router as sql_router from .sql_router import sql_query, sql_ume_query # noqa: F401 — tests import from main -from .security_bootstrap import assert_secure_defaults_or_exit -from .integrations_router import router as integrations_router +from .topology_router import router as topology_router from .ume_router import router as ume_router -from .alarms_router import router as alarms_router from .ume_router import ( # noqa: F401 — tests import from main _extract_ume_raw_group_field, _serialize_ume_alarm_raw_row, @@ -49,89 +35,12 @@ from .ume_support import ( # noqa: F401 — tests import from main _protocol_bucket_label, ) import netx_api.ume_support as ume_support -from .ume_runtime import start_device_schedulers -from .importer import aggregate_alarms, import_alarm_excel, query_alarms -from .models import ( - AiAnalyzeHistory, - AlarmBatch, - AlarmNorm, - ApiToken, - AppUser, - AuditLog, - ImportErrorRow, - ManagedNE, - NeCollectionJob, - NeCollectionRun, - UmeAlarmCurrent, - UmeAlarmHistory, - UmeInventoryNE, - UmeKeyAlertRule, - UmeKeyAlertForwardLog, - UmeSyncJob, -) -from .models import ImportJob -from .parser_config import load_parser_config -from .ume_client import UMEClient -from .ume_alarm_ws import ( - begin_startup_alarm_sync_gate, - cancel_alarm_subscription_manual, - clear_local_alarm_subscription_manual, - complete_startup_alarm_sync_gate, - establish_alarm_subscription_manual, - get_alarms_coordination_status, - get_subscription_status, - get_ws_connection_status, - get_ws_logs, - is_startup_alarm_sync_pending, - is_wss_active_for_current_alarms, - load_persisted_subscription, - request_ws_reconnect, - shutdown_ws_consumer, - start_ume_alarm_ws_consumer, -) -from .ume_sync_service import sync_alarms_current, sync_alarms_history_full, sync_inventory_full -from .runtime_task_messages import ( - RT_ALARMS_SYNC_IN_PROGRESS_SKIP, - RT_OCLAW_FWD_DISABLED, - RT_PULLING_ALARMS_CURRENT, - RT_PULLING_INVENTORY, - RT_RESUMED, - RT_RESUMED_OCLAW_WSS_RECONNECT, - RT_RESUMED_SYNC_SOON, - RT_RESUMED_WSS_RECONNECT, - RT_STARTUP_ALARM_SYNC_BEFORE_WS, - RT_STARTUP_GATE_WAITING, - RT_KEEPALIVE_FAILED, - RT_UME_WS_DISABLED_NO_BASE_URL, - RT_WSS_ACTIVE_SKIP_REST, -) -from .oclaw_alarm_forwarder import ( - forwarder_status, - is_forwarder_enabled, - request_forwarder_reconnect, - configure_oclaw_alarm_forwarder, - shutdown_oclaw_alarm_forwarder, - start_oclaw_alarm_forwarder, -) -from .ume_token_store import ( - clear_shared_token, - load_shared_token, - release_refresh_lock, - save_shared_token, - try_acquire_refresh_lock, - wait_for_token_update, -) -from .schemas import ( - AlarmAggregateBucket, - AlarmAggregateResponse, - AiAnalyzeHistoryItem, - AiAnalyzeHistoryResponse, - AlarmItem, - AlarmQueryResponse, - BatchSummary, - ImportJobItem, - ImportJobListResponse, -) +from .oclaw_alarm_forwarder import shutdown_oclaw_alarm_forwarder +from .ume_alarm_ws import shutdown_ws_consumer +from .webcrt_router import router as webcrt_router + +_schedule_log = logging.getLogger("netx.ume.schedule") +_BOOT_MONO = time.monotonic() app = FastAPI( title="netx ops tool", @@ -156,423 +65,13 @@ app.include_router(integrations_router) app.include_router(ume_router) app.include_router(alarms_router) parser_cfg = load_parser_config() -def _configure_ume_diag_logging() -> None: - """Emit netx.ume.* INFO to stderr so background scripts/.run/*.log and consoles show scheduler lines.""" - fmt = logging.Formatter("%(asctime)s %(levelname)s %(name)s: %(message)s") - for name in ("netx.ume.schedule", "netx.ume.sync"): - lg = logging.getLogger(name) - if lg.handlers: - continue - h = logging.StreamHandler() - h.setFormatter(fmt) - lg.addHandler(h) - lg.setLevel(logging.INFO) - lg.propagate = False @app.on_event("startup") def on_startup() -> None: - assert_secure_defaults_or_exit() - _configure_ume_diag_logging() - Base.metadata.create_all(bind=engine) - from .schema_patches import ( - apply_all_legacy_startup_ddl, - apply_auth_schema_patches, - run_alembic_upgrade_to_head, - ) + from .app_startup import run_api_startup - alembic_ok = True - if bool(getattr(settings, "alembic_upgrade_on_start", True)): - try: - run_alembic_upgrade_to_head() - except Exception: - alembic_ok = False - _schedule_log.exception("startup: alembic upgrade head failed") - - skip_ddl = bool(getattr(settings, "skip_legacy_startup_ddl", True)) - # Auth columns must exist before bootstrap even when legacy DDL is skipped. - try: - with engine.begin() as conn: - apply_auth_schema_patches(conn) - except Exception: - _schedule_log.exception("startup: auth schema patches failed") - if skip_ddl and alembic_ok: - _schedule_log.info("startup: schema via Alembic (legacy inline DDL skipped)") - else: - if skip_ddl and not alembic_ok: - _schedule_log.warning( - "startup: Alembic failed — falling back to legacy schema patches" - ) - try: - apply_all_legacy_startup_ddl(engine) - except Exception: - _schedule_log.exception("startup: legacy schema patches failed") - ume_support._reset_runtime_pause_flags() - ume_support._fail_stale_running_sync_jobs_on_startup() - try: - from .topology_service import bootstrap_topology_tree, reclaim_stale_discover_jobs - - db_topo = SessionLocal() - try: - bootstrap_topology_tree(db_topo) - closed = reclaim_stale_discover_jobs(db_topo, force_all_open=True) - if closed: - _schedule_log.warning( - "startup: closed %s orphaned topology discover jobs", closed - ) - finally: - db_topo.close() - except Exception: - _schedule_log.exception("startup: topology discover job cleanup failed") - if ume_support._needs_startup_alarm_sync_before_ws(): - begin_startup_alarm_sync_gate() - _schedule_log.info( - "startup: WSS blocked until initial REST current-alarm sync completes (delay=%ss)", - ume_support._startup_alarm_pull_delay_s(), - ) - else: - complete_startup_alarm_sync_gate() - db = SessionLocal() - try: - try: - bootstrap_admin_if_needed(db) - except Exception: - _schedule_log.exception("startup: auth bootstrap admin failed") - from .collection_recovery import recover_collection_jobs_on_startup - - resumed = recover_collection_jobs_on_startup(db) - if resumed: - _schedule_log.info("startup: resumed %s pending ne collection runs", resumed) - from .config_sync_recovery import recover_config_sync_on_startup - from .config_sync_service import ensure_policy - from .port_traffic_recovery import recover_port_traffic_on_startup - - ensure_policy(db) - cfg_resumed = recover_config_sync_on_startup(db) - if cfg_resumed: - _schedule_log.info("startup: resumed %s config_sync task(s) from interrupted cycle", cfg_resumed) - try: - from .lldp_collect_service import ensure_policy as ensure_lldp_collect_policy - - ensure_lldp_collect_policy(db) - except Exception: - _schedule_log.exception("startup: lldp_collect policy ensure failed") - pt_cleared = recover_port_traffic_on_startup(db) - if pt_cleared: - _schedule_log.info("startup: cleared %s port_traffic stuck collect_running flag(s)", pt_cleared) - try: - from .port_traffic_migrate import backfill_port_traffic_series - - backfill_port_traffic_series(db) - except Exception: - _schedule_log.exception("startup: port_traffic series backfill failed") - except Exception: - _schedule_log.exception("startup: ne collection / config_sync recovery failed") - finally: - db.close() - if bool(getattr(settings, "run_inline_schedulers", True)): - try: - start_device_schedulers() - except Exception: - _schedule_log.exception("startup: device schedulers init failed") - else: - _schedule_log.info( - "startup: inline schedulers disabled — run `python -m netx_api.worker` for " - "config_sync / lldp_collect / port_traffic" - ) - try: - if bool(getattr(settings, "ume_keepalive_enabled", True)): - interval_keepalive_s = int(getattr(settings, "ume_keepalive_interval_s", 600) or 600) - interval_keepalive_s = max(30, min(interval_keepalive_s, 3600)) - renew_before_s = int(getattr(settings, "ume_keepalive_renew_before_s", 900) or 900) - renew_before_s = max(30, min(renew_before_s, 86400)) - - def _keepalive_loop() -> None: - # Best-effort keepalive: if token exists, periodically handshake to extend TTL. - while True: - try: - if ume_support._runtime_is_paused("token_keepalive"): - time.sleep(1) - continue - client = ume_support._ume_client() - st = client.token_status() - expires_in = int(st.get("expires_in_s") or 0) - # Renew when missing/invalid TTL (0) or nearing expiry — previously 0 skipped renew forever. - if bool(st.get("has_token")) and (expires_in <= 0 or expires_in < renew_before_s): - client.renew_token() - ume_support._set_runtime_task("token_keepalive", status="running", last_run_at=datetime.now(timezone.utc), last_error="") - except Exception: - ume_support._set_runtime_task("token_keepalive", status="error", last_run_at=datetime.now(timezone.utc), last_error=RT_KEEPALIVE_FAILED) - time.sleep(interval_keepalive_s) - - t = threading.Thread(target=_keepalive_loop, name="ume-token-keepalive", daemon=True) - t.start() - except Exception as exc: - _schedule_log.exception("startup: token_keepalive thread init failed: %s", exc) - ume_support._set_runtime_task( - "token_keepalive", - status="error", - last_run_at=datetime.now(timezone.utc), - last_error=f"startup_thread_init_failed: {str(exc)[:180]}", - ) - try: - - def _startup_alarm_sync_worker() -> None: - try: - ume_support._run_startup_alarm_sync_before_ws() - except Exception as exc: - _schedule_log.exception("startup: alarm sync before WSS failed: %s", exc) - complete_startup_alarm_sync_gate() - - # Do not block HTTP /health on slow UME REST pull; WSS waits on startup_alarm_sync_gate. - t_startup_sync = threading.Thread( - target=_startup_alarm_sync_worker, - name="ume-startup-alarm-sync", - daemon=True, - ) - t_startup_sync.start() - except Exception as exc: - _schedule_log.exception("startup: alarm sync thread init failed: %s", exc) - complete_startup_alarm_sync_gate() - try: - if bool(getattr(settings, "ume_sync_alarms_current_enabled", True)): - alarms_interval_s = int(getattr(settings, "ume_sync_alarms_current_interval_s", 18000) or 18000) - alarms_interval_s = max(30, min(alarms_interval_s, 86400)) - - def _alarms_current_sync_loop() -> None: - ume_support._refresh_runtime_task_idle("alarms_current_auto_sync", "alarms_current") - ume_support._wait_until_startup_alarm_pull_allowed("alarms_current_auto_sync") - while True: - try: - _schedule_log.info( - "alarms_current_auto_sync: loop tick paused=%s", - ume_support._runtime_is_paused("alarms_current_auto_sync"), - ) - if ume_support._runtime_is_paused("alarms_current_auto_sync"): - time.sleep(1) - continue - if is_startup_alarm_sync_pending(): - ume_support._refresh_runtime_task_idle( - "alarms_current_auto_sync", - "alarms_current", - last_error=RT_STARTUP_GATE_WAITING, - ) - time.sleep(10) - continue - if ( - bool(getattr(settings, "ume_sync_alarms_current_skip_when_ws", True)) - and is_wss_active_for_current_alarms() - ): - ume_support._refresh_runtime_task_idle( - "alarms_current_auto_sync", - "alarms_current", - last_error=RT_WSS_ACTIVE_SKIP_REST, - ) - time.sleep(max(30, min(alarms_interval_s, 300))) - continue - ume_support._maybe_wait_for_sync_interval( - task_id="alarms_current_auto_sync", - domain="alarms_current", - interval_s=alarms_interval_s, - label="alarms_current_auto_sync", - ) - _schedule_log.info( - "alarms_current_auto_sync: iteration start (interval=%ss)", - alarms_interval_s, - ) - ume_support._set_runtime_task( - "alarms_current_auto_sync", - status="running", - last_run_at=datetime.now(timezone.utc), - last_error=RT_PULLING_ALARMS_CURRENT, - ) - db = SessionLocal() - try: - client = ume_support._ume_client() - sync_alarms_current(db, client, trigger_mode="schedule") - _schedule_log.info("alarms_current_auto_sync: sync finished ok") - ume_support._set_runtime_task( - "alarms_current_auto_sync", - status="running", - last_run_at=datetime.now(timezone.utc), - last_error="", - ) - finally: - db.close() - except RuntimeError as exc: - if str(exc) == "alarms_current_sync_busy": - ume_support._refresh_runtime_task_idle( - "alarms_current_auto_sync", - "alarms_current", - last_error=RT_ALARMS_SYNC_IN_PROGRESS_SKIP, - ) - time.sleep(30) - else: - raise - except Exception as exc: - _schedule_log.exception("alarms_current_auto_sync: sync failed: %s", exc) - ume_support._set_runtime_task( - "alarms_current_auto_sync", - status="error", - last_run_at=datetime.now(timezone.utc), - last_error=str(exc)[:240], - ) - - t2 = threading.Thread(target=_alarms_current_sync_loop, name="ume-alarms-current-sync", daemon=True) - t2.start() - _schedule_log.info("started thread %s alive=%s", t2.name, t2.is_alive()) - if not t2.is_alive(): - _schedule_log.error("ume-alarms-current-sync thread exited immediately (check uncaught errors above)") - except Exception as exc: - _schedule_log.exception("startup: alarms_current_auto_sync thread init failed: %s", exc) - ume_support._set_runtime_task( - "alarms_current_auto_sync", - status="error", - last_run_at=datetime.now(timezone.utc), - last_error=f"startup_thread_init_failed: {str(exc)[:180]}", - ) - try: - if bool(getattr(settings, "ume_sync_inventory_auto_enabled", True)): - hours = int(getattr(settings, "ume_sync_inventory_every_hours", 48) or 48) - hours = max(1, min(hours, 168)) - inventory_interval_s = int(hours * 3600) - ume_support._refresh_runtime_task_idle("inventory_auto_sync", "inventory") - - def _inventory_auto_sync_loop() -> None: - ume_support._refresh_runtime_task_idle("inventory_auto_sync", "inventory") - while True: - try: - _schedule_log.info( - "inventory_auto_sync: loop tick paused=%s", - ume_support._runtime_is_paused("inventory_auto_sync"), - ) - if ume_support._runtime_is_paused("inventory_auto_sync"): - time.sleep(1) - continue - ume_support._maybe_wait_for_sync_interval( - task_id="inventory_auto_sync", - domain="inventory", - interval_s=inventory_interval_s, - label="inventory_auto_sync", - ) - _schedule_log.info( - "inventory_auto_sync: iteration start (interval=%ss)", - inventory_interval_s, - ) - ume_support._set_runtime_task( - "inventory_auto_sync", - status="running", - last_run_at=datetime.now(timezone.utc), - last_error=RT_PULLING_INVENTORY, - ) - db = SessionLocal() - try: - client = ume_support._ume_client() - sync_inventory_full(db, client, trigger_mode="schedule") - _schedule_log.info("inventory_auto_sync: sync finished ok") - ume_support._set_runtime_task( - "inventory_auto_sync", - status="running", - last_run_at=datetime.now(timezone.utc), - last_error="", - ) - finally: - db.close() - except Exception as exc: - _schedule_log.exception("inventory_auto_sync: sync failed: %s", exc) - ume_support._set_runtime_task( - "inventory_auto_sync", - status="error", - last_run_at=datetime.now(timezone.utc), - last_error=str(exc)[:240], - ) - - t3 = threading.Thread(target=_inventory_auto_sync_loop, name="ume-inventory-auto-sync", daemon=True) - t3.start() - _schedule_log.info("started thread %s alive=%s", t3.name, t3.is_alive()) - if not t3.is_alive(): - _schedule_log.error("ume-inventory-auto-sync thread exited immediately (check uncaught errors above)") - except Exception as exc: - _schedule_log.exception("startup: inventory_auto_sync thread init failed: %s", exc) - ume_support._set_runtime_task( - "inventory_auto_sync", - status="error", - last_run_at=datetime.now(timezone.utc), - last_error=f"startup_thread_init_failed: {str(exc)[:180]}", - ) - try: - if bool(getattr(settings, "ume_alarm_ws_enabled", True)) and str(getattr(settings, "ume_base_url", "") or "").strip(): - if load_persisted_subscription(): - _schedule_log.info("startup: loaded persisted UME alarm subscription") - ume_support._UME_WS_STOP_EVENT = threading.Event() - - def _ws_on_status(msg: str) -> None: - ume_support._set_runtime_task( - "alarms_current_ws_consumer", - status="running", - last_run_at=datetime.now(timezone.utc), - last_error=str(msg or "")[:240], - ) - - t_ws = start_ume_alarm_ws_consumer( - ume_support._ume_client(), - on_status=_ws_on_status, - stop_event=ume_support._UME_WS_STOP_EVENT, - is_paused=lambda: ume_support._runtime_is_paused("alarms_current_ws_consumer"), - ) - _schedule_log.info("started thread %s alive=%s", t_ws.name, t_ws.is_alive()) - else: - ume_support._set_runtime_task("alarms_current_ws_consumer", status="paused", last_error=RT_UME_WS_DISABLED_NO_BASE_URL) - except Exception as exc: - _schedule_log.exception("startup: alarms_current_ws_consumer thread init failed: %s", exc) - ume_support._set_runtime_task( - "alarms_current_ws_consumer", - status="error", - last_run_at=datetime.now(timezone.utc), - last_error=f"startup_thread_init_failed: {str(exc)[:180]}", - ) - try: - def _fwd_on_status(msg: str) -> None: - paused = ume_support._runtime_is_paused("oclaw_alarm_forwarder") - fwd = forwarder_status() - if paused: - status = "paused" - elif not bool(fwd.get("enabled")): - status = "paused" - elif bool(fwd.get("connected")): - status = "running" - else: - status = "running" - ume_support._set_runtime_task( - "oclaw_alarm_forwarder", - status=status, - last_run_at=datetime.now(timezone.utc), - last_error=str(msg or "")[:240], - ) - - configure_oclaw_alarm_forwarder( - is_paused=lambda: ume_support._runtime_is_paused("oclaw_alarm_forwarder"), - on_status=_fwd_on_status, - ) - if is_forwarder_enabled(): - ume_support._set_runtime_task("oclaw_alarm_forwarder", status="running", last_error="") - else: - ume_support._set_runtime_task( - "oclaw_alarm_forwarder", - status="paused", - last_error=RT_OCLAW_FWD_DISABLED, - ) - t_fwd = start_oclaw_alarm_forwarder() - if t_fwd is not None: - _schedule_log.info("started thread %s alive=%s", t_fwd.name, t_fwd.is_alive()) - except Exception as exc: - _schedule_log.exception("startup: oclaw_alarm_forwarder thread init failed: %s", exc) - ume_support._set_runtime_task( - "oclaw_alarm_forwarder", - status="error", - last_run_at=datetime.now(timezone.utc), - last_error=f"startup_thread_init_failed: {str(exc)[:180]}", - ) + run_api_startup() @app.on_event("shutdown") @@ -588,70 +87,6 @@ def health() -> dict[str, str]: return {"status": "ok"} - -@app.get("/v1/integrations/status") -def integrations_status(db: Session = Depends(get_db)) -> dict: - # netx api is up if this handler executes; still verify DB + oclaw bridge separately. - netx_api = {"status": "up"} - - db_status: dict = {"status": "unknown"} - try: - t0 = time.monotonic() - db.execute(sql_text("select 1")) - db_status = {"status": "up", "latency_ms": int((time.monotonic() - t0) * 1000)} - except Exception as exc: - db_status = {"status": "down", "error": str(exc)[:240]} - - oclaw_status: dict = {"status": "unknown"} - fwd = forwarder_status() - if not bool(fwd.get("enabled")): - oclaw_status = { - "status": "unknown", - "mode": "ws", - "enabled": False, - "connected": False, - "error_kind": "disabled", - "error": "NETX_OCLAW_ALARM_WS_ENABLED=false or missing token/url", - "forwarder": fwd, - } - elif bool(fwd.get("paused")): - oclaw_status = { - "status": "unknown", - "mode": "ws", - "enabled": True, - "connected": False, - "error_kind": "paused", - "error": "oclaw_alarm_forwarder runtime task paused", - "forwarder": fwd, - } - elif bool(fwd.get("connected")): - oclaw_status = { - "status": "up", - "mode": "ws", - "enabled": True, - "connected": True, - "queue_size": int(fwd.get("queue_size") or 0), - "published_ok": int(fwd.get("published_ok") or 0), - "published_fail": int(fwd.get("published_fail") or 0), - "url": str(fwd.get("url") or ""), - "forwarder": fwd, - } - else: - oclaw_status = { - "status": "down", - "mode": "ws", - "enabled": True, - "connected": False, - "error_kind": "ws_disconnected", - "error": "oclaw netx-bridge WebSocket not connected", - "queue_size": int(fwd.get("queue_size") or 0), - "url": str(fwd.get("url") or ""), - "forwarder": fwd, - } - - return {"netx_api": netx_api, "db": db_status, "oclaw_bridge": oclaw_status} - - @app.get("/") def root() -> dict: return { diff --git a/netx_api/ne_exec.py b/netx_api/ne_exec.py index 8be106f..97fd106 100644 --- a/netx_api/ne_exec.py +++ b/netx_api/ne_exec.py @@ -2,7 +2,6 @@ from __future__ import annotations -import re from typing import Any from fastapi import HTTPException @@ -12,72 +11,24 @@ from .cli_resolve import resolve_cli_target from .config import settings from .ne_collect_runner import _collect_on_device from .ne_crypto import credentials_configured +from .ne_exec_guard import _validate_command, validate_ne_exec_command _EXEC_MAX_COMMANDS_CAP = 50 _EXEC_MAX_OUTPUT = 32_000 _EXEC_READ_TIMEOUT_DEFAULT = 60 _EXEC_READ_TIMEOUT_MAX = 120 +__all__ = [ + "_validate_command", + "execute_managed_ne_commands", + "validate_ne_exec_command", +] + def _exec_max_commands() -> int: raw = int(settings.ne_exec_max_commands or 5) return max(1, min(_EXEC_MAX_COMMANDS_CAP, raw)) -# Block obvious config-change / destructive patterns (case-insensitive). -_BLOCKED_RE = re.compile( - r"(?i)(" - r"configure\s+terminal|conf\s+t\b|" - r"\bwrite\s+(memory|erase)|\bcopy\s+run|\bcopy\s+startup|" - r"\breload\b|\breboot\b|\berase\b|\bformat\b|\bdelete\b|" - r"\bcommit\b|\brollback\b|startup-config|" - r"\bsystem-view\b|\bip\s+address\b|\bvlan\s+\d" - r")" -) - -# Read-only CLI: show/display plus ping/traceroute reachability checks. -_ALLOWED_PREFIX_RE = re.compile( - r"(?i)^(show\s|display\s|ping\s|ping6\s|traceroute\s|tracert\s|trace\s|trace6\s)" -) - -# Unicode / C1 line separators that can smuggle a second CLI after a show prefix. -_FORBIDDEN_LINE_SEPARATORS = ("\u2028", "\u2029", "\x85", "\x0b", "\x0c") - -# Pipe segments allowed after show/display (output filtering only). -_ALLOWED_PIPE_SEGMENT_RE = re.compile( - r"(?i)^(include|exclude|begin|section|count|match|grep|one-line|no-more)(\s|$)" -) -_BLOCKED_PIPE_SEGMENT_RE = re.compile(r"(?i)\b(redirect|append|tee|send)\b") - - -def _validate_pipe_segments(cmd: str) -> None: - if "|" not in cmd: - return - parts = [p.strip() for p in cmd.split("|")] - if len(parts) < 2 or not parts[0] or any(not p for p in parts[1:]): - raise HTTPException(status_code=400, detail="command_pipe_not_allowed") - for segment in parts[1:]: - if _BLOCKED_PIPE_SEGMENT_RE.search(segment): - raise HTTPException(status_code=400, detail="command_pipe_not_allowed") - if not _ALLOWED_PIPE_SEGMENT_RE.match(segment): - raise HTTPException(status_code=400, detail="command_pipe_not_allowed") - - -def _validate_command(command: str) -> None: - cmd = str(command or "").strip() - if not cmd: - raise HTTPException(status_code=400, detail="empty_command") - if len(cmd) > 500: - raise HTTPException(status_code=400, detail="command_too_long") - if any(ch in cmd for ch in (";", "\n", "\r", "`")): - raise HTTPException(status_code=400, detail="command_chars_not_allowed") - if any(sep in cmd for sep in _FORBIDDEN_LINE_SEPARATORS): - raise HTTPException(status_code=400, detail="command_chars_not_allowed") - if _BLOCKED_RE.search(cmd): - raise HTTPException(status_code=400, detail="command_blocked") - if not _ALLOWED_PREFIX_RE.match(cmd): - raise HTTPException(status_code=400, detail="command_not_allowed_prefix") - _validate_pipe_segments(cmd) - def _normalize_read_timeout(sec: int | None) -> int: raw = int(sec if sec is not None else _EXEC_READ_TIMEOUT_DEFAULT) @@ -105,7 +56,7 @@ def execute_managed_ne_commands( if len(cmds) > max_cmds: raise HTTPException(status_code=400, detail=f"too_many_commands (max {max_cmds})") for c in cmds: - _validate_command(c) + validate_ne_exec_command(c) creds, device = resolve_cli_target(db, managed_ne_id=mid or None, ume_ne_id=uid or None) read_timeout = _normalize_read_timeout(read_timeout_sec) diff --git a/netx_api/ne_exec_guard.py b/netx_api/ne_exec_guard.py new file mode 100644 index 0000000..dd408a6 --- /dev/null +++ b/netx_api/ne_exec_guard.py @@ -0,0 +1,74 @@ +"""NE CLI command allow/deny gates (read-only exec for ops tools).""" + +from __future__ import annotations + +import re + +from fastapi import HTTPException + +# Block obvious config-change / destructive patterns (case-insensitive). +_BLOCKED_RE = re.compile( + r"(?i)(" + r"configure\s+terminal|conf\s+t\b|" + r"\bwrite\s+(memory|erase)|\bcopy\s+run|\bcopy\s+startup|" + r"\breload\b|\breboot\b|\berase\b|\bformat\b|\bdelete\b|" + r"\bcommit\b|\brollback\b|startup-config|" + r"\bsystem-view\b|\bip\s+address\b|\bvlan\s+\d|" + # Extra vendor / destructive surface (avoid words that appear in show output filters) + r"\bclear\s+configuration\b|\breset\s+saved-configuration\b|" + r"\bundo\s+|\bsave\s*$|\bsave\s+\S|" + r"\bfile\s+delete\b|\bftp\s+put\b|\btftp\s+put\b|" + r"\bdebug\s+all\b|\bundebug\s+all\b|" + r"\brequest\s+system\s+(reboot|halt|power-off|zeroize)\b|" + r"\bset\s+system\s+reboot\b" + r")" +) + +# Read-only CLI: show/display plus ping/traceroute reachability checks. +_ALLOWED_PREFIX_RE = re.compile( + r"(?i)^(show\s|display\s|ping\s|ping6\s|traceroute\s|tracert\s|trace\s|trace6\s)" +) + +# Unicode / C1 line separators that can smuggle a second CLI after a show prefix. +_FORBIDDEN_LINE_SEPARATORS = ("\u2028", "\u2029", "\x85", "\x0b", "\x0c") + +# Pipe segments allowed after show/display (output filtering only). +_ALLOWED_PIPE_SEGMENT_RE = re.compile( + r"(?i)^(include|exclude|begin|section|count|match|grep|one-line|no-more)(\s|$)" +) +_BLOCKED_PIPE_SEGMENT_RE = re.compile(r"(?i)\b(redirect|append|tee|send)\b") + + +def _validate_pipe_segments(cmd: str) -> None: + if "|" not in cmd: + return + parts = [p.strip() for p in cmd.split("|")] + if len(parts) < 2 or not parts[0] or any(not p for p in parts[1:]): + raise HTTPException(status_code=400, detail="command_pipe_not_allowed") + for segment in parts[1:]: + if _BLOCKED_PIPE_SEGMENT_RE.search(segment): + raise HTTPException(status_code=400, detail="command_pipe_not_allowed") + if not _ALLOWED_PIPE_SEGMENT_RE.match(segment): + raise HTTPException(status_code=400, detail="command_pipe_not_allowed") + + +def validate_ne_exec_command(command: str) -> None: + """Raise HTTPException if command is empty, smuggled, blocked, or not allowlisted.""" + cmd = str(command or "").strip() + if not cmd: + raise HTTPException(status_code=400, detail="empty_command") + if len(cmd) > 500: + raise HTTPException(status_code=400, detail="command_too_long") + if any(ch in cmd for ch in (";", "\n", "\r", "`")): + raise HTTPException(status_code=400, detail="command_chars_not_allowed") + if any(sep in cmd for sep in _FORBIDDEN_LINE_SEPARATORS): + raise HTTPException(status_code=400, detail="command_chars_not_allowed") + if _BLOCKED_RE.search(cmd): + raise HTTPException(status_code=400, detail="command_blocked") + if not _ALLOWED_PREFIX_RE.match(cmd): + raise HTTPException(status_code=400, detail="command_not_allowed_prefix") + _validate_pipe_segments(cmd) + + +# Back-compat alias used by tests / callers. +_validate_command = validate_ne_exec_command diff --git a/netx_api/ume_runtime.py b/netx_api/ume_runtime.py index c6e9745..812407a 100644 --- a/netx_api/ume_runtime.py +++ b/netx_api/ume_runtime.py @@ -1,15 +1,48 @@ """UME / long-task runtime helpers shared by API and optional worker process. -The API process still owns UME token keepalive, alarm WSS, and oclaw forwarder. -Config-sync / LLDP / port-traffic tick loops can run inline (default) or via -``python -m netx_api.worker`` when ``NETX_RUN_INLINE_SCHEDULERS=false``. +Device collectors (config_sync / LLDP / port_traffic) run via ``start_device_schedulers`` +(API inline when ``NETX_RUN_INLINE_SCHEDULERS=true``, otherwise ``python -m netx_api.worker``). +API process also owns UME keepalive, alarm WSS, current-alarm/inventory sync loops, +and oclaw forwarder via ``start_api_sideband_threads``. """ from __future__ import annotations import logging +import threading +import time +from datetime import datetime, timezone + +from .config import settings +from .db import SessionLocal +import netx_api.ume_support as ume_support +from .ume_alarm_ws import ( + complete_startup_alarm_sync_gate, + is_startup_alarm_sync_pending, + is_wss_active_for_current_alarms, + load_persisted_subscription, + start_ume_alarm_ws_consumer, +) +from .ume_sync_service import sync_alarms_current, sync_inventory_full +from .runtime_task_messages import ( + RT_ALARMS_SYNC_IN_PROGRESS_SKIP, + RT_KEEPALIVE_FAILED, + RT_OCLAW_FWD_DISABLED, + RT_PULLING_ALARMS_CURRENT, + RT_PULLING_INVENTORY, + RT_STARTUP_GATE_WAITING, + RT_UME_WS_DISABLED_NO_BASE_URL, + RT_WSS_ACTIVE_SKIP_REST, +) +from .oclaw_alarm_forwarder import ( + configure_oclaw_alarm_forwarder, + forwarder_status, + is_forwarder_enabled, + start_oclaw_alarm_forwarder, +) _log = logging.getLogger("netx.ume.runtime") +_schedule_log = logging.getLogger("netx.ume.schedule") def start_device_schedulers() -> None: @@ -22,3 +55,303 @@ def start_device_schedulers() -> None: start_lldp_collect_scheduler() start_port_traffic_scheduler() _log.info("device schedulers started") + + +def start_api_sideband_threads() -> None: + """UME keepalive / alarm sync / inventory / WSS / oclaw forwarder (API process).""" + try: + if bool(getattr(settings, "ume_keepalive_enabled", True)): + interval_keepalive_s = int(getattr(settings, "ume_keepalive_interval_s", 600) or 600) + interval_keepalive_s = max(30, min(interval_keepalive_s, 3600)) + renew_before_s = int(getattr(settings, "ume_keepalive_renew_before_s", 900) or 900) + renew_before_s = max(30, min(renew_before_s, 86400)) + + def _keepalive_loop() -> None: + # Best-effort keepalive: if token exists, periodically handshake to extend TTL. + while True: + try: + if ume_support._runtime_is_paused("token_keepalive"): + time.sleep(1) + continue + client = ume_support._ume_client() + st = client.token_status() + expires_in = int(st.get("expires_in_s") or 0) + # Renew when missing/invalid TTL (0) or nearing expiry — previously 0 skipped renew forever. + if bool(st.get("has_token")) and (expires_in <= 0 or expires_in < renew_before_s): + client.renew_token() + ume_support._set_runtime_task("token_keepalive", status="running", last_run_at=datetime.now(timezone.utc), last_error="") + except Exception: + ume_support._set_runtime_task("token_keepalive", status="error", last_run_at=datetime.now(timezone.utc), last_error=RT_KEEPALIVE_FAILED) + time.sleep(interval_keepalive_s) + + t = threading.Thread(target=_keepalive_loop, name="ume-token-keepalive", daemon=True) + t.start() + except Exception as exc: + _schedule_log.exception("startup: token_keepalive thread init failed: %s", exc) + ume_support._set_runtime_task( + "token_keepalive", + status="error", + last_run_at=datetime.now(timezone.utc), + last_error=f"startup_thread_init_failed: {str(exc)[:180]}", + ) + try: + + def _startup_alarm_sync_worker() -> None: + try: + ume_support._run_startup_alarm_sync_before_ws() + except Exception as exc: + _schedule_log.exception("startup: alarm sync before WSS failed: %s", exc) + complete_startup_alarm_sync_gate() + + # Do not block HTTP /health on slow UME REST pull; WSS waits on startup_alarm_sync_gate. + t_startup_sync = threading.Thread( + target=_startup_alarm_sync_worker, + name="ume-startup-alarm-sync", + daemon=True, + ) + t_startup_sync.start() + except Exception as exc: + _schedule_log.exception("startup: alarm sync thread init failed: %s", exc) + complete_startup_alarm_sync_gate() + try: + if bool(getattr(settings, "ume_sync_alarms_current_enabled", True)): + alarms_interval_s = int(getattr(settings, "ume_sync_alarms_current_interval_s", 18000) or 18000) + alarms_interval_s = max(30, min(alarms_interval_s, 86400)) + + def _alarms_current_sync_loop() -> None: + ume_support._refresh_runtime_task_idle("alarms_current_auto_sync", "alarms_current") + ume_support._wait_until_startup_alarm_pull_allowed("alarms_current_auto_sync") + while True: + try: + _schedule_log.info( + "alarms_current_auto_sync: loop tick paused=%s", + ume_support._runtime_is_paused("alarms_current_auto_sync"), + ) + if ume_support._runtime_is_paused("alarms_current_auto_sync"): + time.sleep(1) + continue + if is_startup_alarm_sync_pending(): + ume_support._refresh_runtime_task_idle( + "alarms_current_auto_sync", + "alarms_current", + last_error=RT_STARTUP_GATE_WAITING, + ) + time.sleep(10) + continue + if ( + bool(getattr(settings, "ume_sync_alarms_current_skip_when_ws", True)) + and is_wss_active_for_current_alarms() + ): + ume_support._refresh_runtime_task_idle( + "alarms_current_auto_sync", + "alarms_current", + last_error=RT_WSS_ACTIVE_SKIP_REST, + ) + time.sleep(max(30, min(alarms_interval_s, 300))) + continue + ume_support._maybe_wait_for_sync_interval( + task_id="alarms_current_auto_sync", + domain="alarms_current", + interval_s=alarms_interval_s, + label="alarms_current_auto_sync", + ) + _schedule_log.info( + "alarms_current_auto_sync: iteration start (interval=%ss)", + alarms_interval_s, + ) + ume_support._set_runtime_task( + "alarms_current_auto_sync", + status="running", + last_run_at=datetime.now(timezone.utc), + last_error=RT_PULLING_ALARMS_CURRENT, + ) + db = SessionLocal() + try: + client = ume_support._ume_client() + sync_alarms_current(db, client, trigger_mode="schedule") + _schedule_log.info("alarms_current_auto_sync: sync finished ok") + ume_support._set_runtime_task( + "alarms_current_auto_sync", + status="running", + last_run_at=datetime.now(timezone.utc), + last_error="", + ) + finally: + db.close() + except RuntimeError as exc: + if str(exc) == "alarms_current_sync_busy": + ume_support._refresh_runtime_task_idle( + "alarms_current_auto_sync", + "alarms_current", + last_error=RT_ALARMS_SYNC_IN_PROGRESS_SKIP, + ) + time.sleep(30) + else: + raise + except Exception as exc: + _schedule_log.exception("alarms_current_auto_sync: sync failed: %s", exc) + ume_support._set_runtime_task( + "alarms_current_auto_sync", + status="error", + last_run_at=datetime.now(timezone.utc), + last_error=str(exc)[:240], + ) + + t2 = threading.Thread(target=_alarms_current_sync_loop, name="ume-alarms-current-sync", daemon=True) + t2.start() + _schedule_log.info("started thread %s alive=%s", t2.name, t2.is_alive()) + if not t2.is_alive(): + _schedule_log.error("ume-alarms-current-sync thread exited immediately (check uncaught errors above)") + except Exception as exc: + _schedule_log.exception("startup: alarms_current_auto_sync thread init failed: %s", exc) + ume_support._set_runtime_task( + "alarms_current_auto_sync", + status="error", + last_run_at=datetime.now(timezone.utc), + last_error=f"startup_thread_init_failed: {str(exc)[:180]}", + ) + try: + if bool(getattr(settings, "ume_sync_inventory_auto_enabled", True)): + hours = int(getattr(settings, "ume_sync_inventory_every_hours", 48) or 48) + hours = max(1, min(hours, 168)) + inventory_interval_s = int(hours * 3600) + ume_support._refresh_runtime_task_idle("inventory_auto_sync", "inventory") + + def _inventory_auto_sync_loop() -> None: + ume_support._refresh_runtime_task_idle("inventory_auto_sync", "inventory") + while True: + try: + _schedule_log.info( + "inventory_auto_sync: loop tick paused=%s", + ume_support._runtime_is_paused("inventory_auto_sync"), + ) + if ume_support._runtime_is_paused("inventory_auto_sync"): + time.sleep(1) + continue + ume_support._maybe_wait_for_sync_interval( + task_id="inventory_auto_sync", + domain="inventory", + interval_s=inventory_interval_s, + label="inventory_auto_sync", + ) + _schedule_log.info( + "inventory_auto_sync: iteration start (interval=%ss)", + inventory_interval_s, + ) + ume_support._set_runtime_task( + "inventory_auto_sync", + status="running", + last_run_at=datetime.now(timezone.utc), + last_error=RT_PULLING_INVENTORY, + ) + db = SessionLocal() + try: + client = ume_support._ume_client() + sync_inventory_full(db, client, trigger_mode="schedule") + _schedule_log.info("inventory_auto_sync: sync finished ok") + ume_support._set_runtime_task( + "inventory_auto_sync", + status="running", + last_run_at=datetime.now(timezone.utc), + last_error="", + ) + finally: + db.close() + except Exception as exc: + _schedule_log.exception("inventory_auto_sync: sync failed: %s", exc) + ume_support._set_runtime_task( + "inventory_auto_sync", + status="error", + last_run_at=datetime.now(timezone.utc), + last_error=str(exc)[:240], + ) + + t3 = threading.Thread(target=_inventory_auto_sync_loop, name="ume-inventory-auto-sync", daemon=True) + t3.start() + _schedule_log.info("started thread %s alive=%s", t3.name, t3.is_alive()) + if not t3.is_alive(): + _schedule_log.error("ume-inventory-auto-sync thread exited immediately (check uncaught errors above)") + except Exception as exc: + _schedule_log.exception("startup: inventory_auto_sync thread init failed: %s", exc) + ume_support._set_runtime_task( + "inventory_auto_sync", + status="error", + last_run_at=datetime.now(timezone.utc), + last_error=f"startup_thread_init_failed: {str(exc)[:180]}", + ) + try: + if bool(getattr(settings, "ume_alarm_ws_enabled", True)) and str(getattr(settings, "ume_base_url", "") or "").strip(): + if load_persisted_subscription(): + _schedule_log.info("startup: loaded persisted UME alarm subscription") + ume_support._UME_WS_STOP_EVENT = threading.Event() + + def _ws_on_status(msg: str) -> None: + ume_support._set_runtime_task( + "alarms_current_ws_consumer", + status="running", + last_run_at=datetime.now(timezone.utc), + last_error=str(msg or "")[:240], + ) + + t_ws = start_ume_alarm_ws_consumer( + ume_support._ume_client(), + on_status=_ws_on_status, + stop_event=ume_support._UME_WS_STOP_EVENT, + is_paused=lambda: ume_support._runtime_is_paused("alarms_current_ws_consumer"), + ) + _schedule_log.info("started thread %s alive=%s", t_ws.name, t_ws.is_alive()) + else: + ume_support._set_runtime_task("alarms_current_ws_consumer", status="paused", last_error=RT_UME_WS_DISABLED_NO_BASE_URL) + except Exception as exc: + _schedule_log.exception("startup: alarms_current_ws_consumer thread init failed: %s", exc) + ume_support._set_runtime_task( + "alarms_current_ws_consumer", + status="error", + last_run_at=datetime.now(timezone.utc), + last_error=f"startup_thread_init_failed: {str(exc)[:180]}", + ) + try: + def _fwd_on_status(msg: str) -> None: + paused = ume_support._runtime_is_paused("oclaw_alarm_forwarder") + fwd = forwarder_status() + if paused: + status = "paused" + elif not bool(fwd.get("enabled")): + status = "paused" + elif bool(fwd.get("connected")): + status = "running" + else: + status = "running" + ume_support._set_runtime_task( + "oclaw_alarm_forwarder", + status=status, + last_run_at=datetime.now(timezone.utc), + last_error=str(msg or "")[:240], + ) + + configure_oclaw_alarm_forwarder( + is_paused=lambda: ume_support._runtime_is_paused("oclaw_alarm_forwarder"), + on_status=_fwd_on_status, + ) + if is_forwarder_enabled(): + ume_support._set_runtime_task("oclaw_alarm_forwarder", status="running", last_error="") + else: + ume_support._set_runtime_task( + "oclaw_alarm_forwarder", + status="paused", + last_error=RT_OCLAW_FWD_DISABLED, + ) + t_fwd = start_oclaw_alarm_forwarder() + if t_fwd is not None: + _schedule_log.info("started thread %s alive=%s", t_fwd.name, t_fwd.is_alive()) + except Exception as exc: + _schedule_log.exception("startup: oclaw_alarm_forwarder thread init failed: %s", exc) + ume_support._set_runtime_task( + "oclaw_alarm_forwarder", + status="error", + last_run_at=datetime.now(timezone.utc), + last_error=f"startup_thread_init_failed: {str(exc)[:180]}", + ) + + + diff --git a/netx_api/ume_support.py b/netx_api/ume_support.py index d1a88e1..c9b13af 100644 --- a/netx_api/ume_support.py +++ b/netx_api/ume_support.py @@ -47,6 +47,7 @@ from .ume_token_store import ( ) _schedule_log = logging.getLogger("netx.ume.schedule") +_BOOT_MONO = time.monotonic() _UME_CLIENT_SINGLETON = UMEClient( token_loader=lambda: load_shared_token(), diff --git a/netx_api/webcrt_channel.py b/netx_api/webcrt_channel.py new file mode 100644 index 0000000..0b582f0 --- /dev/null +++ b/netx_api/webcrt_channel.py @@ -0,0 +1,425 @@ +"""WebCRT channel helpers: keymap, prompt heuristics, encoding, queues.""" +from __future__ import annotations + +import io +import json +import logging +import queue +import re +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from netmiko import ConnectHandler + +from .config import settings + +_log = logging.getLogger("netx.webcrt") + +_NETWORK_CLI_KEY_SEQS: tuple[tuple[str, str], ...] = ( + ("\x1b[1~", "\x01"), # Home -> Ctrl-A + ("\x1b[3~", "\x04"), # Delete key -> Ctrl-D + ("\x1b[4~", "\x05"), # End -> Ctrl-E + ("\x1b[H", "\x01"), + ("\x1b[F", "\x05"), + ("\x1bOH", "\x01"), + ("\x1bOF", "\x05"), + ("\x1bOA", "\x1b[A"), # App Up -> CSI Up + ("\x1bOB", "\x1b[B"), + ("\x1bOC", "\x1b[C"), + ("\x1bOD", "\x1b[D"), # App Left -> CSI Left + ("\x7f", "\x08"), # DEL -> BS +) + + +def uses_network_cli_keymap(device_type: str = "", vendor: str = "") -> bool: + blob = f"{device_type} {vendor}".strip().lower() + if not blob: + return True + for token in ("linux", "ubuntu", "centos", "debian", "redhat", "unix", "generic_telnet", "generic"): + if token in blob: + return False + return True + + +def map_network_cli_keys( + data: str, + *, + device_type: str = "", + vendor: str = "", + protocol: str = "", +) -> str: + """Rewrite xterm key sequences for network-device CLIs.""" + del device_type, vendor, protocol # protocol kept for call-site compatibility + text = str(data or "") + if not text: + return text + out: list[str] = [] + i = 0 + n = len(text) + while i < n: + matched = False + for seq, repl in _NETWORK_CLI_KEY_SEQS: + if text.startswith(seq, i): + out.append(repl) + i += len(seq) + matched = True + break + if not matched: + out.append(text[i]) + i += 1 + return "".join(out) + + +def channel_return(conn: ConnectHandler | None) -> str: + """Netmiko line ending for this session (SSH usually \\n, Telnet often \\r\\n).""" + if conn is None: + return "\n" + ret = getattr(conn, "RETURN", None) + if isinstance(ret, str) and ret: + return ret + return "\n" + + +def map_network_cli_enter(data: str, conn: ConnectHandler | None) -> str: + """Map xterm Enter (\\r) to the device's Netmiko RETURN.""" + text = str(data or "") + if not text: + return text + ret = channel_return(conn) + if ret == "\r": + return text + # Prefer replacing CRLF first so Telnet RETURN \\r\\n does not double-expand. + return text.replace("\r\n", ret).replace("\r", ret) + + +def _drain_channel(conn: ConnectHandler, *, rounds: int = 6, wait: float = 0.06) -> str: + """Read whatever is already sitting on the channel after login.""" + chunks: list[str] = [] + empty_streak = 0 + for _ in range(max(1, rounds)): + time.sleep(wait) + try: + part = conn.read_channel() + except Exception: + break + if part: + chunks.append(str(part)) + empty_streak = 0 + else: + empty_streak += 1 + if empty_streak >= 2 and chunks: + break + return "".join(chunks) + + +def _session_log_text(buf: io.BytesIO | None) -> str: + """Decode Netmiko session_log buffer into display text.""" + if buf is None: + return "" + try: + raw = buf.getvalue() + except Exception: + return "" + if isinstance(raw, bytes): + return raw.decode("utf-8", errors="replace") + return str(raw or "") + + +def _looks_like_cli_prompt(text: str) -> bool: + s = str(text or "").rstrip() + if not s: + return False + # Buffer races can leave a stray ':' after Huawei ```` (from prior ``[Y/N]:``). + if s.endswith(":") and ">" in s: + s = s[:-1].rstrip() + # Common network CLI prompts: [HUAWEI] Router# Router> + return bool(re.search(r"(?:[>\]]|#)\s*$", s)) or bool(re.search(r"<[^>\r\n]+>\s*$", s)) + + +def _looks_like_login_prompt(text: str) -> bool: + """True when the transcript ends at Username:/Login:/Password: (interactive auth).""" + s = str(text or "").replace("\r\n", "\n").replace("\r", "\n") + lines = [ln.strip() for ln in s.split("\n") if ln.strip()] + if not lines: + return False + last = lines[-1] + return bool(re.search(r"(?i)(user\s*name|login|password)\s*:\s*$", last)) + + +def _looks_like_password_change_prompt(text: str) -> bool: + """Huawei/VRP post-auth ``Change now? [Y/N]:`` (Netmiko already answers N).""" + s = str(text or "").replace("\r\n", "\n").replace("\r", "\n") + lines = [ln.strip() for ln in s.split("\n") if ln.strip()] + if not lines: + return False + last = lines[-1] + return bool(re.search(r"(?i)(change\s*now|please\s*choose|password\s+needs\s+to\s+be\s+changed).{0,80}:\s*$", last)) or bool( + re.search(r"\[Y/N\]\s*:\s*$", last, flags=re.I) + ) + + +# Cisco/Netmiko often yields "R2#R2#" when a sync Enter is appended without a newline. +_GLUED_PROMPT_RE = re.compile(r"(?<=[#>])(?=(?:[A-Za-z0-9][\w.\-:]{0,62})[#>])") + + +def normalize_cli_transcript(text: str) -> str: + """Normalize login transcript for xterm (convertEol) and un-glue prompts.""" + s = str(text or "").replace("\r\n", "\n").replace("\r", "\n") + s = _GLUED_PROMPT_RE.sub("\n", s) + lines = s.split("\n") + while lines and not str(lines[-1]).strip(): + lines.pop() + # Drop blank lines immediately before a final prompt (banner\n\nR2# -> banner\nR2#). + while len(lines) >= 2 and not str(lines[-2]).strip() and _looks_like_cli_prompt(lines[-1]): + lines.pop(-2) + # Collapse trailing duplicate prompt lines (slow VMs often echo R2# several times). + while len(lines) >= 2 and str(lines[-1]).strip() == str(lines[-2]).strip() and _looks_like_cli_prompt(lines[-1]): + lines.pop() + return "\n".join(lines) + + +def prepare_bootstrap_output(text: str) -> str: + """Full login transcript for UI replay; keep final prompt, no trailing newline after it. + + Trailing newline would leave the cursor on a blank line so the first typed line + looks wrong; cursor should sit after the prompt like a real CRT. + """ + s = normalize_cli_transcript(text) + # Drop a stray ':' glued onto Huawei ```` after ``[Y/N]:`` buffer races. + s = re.sub(r"(<[^\r\n>]+>):\s*$", r"\1", s) + return s + + +def _capture_raw_channel(conn: ConnectHandler, *, duration: float = 0.5) -> str: + """Read leftover PTY bytes into text (banner/MOTD after SSH auth). + + Interactive WebCRT skips Netmiko session_preparation, so the post-auth banner + often never lands in ``session_log`` and must be pulled from the live channel. + """ + chunks: list[str] = [] + channel = getattr(conn, "remote_conn", None) + if channel is None: + try: + return _drain_channel(conn, rounds=max(2, int(duration / 0.05)), wait=0.05) + except Exception: + return "" + end = time.time() + max(0.1, float(duration)) + while time.time() < end: + got = False + try: + # Paramiko SSH channel + if hasattr(channel, "recv_ready") and hasattr(channel, "recv"): + if channel.recv_ready(): + raw = channel.recv(65535) + if raw: + got = True + if isinstance(raw, bytes): + chunks.append(raw.decode("utf-8", errors="replace")) + else: + chunks.append(str(raw)) + # telnetlib-style + elif callable(getattr(channel, "read_very_eager", None)): + data = channel.read_very_eager() + if data: + got = True + if isinstance(data, bytes): + chunks.append(data.decode("utf-8", errors="replace")) + else: + chunks.append(str(data)) + else: + part = conn.read_channel() + if part: + got = True + chunks.append(str(part)) + except Exception: + break + if not got: + time.sleep(0.04) + return "".join(chunks) + + +def _drain_raw_channel(conn: ConnectHandler, *, duration: float = 0.5) -> None: + """Discard leftover bytes on the live channel (SSH/Telnet) after login priming.""" + _capture_raw_channel(conn, duration=duration) + + +def _prime_interactive_channel(conn: ConnectHandler, *, already_prompted: bool = False) -> str: + """Sync interactive channel after login; return captured banner/prompt text. + + Skip the sync Enter when the login transcript already ends with a CLI prompt — + otherwise slow Cisco VMs accumulate duplicate ``R2#`` lines in the bootstrap. + """ + parts: list[str] = [] + try: + parts.append(_capture_raw_channel(conn, duration=0.25)) + except Exception: + pass + if not already_prompted: + try: + conn.write_channel(channel_return(conn)) + except Exception: + try: + conn.write_channel("\n") + except Exception: + return "".join(parts) + try: + parts.append(_drain_channel(conn, rounds=6, wait=0.08)) + except Exception: + pass + try: + parts.append(_capture_raw_channel(conn, duration=0.35)) + except Exception: + pass + return "".join(parts) + + +def _is_prompt_only_echo(text: str, prompt_hint: str = "") -> bool: + """True when chunk is only whitespace / CR / a repeated prompt (safe to drop after bootstrap).""" + s = str(text or "").replace("\r\n", "\n").replace("\r", "\n").strip() + if not s: + return True + hint = str(prompt_hint or "").strip() + if hint and s == hint: + return True + # Single-line prompt echo only. + if "\n" not in s and _looks_like_cli_prompt(s): + return True + if hint and all(line.strip() in ("", hint) for line in s.split("\n")): + return True + return False + + +def _normalize_encoding(name: str) -> str: + enc = str(name or "utf-8").strip().lower().replace("_", "-") + if enc in ("gbk", "gb2312", "gb18030", "cp936"): + return "gbk" + return "utf-8" + + +def _decode_bytes(data: bytes, encoding: str) -> str: + enc = _normalize_encoding(encoding) + try: + return data.decode(enc, errors="replace") + except Exception: + return data.decode("utf-8", errors="replace") + + +def _encode_text(text: str, encoding: str) -> bytes: + enc = _normalize_encoding(encoding) + try: + return text.encode(enc, errors="replace") + except Exception: + return text.encode("utf-8", errors="replace") + + +class _BoundedByteQueue: + """Thread-safe queue that drops oldest chunks when full (backpressure).""" + + def __init__(self, maxsize: int = 2000) -> None: + self._q: queue.Queue[bytes | None] = queue.Queue() + self._max = max(8, int(maxsize or 2000)) + self._cond = threading.Condition() + self.dropped = 0 + self._reported = 0 + + def put(self, item: bytes | None) -> None: + with self._cond: + while self._q.qsize() >= self._max: + try: + self._q.get_nowait() + self.dropped += 1 + except queue.Empty: + break + self._q.put(item) + self._cond.notify() + + def put_nowait(self, item: bytes | None) -> None: + self.put(item) + + def get_nowait(self) -> bytes | None: + with self._cond: + return self._q.get_nowait() + + def get(self, timeout: float = 0.25) -> bytes | None: + """Block until a chunk is available or timeout (raises queue.Empty).""" + deadline = time.time() + max(0.0, float(timeout)) + with self._cond: + while self._q.empty(): + remaining = deadline - time.time() + if remaining <= 0: + raise queue.Empty + self._cond.wait(timeout=remaining) + return self._q.get_nowait() + + def qsize(self) -> int: + with self._cond: + return self._q.qsize() + + def take_drop_delta(self) -> int: + """Return newly dropped chunk count since last call (for client notice).""" + with self._cond: + delta = int(self.dropped) - int(self._reported) + if delta <= 0: + return 0 + self._reported = int(self.dropped) + return delta + + +def _utc_now() -> datetime: + return datetime.now(timezone.utc) + + +def _utc_iso() -> str: + return _utc_now().isoformat() + + +def webcrt_data_root() -> Path: + root = Path(str(settings.webcrt_data_dir or "data/webcrt")) + root.mkdir(parents=True, exist_ok=True) + return root.resolve() + + +def _session_log_path(session_id: str) -> Path: + folder = webcrt_data_root() / "sessions" + folder.mkdir(parents=True, exist_ok=True) + return folder / f"{session_id}.log" + + +def read_session_log_tail(session_id: str, *, max_bytes: int = 49152) -> str: + """Best-effort UTF-8 tail of the on-disk session transcript (for WS re-attach).""" + path = _session_log_path(session_id) + try: + if not path.is_file(): + return "" + size = path.stat().st_size + take = max(1024, min(int(max_bytes or 49152), 256 * 1024)) + with path.open("rb") as fh: + if size > take: + fh.seek(size - take) + raw = fh.read() + # Drop partial first line after seek. + nl = raw.find(b"\n") + if 0 <= nl < len(raw) - 1: + raw = raw[nl + 1 :] + else: + raw = fh.read() + text = raw.decode("utf-8", errors="replace") + # Strip header comment lines from the visible replay. + lines = [ln for ln in text.splitlines(keepends=True) if not ln.startswith("# session=")] + return "".join(lines) + except Exception: + _log.debug("webcrt session log tail failed session=%s", session_id, exc_info=True) + return "" + + +def _audit(event: str, **fields: Any) -> None: + record = {"ts": _utc_iso(), "event": event, **fields} + try: + path = webcrt_data_root() / "audit.jsonl" + with path.open("a", encoding="utf-8") as fh: + fh.write(json.dumps(record, ensure_ascii=False) + "\n") + except Exception: + _log.exception("webcrt audit write failed") + _log.info("webcrt.%s %s", event, {k: v for k, v in fields.items() if k != "detail"}) diff --git a/netx_api/webcrt_service.py b/netx_api/webcrt_service.py index 6474c7a..34cae9c 100644 --- a/netx_api/webcrt_service.py +++ b/netx_api/webcrt_service.py @@ -1,1501 +1,54 @@ -"""Interactive WebCRT sessions: bridge browser WebSocket <-> Netmiko device channel.""" - +"""WebCRT facade — re-exports channel helpers and session registry.""" from __future__ import annotations -import io -import json -import logging -import queue -import re -import threading -import time -import uuid -from dataclasses import dataclass, field -from datetime import datetime, timezone -from pathlib import Path -from typing import Any - -from fastapi import HTTPException -from netmiko import ConnectHandler -from sqlalchemy.orm import Session - -from .config import settings -from .ne_crypto import CredentialCryptoError -from .ne_session_factory import ( - close_netmiko_connection, - extract_cli_prompt_marker, - get_cli_hop_guard, - open_netmiko_connection, - should_close_cli_hop_session, +from .webcrt_channel import ( + _decode_bytes, + _encode_text, + _normalize_encoding, + channel_return, + map_network_cli_enter, + map_network_cli_keys, + normalize_cli_transcript, + prepare_bootstrap_output, + read_session_log_tail, + uses_network_cli_keymap, + webcrt_data_root, +) +from .webcrt_session import ( + WebcrtSession, + _webcrt_creds_ready, + active_session_count, + close_session, + create_session, + detach_session, + find_ssh_session_for_ne, + get_session, + list_sessions, + mark_attached, + wait_session_ready, ) -_log = logging.getLogger("netx.webcrt") - -_sessions_lock = threading.Lock() -_sessions: dict[str, "WebcrtSession"] = {} -_reaper_started = False - -# Network device CLI key rewrites (SecureCRT-like). -# - Backspace: DEL(0x7f) -> BS(0x08) -# - Home/End/Delete: emacs controls (widely accepted) -# - Arrows: after login many boxes enable DECCKM (application cursor), so xterm -# sends SS3 forms (\x1bOD) which VRP/IOS ignore; normalize SS3 -> CSI and -# pass CSI through. Do not rewrite arrows to Ctrl-B/F — that leaves the -# device cursor stuck at EOL when SS3 was what actually arrived. -_NETWORK_CLI_KEY_SEQS: tuple[tuple[str, str], ...] = ( - ("\x1b[1~", "\x01"), # Home -> Ctrl-A - ("\x1b[3~", "\x04"), # Delete key -> Ctrl-D - ("\x1b[4~", "\x05"), # End -> Ctrl-E - ("\x1b[H", "\x01"), - ("\x1b[F", "\x05"), - ("\x1bOH", "\x01"), - ("\x1bOF", "\x05"), - ("\x1bOA", "\x1b[A"), # App Up -> CSI Up - ("\x1bOB", "\x1b[B"), - ("\x1bOC", "\x1b[C"), - ("\x1bOD", "\x1b[D"), # App Left -> CSI Left - ("\x7f", "\x08"), # DEL -> BS -) - - -def uses_network_cli_keymap(device_type: str = "", vendor: str = "") -> bool: - blob = f"{device_type} {vendor}".strip().lower() - if not blob: - return True - for token in ("linux", "ubuntu", "centos", "debian", "redhat", "unix", "generic_telnet", "generic"): - if token in blob: - return False - return True - - -def map_network_cli_keys( - data: str, - *, - device_type: str = "", - vendor: str = "", - protocol: str = "", -) -> str: - """Rewrite xterm key sequences for network-device CLIs.""" - del device_type, vendor, protocol # protocol kept for call-site compatibility - text = str(data or "") - if not text: - return text - out: list[str] = [] - i = 0 - n = len(text) - while i < n: - matched = False - for seq, repl in _NETWORK_CLI_KEY_SEQS: - if text.startswith(seq, i): - out.append(repl) - i += len(seq) - matched = True - break - if not matched: - out.append(text[i]) - i += 1 - return "".join(out) - - -def channel_return(conn: ConnectHandler | None) -> str: - """Netmiko line ending for this session (SSH usually \\n, Telnet often \\r\\n).""" - if conn is None: - return "\n" - ret = getattr(conn, "RETURN", None) - if isinstance(ret, str) and ret: - return ret - return "\n" - - -def map_network_cli_enter(data: str, conn: ConnectHandler | None) -> str: - """Map xterm Enter (\\r) to the device's Netmiko RETURN.""" - text = str(data or "") - if not text: - return text - ret = channel_return(conn) - if ret == "\r": - return text - # Prefer replacing CRLF first so Telnet RETURN \\r\\n does not double-expand. - return text.replace("\r\n", ret).replace("\r", ret) - - -def _drain_channel(conn: ConnectHandler, *, rounds: int = 6, wait: float = 0.06) -> str: - """Read whatever is already sitting on the channel after login.""" - chunks: list[str] = [] - empty_streak = 0 - for _ in range(max(1, rounds)): - time.sleep(wait) - try: - part = conn.read_channel() - except Exception: - break - if part: - chunks.append(str(part)) - empty_streak = 0 - else: - empty_streak += 1 - if empty_streak >= 2 and chunks: - break - return "".join(chunks) - - -def _session_log_text(buf: io.BytesIO | None) -> str: - """Decode Netmiko session_log buffer into display text.""" - if buf is None: - return "" - try: - raw = buf.getvalue() - except Exception: - return "" - if isinstance(raw, bytes): - return raw.decode("utf-8", errors="replace") - return str(raw or "") - - -def _looks_like_cli_prompt(text: str) -> bool: - s = str(text or "").rstrip() - if not s: - return False - # Buffer races can leave a stray ':' after Huawei ```` (from prior ``[Y/N]:``). - if s.endswith(":") and ">" in s: - s = s[:-1].rstrip() - # Common network CLI prompts: [HUAWEI] Router# Router> - return bool(re.search(r"(?:[>\]]|#)\s*$", s)) or bool(re.search(r"<[^>\r\n]+>\s*$", s)) - - -def _looks_like_login_prompt(text: str) -> bool: - """True when the transcript ends at Username:/Login:/Password: (interactive auth).""" - s = str(text or "").replace("\r\n", "\n").replace("\r", "\n") - lines = [ln.strip() for ln in s.split("\n") if ln.strip()] - if not lines: - return False - last = lines[-1] - return bool(re.search(r"(?i)(user\s*name|login|password)\s*:\s*$", last)) - - -def _looks_like_password_change_prompt(text: str) -> bool: - """Huawei/VRP post-auth ``Change now? [Y/N]:`` (Netmiko already answers N).""" - s = str(text or "").replace("\r\n", "\n").replace("\r", "\n") - lines = [ln.strip() for ln in s.split("\n") if ln.strip()] - if not lines: - return False - last = lines[-1] - return bool(re.search(r"(?i)(change\s*now|please\s*choose|password\s+needs\s+to\s+be\s+changed).{0,80}:\s*$", last)) or bool( - re.search(r"\[Y/N\]\s*:\s*$", last, flags=re.I) - ) - - -# Cisco/Netmiko often yields "R2#R2#" when a sync Enter is appended without a newline. -_GLUED_PROMPT_RE = re.compile(r"(?<=[#>])(?=(?:[A-Za-z0-9][\w.\-:]{0,62})[#>])") - - -def normalize_cli_transcript(text: str) -> str: - """Normalize login transcript for xterm (convertEol) and un-glue prompts.""" - s = str(text or "").replace("\r\n", "\n").replace("\r", "\n") - s = _GLUED_PROMPT_RE.sub("\n", s) - lines = s.split("\n") - while lines and not str(lines[-1]).strip(): - lines.pop() - # Drop blank lines immediately before a final prompt (banner\n\nR2# -> banner\nR2#). - while len(lines) >= 2 and not str(lines[-2]).strip() and _looks_like_cli_prompt(lines[-1]): - lines.pop(-2) - # Collapse trailing duplicate prompt lines (slow VMs often echo R2# several times). - while len(lines) >= 2 and str(lines[-1]).strip() == str(lines[-2]).strip() and _looks_like_cli_prompt(lines[-1]): - lines.pop() - return "\n".join(lines) - - -def prepare_bootstrap_output(text: str) -> str: - """Full login transcript for UI replay; keep final prompt, no trailing newline after it. - - Trailing newline would leave the cursor on a blank line so the first typed line - looks wrong; cursor should sit after the prompt like a real CRT. - """ - s = normalize_cli_transcript(text) - # Drop a stray ':' glued onto Huawei ```` after ``[Y/N]:`` buffer races. - s = re.sub(r"(<[^\r\n>]+>):\s*$", r"\1", s) - return s - - -def _capture_raw_channel(conn: ConnectHandler, *, duration: float = 0.5) -> str: - """Read leftover PTY bytes into text (banner/MOTD after SSH auth). - - Interactive WebCRT skips Netmiko session_preparation, so the post-auth banner - often never lands in ``session_log`` and must be pulled from the live channel. - """ - chunks: list[str] = [] - channel = getattr(conn, "remote_conn", None) - if channel is None: - try: - return _drain_channel(conn, rounds=max(2, int(duration / 0.05)), wait=0.05) - except Exception: - return "" - end = time.time() + max(0.1, float(duration)) - while time.time() < end: - got = False - try: - # Paramiko SSH channel - if hasattr(channel, "recv_ready") and hasattr(channel, "recv"): - if channel.recv_ready(): - raw = channel.recv(65535) - if raw: - got = True - if isinstance(raw, bytes): - chunks.append(raw.decode("utf-8", errors="replace")) - else: - chunks.append(str(raw)) - # telnetlib-style - elif callable(getattr(channel, "read_very_eager", None)): - data = channel.read_very_eager() - if data: - got = True - if isinstance(data, bytes): - chunks.append(data.decode("utf-8", errors="replace")) - else: - chunks.append(str(data)) - else: - part = conn.read_channel() - if part: - got = True - chunks.append(str(part)) - except Exception: - break - if not got: - time.sleep(0.04) - return "".join(chunks) - - -def _drain_raw_channel(conn: ConnectHandler, *, duration: float = 0.5) -> None: - """Discard leftover bytes on the live channel (SSH/Telnet) after login priming.""" - _capture_raw_channel(conn, duration=duration) - - -def _prime_interactive_channel(conn: ConnectHandler, *, already_prompted: bool = False) -> str: - """Sync interactive channel after login; return captured banner/prompt text. - - Skip the sync Enter when the login transcript already ends with a CLI prompt — - otherwise slow Cisco VMs accumulate duplicate ``R2#`` lines in the bootstrap. - """ - parts: list[str] = [] - try: - parts.append(_capture_raw_channel(conn, duration=0.25)) - except Exception: - pass - if not already_prompted: - try: - conn.write_channel(channel_return(conn)) - except Exception: - try: - conn.write_channel("\n") - except Exception: - return "".join(parts) - try: - parts.append(_drain_channel(conn, rounds=6, wait=0.08)) - except Exception: - pass - try: - parts.append(_capture_raw_channel(conn, duration=0.35)) - except Exception: - pass - return "".join(parts) - - -def _is_prompt_only_echo(text: str, prompt_hint: str = "") -> bool: - """True when chunk is only whitespace / CR / a repeated prompt (safe to drop after bootstrap).""" - s = str(text or "").replace("\r\n", "\n").replace("\r", "\n").strip() - if not s: - return True - hint = str(prompt_hint or "").strip() - if hint and s == hint: - return True - # Single-line prompt echo only. - if "\n" not in s and _looks_like_cli_prompt(s): - return True - if hint and all(line.strip() in ("", hint) for line in s.split("\n")): - return True - return False - - -def _normalize_encoding(name: str) -> str: - enc = str(name or "utf-8").strip().lower().replace("_", "-") - if enc in ("gbk", "gb2312", "gb18030", "cp936"): - return "gbk" - return "utf-8" - - -def _decode_bytes(data: bytes, encoding: str) -> str: - enc = _normalize_encoding(encoding) - try: - return data.decode(enc, errors="replace") - except Exception: - return data.decode("utf-8", errors="replace") - - -def _encode_text(text: str, encoding: str) -> bytes: - enc = _normalize_encoding(encoding) - try: - return text.encode(enc, errors="replace") - except Exception: - return text.encode("utf-8", errors="replace") - - -class _BoundedByteQueue: - """Thread-safe queue that drops oldest chunks when full (backpressure).""" - - def __init__(self, maxsize: int = 2000) -> None: - self._q: queue.Queue[bytes | None] = queue.Queue() - self._max = max(8, int(maxsize or 2000)) - self._cond = threading.Condition() - self.dropped = 0 - self._reported = 0 - - def put(self, item: bytes | None) -> None: - with self._cond: - while self._q.qsize() >= self._max: - try: - self._q.get_nowait() - self.dropped += 1 - except queue.Empty: - break - self._q.put(item) - self._cond.notify() - - def put_nowait(self, item: bytes | None) -> None: - self.put(item) - - def get_nowait(self) -> bytes | None: - with self._cond: - return self._q.get_nowait() - - def get(self, timeout: float = 0.25) -> bytes | None: - """Block until a chunk is available or timeout (raises queue.Empty).""" - deadline = time.time() + max(0.0, float(timeout)) - with self._cond: - while self._q.empty(): - remaining = deadline - time.time() - if remaining <= 0: - raise queue.Empty - self._cond.wait(timeout=remaining) - return self._q.get_nowait() - - def qsize(self) -> int: - with self._cond: - return self._q.qsize() - - def take_drop_delta(self) -> int: - """Return newly dropped chunk count since last call (for client notice).""" - with self._cond: - delta = int(self.dropped) - int(self._reported) - if delta <= 0: - return 0 - self._reported = int(self.dropped) - return delta - - -def _utc_now() -> datetime: - return datetime.now(timezone.utc) - - -def _utc_iso() -> str: - return _utc_now().isoformat() - - -def webcrt_data_root() -> Path: - root = Path(str(settings.webcrt_data_dir or "data/webcrt")) - root.mkdir(parents=True, exist_ok=True) - return root.resolve() - - -def _session_log_path(session_id: str) -> Path: - folder = webcrt_data_root() / "sessions" - folder.mkdir(parents=True, exist_ok=True) - return folder / f"{session_id}.log" - - -def read_session_log_tail(session_id: str, *, max_bytes: int = 49152) -> str: - """Best-effort UTF-8 tail of the on-disk session transcript (for WS re-attach).""" - path = _session_log_path(session_id) - try: - if not path.is_file(): - return "" - size = path.stat().st_size - take = max(1024, min(int(max_bytes or 49152), 256 * 1024)) - with path.open("rb") as fh: - if size > take: - fh.seek(size - take) - raw = fh.read() - # Drop partial first line after seek. - nl = raw.find(b"\n") - if 0 <= nl < len(raw) - 1: - raw = raw[nl + 1 :] - else: - raw = fh.read() - text = raw.decode("utf-8", errors="replace") - # Strip header comment lines from the visible replay. - lines = [ln for ln in text.splitlines(keepends=True) if not ln.startswith("# session=")] - return "".join(lines) - except Exception: - _log.debug("webcrt session log tail failed session=%s", session_id, exc_info=True) - return "" - - -def _audit(event: str, **fields: Any) -> None: - record = {"ts": _utc_iso(), "event": event, **fields} - try: - path = webcrt_data_root() / "audit.jsonl" - with path.open("a", encoding="utf-8") as fh: - fh.write(json.dumps(record, ensure_ascii=False) + "\n") - except Exception: - _log.exception("webcrt audit write failed") - _log.info("webcrt.%s %s", event, {k: v for k, v in fields.items() if k != "detail"}) - - -@dataclass -class WebcrtSession: - session_id: str - ne_id: str - ne_name: str - ne_ip: str - protocol: str - cols: int - rows: int - device_type: str = "" - vendor: str = "" - cli_keymap: bool = True - encoding: str = "utf-8" - keepalive_sec: int = 0 - conn: ConnectHandler | None = None - created_at: float = field(default_factory=time.time) - last_activity: float = field(default_factory=time.time) - attached: bool = False - detach_deadline: float | None = None - closed: bool = False - close_reason: str = "" - state: str = "connecting" - connect_error: str = "" - connect_started_at: float = field(default_factory=time.time) - connect_finished_at: float | None = None - bootstrap_output: bytes = b"" - # First WS attach gets login bootstrap; later attaches prefer session-log tail. - bootstrap_replayed: bool = False - needs_live_prompt: bool = True - # React StrictMode remounts open a second WS before the first fully tears down. - # Only the newest attach_gen may consume out_queue / mark detach. - attach_gen: int = 0 - out_queue: _BoundedByteQueue = field( - default_factory=lambda: _BoundedByteQueue(int(getattr(settings, "webcrt_out_queue_max", 2000) or 2000)) - ) - # Vendor CLI hop (Huawei/ZTE/Cisco): close when nested target session returns to hop. - cli_hop_guard: bool = False - cli_hop_prompt: str = "" - post_login_commands: list[str] = field(default_factory=list) - bytes_in: int = 0 - bytes_out: int = 0 - _reader: threading.Thread | None = field(default=None, repr=False) - _write_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) - _stdout_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) - _hop_scan_buf: str = field(default="", repr=False) - _cli_hop_seen_other_prompt: bool = field(default=False, repr=False) - _log_fh: Any = field(default=None, repr=False) - _ready_event: threading.Event = field(default_factory=threading.Event, repr=False) - # SFTP channel on the same SSH transport as the interactive shell (direct SSH only). - sftp_ready: bool = False - _sftp: Any = field(default=None, repr=False) - _sftp_lock: threading.RLock = field(default_factory=threading.RLock, repr=False) - - def touch(self) -> None: - self.last_activity = time.time() - - def close_sftp(self) -> None: - with self._sftp_lock: - sftp = self._sftp - self._sftp = None - self.sftp_ready = False - if sftp is None: - return - try: - sftp.close() - except Exception: - pass - - def _ssh_transport_unlocked(self) -> Any: - """Caller must hold ``_sftp_lock``. Returns an active Paramiko Transport.""" - if self.closed or self.conn is None: - raise RuntimeError("session_closed") - if str(self.protocol or "ssh").lower() != "ssh": - raise RuntimeError("sftp_requires_ssh") - if self.cli_hop_guard: - raise RuntimeError("sftp_hop_not_supported") - channel = getattr(self.conn, "remote_conn", None) - transport = None - if channel is not None and hasattr(channel, "get_transport"): - try: - transport = channel.get_transport() - except Exception: - transport = None - if transport is None or not bool(getattr(transport, "is_active", lambda: False)()): - raise RuntimeError("ssh_transport_unavailable") - return transport - - def _ensure_sftp_unlocked(self) -> Any: - """Caller must hold ``_sftp_lock``. Shared probe client (sftp_ready).""" - import paramiko - - if self._sftp is not None: - sock = getattr(self._sftp, "sock", None) - if sock is not None and not bool(getattr(sock, "closed", False)): - return self._sftp - try: - self._sftp.close() - except Exception: - pass - self._sftp = None - transport = self._ssh_transport_unlocked() - self._sftp = paramiko.SFTPClient.from_transport(transport) - if self._sftp is None: - raise RuntimeError("sftp_open_failed") - self.sftp_ready = True - return self._sftp - - def open_sftp(self) -> Any: - """Open/reuse an SFTP client on this session's SSH transport.""" - with self._sftp_lock: - return self._ensure_sftp_unlocked() - - def open_ephemeral_sftp(self) -> Any: - """Open a dedicated SFTP channel for one operation; caller must ``close()`` it. - - Only holds ``_sftp_lock`` briefly while resolving the SSH transport, so long - list/upload/download work does not block other SFTP ops on the same session. - """ - import paramiko - - with self._sftp_lock: - transport = self._ssh_transport_unlocked() - # Keep probe client warm for UI sftp_ready without sharing it for I/O. - try: - self._ensure_sftp_unlocked() - except Exception: - pass - sftp = paramiko.SFTPClient.from_transport(transport) - if sftp is None: - raise RuntimeError("sftp_open_failed") - return sftp - - def run_sftp(self, fn: Any) -> Any: - """Run ``fn(sftp)`` on an ephemeral channel (does not hold the lock during ``fn``).""" - sftp = self.open_ephemeral_sftp() - try: - return fn(sftp) - finally: - try: - sftp.close() - except Exception: - pass - - def try_attach_sftp(self) -> bool: - """Best-effort SFTP channel open after SSH login (does not fail the shell).""" - if str(self.protocol or "ssh").lower() != "ssh" or self.cli_hop_guard: - self.sftp_ready = False - return False - try: - self.open_sftp() - self.sftp_ready = True - return True - except Exception: - self.sftp_ready = False - _log.debug("webcrt sftp attach skipped session=%s", self.session_id, exc_info=True) - return False - - def open_session_log(self) -> None: - if not bool(getattr(settings, "webcrt_session_log_enabled", True)): - return - if self._log_fh is not None: - return - try: - self._log_fh = _session_log_path(self.session_id).open("a", encoding="utf-8", errors="replace") - self._log_fh.write(f"# session={self.session_id} ne={self.ne_id} ip={self.ne_ip} ts={_utc_iso()}\n") - self._log_fh.flush() - except Exception: - _log.debug("webcrt session log open failed", exc_info=True) - self._log_fh = None - - def append_session_log(self, text: str) -> None: - if not text or self._log_fh is None: - return - try: - self._log_fh.write(text) - self._log_fh.flush() - except Exception: - pass - - def close_session_log(self) -> None: - fh = self._log_fh - self._log_fh = None - if fh is None: - return - try: - fh.write(f"\n# closed reason={self.close_reason} ts={_utc_iso()}\n") - fh.close() - except Exception: - pass - - def take_stdout(self, attach_gen: int, *, timeout: float = 0.25) -> bytes | None | str: - """Exclusive stdout take for one WS attach generation. - - Returns: - bytes — device output chunk - None — device reader closed (session end) - \"stale\" — a newer WebSocket owns this session; caller must stop - \"empty\" — no data within timeout (keep polling) - """ - deadline = time.time() + max(0.05, float(timeout)) - while True: - with self._stdout_lock: - if attach_gen != self.attach_gen: - return "stale" - remaining = deadline - time.time() - if remaining <= 0: - return "empty" - # Slice waits so we can notice attach_gen bumps without busy-spinning. - try: - chunk = self.out_queue.get(timeout=min(0.05, remaining)) - except queue.Empty: - continue - with self._stdout_lock: - if attach_gen != self.attach_gen: - # Put back including EOF sentinel so the new owner still sees close. - self.out_queue.put(chunk) - return "stale" - return chunk # bytes | None - - def write_stdin(self, data: str) -> None: - if self.closed or self.conn is None: - raise RuntimeError("session_closed") - text = str(data or "") - if not text: - return - if self.cli_keymap: - text = map_network_cli_keys( - text, - device_type=self.device_type, - vendor=self.vendor, - protocol=self.protocol, - ) - text = map_network_cli_enter(text, self.conn) - if not text: - return - with self._write_lock: - # Prefer raw channel I/O for interactive typing (char echo / backspace). - channel = getattr(self.conn, "remote_conn", None) - try: - if channel is not None and hasattr(channel, "send") and callable(channel.send): - payload = _encode_text(text, self.encoding) - # Paramiko may write partially when the window is full. - view = memoryview(payload) - while len(view): - n = int(channel.send(view) or 0) - if n <= 0: - time.sleep(0.01) - continue - view = view[n:] - self.bytes_in += len(payload) - elif channel is not None and hasattr(channel, "write") and callable(channel.write): - payload = _encode_text(text, self.encoding) - channel.write(payload) - self.bytes_in += len(payload) - else: - self.conn.write_channel(text) - self.bytes_in += len(text) - except Exception: - self.conn.write_channel(text) - self.bytes_in += len(text) - self.touch() - - def send_break(self) -> None: - """Send SSH break / Telnet IAC BREAK to interrupt paging or hung commands.""" - if self.closed or self.conn is None: - raise RuntimeError("session_closed") - channel = getattr(self.conn, "remote_conn", None) - with self._write_lock: - sent = False - if channel is not None and hasattr(channel, "send_break") and callable(channel.send_break): - try: - channel.send_break(0) - sent = True - except Exception: - _log.debug("send_break failed session=%s", self.session_id, exc_info=True) - if not sent and channel is not None and hasattr(channel, "send") and callable(channel.send): - # Telnet IAC BREAK = 255 243 - try: - channel.send(b"\xff\xf3") - sent = True - except Exception: - pass - if not sent: - # Fallback: Ctrl-C often interrupts device CLI more-pages. - try: - self.conn.write_channel("\x03") - except Exception: - raise RuntimeError("break_failed") - self.touch() - - def resize(self, cols: int, rows: int) -> None: - if self.closed or self.conn is None: - return - c = max(20, min(500, int(cols or 80))) - r = max(5, min(200, int(rows or 24))) - self.cols = c - self.rows = r - channel = getattr(self.conn, "remote_conn", None) - if channel is not None and hasattr(channel, "resize_pty"): - try: - channel.resize_pty(width=c, height=r) - except Exception: - _log.debug("resize_pty failed session=%s", self.session_id, exc_info=True) - self.touch() - - def start_reader(self) -> None: - if self._reader and self._reader.is_alive(): - return - self._reader = threading.Thread( - target=self._reader_loop, - name=f"webcrt-reader-{self.session_id[:8]}", - daemon=True, - ) - self._reader.start() - - def _reader_loop(self) -> None: - conn = self.conn - if conn is None: - self.out_queue.put(None) - return - channel = getattr(conn, "remote_conn", None) - hop_return = False - poll = max(0.002, float(getattr(settings, "webcrt_reader_poll_sec", 0.01) or 0.01)) - try: - while not self.closed: - chunk = b"" - try: - if channel is not None and hasattr(channel, "recv_ready") and hasattr(channel, "recv"): - # Paramiko SSH: prefer short blocking recv over fixed spin-sleep. - ready = False - try: - ready = bool(channel.recv_ready()) - except Exception: - ready = False - if ready: - chunk = channel.recv(16384) - if not chunk: - break - elif hasattr(channel, "exit_status_ready") and channel.exit_status_ready(): - break - else: - # Brief block: settimeout + recv wakes sooner than sleep(0.04). - prev_timeout = None - try: - prev_timeout = channel.gettimeout() - except Exception: - prev_timeout = None - try: - channel.settimeout(poll) - chunk = channel.recv(16384) - except Exception: - chunk = b"" - finally: - try: - channel.settimeout(prev_timeout) - except Exception: - pass - if not chunk: - continue - elif channel is not None and hasattr(channel, "read_very_eager"): - # Telnet: do NOT use conn.read_channel() — Netmiko strips ANSI. - data = channel.read_very_eager() - if data: - chunk = ( - data - if isinstance(data, (bytes, bytearray)) - else _encode_text(str(data), self.encoding) - ) - else: - time.sleep(poll) - continue - else: - text = conn.read_channel() - if text: - chunk = _encode_text(str(text), self.encoding) - else: - time.sleep(poll) - continue - except Exception as exc: - if self.closed: - break - _log.debug("webcrt reader error session=%s: %s", self.session_id, exc) - time.sleep(0.05) - continue - if chunk: - self.touch() - self.bytes_out += len(chunk) - self.out_queue.put(chunk) - try: - self.append_session_log(_decode_bytes(chunk, self.encoding)) - except Exception: - pass - if self.cli_hop_guard and self._note_cli_hop_output(chunk): - hop_return = True - notice = ( - "\r\n*** WebCRT: 目标会话已结束,已断开代理连接 " - "(target session ended; closing hop proxy) ***\r\n" - ) - self.out_queue.put(notice.encode("utf-8", errors="replace")) - break - finally: - if hop_return and not self.closed: - try: - close_session(self.session_id, reason="cli_hop_return") - except Exception: - self.close("cli_hop_return") - self.out_queue.put(None) - - def _note_cli_hop_output(self, chunk: bytes) -> bool: - """Accumulate stdout and return True when nested CLI hop has returned to proxy.""" - try: - text = _decode_bytes(chunk, self.encoding) - except Exception: - text = str(chunk) - self._hop_scan_buf = (self._hop_scan_buf + text)[-12000:] - marker = str(self.cli_hop_prompt or "").strip() - last = extract_cli_prompt_marker(self._hop_scan_buf) - if last and (not marker or last != marker): - self._cli_hop_seen_other_prompt = True - return should_close_cli_hop_session( - self._hop_scan_buf, - self.cli_hop_prompt, - seen_other_prompt=self._cli_hop_seen_other_prompt, - ) - - def run_post_login_commands(self) -> None: - cmds = [str(c).rstrip("\r\n") for c in (self.post_login_commands or []) if str(c).strip()] - if not cmds or self.closed or self.conn is None: - return - for cmd in cmds[:20]: - try: - self.write_stdin(cmd + "\r") - time.sleep(0.15) - except Exception: - _log.debug("post_login command failed session=%s", self.session_id, exc_info=True) - break - - def close(self, reason: str = "closed") -> None: - if self.closed: - return - self.closed = True - self.state = "closed" - self.close_reason = reason or "closed" - self._ready_event.set() - self.close_sftp() - try: - close_netmiko_connection(self.conn) - except Exception: - pass - self.conn = None - try: - self.out_queue.put_nowait(None) - except Exception: - pass - self.close_session_log() - - -def _ensure_reaper() -> None: - global _reaper_started - with _sessions_lock: - if _reaper_started: - return - _reaper_started = True - t = threading.Thread(target=_reaper_loop, name="webcrt-reaper", daemon=True) - t.start() - - -def _reaper_loop() -> None: - while True: - try: - _reap_sessions() - except Exception: - _log.exception("webcrt reaper failed") - time.sleep(2) - - -def _reap_sessions() -> None: - idle = max(60, int(settings.webcrt_idle_timeout_sec or 1800)) - attach = max(10, int(settings.webcrt_attach_timeout_sec or 60)) - anti_idle = max(0, int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0)) - anti_payload = str(getattr(settings, "webcrt_anti_idle_payload", " ") or " ") - now = time.time() - to_close: list[tuple[WebcrtSession, str]] = [] - to_nudge: list[WebcrtSession] = [] - with _sessions_lock: - for sess in list(_sessions.values()): - if sess.closed: - _sessions.pop(sess.session_id, None) - continue - if sess.state == "connecting": - # Connecting sessions use connect timeout, not attach timeout alone. - connect_budget = max(30, int(settings.webcrt_connect_timeout_sec or 90)) + 30 - if (now - sess.connect_started_at) > connect_budget: - to_close.append((sess, "connect_timeout")) - continue - if sess.attached: - if (now - sess.last_activity) > idle: - to_close.append((sess, "idle_timeout")) - elif ( - anti_idle > 0 - and sess.state == "ready" - and sess.conn is not None - and (now - sess.last_activity) >= anti_idle - ): - to_nudge.append(sess) - continue - # Not attached: either never attached, or briefly detached for reconnect. - if sess.detach_deadline is not None: - if now >= sess.detach_deadline: - to_close.append((sess, "detach_timeout")) - else: - # Start attach clock after connect finishes (not HTTP create time), - # so slow auth + UI mount does not race attach_timeout. - anchor = float(sess.connect_finished_at or sess.created_at or now) - if (now - anchor) > attach: - to_close.append((sess, "attach_timeout")) - elif (now - sess.last_activity) > idle: - to_close.append((sess, "idle_timeout")) - for sess in to_nudge: - try: - # Touch without changing visible prompt when payload is empty/null-ish. - payload = anti_payload - if payload == "\\0": - payload = "\x00" - if payload: - sess.write_stdin(payload) - else: - sess.touch() - except Exception: - _log.debug("webcrt anti-idle failed session=%s", sess.session_id, exc_info=True) - for sess, reason in to_close: - close_session(sess.session_id, reason=reason) - - -def active_session_count() -> int: - with _sessions_lock: - return sum(1 for s in _sessions.values() if not s.closed) - - -def get_session(session_id: str) -> WebcrtSession | None: - with _sessions_lock: - sess = _sessions.get(session_id) - if sess is None or sess.closed: - return None - return sess - - -def find_ssh_session_for_ne(ne_id: str) -> WebcrtSession | None: - """Prefer a ready/attached interactive SSH session for SFTP channel reuse.""" - nid = str(ne_id or "").strip() - if not nid: - return None - with _sessions_lock: - candidates = [ - s - for s in _sessions.values() - if (not s.closed) - and str(s.ne_id) == nid - and str(s.protocol or "ssh").lower() == "ssh" - and s.conn is not None - and s.state in ("ready", "connecting") - and not s.cli_hop_guard - ] - if not candidates: - return None - # Prefer attached + ready sessions. - candidates.sort(key=lambda s: (0 if s.attached and s.state == "ready" else 1, -s.last_activity)) - return candidates[0] - - -def wait_session_ready(session_id: str, *, timeout: float = 120.0) -> WebcrtSession: - """Block until async connect finishes (ready or error). Used by tests and WS.""" - deadline = time.time() + max(1.0, float(timeout)) - while time.time() < deadline: - sess = get_session(session_id) - if sess is None: - raise HTTPException(status_code=404, detail="webcrt_session_not_found") - if sess.state == "ready": - return sess - if sess.state == "error": - raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") - sess._ready_event.wait(timeout=0.25) - raise HTTPException(status_code=504, detail="connect_timeout") - - -def _webcrt_creds_ready(creds: dict[str, Any]) -> bool: - """True when WebCRT can open a session with the resolved credentials. - - Bastion-managed hops store the target password on the bastion side, so an empty - NE password is valid (same as connectivity test). Direct / manual / Linux hops - still require a target password for SSH. - - Telnet (no hop) allows empty username/password so the user can authenticate - interactively in the terminal (SecureCRT-style). - """ - hop_enabled = bool(creds.get("hop_enabled")) - hop_vendor = str(creds.get("hop_vendor") or "").strip().lower() - auth_mode = str(creds.get("hop_target_auth_mode") or "bastion_managed").strip().lower() - protocol = str(creds.get("protocol") or "ssh").strip().lower() - if hop_enabled and hop_vendor == "bastion" and auth_mode == "bastion_managed": - return bool( - str(creds.get("hop_host") or "").strip() - and str(creds.get("hop_username") or "").strip() - and str(creds.get("hop_password") or "") - ) - if protocol == "telnet" and not hop_enabled: - return True - if not str(creds.get("username") or "").strip(): - return False - return bool(str(creds.get("password") or "")) - - -def _finish_connect( - sess: WebcrtSession, - *, - creds: dict[str, Any], - device: dict[str, Any], - connect_timeout: int, - client: str, -) -> None: - log_buf = io.BytesIO() - try: - conn = open_netmiko_connection( - creds, - session_timeout=connect_timeout, - session_log=log_buf, - cols=sess.cols, - rows=sess.rows, - interactive=True, - keepalive=int(sess.keepalive_sec or 0), - ) - except Exception as exc: - partial = _session_log_text(log_buf).strip() - from .ne_cli_errors import format_cli_failure - - classified = format_cli_failure(exc, partial) - detail = f"connect_failed:{classified}" - if partial: - detail = f"{detail}\n--- device transcript ---\n{partial[-4000:]}" - sess.state = "error" - sess.connect_error = detail - sess.connect_finished_at = time.time() - sess._ready_event.set() - _audit( - "session_open_failed", - session_id=sess.session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - source=str(device.get("source") or ""), - client=client or "", - error=str(exc)[:500], - transcript_len=len(partial), - ) - return - - channel = getattr(conn, "remote_conn", None) - if channel is not None and hasattr(channel, "resize_pty"): - try: - channel.resize_pty(width=sess.cols, height=sess.rows) - except Exception: - pass - - pre_log = _session_log_text(log_buf) - # Pull post-auth banner/MOTD from the PTY. With interactive no-op session_preparation - # (generic_termserver), Netmiko session_log is often empty — do not discard these bytes. - try: - early = _capture_raw_channel(conn, duration=0.35) - except Exception: - early = "" - seed = f"{pre_log}{early}" - already_prompted = _looks_like_cli_prompt(seed) - primed = "" - # Do not send Enter at Username:/Password: or Huawei password-change [Y/N]: - # (Netmiko telnet_login already answers password-change with "N"). - if _looks_like_login_prompt(seed) or _looks_like_password_change_prompt(seed): - try: - primed = _capture_raw_channel(conn, duration=0.9) - except Exception: - primed = "" - else: - try: - primed = _prime_interactive_channel(conn, already_prompted=already_prompted) - except Exception: - primed = "" - combined = f"{seed}{primed}" - # Final settle: keep stragglers in bootstrap (normalize collapses duplicate prompts). - try: - combined += _capture_raw_channel(conn, duration=0.35) - except Exception: - pass - if not str(combined).strip(): - try: - combined = _drain_channel(conn, rounds=6, wait=0.08) - except Exception: - combined = "" - bootstrap = prepare_bootstrap_output(combined) - # Discard lone punctuation left on the wire (would glue onto ```` in xterm). - try: - leftover = _capture_raw_channel(conn, duration=0.12) - except Exception: - leftover = "" - if leftover and leftover.strip() not in {":", ">", "#", "]", "$"}: - bootstrap = prepare_bootstrap_output(f"{bootstrap}{leftover}") - - hop_guard = get_cli_hop_guard(conn) - sess.conn = conn - sess.cli_hop_guard = bool(hop_guard) - sess.cli_hop_prompt = str((hop_guard or {}).get("hop_prompt") or "") - sess.bootstrap_output = _encode_text(str(bootstrap or ""), sess.encoding) - # Nudge Enter on WS attach only when we still need a shell prompt. - # Never when already at CLI prompt or Username:/Password: (would empty-submit login). - sess.needs_live_prompt = ( - not _looks_like_cli_prompt(bootstrap) and not _looks_like_login_prompt(bootstrap) - ) - sess.open_session_log() - if bootstrap: - sess.append_session_log(bootstrap if bootstrap.endswith("\n") else bootstrap + "\n") - sess.start_reader() - # Drop late prompt echoes that race into the queue right after reader start. - prompt_hint = "" - if bootstrap: - prompt_hint = str(bootstrap).replace("\r\n", "\n").replace("\r", "\n").strip().split("\n")[-1].strip() - settle_deadline = time.time() + 0.45 - while time.time() < settle_deadline: - try: - chunk = sess.out_queue.get_nowait() - except queue.Empty: - time.sleep(0.02) - continue - if chunk is None: - sess.out_queue.put(None) - break - try: - text = _decode_bytes(chunk, sess.encoding) - except Exception: - text = "" - if _is_prompt_only_echo(text, prompt_hint): - continue - # Non-prompt data: put back and stop settling. - sess.out_queue.put(chunk) - break - try: - sess.run_post_login_commands() - except Exception: - _log.debug("post_login failed session=%s", sess.session_id, exc_info=True) - # Same SSH transport: open SFTP channel when the device supports it. - sftp_ok = sess.try_attach_sftp() - sess.state = "ready" - sess.connect_finished_at = time.time() - sess._ready_event.set() - elapsed_ms = int((sess.connect_finished_at - sess.connect_started_at) * 1000) - _audit( - "session_created", - session_id=sess.session_id, - ne_id=sess.ne_id, - ne_name=sess.ne_name, - ne_ip=sess.ne_ip, - protocol=sess.protocol, - encoding=sess.encoding, - source=str(device.get("source") or ""), - hop_enabled=bool(creds.get("hop_enabled")), - hop_vendor=str(creds.get("hop_vendor") or "") if creds.get("hop_enabled") else "", - cli_hop_guard=bool(hop_guard), - cli_hop_prompt=str((hop_guard or {}).get("hop_prompt") or ""), - sftp_ready=bool(sftp_ok), - client=client or "", - connect_ms=elapsed_ms, - active=active_session_count(), - ) - - -def create_session( - db: Session, - *, - ne_id: str | None = None, - ume_ne_id: str | None = None, - cols: int = 80, - rows: int = 24, - client: str = "", - encoding: str = "utf-8", - keepalive_sec: int | None = None, - post_login_commands: list[str] | None = None, - async_connect: bool = True, - username_override: str | None = None, - password_override: str | None = None, -) -> dict[str, Any]: - from .cli_resolve import resolve_cli_target - - _ensure_reaper() - max_sessions = max(1, int(settings.webcrt_max_sessions or 20)) - if active_session_count() >= max_sessions: - raise HTTPException(status_code=429, detail="webcrt_session_limit") - - mid = str(ne_id or "").strip() - uid = str(ume_ne_id or "").strip() - try: - creds, device = resolve_cli_target(db, managed_ne_id=mid or None, ume_ne_id=uid or None) - except HTTPException: - raise - except CredentialCryptoError as exc: - raise HTTPException(status_code=400, detail=str(exc) or "credential_crypto_error") from exc - except Exception as exc: - raise HTTPException(status_code=400, detail=f"credential_error:{exc}") from exc - - # One-shot credentials for SecureCRT-style "do not save password" / retry. - if username_override is not None and str(username_override).strip(): - creds["username"] = str(username_override).strip() - if password_override is not None: - creds["password"] = str(password_override) - - protocol = str(device.get("protocol") or creds.get("protocol") or "ssh").strip().lower() - creds["protocol"] = protocol - # Netmiko telnet drivers dislike a completely missing username; use a placeholder - # for the wire only (interactive login still happens in the terminal). - if protocol == "telnet" and not bool(creds.get("hop_enabled")) and not str(creds.get("username") or "").strip(): - creds["username"] = "telnet" - - if not _webcrt_creds_ready(creds): - raise HTTPException(status_code=400, detail="credentials_incomplete") - - session_id = str(uuid.uuid4()) - c = max(20, min(500, int(cols or 80))) - r = max(5, min(200, int(rows or 24))) - connect_timeout = max(30, int(settings.webcrt_connect_timeout_sec or 90)) - target_id = str(device.get("id") or mid or uid) - target_ip = str(device.get("ip_address") or "") - target_name = str(device.get("name") or target_ip) - device_type = str(device.get("device_type") or creds.get("device_type") or "") - vendor = str(device.get("vendor") or creds.get("vendor") or "") - cli_keymap = uses_network_cli_keymap(device_type, vendor) - enc = _normalize_encoding(encoding) - if keepalive_sec is None: - ka = max(0, int(getattr(settings, "webcrt_keepalive_sec", 0) or 0)) - else: - ka = max(0, min(600, int(keepalive_sec))) - - sess = WebcrtSession( - session_id=session_id, - ne_id=target_id, - ne_name=target_name, - ne_ip=target_ip, - protocol=protocol, - cols=c, - rows=r, - device_type=device_type, - vendor=vendor, - cli_keymap=cli_keymap, - encoding=enc, - keepalive_sec=ka, - state="connecting", - post_login_commands=list(post_login_commands or [])[:20], - ) - with _sessions_lock: - _sessions[session_id] = sess - - _audit( - "session_connecting", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - protocol=sess.protocol, - encoding=enc, - client=client or "", - async_connect=bool(async_connect), - ) - - if async_connect: - t = threading.Thread( - target=_finish_connect, - kwargs={ - "sess": sess, - "creds": creds, - "device": device, - "connect_timeout": connect_timeout, - "client": client or "", - }, - name=f"webcrt-connect-{session_id[:8]}", - daemon=True, - ) - t.start() - else: - _finish_connect( - sess, - creds=creds, - device=device, - connect_timeout=connect_timeout, - client=client or "", - ) - if sess.state == "error": - with _sessions_lock: - _sessions.pop(session_id, None) - raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") - - return { - "session_id": session_id, - "ne_id": sess.ne_id, - "ne_name": sess.ne_name, - "ne_ip": sess.ne_ip, - "source": str(device.get("source") or ""), - "protocol": sess.protocol, - "cols": sess.cols, - "rows": sess.rows, - "encoding": enc, - "keepalive_sec": ka, - "state": sess.state, - "ws_path": f"/v1/webcrt/sessions/{session_id}/ws", - "cli_hop": bool(sess.cli_hop_guard), - "sftp_ready": bool(sess.sftp_ready), - } - - -def mark_attached(session_id: str) -> tuple[WebcrtSession, int]: - sess = get_session(session_id) - if sess is None: - raise HTTPException(status_code=404, detail="webcrt_session_not_found") - # Allow re-attach after brief WS drop (React StrictMode remount / network blip). - # Bump generation so the previous WS pump stops and does not steal echo bytes. - sess.attach_gen += 1 - attach_gen = sess.attach_gen - sess.attached = True - sess.detach_deadline = None - sess.touch() - _audit( - "session_attached", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - attach_gen=attach_gen, - state=sess.state, - ) - return sess, attach_gen - - -def detach_session( - session_id: str, - *, - grace_sec: float = 8.0, - client: str = "", - attach_gen: int | None = None, -) -> dict[str, Any]: - """Mark session unattached but keep device channel open briefly for reconnect.""" - sess = get_session(session_id) - if sess is None: - return {"ok": True, "session_id": session_id, "detached": False} - if sess.closed: - return {"ok": True, "session_id": session_id, "detached": False} - # Ignore detach from an older StrictMode WS once a newer attach owns the session. - if attach_gen is not None and attach_gen != sess.attach_gen: - return { - "ok": True, - "session_id": session_id, - "detached": False, - "ignored_stale_attach": True, - } - sess.attached = False - sess.detach_deadline = time.time() + max(1.0, float(grace_sec)) - sess.touch() - _audit( - "session_detached", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - grace_sec=grace_sec, - client=client or "", - attach_gen=attach_gen, - ) - return {"ok": True, "session_id": session_id, "detached": True} - - -def close_session(session_id: str, *, reason: str = "closed", client: str = "") -> dict[str, Any]: - with _sessions_lock: - sess = _sessions.pop(session_id, None) - if sess is None: - return {"ok": True, "session_id": session_id, "closed": False} - if not sess.closed: - sess.close(reason) - _audit( - "session_closed", - session_id=session_id, - ne_id=sess.ne_id, - ne_ip=sess.ne_ip, - reason=reason, - client=client or "", - bytes_in=sess.bytes_in, - bytes_out=sess.bytes_out, - queue_dropped=getattr(sess.out_queue, "dropped", 0), - active=active_session_count(), - ) - return {"ok": True, "session_id": session_id, "closed": True, "reason": reason} - - -def list_sessions() -> dict[str, Any]: - with _sessions_lock: - items = [] - for s in _sessions.values(): - if s.closed: - continue - state = str(s.state or "unknown") - attached = bool(s.attached) - # Lifecycle for ops UI: distinguish login vs live vs grace-period detach. - if state == "connecting": - lifecycle = "connecting" - elif state == "error": - lifecycle = "error" - elif state == "ready" and attached: - lifecycle = "ready" - elif state == "ready" and not attached: - lifecycle = "detached" - else: - lifecycle = state - elapsed_ms = None - if state == "connecting": - elapsed_ms = int(max(0.0, time.time() - float(s.connect_started_at or time.time())) * 1000) - items.append( - { - "session_id": s.session_id, - "ne_id": s.ne_id, - "ne_name": s.ne_name, - "ne_ip": s.ne_ip, - "protocol": s.protocol, - "encoding": s.encoding, - "keepalive_sec": int(s.keepalive_sec or 0), - "state": state, - "lifecycle": lifecycle, - "attached": attached, - "detach_deadline": s.detach_deadline, - "connect_error": str(s.connect_error or "")[:500], - "elapsed_ms": elapsed_ms, - "created_at": datetime.fromtimestamp(s.created_at, tz=timezone.utc).isoformat(), - "last_activity": datetime.fromtimestamp(s.last_activity, tz=timezone.utc).isoformat(), - "bytes_in": s.bytes_in, - "bytes_out": s.bytes_out, - "queue_depth": s.out_queue.qsize(), - "queue_dropped": getattr(s.out_queue, "dropped", 0), - "connect_ms": ( - int((s.connect_finished_at - s.connect_started_at) * 1000) - if s.connect_finished_at - else None - ), - } - ) - return { - "total": len(items), - "max_sessions": max(1, int(settings.webcrt_max_sessions or 20)), - "idle_timeout_sec": max(60, int(settings.webcrt_idle_timeout_sec or 1800)), - "keepalive_sec": int(getattr(settings, "webcrt_keepalive_sec", 0) or 0), - "anti_idle_sec": int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0), - "items": items, - } +__all__ = [ + "WebcrtSession", + "_decode_bytes", + "_encode_text", + "_normalize_encoding", + "_webcrt_creds_ready", + "active_session_count", + "channel_return", + "close_session", + "create_session", + "detach_session", + "find_ssh_session_for_ne", + "get_session", + "list_sessions", + "map_network_cli_enter", + "map_network_cli_keys", + "mark_attached", + "normalize_cli_transcript", + "prepare_bootstrap_output", + "read_session_log_tail", + "uses_network_cli_keymap", + "wait_session_ready", + "webcrt_data_root", +] diff --git a/netx_api/webcrt_session.py b/netx_api/webcrt_session.py new file mode 100644 index 0000000..916099e --- /dev/null +++ b/netx_api/webcrt_session.py @@ -0,0 +1,1111 @@ +"""WebCRT interactive sessions and process-local registry.""" +from __future__ import annotations + +import io +import json +import logging +import threading +import time +import uuid +from dataclasses import dataclass, field +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from fastapi import HTTPException +from netmiko import ConnectHandler +from sqlalchemy.orm import Session + +from .config import settings +from .ne_crypto import CredentialCryptoError +from .ne_session_factory import ( + close_netmiko_connection, + extract_cli_prompt_marker, + get_cli_hop_guard, + open_netmiko_connection, + should_close_cli_hop_session, +) +from .webcrt_channel import ( + _BoundedByteQueue, + _audit, + _capture_raw_channel, + _decode_bytes, + _drain_channel, + _drain_raw_channel, + _encode_text, + _is_prompt_only_echo, + _looks_like_cli_prompt, + _looks_like_login_prompt, + _looks_like_password_change_prompt, + _normalize_encoding, + _prime_interactive_channel, + _session_log_path, + _session_log_text, + _utc_iso, + _utc_now, + channel_return, + map_network_cli_enter, + map_network_cli_keys, + normalize_cli_transcript, + prepare_bootstrap_output, + read_session_log_tail, + uses_network_cli_keymap, + webcrt_data_root, +) + +_log = logging.getLogger("netx.webcrt") + +_sessions_lock = threading.Lock() +_sessions: dict[str, "WebcrtSession"] = {} +_reaper_started = False + + +@dataclass +class WebcrtSession: + session_id: str + ne_id: str + ne_name: str + ne_ip: str + protocol: str + cols: int + rows: int + device_type: str = "" + vendor: str = "" + cli_keymap: bool = True + encoding: str = "utf-8" + keepalive_sec: int = 0 + conn: ConnectHandler | None = None + created_at: float = field(default_factory=time.time) + last_activity: float = field(default_factory=time.time) + attached: bool = False + detach_deadline: float | None = None + closed: bool = False + close_reason: str = "" + state: str = "connecting" + connect_error: str = "" + connect_started_at: float = field(default_factory=time.time) + connect_finished_at: float | None = None + bootstrap_output: bytes = b"" + # First WS attach gets login bootstrap; later attaches prefer session-log tail. + bootstrap_replayed: bool = False + needs_live_prompt: bool = True + # React StrictMode remounts open a second WS before the first fully tears down. + # Only the newest attach_gen may consume out_queue / mark detach. + attach_gen: int = 0 + out_queue: _BoundedByteQueue = field( + default_factory=lambda: _BoundedByteQueue(int(getattr(settings, "webcrt_out_queue_max", 2000) or 2000)) + ) + # Vendor CLI hop (Huawei/ZTE/Cisco): close when nested target session returns to hop. + cli_hop_guard: bool = False + cli_hop_prompt: str = "" + post_login_commands: list[str] = field(default_factory=list) + bytes_in: int = 0 + bytes_out: int = 0 + _reader: threading.Thread | None = field(default=None, repr=False) + _write_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) + _stdout_lock: threading.Lock = field(default_factory=threading.Lock, repr=False) + _hop_scan_buf: str = field(default="", repr=False) + _cli_hop_seen_other_prompt: bool = field(default=False, repr=False) + _log_fh: Any = field(default=None, repr=False) + _ready_event: threading.Event = field(default_factory=threading.Event, repr=False) + # SFTP channel on the same SSH transport as the interactive shell (direct SSH only). + sftp_ready: bool = False + _sftp: Any = field(default=None, repr=False) + _sftp_lock: threading.RLock = field(default_factory=threading.RLock, repr=False) + + def touch(self) -> None: + self.last_activity = time.time() + + def close_sftp(self) -> None: + with self._sftp_lock: + sftp = self._sftp + self._sftp = None + self.sftp_ready = False + if sftp is None: + return + try: + sftp.close() + except Exception: + pass + + def _ssh_transport_unlocked(self) -> Any: + """Caller must hold ``_sftp_lock``. Returns an active Paramiko Transport.""" + if self.closed or self.conn is None: + raise RuntimeError("session_closed") + if str(self.protocol or "ssh").lower() != "ssh": + raise RuntimeError("sftp_requires_ssh") + if self.cli_hop_guard: + raise RuntimeError("sftp_hop_not_supported") + channel = getattr(self.conn, "remote_conn", None) + transport = None + if channel is not None and hasattr(channel, "get_transport"): + try: + transport = channel.get_transport() + except Exception: + transport = None + if transport is None or not bool(getattr(transport, "is_active", lambda: False)()): + raise RuntimeError("ssh_transport_unavailable") + return transport + + def _ensure_sftp_unlocked(self) -> Any: + """Caller must hold ``_sftp_lock``. Shared probe client (sftp_ready).""" + import paramiko + + if self._sftp is not None: + sock = getattr(self._sftp, "sock", None) + if sock is not None and not bool(getattr(sock, "closed", False)): + return self._sftp + try: + self._sftp.close() + except Exception: + pass + self._sftp = None + transport = self._ssh_transport_unlocked() + self._sftp = paramiko.SFTPClient.from_transport(transport) + if self._sftp is None: + raise RuntimeError("sftp_open_failed") + self.sftp_ready = True + return self._sftp + + def open_sftp(self) -> Any: + """Open/reuse an SFTP client on this session's SSH transport.""" + with self._sftp_lock: + return self._ensure_sftp_unlocked() + + def open_ephemeral_sftp(self) -> Any: + """Open a dedicated SFTP channel for one operation; caller must ``close()`` it. + + Only holds ``_sftp_lock`` briefly while resolving the SSH transport, so long + list/upload/download work does not block other SFTP ops on the same session. + """ + import paramiko + + with self._sftp_lock: + transport = self._ssh_transport_unlocked() + # Keep probe client warm for UI sftp_ready without sharing it for I/O. + try: + self._ensure_sftp_unlocked() + except Exception: + pass + sftp = paramiko.SFTPClient.from_transport(transport) + if sftp is None: + raise RuntimeError("sftp_open_failed") + return sftp + + def run_sftp(self, fn: Any) -> Any: + """Run ``fn(sftp)`` on an ephemeral channel (does not hold the lock during ``fn``).""" + sftp = self.open_ephemeral_sftp() + try: + return fn(sftp) + finally: + try: + sftp.close() + except Exception: + pass + + def try_attach_sftp(self) -> bool: + """Best-effort SFTP channel open after SSH login (does not fail the shell).""" + if str(self.protocol or "ssh").lower() != "ssh" or self.cli_hop_guard: + self.sftp_ready = False + return False + try: + self.open_sftp() + self.sftp_ready = True + return True + except Exception: + self.sftp_ready = False + _log.debug("webcrt sftp attach skipped session=%s", self.session_id, exc_info=True) + return False + + def open_session_log(self) -> None: + if not bool(getattr(settings, "webcrt_session_log_enabled", True)): + return + if self._log_fh is not None: + return + try: + self._log_fh = _session_log_path(self.session_id).open("a", encoding="utf-8", errors="replace") + self._log_fh.write(f"# session={self.session_id} ne={self.ne_id} ip={self.ne_ip} ts={_utc_iso()}\n") + self._log_fh.flush() + except Exception: + _log.debug("webcrt session log open failed", exc_info=True) + self._log_fh = None + + def append_session_log(self, text: str) -> None: + if not text or self._log_fh is None: + return + try: + self._log_fh.write(text) + self._log_fh.flush() + except Exception: + pass + + def close_session_log(self) -> None: + fh = self._log_fh + self._log_fh = None + if fh is None: + return + try: + fh.write(f"\n# closed reason={self.close_reason} ts={_utc_iso()}\n") + fh.close() + except Exception: + pass + + def take_stdout(self, attach_gen: int, *, timeout: float = 0.25) -> bytes | None | str: + """Exclusive stdout take for one WS attach generation. + + Returns: + bytes — device output chunk + None — device reader closed (session end) + \"stale\" — a newer WebSocket owns this session; caller must stop + \"empty\" — no data within timeout (keep polling) + """ + deadline = time.time() + max(0.05, float(timeout)) + while True: + with self._stdout_lock: + if attach_gen != self.attach_gen: + return "stale" + remaining = deadline - time.time() + if remaining <= 0: + return "empty" + # Slice waits so we can notice attach_gen bumps without busy-spinning. + try: + chunk = self.out_queue.get(timeout=min(0.05, remaining)) + except queue.Empty: + continue + with self._stdout_lock: + if attach_gen != self.attach_gen: + # Put back including EOF sentinel so the new owner still sees close. + self.out_queue.put(chunk) + return "stale" + return chunk # bytes | None + + def write_stdin(self, data: str) -> None: + if self.closed or self.conn is None: + raise RuntimeError("session_closed") + text = str(data or "") + if not text: + return + if self.cli_keymap: + text = map_network_cli_keys( + text, + device_type=self.device_type, + vendor=self.vendor, + protocol=self.protocol, + ) + text = map_network_cli_enter(text, self.conn) + if not text: + return + with self._write_lock: + # Prefer raw channel I/O for interactive typing (char echo / backspace). + channel = getattr(self.conn, "remote_conn", None) + try: + if channel is not None and hasattr(channel, "send") and callable(channel.send): + payload = _encode_text(text, self.encoding) + # Paramiko may write partially when the window is full. + view = memoryview(payload) + while len(view): + n = int(channel.send(view) or 0) + if n <= 0: + time.sleep(0.01) + continue + view = view[n:] + self.bytes_in += len(payload) + elif channel is not None and hasattr(channel, "write") and callable(channel.write): + payload = _encode_text(text, self.encoding) + channel.write(payload) + self.bytes_in += len(payload) + else: + self.conn.write_channel(text) + self.bytes_in += len(text) + except Exception: + self.conn.write_channel(text) + self.bytes_in += len(text) + self.touch() + + def send_break(self) -> None: + """Send SSH break / Telnet IAC BREAK to interrupt paging or hung commands.""" + if self.closed or self.conn is None: + raise RuntimeError("session_closed") + channel = getattr(self.conn, "remote_conn", None) + with self._write_lock: + sent = False + if channel is not None and hasattr(channel, "send_break") and callable(channel.send_break): + try: + channel.send_break(0) + sent = True + except Exception: + _log.debug("send_break failed session=%s", self.session_id, exc_info=True) + if not sent and channel is not None and hasattr(channel, "send") and callable(channel.send): + # Telnet IAC BREAK = 255 243 + try: + channel.send(b"\xff\xf3") + sent = True + except Exception: + pass + if not sent: + # Fallback: Ctrl-C often interrupts device CLI more-pages. + try: + self.conn.write_channel("\x03") + except Exception: + raise RuntimeError("break_failed") + self.touch() + + def resize(self, cols: int, rows: int) -> None: + if self.closed or self.conn is None: + return + c = max(20, min(500, int(cols or 80))) + r = max(5, min(200, int(rows or 24))) + self.cols = c + self.rows = r + channel = getattr(self.conn, "remote_conn", None) + if channel is not None and hasattr(channel, "resize_pty"): + try: + channel.resize_pty(width=c, height=r) + except Exception: + _log.debug("resize_pty failed session=%s", self.session_id, exc_info=True) + self.touch() + + def start_reader(self) -> None: + if self._reader and self._reader.is_alive(): + return + self._reader = threading.Thread( + target=self._reader_loop, + name=f"webcrt-reader-{self.session_id[:8]}", + daemon=True, + ) + self._reader.start() + + def _reader_loop(self) -> None: + conn = self.conn + if conn is None: + self.out_queue.put(None) + return + channel = getattr(conn, "remote_conn", None) + hop_return = False + poll = max(0.002, float(getattr(settings, "webcrt_reader_poll_sec", 0.01) or 0.01)) + try: + while not self.closed: + chunk = b"" + try: + if channel is not None and hasattr(channel, "recv_ready") and hasattr(channel, "recv"): + # Paramiko SSH: prefer short blocking recv over fixed spin-sleep. + ready = False + try: + ready = bool(channel.recv_ready()) + except Exception: + ready = False + if ready: + chunk = channel.recv(16384) + if not chunk: + break + elif hasattr(channel, "exit_status_ready") and channel.exit_status_ready(): + break + else: + # Brief block: settimeout + recv wakes sooner than sleep(0.04). + prev_timeout = None + try: + prev_timeout = channel.gettimeout() + except Exception: + prev_timeout = None + try: + channel.settimeout(poll) + chunk = channel.recv(16384) + except Exception: + chunk = b"" + finally: + try: + channel.settimeout(prev_timeout) + except Exception: + pass + if not chunk: + continue + elif channel is not None and hasattr(channel, "read_very_eager"): + # Telnet: do NOT use conn.read_channel() — Netmiko strips ANSI. + data = channel.read_very_eager() + if data: + chunk = ( + data + if isinstance(data, (bytes, bytearray)) + else _encode_text(str(data), self.encoding) + ) + else: + time.sleep(poll) + continue + else: + text = conn.read_channel() + if text: + chunk = _encode_text(str(text), self.encoding) + else: + time.sleep(poll) + continue + except Exception as exc: + if self.closed: + break + _log.debug("webcrt reader error session=%s: %s", self.session_id, exc) + time.sleep(0.05) + continue + if chunk: + self.touch() + self.bytes_out += len(chunk) + self.out_queue.put(chunk) + try: + self.append_session_log(_decode_bytes(chunk, self.encoding)) + except Exception: + pass + if self.cli_hop_guard and self._note_cli_hop_output(chunk): + hop_return = True + notice = ( + "\r\n*** WebCRT: 目标会话已结束,已断开代理连接 " + "(target session ended; closing hop proxy) ***\r\n" + ) + self.out_queue.put(notice.encode("utf-8", errors="replace")) + break + finally: + if hop_return and not self.closed: + try: + close_session(self.session_id, reason="cli_hop_return") + except Exception: + self.close("cli_hop_return") + self.out_queue.put(None) + + def _note_cli_hop_output(self, chunk: bytes) -> bool: + """Accumulate stdout and return True when nested CLI hop has returned to proxy.""" + try: + text = _decode_bytes(chunk, self.encoding) + except Exception: + text = str(chunk) + self._hop_scan_buf = (self._hop_scan_buf + text)[-12000:] + marker = str(self.cli_hop_prompt or "").strip() + last = extract_cli_prompt_marker(self._hop_scan_buf) + if last and (not marker or last != marker): + self._cli_hop_seen_other_prompt = True + return should_close_cli_hop_session( + self._hop_scan_buf, + self.cli_hop_prompt, + seen_other_prompt=self._cli_hop_seen_other_prompt, + ) + + def run_post_login_commands(self) -> None: + cmds = [str(c).rstrip("\r\n") for c in (self.post_login_commands or []) if str(c).strip()] + if not cmds or self.closed or self.conn is None: + return + for cmd in cmds[:20]: + try: + self.write_stdin(cmd + "\r") + time.sleep(0.15) + except Exception: + _log.debug("post_login command failed session=%s", self.session_id, exc_info=True) + break + + def close(self, reason: str = "closed") -> None: + if self.closed: + return + self.closed = True + self.state = "closed" + self.close_reason = reason or "closed" + self._ready_event.set() + self.close_sftp() + try: + close_netmiko_connection(self.conn) + except Exception: + pass + self.conn = None + try: + self.out_queue.put_nowait(None) + except Exception: + pass + self.close_session_log() + + +def _ensure_reaper() -> None: + global _reaper_started + with _sessions_lock: + if _reaper_started: + return + _reaper_started = True + t = threading.Thread(target=_reaper_loop, name="webcrt-reaper", daemon=True) + t.start() + + +def _reaper_loop() -> None: + while True: + try: + _reap_sessions() + except Exception: + _log.exception("webcrt reaper failed") + time.sleep(2) + + +def _reap_sessions() -> None: + idle = max(60, int(settings.webcrt_idle_timeout_sec or 1800)) + attach = max(10, int(settings.webcrt_attach_timeout_sec or 60)) + anti_idle = max(0, int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0)) + anti_payload = str(getattr(settings, "webcrt_anti_idle_payload", " ") or " ") + now = time.time() + to_close: list[tuple[WebcrtSession, str]] = [] + to_nudge: list[WebcrtSession] = [] + with _sessions_lock: + for sess in list(_sessions.values()): + if sess.closed: + _sessions.pop(sess.session_id, None) + continue + if sess.state == "connecting": + # Connecting sessions use connect timeout, not attach timeout alone. + connect_budget = max(30, int(settings.webcrt_connect_timeout_sec or 90)) + 30 + if (now - sess.connect_started_at) > connect_budget: + to_close.append((sess, "connect_timeout")) + continue + if sess.attached: + if (now - sess.last_activity) > idle: + to_close.append((sess, "idle_timeout")) + elif ( + anti_idle > 0 + and sess.state == "ready" + and sess.conn is not None + and (now - sess.last_activity) >= anti_idle + ): + to_nudge.append(sess) + continue + # Not attached: either never attached, or briefly detached for reconnect. + if sess.detach_deadline is not None: + if now >= sess.detach_deadline: + to_close.append((sess, "detach_timeout")) + else: + # Start attach clock after connect finishes (not HTTP create time), + # so slow auth + UI mount does not race attach_timeout. + anchor = float(sess.connect_finished_at or sess.created_at or now) + if (now - anchor) > attach: + to_close.append((sess, "attach_timeout")) + elif (now - sess.last_activity) > idle: + to_close.append((sess, "idle_timeout")) + for sess in to_nudge: + try: + # Touch without changing visible prompt when payload is empty/null-ish. + payload = anti_payload + if payload == "\\0": + payload = "\x00" + if payload: + sess.write_stdin(payload) + else: + sess.touch() + except Exception: + _log.debug("webcrt anti-idle failed session=%s", sess.session_id, exc_info=True) + for sess, reason in to_close: + close_session(sess.session_id, reason=reason) + + +def active_session_count() -> int: + with _sessions_lock: + return sum(1 for s in _sessions.values() if not s.closed) + + +def get_session(session_id: str) -> WebcrtSession | None: + with _sessions_lock: + sess = _sessions.get(session_id) + if sess is None or sess.closed: + return None + return sess + + +def find_ssh_session_for_ne(ne_id: str) -> WebcrtSession | None: + """Prefer a ready/attached interactive SSH session for SFTP channel reuse.""" + nid = str(ne_id or "").strip() + if not nid: + return None + with _sessions_lock: + candidates = [ + s + for s in _sessions.values() + if (not s.closed) + and str(s.ne_id) == nid + and str(s.protocol or "ssh").lower() == "ssh" + and s.conn is not None + and s.state in ("ready", "connecting") + and not s.cli_hop_guard + ] + if not candidates: + return None + # Prefer attached + ready sessions. + candidates.sort(key=lambda s: (0 if s.attached and s.state == "ready" else 1, -s.last_activity)) + return candidates[0] + + +def wait_session_ready(session_id: str, *, timeout: float = 120.0) -> WebcrtSession: + """Block until async connect finishes (ready or error). Used by tests and WS.""" + deadline = time.time() + max(1.0, float(timeout)) + while time.time() < deadline: + sess = get_session(session_id) + if sess is None: + raise HTTPException(status_code=404, detail="webcrt_session_not_found") + if sess.state == "ready": + return sess + if sess.state == "error": + raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") + sess._ready_event.wait(timeout=0.25) + raise HTTPException(status_code=504, detail="connect_timeout") + + +def _webcrt_creds_ready(creds: dict[str, Any]) -> bool: + """True when WebCRT can open a session with the resolved credentials. + + Bastion-managed hops store the target password on the bastion side, so an empty + NE password is valid (same as connectivity test). Direct / manual / Linux hops + still require a target password for SSH. + + Telnet (no hop) allows empty username/password so the user can authenticate + interactively in the terminal (SecureCRT-style). + """ + hop_enabled = bool(creds.get("hop_enabled")) + hop_vendor = str(creds.get("hop_vendor") or "").strip().lower() + auth_mode = str(creds.get("hop_target_auth_mode") or "bastion_managed").strip().lower() + protocol = str(creds.get("protocol") or "ssh").strip().lower() + if hop_enabled and hop_vendor == "bastion" and auth_mode == "bastion_managed": + return bool( + str(creds.get("hop_host") or "").strip() + and str(creds.get("hop_username") or "").strip() + and str(creds.get("hop_password") or "") + ) + if protocol == "telnet" and not hop_enabled: + return True + if not str(creds.get("username") or "").strip(): + return False + return bool(str(creds.get("password") or "")) + + +def _finish_connect( + sess: WebcrtSession, + *, + creds: dict[str, Any], + device: dict[str, Any], + connect_timeout: int, + client: str, +) -> None: + log_buf = io.BytesIO() + try: + conn = open_netmiko_connection( + creds, + session_timeout=connect_timeout, + session_log=log_buf, + cols=sess.cols, + rows=sess.rows, + interactive=True, + keepalive=int(sess.keepalive_sec or 0), + ) + except Exception as exc: + partial = _session_log_text(log_buf).strip() + from .ne_cli_errors import format_cli_failure + + classified = format_cli_failure(exc, partial) + detail = f"connect_failed:{classified}" + if partial: + detail = f"{detail}\n--- device transcript ---\n{partial[-4000:]}" + sess.state = "error" + sess.connect_error = detail + sess.connect_finished_at = time.time() + sess._ready_event.set() + _audit( + "session_open_failed", + session_id=sess.session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + source=str(device.get("source") or ""), + client=client or "", + error=str(exc)[:500], + transcript_len=len(partial), + ) + return + + channel = getattr(conn, "remote_conn", None) + if channel is not None and hasattr(channel, "resize_pty"): + try: + channel.resize_pty(width=sess.cols, height=sess.rows) + except Exception: + pass + + pre_log = _session_log_text(log_buf) + # Pull post-auth banner/MOTD from the PTY. With interactive no-op session_preparation + # (generic_termserver), Netmiko session_log is often empty — do not discard these bytes. + try: + early = _capture_raw_channel(conn, duration=0.35) + except Exception: + early = "" + seed = f"{pre_log}{early}" + already_prompted = _looks_like_cli_prompt(seed) + primed = "" + # Do not send Enter at Username:/Password: or Huawei password-change [Y/N]: + # (Netmiko telnet_login already answers password-change with "N"). + if _looks_like_login_prompt(seed) or _looks_like_password_change_prompt(seed): + try: + primed = _capture_raw_channel(conn, duration=0.9) + except Exception: + primed = "" + else: + try: + primed = _prime_interactive_channel(conn, already_prompted=already_prompted) + except Exception: + primed = "" + combined = f"{seed}{primed}" + # Final settle: keep stragglers in bootstrap (normalize collapses duplicate prompts). + try: + combined += _capture_raw_channel(conn, duration=0.35) + except Exception: + pass + if not str(combined).strip(): + try: + combined = _drain_channel(conn, rounds=6, wait=0.08) + except Exception: + combined = "" + bootstrap = prepare_bootstrap_output(combined) + # Discard lone punctuation left on the wire (would glue onto ```` in xterm). + try: + leftover = _capture_raw_channel(conn, duration=0.12) + except Exception: + leftover = "" + if leftover and leftover.strip() not in {":", ">", "#", "]", "$"}: + bootstrap = prepare_bootstrap_output(f"{bootstrap}{leftover}") + + hop_guard = get_cli_hop_guard(conn) + sess.conn = conn + sess.cli_hop_guard = bool(hop_guard) + sess.cli_hop_prompt = str((hop_guard or {}).get("hop_prompt") or "") + sess.bootstrap_output = _encode_text(str(bootstrap or ""), sess.encoding) + # Nudge Enter on WS attach only when we still need a shell prompt. + # Never when already at CLI prompt or Username:/Password: (would empty-submit login). + sess.needs_live_prompt = ( + not _looks_like_cli_prompt(bootstrap) and not _looks_like_login_prompt(bootstrap) + ) + sess.open_session_log() + if bootstrap: + sess.append_session_log(bootstrap if bootstrap.endswith("\n") else bootstrap + "\n") + sess.start_reader() + # Drop late prompt echoes that race into the queue right after reader start. + prompt_hint = "" + if bootstrap: + prompt_hint = str(bootstrap).replace("\r\n", "\n").replace("\r", "\n").strip().split("\n")[-1].strip() + settle_deadline = time.time() + 0.45 + while time.time() < settle_deadline: + try: + chunk = sess.out_queue.get_nowait() + except queue.Empty: + time.sleep(0.02) + continue + if chunk is None: + sess.out_queue.put(None) + break + try: + text = _decode_bytes(chunk, sess.encoding) + except Exception: + text = "" + if _is_prompt_only_echo(text, prompt_hint): + continue + # Non-prompt data: put back and stop settling. + sess.out_queue.put(chunk) + break + try: + sess.run_post_login_commands() + except Exception: + _log.debug("post_login failed session=%s", sess.session_id, exc_info=True) + # Same SSH transport: open SFTP channel when the device supports it. + sftp_ok = sess.try_attach_sftp() + sess.state = "ready" + sess.connect_finished_at = time.time() + sess._ready_event.set() + elapsed_ms = int((sess.connect_finished_at - sess.connect_started_at) * 1000) + _audit( + "session_created", + session_id=sess.session_id, + ne_id=sess.ne_id, + ne_name=sess.ne_name, + ne_ip=sess.ne_ip, + protocol=sess.protocol, + encoding=sess.encoding, + source=str(device.get("source") or ""), + hop_enabled=bool(creds.get("hop_enabled")), + hop_vendor=str(creds.get("hop_vendor") or "") if creds.get("hop_enabled") else "", + cli_hop_guard=bool(hop_guard), + cli_hop_prompt=str((hop_guard or {}).get("hop_prompt") or ""), + sftp_ready=bool(sftp_ok), + client=client or "", + connect_ms=elapsed_ms, + active=active_session_count(), + ) + + +def create_session( + db: Session, + *, + ne_id: str | None = None, + ume_ne_id: str | None = None, + cols: int = 80, + rows: int = 24, + client: str = "", + encoding: str = "utf-8", + keepalive_sec: int | None = None, + post_login_commands: list[str] | None = None, + async_connect: bool = True, + username_override: str | None = None, + password_override: str | None = None, +) -> dict[str, Any]: + from .cli_resolve import resolve_cli_target + + _ensure_reaper() + max_sessions = max(1, int(settings.webcrt_max_sessions or 20)) + if active_session_count() >= max_sessions: + raise HTTPException(status_code=429, detail="webcrt_session_limit") + + mid = str(ne_id or "").strip() + uid = str(ume_ne_id or "").strip() + try: + creds, device = resolve_cli_target(db, managed_ne_id=mid or None, ume_ne_id=uid or None) + except HTTPException: + raise + except CredentialCryptoError as exc: + raise HTTPException(status_code=400, detail=str(exc) or "credential_crypto_error") from exc + except Exception as exc: + raise HTTPException(status_code=400, detail=f"credential_error:{exc}") from exc + + # One-shot credentials for SecureCRT-style "do not save password" / retry. + if username_override is not None and str(username_override).strip(): + creds["username"] = str(username_override).strip() + if password_override is not None: + creds["password"] = str(password_override) + + protocol = str(device.get("protocol") or creds.get("protocol") or "ssh").strip().lower() + creds["protocol"] = protocol + # Netmiko telnet drivers dislike a completely missing username; use a placeholder + # for the wire only (interactive login still happens in the terminal). + if protocol == "telnet" and not bool(creds.get("hop_enabled")) and not str(creds.get("username") or "").strip(): + creds["username"] = "telnet" + + if not _webcrt_creds_ready(creds): + raise HTTPException(status_code=400, detail="credentials_incomplete") + + session_id = str(uuid.uuid4()) + c = max(20, min(500, int(cols or 80))) + r = max(5, min(200, int(rows or 24))) + connect_timeout = max(30, int(settings.webcrt_connect_timeout_sec or 90)) + target_id = str(device.get("id") or mid or uid) + target_ip = str(device.get("ip_address") or "") + target_name = str(device.get("name") or target_ip) + device_type = str(device.get("device_type") or creds.get("device_type") or "") + vendor = str(device.get("vendor") or creds.get("vendor") or "") + cli_keymap = uses_network_cli_keymap(device_type, vendor) + enc = _normalize_encoding(encoding) + if keepalive_sec is None: + ka = max(0, int(getattr(settings, "webcrt_keepalive_sec", 0) or 0)) + else: + ka = max(0, min(600, int(keepalive_sec))) + + sess = WebcrtSession( + session_id=session_id, + ne_id=target_id, + ne_name=target_name, + ne_ip=target_ip, + protocol=protocol, + cols=c, + rows=r, + device_type=device_type, + vendor=vendor, + cli_keymap=cli_keymap, + encoding=enc, + keepalive_sec=ka, + state="connecting", + post_login_commands=list(post_login_commands or [])[:20], + ) + with _sessions_lock: + _sessions[session_id] = sess + + _audit( + "session_connecting", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + protocol=sess.protocol, + encoding=enc, + client=client or "", + async_connect=bool(async_connect), + ) + + if async_connect: + t = threading.Thread( + target=_finish_connect, + kwargs={ + "sess": sess, + "creds": creds, + "device": device, + "connect_timeout": connect_timeout, + "client": client or "", + }, + name=f"webcrt-connect-{session_id[:8]}", + daemon=True, + ) + t.start() + else: + _finish_connect( + sess, + creds=creds, + device=device, + connect_timeout=connect_timeout, + client=client or "", + ) + if sess.state == "error": + with _sessions_lock: + _sessions.pop(session_id, None) + raise HTTPException(status_code=502, detail=sess.connect_error or "connect_failed") + + return { + "session_id": session_id, + "ne_id": sess.ne_id, + "ne_name": sess.ne_name, + "ne_ip": sess.ne_ip, + "source": str(device.get("source") or ""), + "protocol": sess.protocol, + "cols": sess.cols, + "rows": sess.rows, + "encoding": enc, + "keepalive_sec": ka, + "state": sess.state, + "ws_path": f"/v1/webcrt/sessions/{session_id}/ws", + "cli_hop": bool(sess.cli_hop_guard), + "sftp_ready": bool(sess.sftp_ready), + } + + +def mark_attached(session_id: str) -> tuple[WebcrtSession, int]: + sess = get_session(session_id) + if sess is None: + raise HTTPException(status_code=404, detail="webcrt_session_not_found") + # Allow re-attach after brief WS drop (React StrictMode remount / network blip). + # Bump generation so the previous WS pump stops and does not steal echo bytes. + sess.attach_gen += 1 + attach_gen = sess.attach_gen + sess.attached = True + sess.detach_deadline = None + sess.touch() + _audit( + "session_attached", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + attach_gen=attach_gen, + state=sess.state, + ) + return sess, attach_gen + + +def detach_session( + session_id: str, + *, + grace_sec: float = 8.0, + client: str = "", + attach_gen: int | None = None, +) -> dict[str, Any]: + """Mark session unattached but keep device channel open briefly for reconnect.""" + sess = get_session(session_id) + if sess is None: + return {"ok": True, "session_id": session_id, "detached": False} + if sess.closed: + return {"ok": True, "session_id": session_id, "detached": False} + # Ignore detach from an older StrictMode WS once a newer attach owns the session. + if attach_gen is not None and attach_gen != sess.attach_gen: + return { + "ok": True, + "session_id": session_id, + "detached": False, + "ignored_stale_attach": True, + } + sess.attached = False + sess.detach_deadline = time.time() + max(1.0, float(grace_sec)) + sess.touch() + _audit( + "session_detached", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + grace_sec=grace_sec, + client=client or "", + attach_gen=attach_gen, + ) + return {"ok": True, "session_id": session_id, "detached": True} + + +def close_session(session_id: str, *, reason: str = "closed", client: str = "") -> dict[str, Any]: + with _sessions_lock: + sess = _sessions.pop(session_id, None) + if sess is None: + return {"ok": True, "session_id": session_id, "closed": False} + if not sess.closed: + sess.close(reason) + _audit( + "session_closed", + session_id=session_id, + ne_id=sess.ne_id, + ne_ip=sess.ne_ip, + reason=reason, + client=client or "", + bytes_in=sess.bytes_in, + bytes_out=sess.bytes_out, + queue_dropped=getattr(sess.out_queue, "dropped", 0), + active=active_session_count(), + ) + return {"ok": True, "session_id": session_id, "closed": True, "reason": reason} + + +def list_sessions() -> dict[str, Any]: + with _sessions_lock: + items = [] + for s in _sessions.values(): + if s.closed: + continue + state = str(s.state or "unknown") + attached = bool(s.attached) + # Lifecycle for ops UI: distinguish login vs live vs grace-period detach. + if state == "connecting": + lifecycle = "connecting" + elif state == "error": + lifecycle = "error" + elif state == "ready" and attached: + lifecycle = "ready" + elif state == "ready" and not attached: + lifecycle = "detached" + else: + lifecycle = state + elapsed_ms = None + if state == "connecting": + elapsed_ms = int(max(0.0, time.time() - float(s.connect_started_at or time.time())) * 1000) + items.append( + { + "session_id": s.session_id, + "ne_id": s.ne_id, + "ne_name": s.ne_name, + "ne_ip": s.ne_ip, + "protocol": s.protocol, + "encoding": s.encoding, + "keepalive_sec": int(s.keepalive_sec or 0), + "state": state, + "lifecycle": lifecycle, + "attached": attached, + "detach_deadline": s.detach_deadline, + "connect_error": str(s.connect_error or "")[:500], + "elapsed_ms": elapsed_ms, + "created_at": datetime.fromtimestamp(s.created_at, tz=timezone.utc).isoformat(), + "last_activity": datetime.fromtimestamp(s.last_activity, tz=timezone.utc).isoformat(), + "bytes_in": s.bytes_in, + "bytes_out": s.bytes_out, + "queue_depth": s.out_queue.qsize(), + "queue_dropped": getattr(s.out_queue, "dropped", 0), + "connect_ms": ( + int((s.connect_finished_at - s.connect_started_at) * 1000) + if s.connect_finished_at + else None + ), + } + ) + return { + "total": len(items), + "max_sessions": max(1, int(settings.webcrt_max_sessions or 20)), + "idle_timeout_sec": max(60, int(settings.webcrt_idle_timeout_sec or 1800)), + "keepalive_sec": int(getattr(settings, "webcrt_keepalive_sec", 0) or 0), + "anti_idle_sec": int(getattr(settings, "webcrt_anti_idle_sec", 0) or 0), + "items": items, + } diff --git a/netx_api/worker.py b/netx_api/worker.py index 1bcb63d..56b1abc 100644 --- a/netx_api/worker.py +++ b/netx_api/worker.py @@ -1,6 +1,7 @@ """Background worker process for long-running schedulers. -Run separately from the API when NETX_RUN_INLINE_SCHEDULERS=false: +Default deployment: API has ``NETX_RUN_INLINE_SCHEDULERS=false``; run this +alongside the API: python -m netx_api.worker diff --git a/tests/test_ne_exec.py b/tests/test_ne_exec.py index 4db654a..650b111 100644 --- a/tests/test_ne_exec.py +++ b/tests/test_ne_exec.py @@ -106,6 +106,12 @@ class NeExecValidationTests(unittest.TestCase): _validate_command("ip address 1.1.1.1 255.255.255.0") self.assertEqual(ctx.exception.detail, "command_blocked") + def test_blocks_vendor_destructive(self) -> None: + for cmd in ("save", "request system reboot", "clear configuration", "file delete flash:/x"): + with self.assertRaises(HTTPException) as ctx: + _validate_command(cmd) + self.assertEqual(ctx.exception.detail, "command_blocked", cmd) + def test_agent_batch_rejects_configure_before_connect(self) -> None: cmds = [ "show interface", diff --git a/tests/test_schema_patches.py b/tests/test_schema_patches.py index 281b817..1996f01 100644 --- a/tests/test_schema_patches.py +++ b/tests/test_schema_patches.py @@ -38,12 +38,13 @@ class SchemaPatchesTests(unittest.TestCase): apply_domain_schema_patches(conn) apply_domain_schema_patches(conn) - def test_alembic_auto_defaults(self) -> None: + def test_worker_default_off_inline(self) -> None: from netx_api.config import Settings s = Settings(_env_file=None) + self.assertFalse(s.run_inline_schedulers) self.assertTrue(s.alembic_upgrade_on_start) - self.assertTrue(s.skip_legacy_startup_ddl) + versions = Path(__file__).resolve().parents[1] / "alembic" / "versions" files = sorted(p.name for p in versions.glob("*.py") if p.name != "__init__.py") self.assertIn("20260802_scopes.py", files)