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:
oliver 2026-09-22 17:00:37 +08:00
parent 0eb65a839e
commit 3a0de0d7fb
22 changed files with 1210 additions and 165 deletions

View file

@ -0,0 +1,125 @@
"""biz_state enqueue / claim (NE mutex) tests."""
from __future__ import annotations
import unittest
from unittest.mock import patch
from uuid import uuid4
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool
from netx_api.biz_state import claim as claim_mod
from netx_api.db import Base
from netx_api.models import BizStateBatch, BizStateTask, BizStateTaskItem
class BizStateClaimTests(unittest.TestCase):
def setUp(self) -> None:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
future=True,
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
TestingSession = sessionmaker(
bind=engine, autoflush=False, autocommit=False, expire_on_commit=False
)
Base.metadata.create_all(bind=engine)
self.Session = TestingSession
self.db = TestingSession()
self._session_patch = patch.object(claim_mod, "SessionLocal", TestingSession)
self._session_patch.start()
self.task_a = BizStateTask(
id="ta",
source="managed",
ne_id="ne-a",
ne_name="A",
status="running",
collect_running=False,
interval_sec=300,
)
self.task_b = BizStateTask(
id="tb",
source="managed",
ne_id="ne-a", # same NE as A
ne_name="A2",
status="running",
collect_running=False,
interval_sec=300,
)
self.task_c = BizStateTask(
id="tc",
source="managed",
ne_id="ne-c",
ne_name="C",
status="running",
collect_running=False,
interval_sec=300,
)
self.db.add_all([self.task_a, self.task_b, self.task_c])
for tid in ("ta", "tb", "tc"):
self.db.add(
BizStateTaskItem(
id=uuid4().hex,
task_id=tid,
source_profile_id="zte.arp",
kind="catalog",
enabled=True,
)
)
self.db.commit()
def tearDown(self) -> None:
self._session_patch.stop()
self.db.close()
def test_enqueue_sets_queued_batch(self) -> None:
r = claim_mod.enqueue_collect("ta", manual=True)
self.assertTrue(r.get("queued"), r)
self.db.expire_all()
task = self.db.get(BizStateTask, "ta")
assert task is not None
self.assertTrue(task.collect_running)
batch = self.db.get(BizStateBatch, r["batch_id"])
assert batch is not None
self.assertEqual(batch.status, "queued")
def test_enqueue_idempotent_while_running(self) -> None:
r1 = claim_mod.enqueue_collect("ta", manual=True)
self.assertTrue(r1.get("queued"), r1)
r2 = claim_mod.enqueue_collect("ta", manual=True)
self.assertFalse(r2.get("queued"))
self.assertEqual(r2.get("reason"), "already_collecting")
def test_claim_ne_mutex_skips_same_ne(self) -> None:
with patch.object(claim_mod, "max_concurrent_tasks", return_value=10):
ra = claim_mod.enqueue_collect("ta", manual=True)
rb = claim_mod.enqueue_collect("tb", manual=True)
rc = claim_mod.enqueue_collect("tc", manual=True)
self.assertTrue(ra["queued"] and rb["queued"] and rc["queued"])
claimed = claim_mod.claim_queued_batches(10)
ids = {c["batch_id"] for c in claimed}
# ta and tc can run; tb same NE as ta must wait
self.assertIn(ra["batch_id"], ids)
self.assertIn(rc["batch_id"], ids)
self.assertNotIn(rb["batch_id"], ids)
self.db.expire_all()
self.assertEqual(self.db.get(BizStateBatch, ra["batch_id"]).status, "running")
self.assertEqual(self.db.get(BizStateBatch, rb["batch_id"]).status, "queued")
self.assertEqual(self.db.get(BizStateBatch, rc["batch_id"]).status, "running")
def test_claim_respects_slot_ceiling(self) -> None:
with patch.object(claim_mod, "max_concurrent_tasks", return_value=1):
claim_mod.enqueue_collect("ta", manual=True)
claim_mod.enqueue_collect("tc", manual=True)
first = claim_mod.claim_queued_batches(10)
self.assertEqual(len(first), 1)
second = claim_mod.claim_queued_batches(10)
self.assertEqual(len(second), 0)
if __name__ == "__main__":
unittest.main()

View file

@ -88,6 +88,21 @@ class BizStateCollectRecoveryTests(unittest.TestCase):
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()

View file

@ -0,0 +1,51 @@
"""Regression: dual-lane absorb must not inflate row_count from cumulative DB."""
from __future__ import annotations
import unittest
from netx_api.biz_state.collect_runner import _absorb_lane_result
class DualLaneAbsorbTests(unittest.TestCase):
def test_absorb_sums_lane_deltas_not_cumulative(self) -> None:
# Each lane should report only its own delta (0 when persist owns rows).
lane_errors: list[str] = []
total_rows, cmd_count, any_fail, any_ok = _absorb_lane_result(
(0, 10, False, True),
total_rows=0,
cmd_count=0,
any_fail=False,
any_ok=False,
lane_errors=lane_errors,
)
total_rows, cmd_count, any_fail, any_ok = _absorb_lane_result(
(0, 5, False, True),
total_rows=total_rows,
cmd_count=cmd_count,
any_fail=any_fail,
any_ok=any_ok,
lane_errors=lane_errors,
)
self.assertEqual(total_rows, 0)
self.assertEqual(cmd_count, 15)
self.assertTrue(any_ok)
self.assertFalse(any_fail)
def test_legacy_bug_pattern_would_inflate(self) -> None:
# Document what NOT to do: feeding cumulative DB totals into absorb.
lane_errors: list[str] = []
# If light returns cumulative 100 and heavy returns cumulative 150...
bad_rows, _, _, _ = _absorb_lane_result(
(150, 5, False, True),
total_rows=100,
cmd_count=10,
any_fail=False,
any_ok=True,
lane_errors=lane_errors,
)
self.assertEqual(bad_rows, 250) # inflated — lanes must return deltas only
if __name__ == "__main__":
unittest.main()

View file

@ -53,23 +53,44 @@ class ManualCollectTests(unittest.TestCase):
dc.assert_called_once_with("t1", manual=True)
def test_dispatch_manual_allows_paused(self) -> None:
task = MagicMock()
task.collect_running = False
task.status = "paused"
task.id = "t1"
task.source = "managed"
task.ne_id = "n1"
task.ne_name = "NE"
task.vendor = "zte"
db = MagicMock()
db.get.return_value = task
# items query → empty so it returns early after setting error
q = MagicMock()
q.filter.return_value.order_by.return_value.all.return_value = []
db.query.return_value = q
with patch("netx_api.biz_state.collect_runner.SessionLocal", return_value=db):
with patch(
"netx_api.biz_state.claim.enqueue_collect",
return_value={
"ok": False,
"queued": False,
"reason": "no_enabled_items",
"task_id": "t1",
},
) as enq:
dispatch_collect("t1", manual=True)
self.assertEqual(task.last_error, "no enabled task items")
enq.assert_called_once_with("t1", manual=True)
def test_dispatch_inline_skips_if_already_claimed(self) -> None:
with (
patch(
"netx_api.biz_state.claim.enqueue_collect",
return_value={
"ok": True,
"queued": True,
"batch_id": "b1",
"task_id": "t1",
},
),
patch(
"netx_api.biz_state.collect_runner._should_execute_inline",
return_value=True,
),
patch(
"netx_api.biz_state.collect_runner._try_claim_batch_for_execute",
return_value=False,
) as claim,
patch(
"netx_api.biz_state.collect_runner.execute_claimed_batch"
) as exe,
):
dispatch_collect("t1", manual=True)
claim.assert_called_once_with("b1")
exe.assert_not_called()
if __name__ == "__main__":