mirror of
https://github.com/hansjone/netx.git
synced 2026-10-09 02:00: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
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.
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue