mirror of
https://github.com/hansjone/netx.git
synced 2026-10-09 03:10:46 +08:00
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:
parent
0eb65a839e
commit
3a0de0d7fb
22 changed files with 1210 additions and 165 deletions
11
.env.example
11
.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
|
||||
|
|
|
|||
321
netx_api/biz_state/claim.py
Normal file
321
netx_api/biz_state/claim.py
Normal 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())
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
126
netx_api/biz_state/persist_pool.py
Normal file
126
netx_api/biz_state/persist_pool.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
),
|
||||
}
|
||||
|
|
|
|||
64
netx_api/biz_state_worker.py
Normal file
64
netx_api/biz_state_worker.py
Normal 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()
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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="")
|
||||
|
|
|
|||
|
|
@ -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 '网元对象详细信息'",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<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
|
||||
|
||||
if [[ "${API_ONLY}" != "1" ]]; then
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
125
tests/test_biz_state_claim.py
Normal file
125
tests/test_biz_state_claim.py
Normal 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()
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
51
tests/test_biz_state_dual_lane_absorb.py
Normal file
51
tests/test_biz_state_dual_lane_absorb.py
Normal 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()
|
||||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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<TaskRow[]>([]);
|
||||
const [busy, setBusy] = useState(false);
|
||||
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 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));
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue