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 <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-08-02 18:02:43 +08:00
parent 2068ce3d04
commit d56020d84c
23 changed files with 824 additions and 102 deletions

88
netx_api/app_shutdown.py Normal file
View file

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

View file

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

View file

@ -24,6 +24,8 @@ _PUBLIC_EXACT = frozenset(
"/health",
"/health/live",
"/health/ready",
"/metrics",
"/metrics/json",
"/favicon.ico",
"/v1/auth/login",
}

68
netx_api/cli_budget.py Normal file
View file

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

78
netx_api/cli_timeout.py Normal file
View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

113
netx_api/metrics_router.py Normal file
View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 = []

View file

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