netx/tests/test_biz_state_collect_recovery.py
oliver 3a0de0d7fb 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>
2026-09-22 17:00:37 +08:00

108 lines
3.4 KiB
Python

"""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")
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()