mirror of
https://github.com/hansjone/netx.git
synced 2026-10-08 23:33:21 +08:00
Mark interrupted biz-state batches as partial on API restart.
Clear stuck collect_running without re-dispatching; finalize in-flight commands as failed with interrupted_by_restart. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
586ca69d3f
commit
12b60c2a52
3 changed files with 179 additions and 13 deletions
|
|
@ -144,21 +144,18 @@ def run_api_startup() -> None:
|
|||
if pt_cleared:
|
||||
_log.info("startup: cleared %s port_traffic stuck collect_running flag(s)", pt_cleared)
|
||||
try:
|
||||
from .models import BizStateTask
|
||||
from .biz_state.collect_recovery import recover_interrupted_collects_on_startup
|
||||
|
||||
stuck = (
|
||||
db.query(BizStateTask)
|
||||
.filter(BizStateTask.collect_running.is_(True))
|
||||
.all()
|
||||
)
|
||||
for t in stuck:
|
||||
t.collect_running = False
|
||||
t.last_error = (t.last_error or "")[:900] + " | reset_on_startup"
|
||||
if stuck:
|
||||
db.commit()
|
||||
_log.info("startup: cleared %s biz_state stuck collect_running flag(s)", len(stuck))
|
||||
rec = recover_interrupted_collects_on_startup(db)
|
||||
if rec.get("batches") or rec.get("tasks"):
|
||||
_log.info(
|
||||
"startup: biz_state interrupted recover batches=%s cmds=%s tasks=%s",
|
||||
rec.get("batches"),
|
||||
rec.get("commands"),
|
||||
rec.get("tasks"),
|
||||
)
|
||||
except Exception:
|
||||
_log.exception("startup: biz_state collect_running recovery failed")
|
||||
_log.exception("startup: biz_state collect recovery failed")
|
||||
try:
|
||||
from .port_traffic_migrate import backfill_port_traffic_series
|
||||
|
||||
|
|
|
|||
76
netx_api/biz_state/collect_recovery.py
Normal file
76
netx_api/biz_state/collect_recovery.py
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
"""Recover biz-state collects interrupted by process restart."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models import BizStateBatch, BizStateBatchCommand, BizStateTask
|
||||
from ..timeutil import utcnow_naive
|
||||
|
||||
_log = logging.getLogger("netx.biz_state.recovery")
|
||||
|
||||
_INTERRUPT_MARK = "interrupted_by_restart"
|
||||
|
||||
|
||||
def recover_interrupted_collects_on_startup(db: Session) -> dict[str, Any]:
|
||||
"""Mark orphaned running batches as partial; clear stuck collect_running.
|
||||
|
||||
Does **not** resume or re-dispatch collects — process crash mid-batch leaves
|
||||
whatever rows/commands were already persisted.
|
||||
"""
|
||||
now = utcnow_naive()
|
||||
batch_n = 0
|
||||
cmd_n = 0
|
||||
task_n = 0
|
||||
|
||||
stuck_batches = (
|
||||
db.query(BizStateBatch).filter(BizStateBatch.status == "running").all()
|
||||
)
|
||||
for b in stuck_batches:
|
||||
b.status = "partial"
|
||||
if not b.ended_at:
|
||||
b.ended_at = now
|
||||
msg = str(b.message or "").strip()
|
||||
if _INTERRUPT_MARK not in msg:
|
||||
b.message = f"{msg} | {_INTERRUPT_MARK}".strip(" |")[:1020]
|
||||
batch_n += 1
|
||||
|
||||
running_cmds = (
|
||||
db.query(BizStateBatchCommand)
|
||||
.filter(
|
||||
BizStateBatchCommand.batch_id == b.id,
|
||||
BizStateBatchCommand.parse_status == "running",
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for c in running_cmds:
|
||||
c.parse_status = "failed"
|
||||
cmsg = str(c.message or "").strip()
|
||||
if _INTERRUPT_MARK not in cmsg:
|
||||
c.message = f"{cmsg} | {_INTERRUPT_MARK}".strip(" |")[:1020]
|
||||
cmd_n += 1
|
||||
|
||||
stuck_tasks = (
|
||||
db.query(BizStateTask).filter(BizStateTask.collect_running.is_(True)).all()
|
||||
)
|
||||
for t in stuck_tasks:
|
||||
t.collect_running = False
|
||||
if not t.last_collect_ended_at:
|
||||
t.last_collect_ended_at = now
|
||||
err = str(t.last_error or "").strip()
|
||||
if _INTERRUPT_MARK not in err:
|
||||
t.last_error = f"{err} | {_INTERRUPT_MARK}".strip(" |")[:1020]
|
||||
task_n += 1
|
||||
|
||||
if batch_n or task_n or cmd_n:
|
||||
db.commit()
|
||||
_log.info(
|
||||
"startup: biz_state recover interrupted batches=%s cmds=%s tasks=%s",
|
||||
batch_n,
|
||||
cmd_n,
|
||||
task_n,
|
||||
)
|
||||
return {"batches": batch_n, "commands": cmd_n, "tasks": task_n}
|
||||
93
tests/test_biz_state_collect_recovery.py
Normal file
93
tests/test_biz_state_collect_recovery.py
Normal file
|
|
@ -0,0 +1,93 @@
|
|||
"""Startup recovery for interrupted biz-state collects."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from netx_api.biz_state.collect_recovery import recover_interrupted_collects_on_startup
|
||||
from netx_api.db import Base
|
||||
from netx_api.models import BizStateBatch, BizStateBatchCommand, BizStateTask
|
||||
|
||||
|
||||
class BizStateCollectRecoveryTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:", future=True)
|
||||
TestingSession = sessionmaker(
|
||||
bind=engine, autoflush=False, autocommit=False, expire_on_commit=False
|
||||
)
|
||||
Base.metadata.create_all(bind=engine)
|
||||
self.db = TestingSession()
|
||||
self.task = BizStateTask(
|
||||
id="t1",
|
||||
source="managed",
|
||||
ne_id="ne1",
|
||||
ne_name="PE1",
|
||||
status="running",
|
||||
collect_running=True,
|
||||
interval_sec=300,
|
||||
)
|
||||
self.db.add(self.task)
|
||||
self.batch = BizStateBatch(
|
||||
id="b1",
|
||||
task_id="t1",
|
||||
status="running",
|
||||
command_count=1,
|
||||
row_count=2,
|
||||
message="",
|
||||
)
|
||||
self.db.add(self.batch)
|
||||
self.db.add(
|
||||
BizStateBatchCommand(
|
||||
id="c1",
|
||||
batch_id="b1",
|
||||
profile_id="zte.lldp_neighbors",
|
||||
parser_id="lldp_neighbors",
|
||||
metric_id="lldp_neighbor",
|
||||
raw_command="show lldp neighbor brief",
|
||||
parse_status="running",
|
||||
row_count=0,
|
||||
)
|
||||
)
|
||||
self.db.add(
|
||||
BizStateBatch(
|
||||
id="b_ok",
|
||||
task_id="t1",
|
||||
status="success",
|
||||
command_count=1,
|
||||
row_count=1,
|
||||
)
|
||||
)
|
||||
self.db.commit()
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self.db.close()
|
||||
|
||||
def test_running_batch_becomes_partial_no_resume(self) -> None:
|
||||
out = recover_interrupted_collects_on_startup(self.db)
|
||||
self.assertEqual(out["batches"], 1)
|
||||
self.assertEqual(out["tasks"], 1)
|
||||
self.assertEqual(out["commands"], 1)
|
||||
|
||||
self.db.refresh(self.batch)
|
||||
self.db.refresh(self.task)
|
||||
self.assertEqual(self.batch.status, "partial")
|
||||
self.assertIsNotNone(self.batch.ended_at)
|
||||
self.assertIn("interrupted_by_restart", self.batch.message)
|
||||
self.assertFalse(self.task.collect_running)
|
||||
self.assertIn("interrupted_by_restart", self.task.last_error or "")
|
||||
|
||||
cmd = self.db.get(BizStateBatchCommand, "c1")
|
||||
assert cmd is not None
|
||||
self.assertEqual(cmd.parse_status, "failed")
|
||||
self.assertIn("interrupted_by_restart", cmd.message or "")
|
||||
|
||||
ok = self.db.get(BizStateBatch, "b_ok")
|
||||
assert ok is not None
|
||||
self.assertEqual(ok.status, "success")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Loading…
Add table
Add a link
Reference in a new issue