mirror of
https://github.com/hansjone/netx.git
synced 2026-10-12 04:10:47 +08:00
400 lines
14 KiB
Python
400 lines
14 KiB
Python
"""Enqueue / claim biz_state collect batches (Postgres SKIP LOCKED + NE mutex)."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
from datetime import datetime, timedelta
|
|
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 task.collect_group_id:
|
|
group_id = task.collect_group_id
|
|
db.close()
|
|
return enqueue_collect_group(group_id, manual=manual, requested_task_id=tid)
|
|
if str(task.source or "").strip().lower() == "import":
|
|
return {
|
|
"ok": False,
|
|
"queued": False,
|
|
"reason": "import_offline_only",
|
|
"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,
|
|
queued_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 enqueue_collect_group(group_id: str, *, manual: bool = False,
|
|
requested_task_id: str = "") -> dict[str, Any]:
|
|
"""Atomically enqueue both devices with one round identity, or neither."""
|
|
db = SessionLocal()
|
|
try:
|
|
tasks = db.query(BizStateTask).filter(BizStateTask.collect_group_id == group_id).order_by(
|
|
BizStateTask.id).with_for_update().populate_existing().all()
|
|
result: dict[str, Any] = {"ok": False, "queued": False, "task_id": requested_task_id,
|
|
"collect_group_id": group_id}
|
|
if len(tasks) != 2 or len({t.ne_id for t in tasks}) != 2:
|
|
return {**result, "reason": "collect_pair_invalid"}
|
|
if max_concurrent_tasks() < 2:
|
|
return {**result, "reason": "collect_pair_capacity"}
|
|
from ..cli_budget import clamp_cli_workers
|
|
|
|
worker_cap = clamp_cli_workers(int(getattr(settings, "biz_state_worker_collect_threads", None)
|
|
or getattr(settings, "biz_state_dispatch_workers", 8) or 8))
|
|
if worker_cap < 2:
|
|
return {**result, "reason": "collect_pair_capacity"}
|
|
if any(t.collect_running for t in tasks):
|
|
return {**result, "reason": "collect_pair_busy"}
|
|
for task in tasks:
|
|
if task.source == "import" or (not manual and (task.status != "running" or task.collect_manual_only)):
|
|
return {**result, "reason": "not_scheduled"}
|
|
if task.status in ("", "deleted"):
|
|
return {**result, "reason": "bad_status"}
|
|
if not db.query(BizStateTaskItem.id).filter(BizStateTaskItem.task_id == task.id,
|
|
BizStateTaskItem.enabled.is_(True)).first():
|
|
return {**result, "reason": "no_enabled_items"}
|
|
now, round_id = _utcnow(), uuid4().hex
|
|
batches = []
|
|
for task in tasks:
|
|
task.collect_running = True
|
|
task.collect_queued_at = now
|
|
task.last_error = ""
|
|
task.updated_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, queued_at=now, collect_group_id=group_id,
|
|
collect_round_id=round_id, collect_priority=10 if task.collect_manual_only else 20,
|
|
message="queued_manual" if manual else "queued")
|
|
db.add(batch)
|
|
batches.append({"batch_id": batch.id, "task_id": task.id})
|
|
db.commit()
|
|
selected = next((b for b in batches if b["task_id"] == requested_task_id), batches[0])
|
|
return {**result, **selected, "ok": True, "queued": True, "collect_round_id": round_id,
|
|
"batches": batches, "manual": manual}
|
|
except Exception:
|
|
db.rollback()
|
|
_log.exception("enqueue collect pair failed group=%s", group_id)
|
|
return {"ok": False, "queued": False, "reason": "enqueue_error", "task_id": requested_task_id}
|
|
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 and_, case, text
|
|
|
|
global_cap = max_concurrent_tasks()
|
|
db = SessionLocal()
|
|
try:
|
|
pg = _dialect_supports_skip_locked(db)
|
|
if pg:
|
|
# A short transaction lock makes capacity reservation and pair claims
|
|
# atomic across worker processes. Never held during device I/O.
|
|
db.execute(text("SELECT pg_advisory_xact_lock(hashtext('biz-state-claim-capacity'))"))
|
|
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()
|
|
}
|
|
|
|
# Manual route rounds yield to fresh monitoring, but must eventually
|
|
# run even when a busy device always has regular metrics queued.
|
|
priority = case((and_(BizStateBatch.collect_priority == 10,
|
|
BizStateBatch.queued_at <= _utcnow() - timedelta(seconds=300)), 30),
|
|
else_=BizStateBatch.collect_priority)
|
|
q = (
|
|
db.query(BizStateBatch)
|
|
.filter(BizStateBatch.status == "queued")
|
|
.order_by(priority.desc(), BizStateBatch.started_at.asc(), BizStateBatch.id.asc())
|
|
)
|
|
if pg:
|
|
q = q.with_for_update(skip_locked=True)
|
|
else:
|
|
q = q.with_for_update()
|
|
|
|
candidates = q.limit(max(slots * 8, 64)).all()
|
|
claimed_rows: list[BizStateBatch] = []
|
|
visited: set[str] = set()
|
|
for batch in candidates:
|
|
if batch.id in visited:
|
|
continue
|
|
if len(claimed_rows) >= slots:
|
|
break
|
|
group = [batch]
|
|
if batch.collect_round_id:
|
|
peers = db.query(BizStateBatch).filter(
|
|
BizStateBatch.collect_round_id == batch.collect_round_id,
|
|
).with_for_update(skip_locked=pg).all()
|
|
if len(peers) != 2 or any(p.status != "queued" for p in peers):
|
|
continue
|
|
group = peers
|
|
visited.update(p.id for p in group)
|
|
if len(group) > slots - len(claimed_rows):
|
|
continue
|
|
nes = {str(p.ne_id or "").strip() for p in group}
|
|
if nes & busy_nes:
|
|
continue
|
|
now = _utcnow()
|
|
for member in group:
|
|
member.status = "running"
|
|
member.message = ""
|
|
if member.collect_round_id:
|
|
member.started_at = now
|
|
task = db.get(BizStateTask, member.task_id)
|
|
if task:
|
|
task.last_collect_started_at = now
|
|
claimed_rows.append(member)
|
|
busy_nes.update(nes - {""})
|
|
|
|
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:
|
|
last = claimed_rows[-1]
|
|
demoted = [b for b in claimed_rows if b.collect_round_id == last.collect_round_id] if last.collect_round_id else [last]
|
|
for demote in demoted:
|
|
claimed_rows.remove(demote)
|
|
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())
|