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

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

View file

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

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

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

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
# 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.

View file

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

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 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 '网元对象详细信息'",

View file

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

View file

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