From d56020d84c2d8c6960fbc9b50345bc413e7f0181 Mon Sep 17 00:00:00 2001 From: oliver Date: Sun, 2 Aug 2026 18:02:43 +0800 Subject: [PATCH] Harden runtime budgets: DB pool, CLI concurrency, timeouts, and metrics. Add shared CLI budget/timeout with force-close, parallel port-traffic dispatch, unified shutdown, bounded audit queue, output/log caps, and /metrics probes. Co-authored-by: Cursor --- .env.example | 12 +++ netx_api/app_shutdown.py | 88 +++++++++++++++++++++ netx_api/audit_async.py | 52 ++++++++++++- netx_api/auth_middleware.py | 2 + netx_api/cli_budget.py | 68 ++++++++++++++++ netx_api/cli_timeout.py | 78 +++++++++++++++++++ netx_api/config.py | 18 +++++ netx_api/config_sync_runner.py | 48 +++++++++--- netx_api/db.py | 47 +++++++++++- netx_api/integrations_router.py | 5 ++ netx_api/main.py | 11 +-- netx_api/metrics_router.py | 113 +++++++++++++++++++++++++++ netx_api/ne_collect_runner.py | 50 +++++++++--- netx_api/ne_connect.py | 14 +++- netx_api/port_traffic_runner.py | 87 ++++++++++----------- netx_api/port_traffic_scheduler.py | 115 ++++++++++++++++++++++++---- netx_api/topology_discover_jobs.py | 9 ++- netx_api/topology_discover_scan.py | 7 +- netx_api/ume_alarm_apply.py | 14 +++- netx_api/webcrt_session_model.py | 24 +++++- netx_api/webcrt_session_registry.py | 15 ++++ netx_api/worker.py | 4 + tests/test_stability_hardening.py | 45 +++++++++++ 23 files changed, 824 insertions(+), 102 deletions(-) create mode 100644 netx_api/app_shutdown.py create mode 100644 netx_api/cli_budget.py create mode 100644 netx_api/cli_timeout.py create mode 100644 netx_api/metrics_router.py create mode 100644 tests/test_stability_hardening.py diff --git a/.env.example b/.env.example index dbecbc0..f46e737 100644 --- a/.env.example +++ b/.env.example @@ -59,3 +59,15 @@ NETX_UME_NOTIFICATION_TOPIC=ALARM # NETX_RUN_INLINE_SCHEDULERS=false # NETX_AUDIT_ASYNC=true # NETX_AUDIT_SAMPLE_N=1 +# DB pool (QueuePool). Aim: pool_size + max_overflow >= HTTP peak + NETX_CLI_MAX_CONCURRENT + UME WS burst. +# NETX_DB_POOL_SIZE=20 +# NETX_DB_MAX_OVERFLOW=20 +# NETX_DB_POOL_RECYCLE_SEC=1800 +# NETX_DB_POOL_TIMEOUT_SEC=30 +# Global SSH/Netmiko concurrency across discover / collect / config_sync / port_traffic. +# NETX_CLI_MAX_CONCURRENT=20 +# NETX_CLI_TIMEOUT_POOL_WORKERS=16 +# NETX_PORT_TRAFFIC_DISPATCH_WORKERS=4 +# NETX_AUDIT_QUEUE_MAX=5000 +# NETX_NE_COLLECT_MAX_OUTPUT_BYTES=8388608 +# NETX_WEBCRT_SESSION_LOG_MAX_BYTES=4194304 diff --git a/netx_api/app_shutdown.py b/netx_api/app_shutdown.py new file mode 100644 index 0000000..17fb593 --- /dev/null +++ b/netx_api/app_shutdown.py @@ -0,0 +1,88 @@ +"""Centralized runtime shutdown for API / worker processes.""" +from __future__ import annotations + +import logging + +_log = logging.getLogger("netx.shutdown") + + +def shutdown_runtime(*, reason: str = "lifespan") -> None: + """Best-effort stop of schedulers, sidebands, pools, and sessions.""" + _log.info("shutdown_runtime begin reason=%s", reason) + + try: + from .config_sync_scheduler import stop_config_sync_scheduler + + stop_config_sync_scheduler() + except Exception: # noqa: BLE001 + _log.exception("stop_config_sync_scheduler failed") + + try: + from .lldp_collect_scheduler import stop_lldp_collect_scheduler + + stop_lldp_collect_scheduler() + except Exception: # noqa: BLE001 + _log.exception("stop_lldp_collect_scheduler failed") + + try: + from .port_traffic_scheduler import stop_port_traffic_scheduler + + stop_port_traffic_scheduler() + except Exception: # noqa: BLE001 + _log.exception("stop_port_traffic_scheduler failed") + + try: + import netx_api.ume_support as ume_support + + if ume_support._UME_WS_STOP_EVENT is not None: + ume_support._UME_WS_STOP_EVENT.set() + except Exception: # noqa: BLE001 + _log.exception("set UME_WS_STOP_EVENT failed") + + try: + from .ume_alarm_ws import shutdown_ws_consumer + + shutdown_ws_consumer() + except Exception: # noqa: BLE001 + _log.exception("shutdown_ws_consumer failed") + + try: + from .oclaw_alarm_forwarder import shutdown_oclaw_alarm_forwarder + + shutdown_oclaw_alarm_forwarder() + except Exception: # noqa: BLE001 + _log.exception("shutdown_oclaw_alarm_forwarder failed") + + try: + from .webcrt_session_registry import close_all_sessions + + close_all_sessions() + except Exception: # noqa: BLE001 + _log.exception("close_all_webcrt_sessions failed") + + try: + from .cli_timeout import shutdown_cli_timeout_pool + + shutdown_cli_timeout_pool(wait=False) + except Exception: # noqa: BLE001 + _log.exception("shutdown_cli_timeout_pool failed") + + try: + from .audit_async import shutdown_audit_worker + + shutdown_audit_worker(timeout_sec=2.0) + except Exception: # noqa: BLE001 + _log.exception("shutdown_audit_worker failed") + + try: + from .ne_collect_runner import shutdown_ne_collect_executor + from .ne_connect import shutdown_ne_connect_executor + from .config_sync_runner import shutdown_config_sync_pools + + shutdown_ne_collect_executor() + shutdown_ne_connect_executor() + shutdown_config_sync_pools() + except Exception: # noqa: BLE001 + _log.exception("shutdown executors failed") + + _log.info("shutdown_runtime done reason=%s", reason) diff --git a/netx_api/audit_async.py b/netx_api/audit_async.py index a15abe5..dabf289 100644 --- a/netx_api/audit_async.py +++ b/netx_api/audit_async.py @@ -12,10 +12,11 @@ from .db import SessionLocal _log = logging.getLogger("netx.audit.async") -_q: queue.SimpleQueue[dict[str, Any]] | None = None +_q: queue.Queue[dict[str, Any] | None] | None = None _worker: threading.Thread | None = None _lock = threading.Lock() _counter = 0 +_dropped = 0 def _sample_ok() -> bool: @@ -28,6 +29,15 @@ def _sample_ok() -> bool: return (_counter % n) == 0 +def audit_queue_status() -> dict[str, int]: + q = _q + return { + "depth": int(q.qsize()) if q is not None else 0, + "dropped": int(_dropped), + "maxsize": max(100, int(getattr(settings, "audit_queue_max", 5000) or 5000)), + } + + def _worker_loop() -> None: assert _q is not None while True: @@ -46,16 +56,35 @@ def _worker_loop() -> None: _log.exception("async audit write failed") -def _ensure_worker() -> queue.SimpleQueue: +def _ensure_worker() -> queue.Queue: global _q, _worker with _lock: if _q is None: - _q = queue.SimpleQueue() + maxsize = max(100, int(getattr(settings, "audit_queue_max", 5000) or 5000)) + _q = queue.Queue(maxsize=maxsize) _worker = threading.Thread(target=_worker_loop, name="netx-audit-writer", daemon=True) _worker.start() return _q +def shutdown_audit_worker(*, timeout_sec: float = 2.0) -> None: + global _q, _worker + with _lock: + q = _q + worker = _worker + if q is None: + return + try: + q.put_nowait(None) + except queue.Full: + pass + if worker is not None: + worker.join(timeout=max(0.1, float(timeout_sec))) + with _lock: + _q = None + _worker = None + + def enqueue_audit( *, action: str, @@ -68,6 +97,7 @@ def enqueue_audit( user_agent: str = "", detail: dict[str, Any] | None = None, ) -> None: + global _dropped act = str(action or "") # Always persist auth / security-relevant events. if act.startswith("auth.") or act.startswith("users.") or act.startswith("api_tokens.") or act.startswith("webcrt."): @@ -95,7 +125,21 @@ def enqueue_audit( db.close() return try: - _ensure_worker().put(payload) + q = _ensure_worker() + try: + q.put_nowait(payload) + except queue.Full: + # Drop oldest then retry once to keep newest events. + try: + q.get_nowait() + _dropped += 1 + except queue.Empty: + pass + try: + q.put_nowait(payload) + except queue.Full: + _dropped += 1 + _log.warning("audit queue full; dropped event action=%s", act) except Exception: _log.exception("audit enqueue failed; falling back to sync") from .auth_service import write_audit diff --git a/netx_api/auth_middleware.py b/netx_api/auth_middleware.py index b8876a5..778e988 100644 --- a/netx_api/auth_middleware.py +++ b/netx_api/auth_middleware.py @@ -24,6 +24,8 @@ _PUBLIC_EXACT = frozenset( "/health", "/health/live", "/health/ready", + "/metrics", + "/metrics/json", "/favicon.ico", "/v1/auth/login", } diff --git a/netx_api/cli_budget.py b/netx_api/cli_budget.py new file mode 100644 index 0000000..18c35bf --- /dev/null +++ b/netx_api/cli_budget.py @@ -0,0 +1,68 @@ +"""Global CLI concurrency budget — shared across discover / collect / config_sync / port_traffic.""" +from __future__ import annotations + +import threading +from contextlib import contextmanager +from typing import Iterator + +from .config import settings + +_lock = threading.Lock() +_sem: threading.BoundedSemaphore | None = None +_limit = 0 +_in_use = 0 +_in_use_lock = threading.Lock() + + +def _ensure_sem() -> threading.BoundedSemaphore: + global _sem, _limit + with _lock: + want = max(1, int(getattr(settings, "cli_max_concurrent", 20) or 20)) + if _sem is None or want != _limit: + _sem = threading.BoundedSemaphore(want) + _limit = want + return _sem + + +def cli_budget_status() -> dict[str, int]: + """Snapshot for health/metrics.""" + _ensure_sem() + with _in_use_lock: + used = _in_use + return {"limit": _limit, "in_use": used, "available": max(0, _limit - used)} + + +def clamp_cli_workers(requested: int, *, hard_cap: int) -> int: + """Clamp a feature concurrency against the global CLI budget and a hard cap.""" + budget = max(1, int(getattr(settings, "cli_max_concurrent", 20) or 20)) + return max(1, min(int(hard_cap), int(requested or 1), budget)) + + +@contextmanager +def acquire_cli_slot(*, blocking: bool = True, timeout: float | None = None) -> Iterator[bool]: + """Acquire one global CLI slot for an SSH/Netmiko session. + + Yields True when acquired. When ``blocking=False`` and the budget is full, + yields False without waiting. + """ + global _in_use + sem = _ensure_sem() + acquired = False + if blocking: + if timeout is None: + sem.acquire() + acquired = True + else: + acquired = bool(sem.acquire(timeout=float(timeout))) + else: + acquired = bool(sem.acquire(blocking=False)) + if acquired: + with _in_use_lock: + _in_use += 1 + try: + yield acquired + finally: + if acquired: + with _in_use_lock: + _in_use = max(0, _in_use - 1) + sem.release() diff --git a/netx_api/cli_timeout.py b/netx_api/cli_timeout.py new file mode 100644 index 0000000..7a0dd37 --- /dev/null +++ b/netx_api/cli_timeout.py @@ -0,0 +1,78 @@ +"""Shared CLI timeout runner — one pool, force-close Netmiko on timeout.""" +from __future__ import annotations + +import logging +import threading +from concurrent.futures import Future, ThreadPoolExecutor, TimeoutError as FuturesTimeout +from contextlib import nullcontext +from typing import Any, Callable, TypeVar + +from .cli_budget import acquire_cli_slot +from .config import settings +from .ne_session_factory import close_netmiko_connection + +_log = logging.getLogger("netx.cli.timeout") +_pool_lock = threading.Lock() +_pool: ThreadPoolExecutor | None = None + +T = TypeVar("T") + + +def _timeout_pool() -> ThreadPoolExecutor: + global _pool + with _pool_lock: + if _pool is None: + # Shared watchdog pool: enough for concurrent timed CLI jobs, not per-task. + workers = max(4, int(getattr(settings, "cli_timeout_pool_workers", 16) or 16)) + _pool = ThreadPoolExecutor(max_workers=workers, thread_name_prefix="cli-timeout") + return _pool + + +def shutdown_cli_timeout_pool(*, wait: bool = False) -> None: + global _pool + with _pool_lock: + if _pool is not None: + _pool.shutdown(wait=wait, cancel_futures=True) + _pool = None + + +def run_cli_with_timeout( + fn: Callable[[], T], + *, + timeout_sec: float, + conn_holder: dict[str, Any] | None = None, + label: str = "cli", + acquire_budget: bool = True, +) -> T: + """Run ``fn`` with a wall-clock timeout. + + If ``conn_holder`` is provided and ``fn`` stores the live connection under + key ``\"conn\"``, a timeout will force-close that connection so the worker + thread can exit instead of leaking SSH/FD until Netmiko finishes. + + CLI budget is acquired *outside* the timed window so queue wait is not + counted against the SSH timeout. + """ + budget = max(1.0, float(timeout_sec)) + slot_cm = acquire_cli_slot() if acquire_budget else nullcontext(True) + + with slot_cm as ok: + if acquire_budget and not ok: + raise TimeoutError(f"{label}_cli_budget_unavailable") + fut: Future[T] = _timeout_pool().submit(fn) + try: + return fut.result(timeout=budget) + except FuturesTimeout as exc: + conn = None + if conn_holder is not None: + conn = conn_holder.get("conn") + conn_holder["timed_out"] = True + if conn is not None: + _log.warning("%s timeout after %.0fs — force-closing netmiko connection", label, budget) + try: + close_netmiko_connection(conn) + except Exception: # noqa: BLE001 + _log.exception("%s force-close failed", label) + else: + _log.warning("%s timeout after %.0fs (no live connection to close yet)", label, budget) + raise TimeoutError(f"{label}_timeout ({int(budget)}s)") from exc diff --git a/netx_api/config.py b/netx_api/config.py index a05752c..ca4852a 100644 --- a/netx_api/config.py +++ b/netx_api/config.py @@ -143,6 +143,24 @@ class Settings(BaseSettings): # When true (default), API also runs config_sync / lldp / port_traffic schedulers. # Production split: set false and run `python -m netx_api.worker` beside the API. run_inline_schedulers: bool = True + # SQLAlchemy QueuePool (API + background workers share one engine). + # Rule of thumb: pool_size + max_overflow >= HTTP peak + cli_max_concurrent + UME WS burst. + db_pool_size: int = 20 + db_max_overflow: int = 20 + db_pool_recycle_sec: int = 1800 + db_pool_timeout_sec: int = 30 + # Global Netmiko/SSH concurrency across discover / collect / config_sync / port_traffic. + cli_max_concurrent: int = 20 + # Shared timeout watchdog pool (not per-task executors). + cli_timeout_pool_workers: int = 16 + # Port-traffic: how many devices may collect in parallel on the scheduler tick. + port_traffic_dispatch_workers: int = 4 + # Bound async audit queue; drop oldest when full to protect RSS. + audit_queue_max: int = 5000 + # Cap NE collection output files (bytes); 0 = unlimited (not recommended). + ne_collect_max_output_bytes: int = 8 * 1024 * 1024 + # WebCRT session transcript rotate size (bytes); 0 disables rotate. + webcrt_session_log_max_bytes: int = 4 * 1024 * 1024 settings = Settings() diff --git a/netx_api/config_sync_runner.py b/netx_api/config_sync_runner.py index 503968a..3a064dd 100644 --- a/netx_api/config_sync_runner.py +++ b/netx_api/config_sync_runner.py @@ -5,7 +5,7 @@ from __future__ import annotations import io import logging import time -from concurrent.futures import ThreadPoolExecutor, TimeoutError as FuturesTimeout +from concurrent.futures import ThreadPoolExecutor from datetime import datetime from threading import Lock from typing import Any @@ -38,10 +38,12 @@ def _format_error(exc: BaseException) -> str: def _pool_for_cycle(cycle_id: str, concurrency: int) -> ThreadPoolExecutor: + from .cli_budget import clamp_cli_workers + with _pools_lock: pool = _pools.get(cycle_id) if pool is None: - workers = max(1, min(30, int(concurrency or 5))) + workers = clamp_cli_workers(int(concurrency or 5), hard_cap=30) pool = ThreadPoolExecutor(max_workers=workers, thread_name_prefix=f"cfg-sync-{cycle_id[:8]}") _pools[cycle_id] = pool return pool @@ -59,6 +61,19 @@ def _release_pool(cycle_id: str) -> None: _log.exception("config_sync pool shutdown failed cycle=%s", cycle_id) +def shutdown_config_sync_pools(*, wait: bool = False) -> None: + with _pools_lock: + pools = list(_pools.items()) + _pools.clear() + for cycle_id, pool in pools: + try: + pool.shutdown(wait=wait, cancel_futures=True) + except TypeError: + pool.shutdown(wait=wait) + except Exception: + _log.exception("config_sync pool shutdown failed cycle=%s", cycle_id) + + def _update_task(task_id: str, **fields: Any) -> None: db = SessionLocal() try: @@ -102,7 +117,12 @@ def _claim_task(cycle_id: str, task_id: str) -> bool: return False -def _collect_commands(creds: dict[str, Any], commands: list[str]) -> list[str]: +def _collect_commands( + creds: dict[str, Any], + commands: list[str], + *, + conn_holder: dict[str, Any] | None = None, +) -> list[str]: per_cmd = int(settings.ne_collect_read_timeout_sec or 120) session_timeout = per_cmd * max(1, len(commands)) + 60 log_buf = io.BytesIO() @@ -110,9 +130,13 @@ def _collect_commands(creds: dict[str, Any], commands: list[str]) -> list[str]: conn = open_netmiko_connection(creds, session_timeout=session_timeout, session_log=log_buf) except Exception as exc: raise RuntimeError(format_cli_failure(exc, session_log_text(log_buf))) from exc + if conn_holder is not None: + conn_holder["conn"] = conn try: outputs: list[str] = [] for command in commands: + if conn_holder is not None and conn_holder.get("timed_out"): + raise TimeoutError("config_sync_aborted") try: out = send_show_command(conn, command, read_timeout=per_cmd) except Exception as exc: @@ -120,19 +144,25 @@ def _collect_commands(creds: dict[str, Any], commands: list[str]) -> list[str]: outputs.append(str(out or "")) return outputs finally: + if conn_holder is not None: + conn_holder.pop("conn", None) close_netmiko_connection(conn) def _collect_with_timeout(creds: dict[str, Any], commands: list[str]) -> list[str]: + from .cli_timeout import run_cli_with_timeout + per_cmd = int(settings.ne_collect_read_timeout_sec or 120) cap = int(settings.ne_collect_run_timeout_cap_sec or 600) budget = min(cap, per_cmd * max(1, len(commands)) + 90) - with ThreadPoolExecutor(max_workers=1) as pool: - fut = pool.submit(_collect_commands, creds, commands) - try: - return fut.result(timeout=budget) - except FuturesTimeout as exc: - raise TimeoutError(f"config_sync_timeout ({budget}s)") from exc + holder: dict[str, Any] = {} + return run_cli_with_timeout( + lambda: _collect_commands(creds, commands, conn_holder=holder), + timeout_sec=budget, + conn_holder=holder, + label="config_sync", + acquire_budget=True, + ) def _history_keep(db) -> int: diff --git a/netx_api/db.py b/netx_api/db.py index 4e4f544..cccfccb 100644 --- a/netx_api/db.py +++ b/netx_api/db.py @@ -1,7 +1,10 @@ from __future__ import annotations +from typing import Any + from sqlalchemy import create_engine from sqlalchemy.orm import DeclarativeBase, sessionmaker +from sqlalchemy.pool import QueuePool from .config import settings @@ -10,7 +13,26 @@ class Base(DeclarativeBase): pass -engine = create_engine(settings.database_url, future=True, pool_pre_ping=True) +def _engine_kwargs() -> dict[str, Any]: + url = str(settings.database_url or "") + kwargs: dict[str, Any] = {"future": True, "pool_pre_ping": True} + # SQLite (tests / local) does not use QueuePool the same way. + if url.startswith("sqlite"): + kwargs["connect_args"] = {"check_same_thread": False} + return kwargs + kwargs.update( + { + "poolclass": QueuePool, + "pool_size": max(1, int(getattr(settings, "db_pool_size", 20) or 20)), + "max_overflow": max(0, int(getattr(settings, "db_max_overflow", 20) or 20)), + "pool_recycle": max(60, int(getattr(settings, "db_pool_recycle_sec", 1800) or 1800)), + "pool_timeout": max(1, int(getattr(settings, "db_pool_timeout_sec", 30) or 30)), + } + ) + return kwargs + + +engine = create_engine(settings.database_url, **_engine_kwargs()) SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False, expire_on_commit=False) @@ -20,3 +42,26 @@ def get_db(): yield db finally: db.close() + + +def db_pool_status() -> dict[str, Any]: + """Best-effort QueuePool snapshot for readiness / metrics.""" + pool = getattr(engine, "pool", None) + if pool is None: + return {"backend": "none"} + out: dict[str, Any] = {"backend": type(pool).__name__} + for key, meth in ( + ("size", "size"), + ("checked_in", "checkedin"), + ("checked_out", "checkedout"), + ("overflow", "overflow"), + ): + fn = getattr(pool, meth, None) + if callable(fn): + try: + out[key] = int(fn()) + except Exception: # noqa: BLE001 + pass + out["pool_size_cfg"] = int(getattr(settings, "db_pool_size", 20) or 20) + out["max_overflow_cfg"] = int(getattr(settings, "db_max_overflow", 20) or 20) + return out diff --git a/netx_api/integrations_router.py b/netx_api/integrations_router.py index c60e88a..a0d7174 100644 --- a/netx_api/integrations_router.py +++ b/netx_api/integrations_router.py @@ -25,6 +25,9 @@ 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 — DB plus scheduler deployment hint.""" + from .cli_budget import cli_budget_status + from .db import db_pool_status + out: dict[str, Any] = {"status": "ok", "probe": "ready"} try: db.execute(sql_text("select 1")) @@ -36,6 +39,8 @@ def health_ready(db: Session = Depends(get_db)) -> dict[str, Any]: "db": "down", "error": str(exc)[:240], } + out["db_pool"] = db_pool_status() + out["cli_budget"] = cli_budget_status() inline = bool(getattr(settings, "run_inline_schedulers", True)) out["schedulers"] = { "inline": inline, diff --git a/netx_api/main.py b/netx_api/main.py index 563158c..74946a1 100644 --- a/netx_api/main.py +++ b/netx_api/main.py @@ -36,10 +36,8 @@ from .ume_support import ( # noqa: F401 — tests import from main _classify_protocol_bucket, _protocol_bucket_label, ) -import netx_api.ume_support as ume_support -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 +from .metrics_router import router as metrics_router _schedule_log = logging.getLogger("netx.ume.schedule") _BOOT_MONO = time.monotonic() @@ -48,15 +46,13 @@ _BOOT_MONO = time.monotonic() @asynccontextmanager async def lifespan(_app: FastAPI) -> AsyncIterator[None]: from .app_startup import run_api_startup + from .app_shutdown import shutdown_runtime run_api_startup() try: yield finally: - shutdown_oclaw_alarm_forwarder() - if ume_support._UME_WS_STOP_EVENT is not None: - ume_support._UME_WS_STOP_EVENT.set() - shutdown_ws_consumer() + shutdown_runtime(reason="api_lifespan") app = FastAPI( @@ -80,6 +76,7 @@ app.include_router(lldp_collect_router) app.include_router(ops_router) app.include_router(sql_router) app.include_router(integrations_router) +app.include_router(metrics_router) app.include_router(ume_router) app.include_router(alarms_router) parser_cfg = load_parser_config() diff --git a/netx_api/metrics_router.py b/netx_api/metrics_router.py new file mode 100644 index 0000000..05e430d --- /dev/null +++ b/netx_api/metrics_router.py @@ -0,0 +1,113 @@ +"""Minimal process metrics for ops (Prometheus text + JSON).""" +from __future__ import annotations + +import os +import threading +import time +from typing import Any + +from fastapi import APIRouter, Response + +from .audit_async import audit_queue_status +from .cli_budget import cli_budget_status +from .config import settings +from .db import db_pool_status +from .oclaw_alarm_forwarder import forwarder_status + +router = APIRouter(tags=["metrics"]) + +_BOOT_MONO = time.monotonic() + + +def collect_runtime_metrics() -> dict[str, Any]: + out: dict[str, Any] = { + "uptime_sec": round(time.monotonic() - _BOOT_MONO, 1), + "pid": os.getpid(), + "thread_count": threading.active_count(), + "db_pool": db_pool_status(), + "cli_budget": cli_budget_status(), + "audit_queue": audit_queue_status(), + "oclaw_forwarder": forwarder_status(), + "schedulers_inline": bool(getattr(settings, "run_inline_schedulers", True)), + } + try: + import resource + + usage = resource.getrusage(resource.RUSAGE_SELF) + # ru_maxrss is KB on Linux, bytes on macOS; report raw + note. + out["ru_maxrss"] = int(usage.ru_maxrss) + except Exception: # noqa: BLE001 + pass + try: + from .port_traffic_scheduler import port_traffic_scheduler_status + + out["port_traffic"] = port_traffic_scheduler_status() + except Exception: # noqa: BLE001 + pass + try: + from .webcrt_session_registry import active_session_count, list_sessions + + out["webcrt"] = { + "active_sessions": active_session_count(), + "max_sessions": int(getattr(settings, "webcrt_max_sessions", 20) or 20), + } + # Avoid dumping full session list into metrics; depth aggregate only. + sess = list_sessions() + dropped = 0 + for item in sess.get("items") or []: + dropped += int(item.get("queue_dropped") or 0) + out["webcrt"]["queue_dropped_total"] = dropped + except Exception: # noqa: BLE001 + pass + return out + + +def _prom_lines(metrics: dict[str, Any]) -> str: + lines: list[str] = [] + lines.append(f'netx_uptime_seconds {metrics.get("uptime_sec") or 0}') + lines.append(f'netx_thread_count {metrics.get("thread_count") or 0}') + pool = metrics.get("db_pool") or {} + for key in ("checked_out", "checked_in", "overflow", "size"): + if key in pool: + lines.append(f"netx_db_pool_{key} {pool[key]}") + budget = metrics.get("cli_budget") or {} + for key in ("limit", "in_use", "available"): + if key in budget: + lines.append(f"netx_cli_budget_{key} {budget[key]}") + audit = metrics.get("audit_queue") or {} + for key in ("depth", "dropped", "maxsize"): + if key in audit: + lines.append(f"netx_audit_queue_{key} {audit[key]}") + fwd = metrics.get("oclaw_forwarder") or {} + for key, prom in ( + ("queue_size", "netx_oclaw_forwarder_queue_size"), + ("published_ok", "netx_oclaw_forwarder_published_ok"), + ("published_fail", "netx_oclaw_forwarder_published_fail"), + ): + if key in fwd: + lines.append(f"{prom} {int(fwd.get(key) or 0)}") + pt = metrics.get("port_traffic") or {} + if pt.get("last_tick_age_sec") is not None: + lines.append(f'netx_port_traffic_tick_age_seconds {pt["last_tick_age_sec"]}') + if pt.get("last_purge_age_sec") is not None: + lines.append(f'netx_port_traffic_purge_age_seconds {pt["last_purge_age_sec"]}') + web = metrics.get("webcrt") or {} + if "active_sessions" in web: + lines.append(f'netx_webcrt_active_sessions {web["active_sessions"]}') + if "queue_dropped_total" in web: + lines.append(f'netx_webcrt_queue_dropped_total {web["queue_dropped_total"]}') + if "ru_maxrss" in metrics: + lines.append(f'netx_ru_maxrss {metrics["ru_maxrss"]}') + lines.append("") + return "\n".join(lines) + + +@router.get("/metrics") +def metrics_prometheus() -> Response: + body = _prom_lines(collect_runtime_metrics()) + return Response(content=body, media_type="text/plain; version=0.0.4; charset=utf-8") + + +@router.get("/metrics/json") +def metrics_json() -> dict[str, Any]: + return collect_runtime_metrics() diff --git a/netx_api/ne_collect_runner.py b/netx_api/ne_collect_runner.py index 94d05f9..6d066bb 100644 --- a/netx_api/ne_collect_runner.py +++ b/netx_api/ne_collect_runner.py @@ -4,7 +4,7 @@ import logging import re import time import traceback -from concurrent.futures import ThreadPoolExecutor, TimeoutError as FuturesTimeout +from concurrent.futures import ThreadPoolExecutor from datetime import datetime from pathlib import Path from typing import Any @@ -14,6 +14,7 @@ from .config import settings from .db import SessionLocal from fastapi import HTTPException +from .cli_budget import clamp_cli_workers from .cli_resolve import resolve_cli_target from .models import NeCollectionJob, NeCollectionRun from .ne_collection_paths import clear_run_output_files, run_output_dir @@ -28,11 +29,21 @@ _executor: ThreadPoolExecutor | None = None def _executor_pool() -> ThreadPoolExecutor: global _executor if _executor is None: - workers = max(1, int(settings.ne_collect_max_workers or 5)) + workers = clamp_cli_workers(int(settings.ne_collect_max_workers or 5), hard_cap=32) _executor = ThreadPoolExecutor(max_workers=workers, thread_name_prefix="ne-collect") return _executor +def shutdown_ne_collect_executor(*, wait: bool = False) -> None: + global _executor + if _executor is not None: + try: + _executor.shutdown(wait=wait, cancel_futures=True) + except TypeError: + _executor.shutdown(wait=wait) + _executor = None + + def _format_run_error(exc: BaseException) -> str: head = f"{type(exc).__name__}: {exc}" tb = traceback.format_exc().strip() @@ -50,14 +61,19 @@ def _collect_on_device( commands: list[str], *, read_timeout_sec: int | None = None, + conn_holder: dict[str, Any] | None = None, ) -> str: per_cmd = int(read_timeout_sec if read_timeout_sec is not None else (settings.ne_collect_read_timeout_sec or 120)) session_timeout = per_cmd * max(1, len(commands)) + 60 conn = open_netmiko_connection(creds, session_timeout=session_timeout) + if conn_holder is not None: + conn_holder["conn"] = conn try: prompt = str(conn.find_prompt() or "") chunks: list[str] = [] for command in commands: + if conn_holder is not None and conn_holder.get("timed_out"): + raise TimeoutError("collection_aborted") ts = datetime.now().isoformat(timespec="seconds") chunks.append(f'>>> [{ts}] {{"String":"{command}", "Match":"{prompt}", "Timeout":0}}\n') out = send_show_command(conn, command, read_timeout=per_cmd) @@ -65,6 +81,8 @@ def _collect_on_device( chunks.append("\n") return "".join(chunks) finally: + if conn_holder is not None: + conn_holder.pop("conn", None) close_netmiko_connection(conn) @@ -74,15 +92,21 @@ def _collect_with_timeout( *, read_timeout_sec: int | None = None, ) -> str: + from .cli_timeout import run_cli_with_timeout + per_cmd = int(read_timeout_sec if read_timeout_sec is not None else (settings.ne_collect_read_timeout_sec or 120)) cap = int(settings.ne_collect_run_timeout_cap_sec or 600) budget = min(cap, per_cmd * max(1, len(commands)) + 90) - with ThreadPoolExecutor(max_workers=1) as pool: - fut = pool.submit(_collect_on_device, creds, commands, read_timeout_sec=per_cmd) - try: - return fut.result(timeout=budget) - except FuturesTimeout as exc: - raise TimeoutError(f"collection_timeout ({budget}s)") from exc + holder: dict[str, Any] = {} + return run_cli_with_timeout( + lambda: _collect_on_device( + creds, commands, read_timeout_sec=per_cmd, conn_holder=holder + ), + timeout_sec=budget, + conn_holder=holder, + label="collection", + acquire_budget=True, + ) def _update_run(run_id: str, **fields: Any) -> None: @@ -192,7 +216,15 @@ def _run_single(job_id: str, run_id: str, commands: list[str]) -> None: out_dir.mkdir(parents=True, exist_ok=True) filename = f"{name_part}-{ip_part}-{ts}.txt" full_path = out_dir / filename - full_path.write_text(output, encoding="utf-8", errors="replace") + max_bytes = int(getattr(settings, "ne_collect_max_output_bytes", 0) or 0) + if max_bytes > 0 and len(output.encode("utf-8", errors="replace")) > max_bytes: + # Truncate by characters roughly under byte budget. + truncated = output.encode("utf-8", errors="replace")[:max_bytes] + text = truncated.decode("utf-8", errors="replace") + text += f"\n...[truncated {max_bytes} bytes cap]\n" + full_path.write_text(text, encoding="utf-8", errors="replace") + else: + full_path.write_text(output, encoding="utf-8", errors="replace") rel_path = str(rel_dir / filename).replace("\\", "/") _update_run( run_id, diff --git a/netx_api/ne_connect.py b/netx_api/ne_connect.py index ff9a107..08287d2 100644 --- a/netx_api/ne_connect.py +++ b/netx_api/ne_connect.py @@ -30,11 +30,23 @@ _DETAIL_MAX = 8000 def _executor_pool() -> ThreadPoolExecutor: global _executor if _executor is None: - workers = max(1, int(settings.ne_connect_max_workers or 5)) + from .cli_budget import clamp_cli_workers + + workers = clamp_cli_workers(int(settings.ne_connect_max_workers or 5), hard_cap=32) _executor = ThreadPoolExecutor(max_workers=workers, thread_name_prefix="ne-connect") return _executor +def shutdown_ne_connect_executor(*, wait: bool = False) -> None: + global _executor + if _executor is not None: + try: + _executor.shutdown(wait=wait, cancel_futures=True) + except TypeError: + _executor.shutdown(wait=wait) + _executor = None + + def _truncate_detail(text: str) -> str: return str(text or "")[:_DETAIL_MAX] diff --git a/netx_api/port_traffic_runner.py b/netx_api/port_traffic_runner.py index fb1aac5..adfdd37 100644 --- a/netx_api/port_traffic_runner.py +++ b/netx_api/port_traffic_runner.py @@ -4,15 +4,14 @@ from __future__ import annotations import logging import time -from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import datetime -from threading import Lock from typing import Any from uuid import uuid4 from fastapi import HTTPException from .cli_resolve import resolve_cli_target +from .cli_timeout import run_cli_with_timeout from .config import settings from .db import SessionLocal from .models import PortTrafficDevice, PortTrafficEvent, PortTrafficSample, PortTrafficTarget @@ -22,8 +21,6 @@ from .port_traffic_commands import commands_for_vendor, detail_command from .port_traffic_parsers import parse_interface_detail, resolve_util_pct _log = logging.getLogger("netx.port_traffic.runner") -_pools: dict[str, ThreadPoolExecutor] = {} -_pools_lock = Lock() def _utcnow() -> datetime: @@ -93,29 +90,6 @@ def _finish_collect_round(device_id: str, *, error: str = "") -> None: db.commit() finally: db.close() - _release_pool(device_id) - - -def _pool_for_device(device_id: str, concurrency: int) -> ThreadPoolExecutor: - with _pools_lock: - pool = _pools.get(device_id) - if pool is None: - workers = max(1, min(5, int(concurrency or 1))) - pool = ThreadPoolExecutor(max_workers=workers, thread_name_prefix=f"pt-{device_id[:8]}") - _pools[device_id] = pool - return pool - - -def _release_pool(device_id: str) -> None: - with _pools_lock: - pool = _pools.pop(device_id, None) - if pool is not None: - try: - pool.shutdown(wait=False, cancel_futures=False) - except TypeError: - pool.shutdown(wait=False) - except Exception: - _log.exception("port_traffic pool shutdown failed device=%s", device_id) def _claim_collect_round(device_id: str) -> list[str] | None: @@ -225,7 +199,6 @@ def _sample_targets_shared_session(device_id: str, target_ids: list[str]) -> tup """ per_cmd = int(settings.ne_collect_read_timeout_sec or 120) cap = int(settings.ne_collect_run_timeout_cap_sec or 600) - errors = 0 db = SessionLocal() try: @@ -277,30 +250,52 @@ def _sample_targets_shared_session(device_id: str, target_ids: list[str]) -> tup return 0, "" budget = min(cap, per_cmd * max(1, len(ifaces)) + 90) - conn = None - try: + holder: dict[str, Any] = {} + + def _run_session() -> tuple[int, str]: conn = open_netmiko_connection(creds, session_timeout=budget) - for tid, ifname in ifaces: - try: - cmd = detail_command(cmds, ifname) - raw = send_show_command(conn, cmd, read_timeout=per_cmd) - parsed = parse_interface_detail(raw, vendor_key) - _save_sample(tid, parsed) - except Exception as exc: - errors += 1 - _set_target_error(tid, _format_error(exc)) - _log.exception("port_traffic iface sample failed device=%s if=%s", device_id, ifname) + holder["conn"] = conn + local_errors = 0 + try: + for tid, ifname in ifaces: + if holder.get("timed_out"): + raise TimeoutError("port_traffic_aborted") + try: + cmd = detail_command(cmds, ifname) + raw = send_show_command(conn, cmd, read_timeout=per_cmd) + parsed = parse_interface_detail(raw, vendor_key) + _save_sample(tid, parsed) + except Exception as exc: + local_errors += 1 + _set_target_error(tid, _format_error(exc)) + _log.exception( + "port_traffic iface sample failed device=%s if=%s", device_id, ifname + ) + return local_errors, "" + finally: + holder.pop("conn", None) + close_netmiko_connection(conn) + + try: + # Budget acquired inside run_cli_with_timeout; avoid double-acquire. + return run_cli_with_timeout( + _run_session, + timeout_sec=budget, + conn_holder=holder, + label="port_traffic", + acquire_budget=True, + ) + except TimeoutError as exc: + msg = str(exc)[:1020] + for tid, _ in ifaces: + _set_target_error(tid, msg) + return len(ifaces), msg except Exception as exc: msg = _format_error(exc) for tid, _ in ifaces: _set_target_error(tid, msg) - errors = len(ifaces) _log.exception("port_traffic session failed device=%s", device_id) - return errors, msg - finally: - if conn is not None: - close_netmiko_connection(conn) - return errors, "" + return len(ifaces), msg def dispatch_collect(device_id: str) -> int: diff --git a/netx_api/port_traffic_scheduler.py b/netx_api/port_traffic_scheduler.py index 2622551..deaa031 100644 --- a/netx_api/port_traffic_scheduler.py +++ b/netx_api/port_traffic_scheduler.py @@ -4,8 +4,10 @@ from __future__ import annotations import logging import threading +from concurrent.futures import ThreadPoolExecutor from datetime import datetime +from .cli_budget import clamp_cli_workers from .config import settings from .db import SessionLocal from .models import PortTrafficDevice @@ -15,15 +17,49 @@ from .port_traffic_service import purge_expired_samples _log = logging.getLogger("netx.port_traffic.scheduler") _stop = threading.Event() _thread: threading.Thread | None = None -_purge_counter = 0 +_purge_thread: threading.Thread | None = None +_dispatch_pool: ThreadPoolExecutor | None = None +_pool_lock = threading.Lock() +_last_tick_mono: float = 0.0 +_last_purge_mono: float = 0.0 def _utcnow() -> datetime: return datetime.utcnow() +def _dispatch_pool_get() -> ThreadPoolExecutor: + global _dispatch_pool + with _pool_lock: + if _dispatch_pool is None: + workers = clamp_cli_workers( + int(getattr(settings, "port_traffic_dispatch_workers", 4) or 4), + hard_cap=16, + ) + _dispatch_pool = ThreadPoolExecutor( + max_workers=workers, thread_name_prefix="pt-dispatch" + ) + return _dispatch_pool + + +def shutdown_port_traffic_dispatch_pool(*, wait: bool = False) -> None: + global _dispatch_pool + with _pool_lock: + if _dispatch_pool is not None: + try: + _dispatch_pool.shutdown(wait=wait, cancel_futures=True) + except TypeError: + _dispatch_pool.shutdown(wait=wait) + _dispatch_pool = None + + def try_dispatch_due_devices() -> int: - """Dispatch collect rounds for due running devices. Returns number started.""" + """Claim due devices and sample them on a bounded pool (non-blocking submit).""" + import time as _time + + global _last_tick_mono + _last_tick_mono = _time.monotonic() + db = SessionLocal() try: devices = ( @@ -47,46 +83,76 @@ def try_dispatch_due_devices() -> int: finally: db.close() + if not due_ids: + return 0 + + pool = _dispatch_pool_get() started = 0 for did in due_ids: try: - n = dispatch_collect(did) - if n: - started += 1 - _log.info("port_traffic collect started device=%s targets=%s", did, n) + pool.submit(_run_device_collect, did) + started += 1 except Exception: - _log.exception("port_traffic dispatch failed device=%s", did) + _log.exception("port_traffic dispatch submit failed device=%s", did) return started +def _run_device_collect(device_id: str) -> None: + try: + n = dispatch_collect(device_id) + if n: + _log.info("port_traffic collect finished device=%s targets=%s", device_id, n) + except Exception: + _log.exception("port_traffic dispatch failed device=%s", device_id) + + def try_dispatch_due_tasks() -> int: return try_dispatch_due_devices() +def _purge_once() -> None: + import time as _time + + global _last_purge_mono + db = SessionLocal() + try: + purge_expired_samples(db) + _last_purge_mono = _time.monotonic() + except Exception: + _log.exception("port_traffic purge failed") + finally: + db.close() + + def _loop() -> None: - global _purge_counter tick = max(5, int(settings.port_traffic_scheduler_tick_sec or 15)) _log.info("port_traffic scheduler started tick=%ss", tick) while not _stop.is_set(): try: if bool(settings.port_traffic_scheduler_enabled): try_dispatch_due_devices() - _purge_counter += 1 - if _purge_counter >= 20: - _purge_counter = 0 - db = SessionLocal() - try: - purge_expired_samples(db) - finally: - db.close() except Exception: _log.exception("port_traffic scheduler tick failed") _stop.wait(tick) _log.info("port_traffic scheduler stopped") +def _purge_loop() -> None: + # Independent of collect latency — every ~5 minutes. + interval = 300 + _log.info("port_traffic purge loop started interval=%ss", interval) + while not _stop.is_set(): + try: + if bool(settings.port_traffic_scheduler_enabled): + _purge_once() + except Exception: + _log.exception("port_traffic purge loop failed") + _stop.wait(interval) + _log.info("port_traffic purge loop stopped") + + def start_port_traffic_scheduler() -> None: - global _thread + global _thread, _purge_thread if not bool(settings.port_traffic_scheduler_enabled): _log.info("port_traffic scheduler disabled by settings") return @@ -95,7 +161,22 @@ def start_port_traffic_scheduler() -> None: _stop.clear() _thread = threading.Thread(target=_loop, name="port-traffic-scheduler", daemon=True) _thread.start() + if not (_purge_thread and _purge_thread.is_alive()): + _purge_thread = threading.Thread(target=_purge_loop, name="port-traffic-purge", daemon=True) + _purge_thread.start() def stop_port_traffic_scheduler() -> None: _stop.set() + shutdown_port_traffic_dispatch_pool(wait=False) + + +def port_traffic_scheduler_status() -> dict: + import time as _time + + now = _time.monotonic() + return { + "last_tick_age_sec": (now - _last_tick_mono) if _last_tick_mono else None, + "last_purge_age_sec": (now - _last_purge_mono) if _last_purge_mono else None, + "running": bool(_thread and _thread.is_alive()), + } diff --git a/netx_api/topology_discover_jobs.py b/netx_api/topology_discover_jobs.py index 56967a8..873929e 100644 --- a/netx_api/topology_discover_jobs.py +++ b/netx_api/topology_discover_jobs.py @@ -1,6 +1,7 @@ """Discover job lifecycle: start, background run, stale reclaim.""" from __future__ import annotations +import logging import threading from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import datetime, timedelta @@ -26,6 +27,8 @@ from .topology_fabric import ( ) from .topology_schemas import FabricDiscoverJobOut, FabricDiscoverRequest +_log = logging.getLogger("netx.topology.discover") + def _run_discover_job(job_id: str, body: FabricDiscoverRequest) -> None: db = SessionLocal() @@ -54,7 +57,9 @@ def _run_discover_job(job_id: str, body: FabricDiscoverRequest) -> None: except Exception: # noqa: BLE001 db.rollback() - concurrency = max(1, min(32, int(body.concurrency or 4))) + from .cli_budget import clamp_cli_workers + + concurrency = clamp_cli_workers(int(body.concurrency or 4), hard_cap=32) added = 0 updated = 0 stale = 0 @@ -142,7 +147,7 @@ def _run_discover_job(job_id: str, body: FabricDiscoverRequest) -> None: keep = int(getattr(ensure_policy(db), "history_keep", DEFAULT_HISTORY_KEEP) or 0) prune_discover_jobs(db, keep=keep) except Exception: # noqa: BLE001 - pass + _log.warning("prune_discover_jobs failed job=%s", job_id, exc_info=True) except Exception as exc: # noqa: BLE001 db.rollback() job = db.get(TopoDiscoverJob, job_id) diff --git a/netx_api/topology_discover_scan.py b/netx_api/topology_discover_scan.py index 06fcb7d..a647645 100644 --- a/netx_api/topology_discover_scan.py +++ b/netx_api/topology_discover_scan.py @@ -80,7 +80,12 @@ def _discover_one_target( else: exec_kwargs["ne_id"] = target["ne_id"] try: - exec_out = execute_managed_ne_commands(db, [cmd], **exec_kwargs) + from .cli_budget import acquire_cli_slot + + with acquire_cli_slot() as ok: + if not ok: + return {**base, "ok": False, "command": cmd, "error": "cli_budget_unavailable"} + exec_out = execute_managed_ne_commands(db, [cmd], **exec_kwargs) except HTTPException as exc: return { **base, diff --git a/netx_api/ume_alarm_apply.py b/netx_api/ume_alarm_apply.py index 55e15fa..c01f77d 100644 --- a/netx_api/ume_alarm_apply.py +++ b/netx_api/ume_alarm_apply.py @@ -129,10 +129,16 @@ def _mark_alarm_cleared_tombstone(alarm_key: str) -> None: expires = time.time() + ttl_s with _cleared_tombstone_lock: _cleared_tombstones[key] = expires - if len(_cleared_tombstones) > 50000: - now = time.time() - stale = [k for k, exp in _cleared_tombstones.items() if exp <= now] - for k in stale[:10000]: + now = time.time() + # Always sweep expired keys; hard-cap with oldest-first eviction. + stale = [k for k, exp in _cleared_tombstones.items() if exp <= now] + for k in stale: + _cleared_tombstones.pop(k, None) + max_keys = 50000 + if len(_cleared_tombstones) > max_keys: + overflow = len(_cleared_tombstones) - max_keys + oldest = sorted(_cleared_tombstones.items(), key=lambda kv: kv[1])[:overflow] + for k, _ in oldest: _cleared_tombstones.pop(k, None) diff --git a/netx_api/webcrt_session_model.py b/netx_api/webcrt_session_model.py index cfb0e93..2cb25b6 100644 --- a/netx_api/webcrt_session_model.py +++ b/netx_api/webcrt_session_model.py @@ -204,8 +204,30 @@ class WebcrtSession: try: self._log_fh.write(text) self._log_fh.flush() + max_bytes = int(getattr(settings, "webcrt_session_log_max_bytes", 0) or 0) + if max_bytes > 0: + try: + pos = int(self._log_fh.tell()) + except Exception: # noqa: BLE001 + pos = 0 + if pos >= max_bytes: + self._log_fh.write( + f"\n# rotated at {pos} bytes (cap={max_bytes}) ts={_utc_iso()}\n" + ) + self._log_fh.flush() + self._log_fh.close() + self._log_fh = None + # Re-open fresh file (append mode continues same path; rotate by rename). + try: + path = _session_log_path(self.session_id) + rotated = path.with_suffix(path.suffix + f".{int(pos)}.old") + if path.exists(): + path.replace(rotated) + except Exception: # noqa: BLE001 + _log.debug("webcrt session log rotate rename failed", exc_info=True) + self.open_session_log() except Exception: - pass + _log.warning("webcrt session log append failed session=%s", self.session_id, exc_info=True) def close_session_log(self) -> None: fh = self._log_fh diff --git a/netx_api/webcrt_session_registry.py b/netx_api/webcrt_session_registry.py index 4c4918d..fb67463 100644 --- a/netx_api/webcrt_session_registry.py +++ b/netx_api/webcrt_session_registry.py @@ -577,6 +577,21 @@ def close_session(session_id: str, *, reason: str = "closed", client: str = "") return {"ok": True, "session_id": session_id, "closed": True, "reason": reason} +def close_all_sessions(*, reason: str = "shutdown") -> int: + """Close every active WebCRT session (API shutdown). Returns count closed.""" + with _sessions_lock: + ids = [sid for sid, s in _sessions.items() if not s.closed] + closed = 0 + for sid in ids: + try: + out = close_session(sid, reason=reason, client="shutdown") + if out.get("closed"): + closed += 1 + except Exception: # noqa: BLE001 + _log.exception("webcrt close_all failed session=%s", sid) + return closed + + def list_sessions() -> dict[str, Any]: with _sessions_lock: items = [] diff --git a/netx_api/worker.py b/netx_api/worker.py index 3fcd6e5..f2c97dc 100644 --- a/netx_api/worker.py +++ b/netx_api/worker.py @@ -43,6 +43,10 @@ def main() -> None: while not stop.is_set(): time.sleep(1.0) + + from .app_shutdown import shutdown_runtime + + shutdown_runtime(reason="worker_signal") _log.info("netx worker exiting") diff --git a/tests/test_stability_hardening.py b/tests/test_stability_hardening.py new file mode 100644 index 0000000..bba42c3 --- /dev/null +++ b/tests/test_stability_hardening.py @@ -0,0 +1,45 @@ +"""Smoke tests for DB pool / CLI budget / metrics wiring.""" + +from __future__ import annotations + +import unittest + +from netx_api.cli_budget import acquire_cli_slot, clamp_cli_workers, cli_budget_status +from netx_api.config import settings +from netx_api.db import db_pool_status +from netx_api.metrics_router import _prom_lines, collect_runtime_metrics + + +class StabilityHardeningTests(unittest.TestCase): + def test_settings_have_pool_and_budget_knobs(self) -> None: + self.assertGreaterEqual(int(settings.db_pool_size), 1) + self.assertGreaterEqual(int(settings.cli_max_concurrent), 1) + self.assertGreaterEqual(int(settings.audit_queue_max), 100) + + def test_clamp_cli_workers(self) -> None: + n = clamp_cli_workers(999, hard_cap=32) + self.assertLessEqual(n, int(settings.cli_max_concurrent)) + self.assertLessEqual(n, 32) + self.assertGreaterEqual(n, 1) + + def test_cli_budget_acquire_release(self) -> None: + before = cli_budget_status()["in_use"] + with acquire_cli_slot() as ok: + self.assertTrue(ok) + self.assertEqual(cli_budget_status()["in_use"], before + 1) + self.assertEqual(cli_budget_status()["in_use"], before) + + def test_db_pool_status_shape(self) -> None: + st = db_pool_status() + self.assertIn("backend", st) + self.assertIn("pool_size_cfg", st) + + def test_metrics_prometheus_text(self) -> None: + body = _prom_lines(collect_runtime_metrics()) + self.assertIn("netx_uptime_seconds", body) + self.assertIn("netx_thread_count", body) + self.assertIn("netx_cli_budget_limit", body) + + +if __name__ == "__main__": + unittest.main()