From 12b60c2a5234753fb2e26a32931487877855992e Mon Sep 17 00:00:00 2001 From: oliver Date: Tue, 22 Sep 2026 15:44:16 +0800 Subject: [PATCH] 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 --- netx_api/app_startup.py | 23 +++--- netx_api/biz_state/collect_recovery.py | 76 +++++++++++++++++++ tests/test_biz_state_collect_recovery.py | 93 ++++++++++++++++++++++++ 3 files changed, 179 insertions(+), 13 deletions(-) create mode 100644 netx_api/biz_state/collect_recovery.py create mode 100644 tests/test_biz_state_collect_recovery.py diff --git a/netx_api/app_startup.py b/netx_api/app_startup.py index bd7c397..ea16cd6 100644 --- a/netx_api/app_startup.py +++ b/netx_api/app_startup.py @@ -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 diff --git a/netx_api/biz_state/collect_recovery.py b/netx_api/biz_state/collect_recovery.py new file mode 100644 index 0000000..c2e6add --- /dev/null +++ b/netx_api/biz_state/collect_recovery.py @@ -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} diff --git a/tests/test_biz_state_collect_recovery.py b/tests/test_biz_state_collect_recovery.py new file mode 100644 index 0000000..bac608a --- /dev/null +++ b/tests/test_biz_state_collect_recovery.py @@ -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()