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 <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-09-22 17:00:37 +08:00
parent 0eb65a839e
commit 3a0de0d7fb
22 changed files with 1210 additions and 165 deletions

View file

@ -66,6 +66,15 @@ NETX_UME_NOTIFICATION_TOPIC=ALARM
# (start_netx.ps1/.sh do this automatically). Worker writes heartbeat for API /metrics. # (start_netx.ps1/.sh do this automatically). Worker writes heartbeat for API /metrics.
# NETX_RUN_INLINE_SCHEDULERS=false # NETX_RUN_INLINE_SCHEDULERS=false
# NETX_SCHEDULER_HEARTBEAT_PATH=data/runtime/scheduler_heartbeat.json # 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) --- # --- Multi-user shared-server capacity (defaults in Settings already match these) ---
# NETX_DB_POOL_SIZE=40 # NETX_DB_POOL_SIZE=40
# NETX_DB_MAX_OVERFLOW=40 # NETX_DB_MAX_OVERFLOW=40
@ -94,6 +103,6 @@ NETX_UME_NOTIFICATION_TOPIC=ALARM
# NETX_NE_COLLECTION_KEEP_DAYS=14 # NETX_NE_COLLECTION_KEEP_DAYS=14
# Heavier fleets: raise CLI/DB together; also ensure Postgres max_connections and bastion session limits. # 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 # 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 # Legacy single-process: add -InlineSchedulers
# Stop all: .\scripts\stop_netx.ps1 # Stop all: .\scripts\stop_netx.ps1

321
netx_api/biz_state/claim.py Normal file
View file

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

View file

@ -27,7 +27,9 @@ def recover_interrupted_collects_on_startup(db: Session) -> dict[str, Any]:
task_n = 0 task_n = 0
stuck_batches = ( 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: for b in stuck_batches:
b.status = "partial" b.status = "partial"
@ -58,6 +60,8 @@ def recover_interrupted_collects_on_startup(db: Session) -> dict[str, Any]:
) )
for t in stuck_tasks: for t in stuck_tasks:
t.collect_running = False t.collect_running = False
if hasattr(t, "collect_queued_at"):
t.collect_queued_at = None
if not t.last_collect_ended_at: if not t.last_collect_ended_at:
t.last_collect_ended_at = now t.last_collect_ended_at = now
err = str(t.last_error or "").strip() err = str(t.last_error or "").strip()

View file

@ -346,6 +346,8 @@ def _finish_task(task_id: str, *, error: str = "") -> None:
if not task: if not task:
return return
task.collect_running = False task.collect_running = False
if hasattr(task, "collect_queued_at"):
task.collect_queued_at = None
task.last_collect_ended_at = _utcnow() task.last_collect_ended_at = _utcnow()
task.last_error = str(error or "")[:1020] task.last_error = str(error or "")[:1020]
task.updated_at = _utcnow() 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: 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``). Dedicated worker mode: only enqueue (claim loop runs the batch).
Collect-now calls with ``manual=True`` (any status, as long as not already collecting). 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() db = SessionLocal()
batch_id = ""
try: try:
task = db.get(BizStateTask, task_id) batch = (
if not task: db.query(BizStateBatch)
return .filter(BizStateBatch.id == batch_id)
if task.collect_running: .with_for_update()
return .one_or_none()
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()
) )
if not items: if not batch or str(batch.status or "") != "queued":
task.last_error = "no enabled task items" return False
task.updated_at = _utcnow() batch.status = "running"
db.commit() batch.message = ""
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)
db.commit() db.commit()
batch_id = batch.id return True
vendor = str(task.vendor or "") except Exception:
device_type = str(task.device_type or "") _log.exception("biz_state claim-for-execute failed batch=%s", batch_id)
source = str(task.source or "managed").strip().lower() try:
ne_id = str(task.ne_id or "").strip() db.rollback()
except Exception:
pass
return False
finally: finally:
db.close() 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 = "" error = ""
try: try:
@ -433,12 +466,14 @@ def dispatch_collect(task_id: str, *, manual: bool = False) -> None:
device_type=device_type, device_type=device_type,
) )
except Exception as exc: 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) error = _format_error(exc)
try: try:
_fail_batch_status(batch_id, error) _fail_batch_status(batch_id, error)
except Exception: 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: finally:
_finish_task(task_id, error=error) _finish_task(task_id, error=error)
@ -486,9 +521,20 @@ def _run_collect_lane(
any_ok = False any_ok = False
pending: list[SpooledCommand] = [] pending: list[SpooledCommand] = []
task_id = "" 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: 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: if records is not None and item.persist_kind:
item.records_rel_path = write_records(batch_id, item.id, records) item.records_rel_path = write_records(batch_id, item.id, records)
item.row_count = len(records) item.row_count = len(records)
@ -499,16 +545,7 @@ def _run_collect_lane(
pending.append(item) pending.append(item)
cmd_count += 1 cmd_count += 1
if len(pending) >= flush_every: if len(pending) >= flush_every:
try: _submit_pending()
_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),
)
try: try:
try: try:
@ -853,22 +890,26 @@ def _run_collect_lane(
primary.message = f"parse: {_format_error(exc)}" primary.message = f"parse: {_format_error(exc)}"
_queue(primary) _queue(primary)
# Final flush for this lane. # Final flush for this lane → persist pool; wait so rows land before return.
if pending: _submit_pending()
_c, _r = _flush_spooled_commands(batch_id, pending) if not persist.wait_idle(timeout=max(30.0, float(budget))):
total_rows += int(_r or 0) _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 return total_rows, cmd_count, any_fail, any_ok
finally: finally:
# Best-effort: persist whatever was collected before timeout/abort. # Best-effort: enqueue leftover spool before connection teardown.
if pending: try:
try: _submit_pending()
_c, _r = _flush_spooled_commands(batch_id, pending) persist.wait_idle(timeout=60.0)
total_rows += int(_r or 0) except Exception:
except Exception: _log.exception(
_log.exception( "biz_state persist drain on lane exit failed batch=%s", batch_id
"biz_state flush on lane exit failed batch=%s", batch_id )
)
holder.pop("conn", None) holder.pop("conn", None)
close_netmiko_connection(conn) close_netmiko_connection(conn)
@ -1097,7 +1138,7 @@ def _run_collect_session(
task = db.get(BizStateTask, task_id) task = db.get(BizStateTask, task_id)
batch = db.get(BizStateBatch, batch_id) batch = db.get(BizStateBatch, batch_id)
if not task or not batch: if not task or not batch:
return raise RuntimeError(f"batch_or_task_missing batch={batch_id} task={task_id}")
try: try:
if source == "managed": if source == "managed":
@ -1302,6 +1343,21 @@ def _run_collect_session(
raise RuntimeError("; ".join(lane_errors)[:1020]) raise RuntimeError("; ".join(lane_errors)[:1020])
any_fail = True 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( _finalize_batch_status(
batch_id=batch_id, batch_id=batch_id,
task_id=task_id, task_id=task_id,

View file

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

View file

@ -245,12 +245,11 @@ def api_collect_now(
raise HTTPException(status_code=404, detail="task not found") raise HTTPException(status_code=404, detail="task not found")
if bool(task.collect_running): if bool(task.collect_running):
return {"ok": True, "started": False, "reason": "already_collecting", "task_id": task_id} 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 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)) 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") @router.get("/tasks/{task_id}/batches")

View file

@ -1,4 +1,4 @@
"""Background scheduler for biz_state collection.""" """Background scheduler for biz_state collection (enqueue + claim loop)."""
from __future__ import annotations from __future__ import annotations
@ -6,12 +6,14 @@ import logging
import threading import threading
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from datetime import datetime from datetime import datetime
from typing import Any
from .cli_budget import clamp_cli_workers from .cli_budget import clamp_cli_workers
from .config import settings from .config import settings
from .db import SessionLocal from .db import SessionLocal
from .models import BizStateTask 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") _log = logging.getLogger("netx.biz_state.scheduler")
_stop = threading.Event() _stop = threading.Event()
@ -20,6 +22,8 @@ _dispatch_pool: ThreadPoolExecutor | None = None
_pool_lock = threading.Lock() _pool_lock = threading.Lock()
_last_tick_mono: float = 0.0 _last_tick_mono: float = 0.0
_last_purge_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 _PURGE_INTERVAL_SEC = 3600.0
@ -57,11 +61,16 @@ def _dispatch_pool_get() -> ThreadPoolExecutor:
global _dispatch_pool global _dispatch_pool
with _pool_lock: with _pool_lock:
if _dispatch_pool is None: if _dispatch_pool is None:
workers = clamp_cli_workers( n = clamp_cli_workers(
int(getattr(settings, "biz_state_dispatch_workers", 2) or 2), 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( _dispatch_pool = ThreadPoolExecutor(
max_workers=workers, thread_name_prefix="biz-dispatch" max_workers=n, thread_name_prefix="biz-collect"
) )
return _dispatch_pool return _dispatch_pool
@ -75,6 +84,12 @@ def shutdown_biz_state_dispatch_pool(*, wait: bool = False) -> None:
except TypeError: except TypeError:
_dispatch_pool.shutdown(wait=wait) _dispatch_pool.shutdown(wait=wait)
_dispatch_pool = None _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: def _sync_cutover_hf_windows() -> None:
@ -122,17 +137,7 @@ def _sync_cutover_hf_windows() -> None:
db.close() db.close()
def try_dispatch_due_tasks() -> int: def _enqueue_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")
db = SessionLocal() db = SessionLocal()
try: try:
tasks = ( tasks = (
@ -156,15 +161,108 @@ def try_dispatch_due_tasks() -> int:
finally: finally:
db.close() db.close()
if not due_ids: n = 0
return 0
pool = _dispatch_pool_get()
for tid in due_ids: for tid in due_ids:
try: try:
pool.submit(dispatch_collect, tid) r = enqueue_collect(tid, manual=False)
if r.get("queued"):
n += 1
except Exception: except Exception:
_log.exception("biz_state submit failed task=%s", tid) _log.exception("biz_state enqueue failed task=%s", tid)
return len(due_ids) 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: def _loop() -> None:
@ -192,7 +290,10 @@ def start_biz_state_scheduler() -> None:
_stop.clear() _stop.clear()
_thread = threading.Thread(target=_loop, name="biz-state-scheduler", daemon=True) _thread = threading.Thread(target=_loop, name="biz-state-scheduler", daemon=True)
_thread.start() _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: def stop_biz_state_scheduler() -> None:
@ -213,8 +314,15 @@ def biz_state_scheduler_status() -> dict:
age = None age = None
if _last_tick_mono: if _last_tick_mono:
age = max(0.0, _time.monotonic() - _last_tick_mono) age = max(0.0, _time.monotonic() - _last_tick_mono)
with _in_flight_lock:
inflight = len(_in_flight)
return { return {
"running": alive, "running": alive,
"enabled": bool(getattr(settings, "biz_state_scheduler_enabled", True)), "enabled": bool(getattr(settings, "biz_state_scheduler_enabled", True)),
"last_tick_age_sec": age, "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)
),
} }

View file

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

View file

@ -127,6 +127,14 @@ class Settings(BaseSettings):
biz_state_persist_every_cmds: int = 8 biz_state_persist_every_cmds: int = 8
# Cap raw_text loaded into Postgres from spool (0 = unlimited). # Cap raw_text loaded into Postgres from spool (0 = unlimited).
biz_state_raw_max_bytes: int = 8 * 1024 * 1024 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). # Managed NE exec: max CLI commands per request (lab can raise; hard-capped in ne_exec).
ne_exec_max_commands: int = 5 ne_exec_max_commands: int = 5
# Opt-in: allow per-NE exec_policy (linux_shell/unrestricted). Default off — UI hidden. # Opt-in: allow per-NE exec_policy (linux_shell/unrestricted). Default off — UI hidden.

View file

@ -42,6 +42,8 @@ class BizStateTask(Base):
# Legacy column kept for brownfield reads; unused by new purge path # Legacy column kept for brownfield reads; unused by new purge path
retention_batches: Mapped[int] = mapped_column(Integer, default=30) retention_batches: Mapped[int] = mapped_column(Integer, default=30)
collect_running: Mapped[bool] = mapped_column(Boolean, default=False) 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_started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
last_collect_ended_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="") 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_id: Mapped[str] = mapped_column(String(128), default="", index=True)
ne_name: Mapped[str] = mapped_column(String(256), default="") ne_name: Mapped[str] = mapped_column(String(256), default="")
vendor: Mapped[str] = mapped_column(String(64), 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) command_count: Mapped[int] = mapped_column(Integer, default=0)
row_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="") message: Mapped[str] = mapped_column(String(1024), default="")

View file

@ -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 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_x DOUBLE PRECISION",
"ALTER TABLE topo_fabric_node ADD COLUMN IF NOT EXISTS world_y 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_x ON topo_fabric_node (world_x)",
"CREATE INDEX IF NOT EXISTS ix_topo_fabric_node_world_y ON topo_fabric_node (world_y)", "CREATE INDEX IF NOT EXISTS ix_topo_fabric_node_world_y ON topo_fabric_node (world_y)",
"COMMENT ON TABLE ume_inventory_ne IS '网元对象详细信息'", "COMMENT ON TABLE ume_inventory_ne IS '网元对象详细信息'",

View file

@ -55,10 +55,17 @@ def start_device_schedulers() -> None:
start_lldp_collect_scheduler() start_lldp_collect_scheduler()
start_ne_collect_scheduler() start_ne_collect_scheduler()
start_port_traffic_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() start_fabric_reconcile_scheduler()
# Publish status so API /metrics can see collectors when run in a split worker. # 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) start_scheduler_heartbeat_publisher(role=role)
_log.info("device schedulers started") _log.info("device schedulers started")

View file

@ -5,6 +5,8 @@ Optional when ``NETX_RUN_INLINE_SCHEDULERS=false`` (API does not start collector
python -m netx_api.worker python -m netx_api.worker
Starts: config_sync, lldp_collect, port_traffic, fabric_reconcile tick loops. 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). UME WS / keepalive remain in the API process (token + alarm coordination).
By default the API runs collectors inline — no separate worker needed. By default the API runs collectors inline — no separate worker needed.
""" """

View file

@ -34,6 +34,7 @@ $errFile = Join-Path $runDir "netx.err.log"
$workerPidFile = Join-Path $runDir "worker.pid" $workerPidFile = Join-Path $runDir "worker.pid"
$workerLogFile = Join-Path $runDir "worker.out.log" $workerLogFile = Join-Path $runDir "worker.out.log"
$workerErrFile = Join-Path $runDir "worker.err.log" $workerErrFile = Join-Path $runDir "worker.err.log"
$bizStateWorkerPidDir = Join-Path $runDir "biz_state_workers"
$webPidFile = Join-Path $runDir "web.pid" $webPidFile = Join-Path $runDir "web.pid"
$webLogFile = Join-Path $runDir "web.out.log" $webLogFile = Join-Path $runDir "web.out.log"
$webErrFile = Join-Path $runDir "web.err.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) { if ($Background) {
# Truncate logs so a failed start is not confused with an old run. # Truncate logs so a failed start is not confused with an old run.
Set-Content -Path $logFile -Value "" -Encoding utf8 Set-Content -Path $logFile -Value "" -Encoding utf8
@ -226,6 +262,7 @@ if ($Background) {
} }
Write-Host "==> netx API ready: http://${BindHost}:${Port}/health" -ForegroundColor Green Write-Host "==> netx API ready: http://${BindHost}:${Port}/health" -ForegroundColor Green
Start-NetxWorker Start-NetxWorker
Start-BizStateWorkers
if ($WithWeb) { if ($WithWeb) {
Write-Host "==> Starting Vite dev server in background" Write-Host "==> Starting Vite dev server in background"
$webRoot = Join-Path $projectRoot "web" $webRoot = Join-Path $projectRoot "web"
@ -283,5 +320,6 @@ if ($WithWeb) {
} }
Start-NetxWorker Start-NetxWorker
Start-BizStateWorkers
Write-Host "==> Starting netx API in foreground" Write-Host "==> Starting netx API in foreground"
& $pythonExe -m netx_api.main & $pythonExe -m netx_api.main

View file

@ -239,6 +239,21 @@ if [[ "${INLINE_SCHEDULERS}" != "1" ]]; then
echo "PID = ${WORKER_PID}" echo "PID = ${WORKER_PID}"
echo "Log = ${WORKER_LOG_FILE}" echo "Log = ${WORKER_LOG_FILE}"
echo "Err = ${WORKER_ERR_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<REPLICAS; i++)); do
(
cd "${ROOT_DIR}"
nohup "${PYTHON_CMD}" -m netx_api.biz_state_worker \
>"${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 fi
if [[ "${API_ONLY}" != "1" ]]; then if [[ "${API_ONLY}" != "1" ]]; then

View file

@ -13,6 +13,7 @@ $runDir = Join-Path $PSScriptRoot ".run"
$pidFile = Join-Path $runDir "netx.pid" $pidFile = Join-Path $runDir "netx.pid"
$workerPidFile = Join-Path $runDir "worker.pid" $workerPidFile = Join-Path $runDir "worker.pid"
$webPidFile = Join-Path $runDir "web.pid" $webPidFile = Join-Path $runDir "web.pid"
$bizStateWorkerPidDir = Join-Path $runDir "biz_state_workers"
function Stop-OnePid { function Stop-OnePid {
param([int]$ProcId, [string]$Label) param([int]$ProcId, [string]$Label)
@ -62,13 +63,19 @@ function Get-ListenPids {
function Stop-NetxByCommandLine { function Stop-NetxByCommandLine {
# Orphan workers often have no worker.pid but still hold worker.out.log. # Orphan workers often have no worker.pid but still hold worker.out.log.
$hits = @(Get-CimInstance Win32_Process -ErrorAction SilentlyContinue | $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) { 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 return
} }
foreach ($p in $hits) { 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 Stop-OnePid -ProcId ([int]$p.ProcessId) -Label $kind
} }
} }
@ -99,6 +106,18 @@ if (Test-Path $workerPidFile) {
Write-Host "[INFO] No worker PID file" 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) { if (Test-Path $webPidFile) {
$webPidText = (Get-Content -Path $webPidFile -ErrorAction SilentlyContinue | Select-Object -First 1) $webPidText = (Get-Content -Path $webPidFile -ErrorAction SilentlyContinue | Select-Object -First 1)
$webProcId = 0 $webProcId = 0

View file

@ -102,6 +102,16 @@ else
echo "[INFO] No worker PID file: ${WORKER_PID_FILE}" echo "[INFO] No worker PID file: ${WORKER_PID_FILE}"
fi 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 if [[ -f "${WEB_PID_FILE}" ]]; then
WEB_PID="$(head -n 1 "${WEB_PID_FILE}" | tr -d '[:space:]' || true)" WEB_PID="$(head -n 1 "${WEB_PID_FILE}" | tr -d '[:space:]' || true)"
kill_pid "${WEB_PID}" "web" kill_pid "${WEB_PID}" "web"

View file

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

View file

@ -88,6 +88,21 @@ class BizStateCollectRecoveryTests(unittest.TestCase):
assert ok is not None assert ok is not None
self.assertEqual(ok.status, "success") 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__": if __name__ == "__main__":
unittest.main() unittest.main()

View file

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

View file

@ -53,23 +53,44 @@ class ManualCollectTests(unittest.TestCase):
dc.assert_called_once_with("t1", manual=True) dc.assert_called_once_with("t1", manual=True)
def test_dispatch_manual_allows_paused(self) -> None: def test_dispatch_manual_allows_paused(self) -> None:
task = MagicMock() with patch(
task.collect_running = False "netx_api.biz_state.claim.enqueue_collect",
task.status = "paused" return_value={
task.id = "t1" "ok": False,
task.source = "managed" "queued": False,
task.ne_id = "n1" "reason": "no_enabled_items",
task.ne_name = "NE" "task_id": "t1",
task.vendor = "zte" },
db = MagicMock() ) as enq:
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):
dispatch_collect("t1", manual=True) 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__": if __name__ == "__main__":

View file

@ -1,5 +1,5 @@
import { Button, Input, Modal } from "@heroui/react"; 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 { ListPager } from "../../components/ListPager";
import { AppModalShell } from "../../components/ui/AppModalShell"; import { AppModalShell } from "../../components/ui/AppModalShell";
import { FieldSelect } from "../../components/ui/FieldSelect"; import { FieldSelect } from "../../components/ui/FieldSelect";
@ -238,6 +238,9 @@ export function BizStatePage() {
const [tasks, setTasks] = useState<TaskRow[]>([]); const [tasks, setTasks] = useState<TaskRow[]>([]);
const [busy, setBusy] = useState(false); const [busy, setBusy] = useState(false);
const [collectingIds, setCollectingIds] = useState<Record<string, true>>({}); const [collectingIds, setCollectingIds] = useState<Record<string, true>>({});
/** Track that we observed collect_running=true so we don't clear the chip before enqueue lands. */
const seenCollectRunningRef = useRef<Record<string, boolean>>({});
const collectStartedAtRef = useRef<Record<string, number>>({});
const [listKeyword, setListKeyword] = useState(""); const [listKeyword, setListKeyword] = useState("");
const debouncedListKw = useDebouncedValue(listKeyword, 250); const debouncedListKw = useDebouncedValue(listKeyword, 250);
const [purposeFilter, setPurposeFilter] = useState<"all" | "portrait" | "cutover_hf">("all"); const [purposeFilter, setPurposeFilter] = useState<"all" | "portrait" | "cutover_hf">("all");
@ -308,6 +311,25 @@ export function BizStatePage() {
return items; return items;
}, [purposeFilter]); }, [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(() => { useEffect(() => {
void (async () => { void (async () => {
try { try {
@ -319,26 +341,55 @@ export function BizStatePage() {
// eslint-disable-next-line react-hooks/exhaustive-deps -- refresh when purpose filter / refreshTasks changes // eslint-disable-next-line react-hooks/exhaustive-deps -- refresh when purpose filter / refreshTasks changes
}, [refreshTasks]); }, [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(() => { 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; let cancelled = false;
const tick = async () => { const tick = async () => {
if (cancelled) return;
try { try {
const items = await refreshTasks();
if (cancelled) return; 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 { } catch {
/* ignore transient poll errors */ /* ignore transient poll errors */
} }
}; };
const timer = window.setInterval(() => void tick(), 3000); const timer = window.setInterval(() => void tick(), 4000);
void tick(); void tick();
return () => { return () => {
cancelled = true; cancelled = true;
window.clearInterval(timer); window.clearInterval(timer);
}; };
// eslint-disable-next-line react-hooks/exhaustive-deps -- poll while collect_running }, [collectingIds, taskId, detail?.collect_running, refreshTasks, refreshTaskProgress]);
}, [taskId, detail?.collect_running]);
useEffect(() => { useEffect(() => {
if (!createOpen) return; if (!createOpen) return;
@ -718,9 +769,13 @@ export function BizStatePage() {
const setTaskCollecting = (id: string, on: boolean) => { const setTaskCollecting = (id: string, on: boolean) => {
setCollectingIds((prev) => { setCollectingIds((prev) => {
if (on) { if (on) {
collectStartedAtRef.current[id] = Date.now();
seenCollectRunningRef.current[id] = false;
if (prev[id]) return prev; if (prev[id]) return prev;
return { ...prev, [id]: true }; return { ...prev, [id]: true };
} }
delete seenCollectRunningRef.current[id];
delete collectStartedAtRef.current[id];
if (!prev[id]) return prev; if (!prev[id]) return prev;
const next = { ...prev }; const next = { ...prev };
delete next[id]; delete next[id];
@ -735,34 +790,23 @@ export function BizStatePage() {
showOk(t("bizState.collecting")); showOk(t("bizState.collecting"));
if (fromModal && taskId === id) { if (fromModal && taskId === id) {
setTaskTab("batches"); setTaskTab("batches");
} try {
// Heavy show-interface can take ~20 minutes; poll long enough and refresh batches. await refreshTaskProgress(id);
const deadline = Date.now() + 32 * 60 * 1000; } catch {
while (Date.now() < deadline) { /* progress poll will retry */
await new Promise((r) => setTimeout(r, 3000)); }
if (fromModal && taskId === id) { } else {
try { try {
await loadTask(id); await refreshTasks();
const task = await bizStateGetTask(id); } catch {
setDetail(task); /* list poll will retry */
if (!task.collect_running) break;
} catch {
break;
}
continue;
} }
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) { } catch (e) {
showError(formatErr(e));
} finally {
setTaskCollecting(id, false); setTaskCollecting(id, false);
showError(formatErr(e));
} }
}; };