From 3a0de0d7fb4e361309e161339faefa01a9598138 Mon Sep 17 00:00:00 2001 From: oliver Date: Tue, 22 Sep 2026 17:00:37 +0800 Subject: [PATCH] Scale biz_state collects with dedicated workers and non-blocking UI poll. Add PG claim/NE mutex, persist pool, and biz_state_worker replicas; fix double-SSH and row-count bugs; stop 32m collectNow while-loop from freezing page switches. Co-authored-by: Cursor --- .env.example | 11 +- netx_api/biz_state/claim.py | 321 +++++++++++++++++++++++ netx_api/biz_state/collect_recovery.py | 6 +- netx_api/biz_state/collect_runner.py | 222 ++++++++++------ netx_api/biz_state/persist_pool.py | 126 +++++++++ netx_api/biz_state_router.py | 7 +- netx_api/biz_state_scheduler.py | 154 +++++++++-- netx_api/biz_state_worker.py | 64 +++++ netx_api/config.py | 8 + netx_api/models/biz_state.py | 4 +- netx_api/schema_patches.py | 1 + netx_api/ume_runtime.py | 11 +- netx_api/worker.py | 2 + scripts/start_netx.ps1 | 38 +++ scripts/start_netx.sh | 15 ++ scripts/stop_netx.ps1 | 25 +- scripts/stop_netx.sh | 10 + tests/test_biz_state_claim.py | 125 +++++++++ tests/test_biz_state_collect_recovery.py | 15 ++ tests/test_biz_state_dual_lane_absorb.py | 51 ++++ tests/test_collect_dedupe.py | 53 ++-- web/src/pages/network/BizStatePage.tsx | 106 +++++--- 22 files changed, 1210 insertions(+), 165 deletions(-) create mode 100644 netx_api/biz_state/claim.py create mode 100644 netx_api/biz_state/persist_pool.py create mode 100644 netx_api/biz_state_worker.py create mode 100644 tests/test_biz_state_claim.py create mode 100644 tests/test_biz_state_dual_lane_absorb.py diff --git a/.env.example b/.env.example index c9083fd..98db0cd 100644 --- a/.env.example +++ b/.env.example @@ -66,6 +66,15 @@ NETX_UME_NOTIFICATION_TOPIC=ALARM # (start_netx.ps1/.sh do this automatically). Worker writes heartbeat for API /metrics. # NETX_RUN_INLINE_SCHEDULERS=false # NETX_SCHEDULER_HEARTBEAT_PATH=data/runtime/scheduler_heartbeat.json +# Biz-state dedicated workers (default on with split start): claim queued batches via PG. +# Rule of thumb: NETX_BIZ_STATE_MAX_CONCURRENT_TASKS * 2 <= NETX_CLI_MAX_CONCURRENT +# NETX_BIZ_STATE_DEDICATED_WORKERS=true +# NETX_BIZ_STATE_MAX_CONCURRENT_TASKS=16 +# NETX_BIZ_STATE_WORKER_COLLECT_THREADS=8 +# NETX_BIZ_STATE_PERSIST_WORKERS=4 +# NETX_BIZ_STATE_WORKER_REPLICAS=2 +# NETX_BIZ_STATE_SPOOL_DIR=data/biz_state_spool +# NETX_BIZ_STATE_PERSIST_EVERY_CMDS=8 # --- Multi-user shared-server capacity (defaults in Settings already match these) --- # NETX_DB_POOL_SIZE=40 # NETX_DB_MAX_OVERFLOW=40 @@ -94,6 +103,6 @@ NETX_UME_NOTIFICATION_TOPIC=ALARM # NETX_NE_COLLECTION_KEEP_DAYS=14 # Heavier fleets: raise CLI/DB together; also ensure Postgres max_connections and bastion session limits. # One-click start (recommended): .\scripts\start_netx.ps1 -Background -WithWeb -# -> API with NETX_RUN_INLINE_SCHEDULERS=false + auto-started netx_api.worker + optional Vite +# -> API + netx_api.worker + N× netx_api.biz_state_worker + optional Vite # Legacy single-process: add -InlineSchedulers # Stop all: .\scripts\stop_netx.ps1 diff --git a/netx_api/biz_state/claim.py b/netx_api/biz_state/claim.py new file mode 100644 index 0000000..7eaab36 --- /dev/null +++ b/netx_api/biz_state/claim.py @@ -0,0 +1,321 @@ +"""Enqueue / claim biz_state collect batches (Postgres SKIP LOCKED + NE mutex).""" + +from __future__ import annotations + +import logging +from datetime import datetime +from typing import Any +from uuid import uuid4 + +from sqlalchemy.orm import Session + +from ..config import settings +from ..db import SessionLocal +from ..models import BizStateBatch, BizStateTask, BizStateTaskItem +from ..timeutil import utcnow_naive + +_log = logging.getLogger("netx.biz_state.claim") + + +def _utcnow() -> datetime: + return utcnow_naive() + + +def max_concurrent_tasks() -> int: + return max(1, int(getattr(settings, "biz_state_max_concurrent_tasks", 16) or 16)) + + +def enqueue_collect(task_id: str, *, manual: bool = False) -> dict[str, Any]: + """Create a queued batch and mark task collect_running. Does not open SSH. + + Returns keys: ok, queued, reason?, batch_id?, task_id + """ + tid = str(task_id or "").strip() + if not tid: + return {"ok": False, "queued": False, "reason": "missing_task_id", "task_id": ""} + + db = SessionLocal() + try: + task = db.get(BizStateTask, tid) + if not task: + return {"ok": False, "queued": False, "reason": "task_not_found", "task_id": tid} + if bool(task.collect_running): + return { + "ok": True, + "queued": False, + "reason": "already_collecting", + "task_id": tid, + } + + st = str(task.status or "").strip() + if manual: + if st in ("", "deleted"): + return {"ok": False, "queued": False, "reason": "bad_status", "task_id": tid} + else: + if st != "running": + return {"ok": False, "queued": False, "reason": "not_scheduled", "task_id": tid} + + items = ( + db.query(BizStateTaskItem) + .filter( + BizStateTaskItem.task_id == tid, + BizStateTaskItem.enabled.is_(True), + ) + .limit(1) + .all() + ) + if not items: + task.last_error = "no enabled task items" + task.updated_at = _utcnow() + db.commit() + return { + "ok": False, + "queued": False, + "reason": "no_enabled_items", + "task_id": tid, + } + + now = _utcnow() + task.collect_running = True + task.last_collect_started_at = now + task.last_error = "" + task.updated_at = now + if hasattr(task, "collect_queued_at"): + task.collect_queued_at = now + + batch = BizStateBatch( + id=uuid4().hex, + task_id=task.id, + source=task.source, + ne_id=task.ne_id, + ne_name=task.ne_name, + vendor=task.vendor, + status="queued", + started_at=now, + message="queued" if not manual else "queued_manual", + ) + db.add(batch) + db.commit() + return { + "ok": True, + "queued": True, + "batch_id": batch.id, + "task_id": tid, + "manual": bool(manual), + } + except Exception: + _log.exception("enqueue_collect failed task=%s", tid) + try: + db.rollback() + except Exception: + pass + return {"ok": False, "queued": False, "reason": "enqueue_error", "task_id": tid} + finally: + db.close() + + +def _dialect_supports_skip_locked(db: Session) -> bool: + try: + bind = db.get_bind() + name = str(getattr(getattr(bind, "dialect", None), "name", "") or "").lower() + return name == "postgresql" + except Exception: + return False + + +def count_running_batches(db: Session | None = None) -> int: + own = db is None + if own: + db = SessionLocal() + assert db is not None + try: + return int( + db.query(BizStateBatch).filter(BizStateBatch.status == "running").count() + ) + finally: + if own: + db.close() + + +def claim_queued_batches(limit: int | None = None) -> list[dict[str, Any]]: + """Claim up to ``limit`` queued batches (SKIP LOCKED + same-NE mutex). + + Returns list of {batch_id, task_id, source, ne_id, vendor, device_type}. + Global ceiling is always ``max_concurrent_tasks()``; ``limit`` only caps + how many this caller wants in one call. + """ + from sqlalchemy import text + + global_cap = max_concurrent_tasks() + db = SessionLocal() + try: + running_n = ( + db.query(BizStateBatch).filter(BizStateBatch.status == "running").count() + ) + slots = max(0, global_cap - int(running_n)) + if limit is not None: + slots = min(slots, max(0, int(limit))) + if slots <= 0: + return [] + + busy_nes = { + str(r[0] or "").strip() + for r in db.query(BizStateBatch.ne_id) + .filter(BizStateBatch.status == "running") + .all() + if str(r[0] or "").strip() + } + + q = ( + db.query(BizStateBatch) + .filter(BizStateBatch.status == "queued") + .order_by(BizStateBatch.started_at.asc()) + ) + pg = _dialect_supports_skip_locked(db) + if pg: + q = q.with_for_update(skip_locked=True) + else: + q = q.with_for_update() + + candidates = q.limit(max(slots * 4, slots)).all() + claimed_rows: list[BizStateBatch] = [] + for batch in candidates: + if len(claimed_rows) >= slots: + break + ne = str(batch.ne_id or "").strip() + if ne and ne in busy_nes: + continue + # Serialize claims per NE across workers (Postgres). + if ne and pg: + try: + db.execute( + text("SELECT pg_advisory_xact_lock(hashtext(:ne))"), + {"ne": ne}, + ) + except Exception: + _log.exception("advisory lock failed ne=%s", ne) + conflict = ( + db.query(BizStateBatch.id) + .filter( + BizStateBatch.status == "running", + BizStateBatch.ne_id == ne, + BizStateBatch.id != batch.id, + ) + .first() + ) + if conflict: + continue + batch.status = "running" + batch.message = "" + claimed_rows.append(batch) + if ne: + busy_nes.add(ne) + + if not claimed_rows: + db.rollback() + return [] + + # Autoflush so COUNT sees our running marks; trim if over global cap. + try: + db.flush() + except Exception: + pass + n_running = int( + db.query(BizStateBatch).filter(BizStateBatch.status == "running").count() + ) + while n_running > global_cap and claimed_rows: + demote = claimed_rows.pop() + demote.status = "queued" + demote.message = "queued" + n_running -= 1 + + if not claimed_rows: + db.rollback() + return [] + + out: list[dict[str, Any]] = [] + for batch in claimed_rows: + task = db.get(BizStateTask, batch.task_id) + out.append( + { + "batch_id": batch.id, + "task_id": str(batch.task_id or ""), + "source": str( + batch.source or (task.source if task else "") or "managed" + ), + "ne_id": str(batch.ne_id or (task.ne_id if task else "") or ""), + "vendor": str(batch.vendor or (task.vendor if task else "") or ""), + "device_type": str((task.device_type if task else "") or ""), + } + ) + db.commit() + return out + except Exception: + _log.exception("claim_queued_batches failed") + try: + db.rollback() + except Exception: + pass + return [] + finally: + db.close() + + +def reclaim_stale_queued(*, max_age_sec: int = 3600) -> int: + """Fail queued batches older than max_age_sec (no worker drained them).""" + from datetime import timedelta + + age = max(60, int(max_age_sec)) + cutoff = _utcnow() - timedelta(seconds=age) + db = SessionLocal() + n = 0 + try: + stale = ( + db.query(BizStateBatch) + .filter( + BizStateBatch.status == "queued", + BizStateBatch.started_at < cutoff, + ) + .all() + ) + for b in stale: + b.status = "failed" + b.message = f"stale_queued_timeout ({age}s)"[:1020] + b.ended_at = _utcnow() + task = db.get(BizStateTask, b.task_id) + if task and bool(task.collect_running): + # Only clear if no other running/queued for this task. + other = ( + db.query(BizStateBatch) + .filter( + BizStateBatch.task_id == b.task_id, + BizStateBatch.id != b.id, + BizStateBatch.status.in_(("queued", "running")), + ) + .first() + ) + if not other: + task.collect_running = False + if hasattr(task, "collect_queued_at"): + task.collect_queued_at = None + task.last_collect_ended_at = _utcnow() + task.last_error = str(b.message)[:1020] + n += 1 + if n: + db.commit() + _log.warning("reclaimed stale queued batches n=%s age_sec=%s", n, age) + return n + except Exception: + _log.exception("reclaim_stale_queued failed") + try: + db.rollback() + except Exception: + pass + return 0 + finally: + db.close() + + +def open_slots() -> int: + """How many new running batches we can still start under the ceiling.""" + return max(0, max_concurrent_tasks() - count_running_batches()) diff --git a/netx_api/biz_state/collect_recovery.py b/netx_api/biz_state/collect_recovery.py index c2e6add..b935c04 100644 --- a/netx_api/biz_state/collect_recovery.py +++ b/netx_api/biz_state/collect_recovery.py @@ -27,7 +27,9 @@ def recover_interrupted_collects_on_startup(db: Session) -> dict[str, Any]: task_n = 0 stuck_batches = ( - db.query(BizStateBatch).filter(BizStateBatch.status == "running").all() + db.query(BizStateBatch) + .filter(BizStateBatch.status.in_(("running", "queued"))) + .all() ) for b in stuck_batches: b.status = "partial" @@ -58,6 +60,8 @@ def recover_interrupted_collects_on_startup(db: Session) -> dict[str, Any]: ) for t in stuck_tasks: t.collect_running = False + if hasattr(t, "collect_queued_at"): + t.collect_queued_at = None if not t.last_collect_ended_at: t.last_collect_ended_at = now err = str(t.last_error or "").strip() diff --git a/netx_api/biz_state/collect_runner.py b/netx_api/biz_state/collect_runner.py index 8a361a5..a0605ff 100644 --- a/netx_api/biz_state/collect_runner.py +++ b/netx_api/biz_state/collect_runner.py @@ -346,6 +346,8 @@ def _finish_task(task_id: str, *, error: str = "") -> None: if not task: return task.collect_running = False + if hasattr(task, "collect_queued_at"): + task.collect_queued_at = None task.last_collect_ended_at = _utcnow() task.last_error = str(error or "")[:1020] task.updated_at = _utcnow() @@ -357,70 +359,101 @@ def _finish_task(task_id: str, *, error: str = "") -> None: def dispatch_collect(task_id: str, *, manual: bool = False) -> None: - """Claim and run one collect round. + """Enqueue a collect round; run inline when this process owns execution. - Scheduler calls with ``manual=False`` (only when task status is ``running``). - Collect-now calls with ``manual=True`` (any status, as long as not already collecting). + Dedicated worker mode: only enqueue (claim loop runs the batch). + Inline / non-dedicated: enqueue then atomically promote+execute (skip if + another worker already claimed the batch). """ + from .claim import enqueue_collect + + result = enqueue_collect(task_id, manual=manual) + if not result.get("queued"): + return + batch_id = str(result.get("batch_id") or "") + if not batch_id: + return + if _should_execute_inline(): + # Atomic queued→running; if false, scheduler/worker already owns it. + if not _try_claim_batch_for_execute(batch_id): + return + execute_claimed_batch( + batch_id=batch_id, + task_id=str(result.get("task_id") or task_id), + source="", + ne_id="", + vendor="", + device_type="", + ) + + +def _should_execute_inline() -> bool: + """True when this process should run SSH after enqueue (not dedicated workers).""" + if bool(getattr(settings, "run_inline_schedulers", True)): + return True + if not bool(getattr(settings, "biz_state_dedicated_workers", True)): + return True + return False + + +def _try_claim_batch_for_execute(batch_id: str) -> bool: + """Promote queued→running only if still queued. Returns True iff we won the claim.""" db = SessionLocal() - batch_id = "" try: - task = db.get(BizStateTask, task_id) - if not task: - return - if task.collect_running: - return - st = str(task.status or "").strip() - if manual: - # Idle manual trigger: allow scheduled / paused / draft / stopped - if st in ("", "deleted"): - return - else: - if st != "running": - return - - items = ( - db.query(BizStateTaskItem) - .filter( - BizStateTaskItem.task_id == task_id, - BizStateTaskItem.enabled.is_(True), - ) - .order_by(BizStateTaskItem.sort_order.asc()) - .all() + batch = ( + db.query(BizStateBatch) + .filter(BizStateBatch.id == batch_id) + .with_for_update() + .one_or_none() ) - if not items: - task.last_error = "no enabled task items" - task.updated_at = _utcnow() - db.commit() - return - - task.collect_running = True - task.last_collect_started_at = _utcnow() - task.last_error = "" - task.updated_at = _utcnow() - - batch = BizStateBatch( - id=uuid4().hex, - task_id=task.id, - source=task.source, - ne_id=task.ne_id, - ne_name=task.ne_name, - vendor=task.vendor, - status="running", - started_at=_utcnow(), - ) - db.add(batch) + if not batch or str(batch.status or "") != "queued": + return False + batch.status = "running" + batch.message = "" db.commit() - batch_id = batch.id - vendor = str(task.vendor or "") - device_type = str(task.device_type or "") - source = str(task.source or "managed").strip().lower() - ne_id = str(task.ne_id or "").strip() + return True + except Exception: + _log.exception("biz_state claim-for-execute failed batch=%s", batch_id) + try: + db.rollback() + except Exception: + pass + return False finally: db.close() - if not batch_id: - return + +def _promote_queued_batch(batch_id: str) -> bool: + """Backward-compatible alias for atomic claim. """ + return _try_claim_batch_for_execute(batch_id) + +def execute_claimed_batch( + *, + batch_id: str, + task_id: str, + source: str = "", + ne_id: str = "", + vendor: str = "", + device_type: str = "", +) -> None: + """Run CLI+persist for an already-claimed (status=running) batch.""" + db = SessionLocal() + try: + batch = db.get(BizStateBatch, batch_id) + task = db.get(BizStateTask, task_id) if task_id else None + if not batch: + return + if not task_id: + task_id = str(batch.task_id or "") + task = db.get(BizStateTask, task_id) if task_id else None + source = str(source or batch.source or (task.source if task else "") or "managed") + ne_id = str(ne_id or batch.ne_id or (task.ne_id if task else "") or "") + vendor = str(vendor or batch.vendor or (task.vendor if task else "") or "") + device_type = str(device_type or (task.device_type if task else "") or "") + if not task_id: + return + finally: + db.close() error = "" try: @@ -433,12 +466,14 @@ def dispatch_collect(task_id: str, *, manual: bool = False) -> None: device_type=device_type, ) except Exception as exc: - _log.exception("biz_state collect failed task=%s", task_id) + _log.exception("biz_state collect failed task=%s batch=%s", task_id, batch_id) error = _format_error(exc) try: _fail_batch_status(batch_id, error) except Exception: - _log.exception("biz_state fail-batch after collect error failed batch=%s", batch_id) + _log.exception( + "biz_state fail-batch after collect error failed batch=%s", batch_id + ) finally: _finish_task(task_id, error=error) @@ -486,9 +521,20 @@ def _run_collect_lane( any_ok = False pending: list[SpooledCommand] = [] task_id = "" + from .persist_pool import get_persist_pool + + persist = get_persist_pool() + + def _submit_pending() -> None: + nonlocal pending + if not pending: + return + chunk = list(pending) + pending = [] + persist.submit(batch_id, chunk) def _queue(item: SpooledCommand, *, records: list[dict[str, Any]] | None = None) -> None: - nonlocal cmd_count, total_rows + nonlocal cmd_count if records is not None and item.persist_kind: item.records_rel_path = write_records(batch_id, item.id, records) item.row_count = len(records) @@ -499,16 +545,7 @@ def _run_collect_lane( pending.append(item) cmd_count += 1 if len(pending) >= flush_every: - try: - _c, _r = _flush_spooled_commands(batch_id, pending) - total_rows += int(_r or 0) - except Exception: - # Keep collecting to spool; retry flush at lane end. - _log.exception( - "biz_state mid-lane flush failed batch=%s pending=%s", - batch_id, - len(pending), - ) + _submit_pending() try: try: @@ -853,22 +890,26 @@ def _run_collect_lane( primary.message = f"parse: {_format_error(exc)}" _queue(primary) - # Final flush for this lane. - if pending: - _c, _r = _flush_spooled_commands(batch_id, pending) - total_rows += int(_r or 0) - + # Final flush for this lane → persist pool; wait so rows land before return. + _submit_pending() + if not persist.wait_idle(timeout=max(30.0, float(budget))): + _log.warning( + "biz_state persist barrier timed out batch=%s lane=%s", + batch_id, + label, + ) + # Do NOT read batch.row_count here — dual light+heavy lanes would each + # see the cumulative DB total and _absorb would double-count. return total_rows, cmd_count, any_fail, any_ok finally: - # Best-effort: persist whatever was collected before timeout/abort. - if pending: - try: - _c, _r = _flush_spooled_commands(batch_id, pending) - total_rows += int(_r or 0) - except Exception: - _log.exception( - "biz_state flush on lane exit failed batch=%s", batch_id - ) + # Best-effort: enqueue leftover spool before connection teardown. + try: + _submit_pending() + persist.wait_idle(timeout=60.0) + except Exception: + _log.exception( + "biz_state persist drain on lane exit failed batch=%s", batch_id + ) holder.pop("conn", None) close_netmiko_connection(conn) @@ -1097,7 +1138,7 @@ def _run_collect_session( task = db.get(BizStateTask, task_id) batch = db.get(BizStateBatch, batch_id) if not task or not batch: - return + raise RuntimeError(f"batch_or_task_missing batch={batch_id} task={task_id}") try: if source == "managed": @@ -1302,6 +1343,21 @@ def _run_collect_session( raise RuntimeError("; ".join(lane_errors)[:1020]) any_fail = True + # Ensure persist pool drained before terminal status write. + try: + from .persist_pool import get_persist_pool + + if not get_persist_pool().wait_idle(timeout=120.0): + _log.warning( + "biz_state persist barrier before finalize timed out batch=%s", + batch_id, + ) + any_fail = True + if "persist_barrier_timeout" not in lane_errors: + lane_errors.append("RuntimeError: persist_barrier_timeout") + except Exception: + _log.exception("biz_state persist barrier before finalize failed batch=%s", batch_id) + _finalize_batch_status( batch_id=batch_id, task_id=task_id, diff --git a/netx_api/biz_state/persist_pool.py b/netx_api/biz_state/persist_pool.py new file mode 100644 index 0000000..69d2df6 --- /dev/null +++ b/netx_api/biz_state/persist_pool.py @@ -0,0 +1,126 @@ +"""Persist pool: flush spooled commands off the collect/SSH threads.""" + +from __future__ import annotations + +import logging +import queue +import threading +import time +from typing import Any + +from ..config import settings + +_log = logging.getLogger("netx.biz_state.persist") + +_SENTINEL = object() + + +class PersistPool: + """Background workers that call ``_flush_spooled_commands``.""" + + def __init__(self, *, workers: int | None = None) -> None: + n = max( + 1, + int( + workers + if workers is not None + else (getattr(settings, "biz_state_persist_workers", 4) or 4) + ), + ) + self._q: queue.Queue[Any] = queue.Queue() + self._inflight = 0 + self._lock = threading.Lock() + self._cv = threading.Condition(self._lock) + self._workers: list[threading.Thread] = [] + self._stopped = False + for i in range(n): + t = threading.Thread( + target=self._loop, + name=f"biz-persist-{i}", + daemon=True, + ) + t.start() + self._workers.append(t) + + def submit(self, batch_id: str, items: list[Any]) -> None: + if not items or self._stopped: + return + payload = (str(batch_id), list(items)) + with self._cv: + self._inflight += 1 + self._q.put(payload) + + def wait_idle(self, *, timeout: float | None = None) -> bool: + """Block until queue empty and no in-flight flush. Returns False on timeout.""" + end = None + if timeout is not None: + end = time.monotonic() + max(0.0, float(timeout)) + with self._cv: + while self._inflight > 0 or not self._q.empty(): + remaining = None + if end is not None: + remaining = end - time.monotonic() + if remaining <= 0: + return False + self._cv.wait(timeout=remaining) + return True + + def shutdown(self, *, wait: bool = True) -> None: + self._stopped = True + for _ in self._workers: + self._q.put(_SENTINEL) + if wait: + for t in self._workers: + t.join(timeout=5.0) + + def _loop(self) -> None: + from .collect_runner import _flush_spooled_commands + + while True: + job = self._q.get() + if job is _SENTINEL: + self._q.task_done() + break + batch_id, items = job + try: + pending = list(items) + _flush_spooled_commands(batch_id, pending) + if pending: + _log.warning( + "biz_state persist retry leftover=%s batch=%s", + len(pending), + batch_id, + ) + try: + _flush_spooled_commands(batch_id, pending) + except Exception: + _log.exception( + "biz_state persist retry failed batch=%s", batch_id + ) + except Exception: + _log.exception("biz_state persist flush failed batch=%s", batch_id) + finally: + with self._cv: + self._inflight = max(0, self._inflight - 1) + self._cv.notify_all() + self._q.task_done() + + +_pool_lock = threading.Lock() +_pool: PersistPool | None = None + + +def get_persist_pool() -> PersistPool: + global _pool + with _pool_lock: + if _pool is None: + _pool = PersistPool() + return _pool + + +def shutdown_persist_pool(*, wait: bool = False) -> None: + global _pool + with _pool_lock: + if _pool is not None: + _pool.shutdown(wait=wait) + _pool = None diff --git a/netx_api/biz_state_router.py b/netx_api/biz_state_router.py index 7c96164..83dabdf 100644 --- a/netx_api/biz_state_router.py +++ b/netx_api/biz_state_router.py @@ -245,12 +245,11 @@ def api_collect_now( raise HTTPException(status_code=404, detail="task not found") if bool(task.collect_running): return {"ok": True, "started": False, "reason": "already_collecting", "task_id": task_id} - # Allow one-shot even when schedule is enabled (idle only). - task.last_collect_ended_at = None - db.commit() tid = task_id + # Enqueue (and run inline only when this process owns collectors). + # Dedicated worker mode: BackgroundTasks only creates queued batch; workers claim. background_tasks.add_task(lambda: dispatch_collect(tid, manual=True)) - return {"ok": True, "started": True, "task_id": task_id} + return {"ok": True, "started": True, "queued": True, "task_id": task_id} @router.get("/tasks/{task_id}/batches") diff --git a/netx_api/biz_state_scheduler.py b/netx_api/biz_state_scheduler.py index 2bed6e9..c6b13ca 100644 --- a/netx_api/biz_state_scheduler.py +++ b/netx_api/biz_state_scheduler.py @@ -1,4 +1,4 @@ -"""Background scheduler for biz_state collection.""" +"""Background scheduler for biz_state collection (enqueue + claim loop).""" from __future__ import annotations @@ -6,12 +6,14 @@ import logging import threading from concurrent.futures import ThreadPoolExecutor from datetime import datetime +from typing import Any from .cli_budget import clamp_cli_workers from .config import settings from .db import SessionLocal from .models import BizStateTask -from .biz_state.collect_runner import dispatch_collect +from .biz_state.claim import claim_queued_batches, enqueue_collect, max_concurrent_tasks, open_slots +from .biz_state.collect_runner import execute_claimed_batch _log = logging.getLogger("netx.biz_state.scheduler") _stop = threading.Event() @@ -20,6 +22,8 @@ _dispatch_pool: ThreadPoolExecutor | None = None _pool_lock = threading.Lock() _last_tick_mono: float = 0.0 _last_purge_mono: float = 0.0 +_in_flight: set[str] = set() +_in_flight_lock = threading.Lock() _PURGE_INTERVAL_SEC = 3600.0 @@ -57,11 +61,16 @@ def _dispatch_pool_get() -> ThreadPoolExecutor: global _dispatch_pool with _pool_lock: if _dispatch_pool is None: - workers = clamp_cli_workers( - int(getattr(settings, "biz_state_dispatch_workers", 2) or 2), + n = clamp_cli_workers( + int( + getattr(settings, "biz_state_worker_collect_threads", None) + or getattr(settings, "biz_state_dispatch_workers", 8) + or 8 + ), ) + n = max(1, min(n, max_concurrent_tasks())) _dispatch_pool = ThreadPoolExecutor( - max_workers=workers, thread_name_prefix="biz-dispatch" + max_workers=n, thread_name_prefix="biz-collect" ) return _dispatch_pool @@ -75,6 +84,12 @@ def shutdown_biz_state_dispatch_pool(*, wait: bool = False) -> None: except TypeError: _dispatch_pool.shutdown(wait=wait) _dispatch_pool = None + try: + from .biz_state.persist_pool import shutdown_persist_pool + + shutdown_persist_pool(wait=wait) + except Exception: + _log.exception("shutdown persist pool failed") def _sync_cutover_hf_windows() -> None: @@ -122,17 +137,7 @@ def _sync_cutover_hf_windows() -> None: db.close() -def try_dispatch_due_tasks() -> int: - import time as _time - - global _last_tick_mono - _last_tick_mono = _time.monotonic() - - try: - _sync_cutover_hf_windows() - except Exception: - _log.exception("hf window sync tick failed") - +def _enqueue_due_tasks() -> int: db = SessionLocal() try: tasks = ( @@ -156,15 +161,108 @@ def try_dispatch_due_tasks() -> int: finally: db.close() - if not due_ids: - return 0 - pool = _dispatch_pool_get() + n = 0 for tid in due_ids: try: - pool.submit(dispatch_collect, tid) + r = enqueue_collect(tid, manual=False) + if r.get("queued"): + n += 1 except Exception: - _log.exception("biz_state submit failed task=%s", tid) - return len(due_ids) + _log.exception("biz_state enqueue failed task=%s", tid) + return n + + +def _run_claimed(job: dict[str, Any]) -> None: + bid = str(job.get("batch_id") or "") + try: + execute_claimed_batch( + batch_id=bid, + task_id=str(job.get("task_id") or ""), + source=str(job.get("source") or ""), + ne_id=str(job.get("ne_id") or ""), + vendor=str(job.get("vendor") or ""), + device_type=str(job.get("device_type") or ""), + ) + finally: + with _in_flight_lock: + _in_flight.discard(bid) + + +def _claim_and_dispatch() -> int: + slots = open_slots() + with _in_flight_lock: + local_busy = len(_in_flight) + # Don't over-submit beyond local pool either. + pool = _dispatch_pool_get() + local_cap = getattr(pool, "_max_workers", 8) or 8 + want = min(slots, max(0, int(local_cap) - local_busy)) + if want <= 0: + return 0 + jobs = claim_queued_batches(want) + if not jobs: + return 0 + submitted = 0 + for job in jobs: + bid = str(job.get("batch_id") or "") + if not bid: + continue + with _in_flight_lock: + if bid in _in_flight: + continue + _in_flight.add(bid) + try: + pool.submit(_run_claimed, job) + submitted += 1 + except Exception: + with _in_flight_lock: + _in_flight.discard(bid) + _log.exception("biz_state submit claimed batch failed batch=%s", bid) + # Compensate: claimed batch must not stay running forever. + try: + from .biz_state.collect_runner import _fail_batch_status, _finish_task + + _fail_batch_status(bid, "RuntimeError: submit_claimed_batch_failed") + _finish_task( + str(job.get("task_id") or ""), + error="RuntimeError: submit_claimed_batch_failed", + ) + except Exception: + _log.exception("biz_state compensate after submit fail batch=%s", bid) + return submitted + + +def try_dispatch_due_tasks() -> int: + """Enqueue due tasks, then claim+run up to concurrency ceiling.""" + import time as _time + + global _last_tick_mono + _last_tick_mono = _time.monotonic() + + try: + _sync_cutover_hf_windows() + except Exception: + _log.exception("hf window sync tick failed") + + enq = 0 + try: + enq = _enqueue_due_tasks() + except Exception: + _log.exception("biz_state enqueue tick failed") + + try: + from .biz_state.claim import reclaim_stale_queued + + reclaim_stale_queued(max_age_sec=3600) + except Exception: + _log.exception("biz_state stale queued reclaim failed") + + claimed = 0 + try: + claimed = _claim_and_dispatch() + except Exception: + _log.exception("biz_state claim tick failed") + + return enq + claimed def _loop() -> None: @@ -192,7 +290,10 @@ def start_biz_state_scheduler() -> None: _stop.clear() _thread = threading.Thread(target=_loop, name="biz-state-scheduler", daemon=True) _thread.start() - _log.info("biz_state scheduler started") + _log.info( + "biz_state scheduler started max_concurrent=%s", + max_concurrent_tasks(), + ) def stop_biz_state_scheduler() -> None: @@ -213,8 +314,15 @@ def biz_state_scheduler_status() -> dict: age = None if _last_tick_mono: age = max(0.0, _time.monotonic() - _last_tick_mono) + with _in_flight_lock: + inflight = len(_in_flight) return { "running": alive, "enabled": bool(getattr(settings, "biz_state_scheduler_enabled", True)), "last_tick_age_sec": age, + "in_flight_batches": inflight, + "max_concurrent_tasks": max_concurrent_tasks(), + "dedicated_workers": bool( + getattr(settings, "biz_state_dedicated_workers", True) + ), } diff --git a/netx_api/biz_state_worker.py b/netx_api/biz_state_worker.py new file mode 100644 index 0000000..4172264 --- /dev/null +++ b/netx_api/biz_state_worker.py @@ -0,0 +1,64 @@ +"""Dedicated biz_state collect worker process. + +When ``NETX_BIZ_STATE_DEDICATED_WORKERS=true`` (default) and API uses external +schedulers (``NETX_RUN_INLINE_SCHEDULERS=false``), run one or more replicas: + + python -m netx_api.biz_state_worker + +Each process enqueues due tasks, claims queued batches (SKIP LOCKED + NE mutex), +runs SSH collect → spool, and drains the persist pool. Scale by starting N +processes on the same host (start_netx.ps1 / .sh do this via replicas). +""" + +from __future__ import annotations + +import logging +import signal +import time + +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(levelname)s %(name)s: %(message)s", +) +_log = logging.getLogger("netx.biz_state.worker") + + +def main() -> None: + from .biz_state_scheduler import start_biz_state_scheduler, stop_biz_state_scheduler + from .config import settings + from .scheduler_heartbeat import start_scheduler_heartbeat_publisher + + stop = False + + def _handle(_sig: int, _frame: object) -> None: + nonlocal stop + _log.info("shutdown signal received") + stop = True + + for sig in (signal.SIGINT, signal.SIGTERM): + try: + signal.signal(sig, _handle) + except Exception: + pass + + start_biz_state_scheduler() + try: + start_scheduler_heartbeat_publisher(role="biz_state_worker") + except Exception: + _log.exception("biz_state worker heartbeat failed to start") + + _log.info( + "biz_state worker started dedicated=%s max_concurrent=%s", + bool(getattr(settings, "biz_state_dedicated_workers", True)), + int(getattr(settings, "biz_state_max_concurrent_tasks", 16) or 16), + ) + + while not stop: + time.sleep(1.0) + + stop_biz_state_scheduler() + _log.info("biz_state worker exiting") + + +if __name__ == "__main__": + main() diff --git a/netx_api/config.py b/netx_api/config.py index b9bb114..9bfd426 100644 --- a/netx_api/config.py +++ b/netx_api/config.py @@ -127,6 +127,14 @@ class Settings(BaseSettings): biz_state_persist_every_cmds: int = 8 # Cap raw_text loaded into Postgres from spool (0 = unlimited). biz_state_raw_max_bytes: int = 8 * 1024 * 1024 + # Dedicated biz_state worker process(es); general worker skips biz_state scheduler. + biz_state_dedicated_workers: bool = True + # Global ceiling for simultaneous running batches (across all workers). + biz_state_max_concurrent_tasks: int = 16 + biz_state_worker_collect_threads: int = 8 + biz_state_persist_workers: int = 4 + # How many biz_state_worker processes start_netx should launch (same host). + biz_state_worker_replicas: int = 2 # Managed NE exec: max CLI commands per request (lab can raise; hard-capped in ne_exec). ne_exec_max_commands: int = 5 # Opt-in: allow per-NE exec_policy (linux_shell/unrestricted). Default off — UI hidden. diff --git a/netx_api/models/biz_state.py b/netx_api/models/biz_state.py index 8e8b5ef..ee6e2a4 100644 --- a/netx_api/models/biz_state.py +++ b/netx_api/models/biz_state.py @@ -42,6 +42,8 @@ class BizStateTask(Base): # Legacy column kept for brownfield reads; unused by new purge path retention_batches: Mapped[int] = mapped_column(Integer, default=30) collect_running: Mapped[bool] = mapped_column(Boolean, default=False) + # Set when a collect batch is queued (worker claim picks it up). + collect_queued_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) last_collect_started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) last_collect_ended_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) last_error: Mapped[str] = mapped_column(String(1024), default="") @@ -88,7 +90,7 @@ class BizStateBatch(Base): ne_id: Mapped[str] = mapped_column(String(128), default="", index=True) ne_name: Mapped[str] = mapped_column(String(256), default="") vendor: Mapped[str] = mapped_column(String(64), default="") - status: Mapped[str] = mapped_column(String(32), default="running", index=True) # running|success|partial|failed + status: Mapped[str] = mapped_column(String(32), default="running", index=True) # queued|running|success|partial|failed command_count: Mapped[int] = mapped_column(Integer, default=0) row_count: Mapped[int] = mapped_column(Integer, default=0) message: Mapped[str] = mapped_column(String(1024), default="") diff --git a/netx_api/schema_patches.py b/netx_api/schema_patches.py index ab4a31d..1dc4528 100644 --- a/netx_api/schema_patches.py +++ b/netx_api/schema_patches.py @@ -496,6 +496,7 @@ def apply_domain_schema_patches(conn: Connection) -> None: "ALTER TABLE ume_topo_link ADD COLUMN IF NOT EXISTS z_ifname VARCHAR(128) DEFAULT ''", "ALTER TABLE topo_fabric_node ADD COLUMN IF NOT EXISTS world_x DOUBLE PRECISION", "ALTER TABLE topo_fabric_node ADD COLUMN IF NOT EXISTS world_y DOUBLE PRECISION", + "ALTER TABLE biz_state_task ADD COLUMN IF NOT EXISTS collect_queued_at TIMESTAMP", "CREATE INDEX IF NOT EXISTS ix_topo_fabric_node_world_x ON topo_fabric_node (world_x)", "CREATE INDEX IF NOT EXISTS ix_topo_fabric_node_world_y ON topo_fabric_node (world_y)", "COMMENT ON TABLE ume_inventory_ne IS '网元对象详细信息'", diff --git a/netx_api/ume_runtime.py b/netx_api/ume_runtime.py index bc39676..d40fa5c 100644 --- a/netx_api/ume_runtime.py +++ b/netx_api/ume_runtime.py @@ -55,10 +55,17 @@ def start_device_schedulers() -> None: start_lldp_collect_scheduler() start_ne_collect_scheduler() start_port_traffic_scheduler() - start_biz_state_scheduler() + # When dedicated biz_state workers are on and we are not the API-inline path, + # skip biz_state here so only `python -m netx_api.biz_state_worker` owns it. + dedicated = bool(getattr(settings, "biz_state_dedicated_workers", True)) + inline = bool(getattr(settings, "run_inline_schedulers", True)) + if dedicated and not inline: + _log.info("biz_state scheduler skipped (dedicated workers)") + else: + start_biz_state_scheduler() start_fabric_reconcile_scheduler() # Publish status so API /metrics can see collectors when run in a split worker. - role = "api_inline" if bool(getattr(settings, "run_inline_schedulers", True)) else "worker" + role = "api_inline" if inline else "worker" start_scheduler_heartbeat_publisher(role=role) _log.info("device schedulers started") diff --git a/netx_api/worker.py b/netx_api/worker.py index 9455c2a..8904be2 100644 --- a/netx_api/worker.py +++ b/netx_api/worker.py @@ -5,6 +5,8 @@ Optional when ``NETX_RUN_INLINE_SCHEDULERS=false`` (API does not start collector python -m netx_api.worker Starts: config_sync, lldp_collect, port_traffic, fabric_reconcile tick loops. +When ``NETX_BIZ_STATE_DEDICATED_WORKERS=true`` (default), biz_state is NOT started +here — run ``python -m netx_api.biz_state_worker`` instead (start_netx does both). UME WS / keepalive remain in the API process (token + alarm coordination). By default the API runs collectors inline — no separate worker needed. """ diff --git a/scripts/start_netx.ps1 b/scripts/start_netx.ps1 index 8e75164..ed2caec 100644 --- a/scripts/start_netx.ps1 +++ b/scripts/start_netx.ps1 @@ -34,6 +34,7 @@ $errFile = Join-Path $runDir "netx.err.log" $workerPidFile = Join-Path $runDir "worker.pid" $workerLogFile = Join-Path $runDir "worker.out.log" $workerErrFile = Join-Path $runDir "worker.err.log" +$bizStateWorkerPidDir = Join-Path $runDir "biz_state_workers" $webPidFile = Join-Path $runDir "web.pid" $webLogFile = Join-Path $runDir "web.out.log" $webErrFile = Join-Path $runDir "web.err.log" @@ -187,6 +188,41 @@ function Start-NetxWorker { } } +function Start-BizStateWorkers { + if ($InlineSchedulers) { + return + } + $replicas = 2 + if ($env:NETX_BIZ_STATE_WORKER_REPLICAS) { + try { $replicas = [int]$env:NETX_BIZ_STATE_WORKER_REPLICAS } catch { $replicas = 2 } + } + if ($replicas -lt 1) { $replicas = 1 } + if (-not (Test-Path $bizStateWorkerPidDir)) { + New-Item -ItemType Directory -Path $bizStateWorkerPidDir | Out-Null + } + Write-Host "==> Starting biz_state workers (replicas=$replicas)" + for ($i = 0; $i -lt $replicas; $i++) { + $outLog = Reset-LogFile -Path (Join-Path $bizStateWorkerPidDir "worker$i.out.log") + $errLog = Reset-LogFile -Path (Join-Path $bizStateWorkerPidDir "worker$i.err.log") + $pidPath = Join-Path $bizStateWorkerPidDir "worker$i.pid" + $proc = Start-Process -FilePath $pythonExe ` + -ArgumentList @("-m", "netx_api.biz_state_worker") ` + -WorkingDirectory $projectRoot ` + -WindowStyle Hidden ` + -RedirectStandardOutput $outLog ` + -RedirectStandardError $errLog ` + -PassThru + Set-Content -Path $pidPath -Value "$($proc.Id)" + Write-Host "biz_state_worker[$i] PID=$($proc.Id)" + Start-Sleep -Milliseconds 400 + if ($proc.HasExited) { + Write-Host "[ERR] biz_state_worker[$i] exited immediately." -ForegroundColor Red + Show-LogTail -Path $errLog + exit 1 + } + } +} + if ($Background) { # Truncate logs so a failed start is not confused with an old run. Set-Content -Path $logFile -Value "" -Encoding utf8 @@ -226,6 +262,7 @@ if ($Background) { } Write-Host "==> netx API ready: http://${BindHost}:${Port}/health" -ForegroundColor Green Start-NetxWorker + Start-BizStateWorkers if ($WithWeb) { Write-Host "==> Starting Vite dev server in background" $webRoot = Join-Path $projectRoot "web" @@ -283,5 +320,6 @@ if ($WithWeb) { } Start-NetxWorker +Start-BizStateWorkers Write-Host "==> Starting netx API in foreground" & $pythonExe -m netx_api.main diff --git a/scripts/start_netx.sh b/scripts/start_netx.sh index 9d19866..e08286b 100644 --- a/scripts/start_netx.sh +++ b/scripts/start_netx.sh @@ -239,6 +239,21 @@ if [[ "${INLINE_SCHEDULERS}" != "1" ]]; then echo "PID = ${WORKER_PID}" echo "Log = ${WORKER_LOG_FILE}" echo "Err = ${WORKER_ERR_FILE}" + + REPLICAS="${NETX_BIZ_STATE_WORKER_REPLICAS:-2}" + if [[ "${REPLICAS}" -lt 1 ]]; then REPLICAS=1; fi + BIZ_DIR="${RUN_DIR}/biz_state_workers" + mkdir -p "${BIZ_DIR}" + echo "==> Starting biz_state workers (replicas=${REPLICAS})" + for ((i=0; i"${BIZ_DIR}/worker${i}.out.log" 2>"${BIZ_DIR}/worker${i}.err.log" & + echo $! > "${BIZ_DIR}/worker${i}.pid" + ) + echo "biz_state_worker[${i}] PID=$(cat "${BIZ_DIR}/worker${i}.pid")" + done fi if [[ "${API_ONLY}" != "1" ]]; then diff --git a/scripts/stop_netx.ps1 b/scripts/stop_netx.ps1 index 270bb2e..2d91466 100644 --- a/scripts/stop_netx.ps1 +++ b/scripts/stop_netx.ps1 @@ -13,6 +13,7 @@ $runDir = Join-Path $PSScriptRoot ".run" $pidFile = Join-Path $runDir "netx.pid" $workerPidFile = Join-Path $runDir "worker.pid" $webPidFile = Join-Path $runDir "web.pid" +$bizStateWorkerPidDir = Join-Path $runDir "biz_state_workers" function Stop-OnePid { param([int]$ProcId, [string]$Label) @@ -62,13 +63,19 @@ function Get-ListenPids { function Stop-NetxByCommandLine { # Orphan workers often have no worker.pid but still hold worker.out.log. $hits = @(Get-CimInstance Win32_Process -ErrorAction SilentlyContinue | - Where-Object { $_.CommandLine -match 'netx_api\.(main|worker)' }) + Where-Object { $_.CommandLine -match 'netx_api\.(main|worker|biz_state_worker)' }) if ($hits.Count -eq 0) { - Write-Host "[INFO] No netx_api.main/worker process by command line" + Write-Host "[INFO] No netx_api.main/worker/biz_state_worker process by command line" return } foreach ($p in $hits) { - $kind = if ($p.CommandLine -match 'netx_api\.worker') { "worker(cmd)" } else { "api(cmd)" } + $kind = if ($p.CommandLine -match 'biz_state_worker') { + "biz_state_worker(cmd)" + } elseif ($p.CommandLine -match 'netx_api\.worker') { + "worker(cmd)" + } else { + "api(cmd)" + } Stop-OnePid -ProcId ([int]$p.ProcessId) -Label $kind } } @@ -99,6 +106,18 @@ if (Test-Path $workerPidFile) { Write-Host "[INFO] No worker PID file" } +if (Test-Path $bizStateWorkerPidDir) { + Get-ChildItem -Path $bizStateWorkerPidDir -Filter "*.pid" -ErrorAction SilentlyContinue | ForEach-Object { + $t = (Get-Content -Path $_.FullName -ErrorAction SilentlyContinue | Select-Object -First 1) + $id = 0 + [void][int]::TryParse("$t", [ref]$id) + if ($id -gt 0) { + Stop-OnePid -ProcId $id -Label "biz_state_worker" + } + Remove-Item -Path $_.FullName -Force -ErrorAction SilentlyContinue + } +} + if (Test-Path $webPidFile) { $webPidText = (Get-Content -Path $webPidFile -ErrorAction SilentlyContinue | Select-Object -First 1) $webProcId = 0 diff --git a/scripts/stop_netx.sh b/scripts/stop_netx.sh index a3c9da0..95e46b4 100644 --- a/scripts/stop_netx.sh +++ b/scripts/stop_netx.sh @@ -102,6 +102,16 @@ else echo "[INFO] No worker PID file: ${WORKER_PID_FILE}" fi +BIZ_DIR="${RUN_DIR}/biz_state_workers" +if [[ -d "${BIZ_DIR}" ]]; then + for f in "${BIZ_DIR}"/*.pid; do + [[ -f "${f}" ]] || continue + BPID="$(head -n 1 "${f}" | tr -d '[:space:]' || true)" + kill_pid "${BPID}" "biz_state_worker" + rm -f "${f}" || true + done +fi + if [[ -f "${WEB_PID_FILE}" ]]; then WEB_PID="$(head -n 1 "${WEB_PID_FILE}" | tr -d '[:space:]' || true)" kill_pid "${WEB_PID}" "web" diff --git a/tests/test_biz_state_claim.py b/tests/test_biz_state_claim.py new file mode 100644 index 0000000..12beb4f --- /dev/null +++ b/tests/test_biz_state_claim.py @@ -0,0 +1,125 @@ +"""biz_state enqueue / claim (NE mutex) tests.""" + +from __future__ import annotations + +import unittest +from unittest.mock import patch +from uuid import uuid4 + +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker +from sqlalchemy.pool import StaticPool + +from netx_api.biz_state import claim as claim_mod +from netx_api.db import Base +from netx_api.models import BizStateBatch, BizStateTask, BizStateTaskItem + + +class BizStateClaimTests(unittest.TestCase): + def setUp(self) -> None: + engine = create_engine( + "sqlite+pysqlite:///:memory:", + future=True, + connect_args={"check_same_thread": False}, + poolclass=StaticPool, + ) + TestingSession = sessionmaker( + bind=engine, autoflush=False, autocommit=False, expire_on_commit=False + ) + Base.metadata.create_all(bind=engine) + self.Session = TestingSession + self.db = TestingSession() + self._session_patch = patch.object(claim_mod, "SessionLocal", TestingSession) + self._session_patch.start() + + self.task_a = BizStateTask( + id="ta", + source="managed", + ne_id="ne-a", + ne_name="A", + status="running", + collect_running=False, + interval_sec=300, + ) + self.task_b = BizStateTask( + id="tb", + source="managed", + ne_id="ne-a", # same NE as A + ne_name="A2", + status="running", + collect_running=False, + interval_sec=300, + ) + self.task_c = BizStateTask( + id="tc", + source="managed", + ne_id="ne-c", + ne_name="C", + status="running", + collect_running=False, + interval_sec=300, + ) + self.db.add_all([self.task_a, self.task_b, self.task_c]) + for tid in ("ta", "tb", "tc"): + self.db.add( + BizStateTaskItem( + id=uuid4().hex, + task_id=tid, + source_profile_id="zte.arp", + kind="catalog", + enabled=True, + ) + ) + self.db.commit() + + def tearDown(self) -> None: + self._session_patch.stop() + self.db.close() + + def test_enqueue_sets_queued_batch(self) -> None: + r = claim_mod.enqueue_collect("ta", manual=True) + self.assertTrue(r.get("queued"), r) + self.db.expire_all() + task = self.db.get(BizStateTask, "ta") + assert task is not None + self.assertTrue(task.collect_running) + batch = self.db.get(BizStateBatch, r["batch_id"]) + assert batch is not None + self.assertEqual(batch.status, "queued") + + def test_enqueue_idempotent_while_running(self) -> None: + r1 = claim_mod.enqueue_collect("ta", manual=True) + self.assertTrue(r1.get("queued"), r1) + r2 = claim_mod.enqueue_collect("ta", manual=True) + self.assertFalse(r2.get("queued")) + self.assertEqual(r2.get("reason"), "already_collecting") + + def test_claim_ne_mutex_skips_same_ne(self) -> None: + with patch.object(claim_mod, "max_concurrent_tasks", return_value=10): + ra = claim_mod.enqueue_collect("ta", manual=True) + rb = claim_mod.enqueue_collect("tb", manual=True) + rc = claim_mod.enqueue_collect("tc", manual=True) + self.assertTrue(ra["queued"] and rb["queued"] and rc["queued"]) + claimed = claim_mod.claim_queued_batches(10) + ids = {c["batch_id"] for c in claimed} + # ta and tc can run; tb same NE as ta must wait + self.assertIn(ra["batch_id"], ids) + self.assertIn(rc["batch_id"], ids) + self.assertNotIn(rb["batch_id"], ids) + self.db.expire_all() + self.assertEqual(self.db.get(BizStateBatch, ra["batch_id"]).status, "running") + self.assertEqual(self.db.get(BizStateBatch, rb["batch_id"]).status, "queued") + self.assertEqual(self.db.get(BizStateBatch, rc["batch_id"]).status, "running") + + def test_claim_respects_slot_ceiling(self) -> None: + with patch.object(claim_mod, "max_concurrent_tasks", return_value=1): + claim_mod.enqueue_collect("ta", manual=True) + claim_mod.enqueue_collect("tc", manual=True) + first = claim_mod.claim_queued_batches(10) + self.assertEqual(len(first), 1) + second = claim_mod.claim_queued_batches(10) + self.assertEqual(len(second), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_biz_state_collect_recovery.py b/tests/test_biz_state_collect_recovery.py index bac608a..2092300 100644 --- a/tests/test_biz_state_collect_recovery.py +++ b/tests/test_biz_state_collect_recovery.py @@ -88,6 +88,21 @@ class BizStateCollectRecoveryTests(unittest.TestCase): assert ok is not None self.assertEqual(ok.status, "success") + def test_queued_batch_also_marked_partial(self) -> None: + q = BizStateBatch( + id="b_q", + task_id="t1", + status="queued", + message="queued", + ) + self.db.add(q) + self.db.commit() + out = recover_interrupted_collects_on_startup(self.db) + self.assertGreaterEqual(out["batches"], 2) + self.db.refresh(q) + self.assertEqual(q.status, "partial") + self.assertIn("interrupted_by_restart", q.message or "") + if __name__ == "__main__": unittest.main() diff --git a/tests/test_biz_state_dual_lane_absorb.py b/tests/test_biz_state_dual_lane_absorb.py new file mode 100644 index 0000000..7a3d9d6 --- /dev/null +++ b/tests/test_biz_state_dual_lane_absorb.py @@ -0,0 +1,51 @@ +"""Regression: dual-lane absorb must not inflate row_count from cumulative DB.""" + +from __future__ import annotations + +import unittest + +from netx_api.biz_state.collect_runner import _absorb_lane_result + + +class DualLaneAbsorbTests(unittest.TestCase): + def test_absorb_sums_lane_deltas_not_cumulative(self) -> None: + # Each lane should report only its own delta (0 when persist owns rows). + lane_errors: list[str] = [] + total_rows, cmd_count, any_fail, any_ok = _absorb_lane_result( + (0, 10, False, True), + total_rows=0, + cmd_count=0, + any_fail=False, + any_ok=False, + lane_errors=lane_errors, + ) + total_rows, cmd_count, any_fail, any_ok = _absorb_lane_result( + (0, 5, False, True), + total_rows=total_rows, + cmd_count=cmd_count, + any_fail=any_fail, + any_ok=any_ok, + lane_errors=lane_errors, + ) + self.assertEqual(total_rows, 0) + self.assertEqual(cmd_count, 15) + self.assertTrue(any_ok) + self.assertFalse(any_fail) + + def test_legacy_bug_pattern_would_inflate(self) -> None: + # Document what NOT to do: feeding cumulative DB totals into absorb. + lane_errors: list[str] = [] + # If light returns cumulative 100 and heavy returns cumulative 150... + bad_rows, _, _, _ = _absorb_lane_result( + (150, 5, False, True), + total_rows=100, + cmd_count=10, + any_fail=False, + any_ok=True, + lane_errors=lane_errors, + ) + self.assertEqual(bad_rows, 250) # inflated — lanes must return deltas only + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_collect_dedupe.py b/tests/test_collect_dedupe.py index ee12c23..49571fe 100644 --- a/tests/test_collect_dedupe.py +++ b/tests/test_collect_dedupe.py @@ -53,23 +53,44 @@ class ManualCollectTests(unittest.TestCase): dc.assert_called_once_with("t1", manual=True) def test_dispatch_manual_allows_paused(self) -> None: - task = MagicMock() - task.collect_running = False - task.status = "paused" - task.id = "t1" - task.source = "managed" - task.ne_id = "n1" - task.ne_name = "NE" - task.vendor = "zte" - db = MagicMock() - db.get.return_value = task - # items query → empty so it returns early after setting error - q = MagicMock() - q.filter.return_value.order_by.return_value.all.return_value = [] - db.query.return_value = q - with patch("netx_api.biz_state.collect_runner.SessionLocal", return_value=db): + with patch( + "netx_api.biz_state.claim.enqueue_collect", + return_value={ + "ok": False, + "queued": False, + "reason": "no_enabled_items", + "task_id": "t1", + }, + ) as enq: dispatch_collect("t1", manual=True) - self.assertEqual(task.last_error, "no enabled task items") + enq.assert_called_once_with("t1", manual=True) + + def test_dispatch_inline_skips_if_already_claimed(self) -> None: + with ( + patch( + "netx_api.biz_state.claim.enqueue_collect", + return_value={ + "ok": True, + "queued": True, + "batch_id": "b1", + "task_id": "t1", + }, + ), + patch( + "netx_api.biz_state.collect_runner._should_execute_inline", + return_value=True, + ), + patch( + "netx_api.biz_state.collect_runner._try_claim_batch_for_execute", + return_value=False, + ) as claim, + patch( + "netx_api.biz_state.collect_runner.execute_claimed_batch" + ) as exe, + ): + dispatch_collect("t1", manual=True) + claim.assert_called_once_with("b1") + exe.assert_not_called() if __name__ == "__main__": diff --git a/web/src/pages/network/BizStatePage.tsx b/web/src/pages/network/BizStatePage.tsx index 4cf8709..f4235a9 100644 --- a/web/src/pages/network/BizStatePage.tsx +++ b/web/src/pages/network/BizStatePage.tsx @@ -1,5 +1,5 @@ import { Button, Input, Modal } from "@heroui/react"; -import { useCallback, useEffect, useMemo, useState } from "react"; +import { useCallback, useEffect, useMemo, useRef, useState } from "react"; import { ListPager } from "../../components/ListPager"; import { AppModalShell } from "../../components/ui/AppModalShell"; import { FieldSelect } from "../../components/ui/FieldSelect"; @@ -238,6 +238,9 @@ export function BizStatePage() { const [tasks, setTasks] = useState([]); const [busy, setBusy] = useState(false); const [collectingIds, setCollectingIds] = useState>({}); + /** Track that we observed collect_running=true so we don't clear the chip before enqueue lands. */ + const seenCollectRunningRef = useRef>({}); + const collectStartedAtRef = useRef>({}); const [listKeyword, setListKeyword] = useState(""); const debouncedListKw = useDebouncedValue(listKeyword, 250); const [purposeFilter, setPurposeFilter] = useState<"all" | "portrait" | "cutover_hf">("all"); @@ -308,6 +311,25 @@ export function BizStatePage() { return items; }, [purposeFilter]); + /** Lightweight poll: task flags + batch counters only (no profile reload). */ + const refreshTaskProgress = useCallback(async (id: string) => { + const task = await bizStateGetTask(id); + setDetail((prev: any) => { + if (!prev || prev.id !== id) return prev; + return { + ...prev, + collect_running: Boolean(task.collect_running), + last_error: task.last_error, + last_collect_started_at: task.last_collect_started_at, + last_collect_ended_at: task.last_collect_ended_at, + status: task.status, + }; + }); + const b = await bizStateListBatches(id, 50); + setBatches((b.items || []) as BatchRow[]); + return task as TaskRow; + }, []); + useEffect(() => { void (async () => { try { @@ -319,26 +341,55 @@ export function BizStatePage() { // eslint-disable-next-line react-hooks/exhaustive-deps -- refresh when purpose filter / refreshTasks changes }, [refreshTasks]); - // While a collect is running, refresh batch counters so cmd/row progress is visible. + // Progress poll while any collect is running (list chips and/or open task). + // Does NOT block navigation; cleans up on unmount / when nothing is collecting. useEffect(() => { - if (!taskId || !detail?.collect_running) return; + const watching = new Set(Object.keys(collectingIds)); + if (taskId && detail?.collect_running) watching.add(taskId); + if (!watching.size) return; + let cancelled = false; const tick = async () => { + if (cancelled) return; try { + const items = await refreshTasks(); if (cancelled) return; - await loadTask(taskId); + setCollectingIds((prev) => { + let changed = false; + const next = { ...prev }; + const now = Date.now(); + for (const id of Object.keys(next)) { + const row = items.find((x) => x.id === id); + if (row?.collect_running) { + seenCollectRunningRef.current[id] = true; + continue; + } + const seen = Boolean(seenCollectRunningRef.current[id]); + const started = collectStartedAtRef.current[id] || 0; + // Clear after we saw running→idle, or enqueue never landed (~20s). + if (seen || (started && now - started > 20_000)) { + delete next[id]; + delete seenCollectRunningRef.current[id]; + delete collectStartedAtRef.current[id]; + changed = true; + } + } + return changed ? next : prev; + }); + if (taskId && watching.has(taskId)) { + await refreshTaskProgress(taskId); + } } catch { /* ignore transient poll errors */ } }; - const timer = window.setInterval(() => void tick(), 3000); + const timer = window.setInterval(() => void tick(), 4000); void tick(); return () => { cancelled = true; window.clearInterval(timer); }; - // eslint-disable-next-line react-hooks/exhaustive-deps -- poll while collect_running - }, [taskId, detail?.collect_running]); + }, [collectingIds, taskId, detail?.collect_running, refreshTasks, refreshTaskProgress]); useEffect(() => { if (!createOpen) return; @@ -718,9 +769,13 @@ export function BizStatePage() { const setTaskCollecting = (id: string, on: boolean) => { setCollectingIds((prev) => { if (on) { + collectStartedAtRef.current[id] = Date.now(); + seenCollectRunningRef.current[id] = false; if (prev[id]) return prev; return { ...prev, [id]: true }; } + delete seenCollectRunningRef.current[id]; + delete collectStartedAtRef.current[id]; if (!prev[id]) return prev; const next = { ...prev }; delete next[id]; @@ -735,34 +790,23 @@ export function BizStatePage() { showOk(t("bizState.collecting")); if (fromModal && taskId === id) { setTaskTab("batches"); - } - // Heavy show-interface can take ~20 minutes; poll long enough and refresh batches. - const deadline = Date.now() + 32 * 60 * 1000; - while (Date.now() < deadline) { - await new Promise((r) => setTimeout(r, 3000)); - if (fromModal && taskId === id) { - try { - await loadTask(id); - const task = await bizStateGetTask(id); - setDetail(task); - if (!task.collect_running) break; - } catch { - break; - } - continue; + try { + await refreshTaskProgress(id); + } catch { + /* progress poll will retry */ + } + } else { + try { + await refreshTasks(); + } catch { + /* list poll will retry */ } - const items = await refreshTasks(); - const latest = items.find((x) => x.id === id); - if (!latest?.collect_running) break; - } - await refreshTasks(); - if (fromModal && taskId === id) { - await loadTask(id); } + // Do NOT block UI for the full collect duration. Progress is driven by + // the collectingIds / collect_running effect (cleans up on unmount). } catch (e) { - showError(formatErr(e)); - } finally { setTaskCollecting(id, false); + showError(formatErr(e)); } };