feat(cutover): pair metric collection rounds and isolate manual routes

This commit is contained in:
oliver 2026-10-11 08:46:46 +08:00
parent 32d7d969f6
commit d3618deff6
14 changed files with 809 additions and 154 deletions

View file

@ -138,10 +138,13 @@ class EvidencePersistTests(unittest.TestCase):
return t
def _mk_batch(self, bid: str, task_id: str) -> BizStateBatch:
task = self.db.get(BizStateTask, task_id)
now = _utcnow()
b = BizStateBatch(
id=bid,
task_id=task_id,
collect_group_id=task.collect_group_id or "",
collect_round_id="test-" + task.collect_group_id if task.collect_group_id else "",
status="success",
started_at=now,
ended_at=now,

View file

@ -47,6 +47,7 @@ class HfDecoupleFlowTests(unittest.TestCase):
)
Base.metadata.create_all(bind=engine)
self.db = TestingSession()
self.Session = TestingSession
self.mt = BizMonitorTemplate(
id="mt1",
@ -94,10 +95,13 @@ class HfDecoupleFlowTests(unittest.TestCase):
return t
def _mk_batch(self, bid: str, task_id: str) -> BizStateBatch:
task = self.db.get(BizStateTask, task_id)
now = datetime.utcnow()
b = BizStateBatch(
id=bid,
task_id=task_id,
collect_group_id=task.collect_group_id or "",
collect_round_id="test-" + task.collect_group_id if task.collect_group_id else "",
status="success",
started_at=now,
ended_at=now,
@ -140,6 +144,64 @@ class HfDecoupleFlowTests(unittest.TestCase):
self.assertNotEqual(out["old_hf_task_id"], "old_p")
self.assertNotEqual(out["new_hf_task_id"], "new_p")
def test_same_interval_metrics_are_separate_paired_tasks(self) -> None:
out = self._create_project(collect_metric_ids=["interface_brief", "arp"])
old, new = out["old_hf_bindings"], out["new_hf_bindings"]
self.assertEqual(len(old), 2)
self.assertEqual({b["interval_sec"] for b in old}, {60})
for o, n in zip(old, new):
self.assertEqual(len(o["metric_ids"]), 1)
self.assertTrue(o["collect_group_id"])
self.assertEqual(o["collect_group_id"], n["collect_group_id"])
self.assertEqual(o["metric_ids"], n["metric_ids"])
def test_dynamic_routes_manual_only_and_removal_cancels_queued_pairs(self) -> None:
from netx_api.models import BizStateTaskItem, BizStateTaskItemBinding
for tid in ("old_p", "new_p"):
self.db.add(BizStateTaskItem(id="route-item-" + tid, task_id=tid, source_profile_id="zte.ip_route_vrf", kind="catalog", enabled=True))
self.db.add(BizStateTaskItemBinding(id="route-bind-" + tid, item_id="route-item-" + tid, placeholder="vrf", value="customer-A"))
self.db.commit()
out = self._create_project()
with mock.patch.object(biz_svc, "_ne_meta", side_effect=_stub_ne_meta):
updated = mig.patch_project(self.db, out["id"], {"collect_metric_ids": ["interface_brief", "ip_route"]})
routes = [b for side in ("old", "new") for b in updated[f"{side}_hf_bindings"] if "ip_route" in b["metric_ids"]]
self.assertEqual(len(routes), 2)
for idx, binding in enumerate(routes):
task = self.db.get(BizStateTask, binding["task_id"])
self.assertTrue(task.collect_manual_only)
self.assertEqual(task.status, "paused")
task.collect_running = True
self.db.add(BizStateBatch(id=f"queued-route-{idx}", task_id=task.id, status="queued", collect_round_id="route-round"))
self.db.commit()
with mock.patch("netx_api.biz_state.collect_stop.SessionLocal", self.Session), mock.patch.object(biz_svc, "_ne_meta", side_effect=_stub_ne_meta):
shrunk = mig.patch_project(self.db, out["id"], {"collect_metric_ids": ["interface_brief"]})
self.db.expire_all()
self.assertEqual(len(shrunk["old_hf_bindings"]), 1)
for idx in range(2):
self.assertEqual(self.db.get(BizStateBatch, f"queued-route-{idx}").status, "cancelled")
def test_route_profile_preserves_portrait_cli_variant_and_bindings(self) -> None:
from netx_api.models import BizStateTaskItem, BizStateTaskItemBinding
self.db.add(BizStateTaskItem(id="vrf-item", task_id="old_p", source_profile_id="zte.ip_route_vrf", kind="catalog", enabled=True))
self.db.add(BizStateTaskItemBinding(id="vrf-bind", item_id="vrf-item", placeholder="vrf", value="customer-A"))
self.db.commit()
out = self._create_project(collect_metric_ids=["interface_brief", "ip_route"])
binding = next(b for b in out["old_hf_bindings"] if b["metric_ids"] == ["ip_route"])
item = self.db.query(BizStateTaskItem).filter(BizStateTaskItem.task_id == binding["task_id"]).one()
self.assertEqual(item.source_profile_id, "zte.ip_route_vrf")
copied = self.db.query(BizStateTaskItemBinding).filter(BizStateTaskItemBinding.item_id == item.id).one()
self.assertEqual((copied.placeholder, copied.value), ("vrf", "customer-A"))
def test_invalid_new_metric_does_not_save_broken_selection(self) -> None:
out = self._create_project()
with self.assertRaises(Exception):
mig.patch_project(self.db, out["id"], {"collect_metric_ids": ["interface_brief", "not-a-metric"]})
self.db.rollback()
current = self.db.get(BizMigrationProject, out["id"])
self.assertEqual(mig.resolve_collect_metric_ids(self.db, current), ["interface_brief"])
def test_ensure_highfreq_does_not_overwrite_portrait(self) -> None:
proj = self._create_project()
res = self._ensure(proj["id"], interval_sec=60)
@ -213,7 +275,7 @@ class HfDecoupleFlowTests(unittest.TestCase):
called = {c.args[0] for c in disp.call_args_list}
self.assertEqual(
called,
{proj["old_hf_task_id"], proj["new_hf_task_id"]},
{proj["old_hf_task_id"]}, # one dispatch atomically enqueues the pair
)
self.assertNotIn("old_p", called)

View file

@ -112,6 +112,97 @@ def test_collected_empty_baseline_can_pass_acceptance(db):
assert finished["run"]["summary"]["config_snapshot"]["expect_set"]
def paired_samples(db):
for tid in ("oh", "nh"):
db.get(BizStateTask, tid).collect_group_id = "pair"
for bid in ("oc", "nc"):
sample = db.get(BizStateBatch, bid)
sample.collect_group_id = "pair"
sample.collect_round_id = "round"
sample.queued_at = sample.started_at
db.commit()
@pytest.mark.parametrize("problem", ["running", "queued", "failed", "round_mismatch"])
def test_paired_samples_cannot_mix_rounds_or_use_old_success(db, problem):
paired_samples(db)
if problem == "round_mismatch":
db.get(BizStateBatch, "nc").collect_round_id = "another-round"
else:
db.get(BizStateBatch, "nc").status = problem
db.commit()
run = svc.run_evaluate(db, batch_id="wave")
assert not run["summary"]["coverage_complete"]
assert not run["summary"]["anomaly"] # incomplete pair is unknown, not service loss
def test_paired_samples_expose_duration_queue_and_real_start_skew(db):
paired_samples(db)
for bid in ("oc", "nc"):
sample = db.get(BizStateBatch, bid)
sample.queued_at = sample.started_at - timedelta(seconds=15)
sample.started_at -= timedelta(seconds=5)
db.commit()
card = svc.run_evaluate(db, batch_id="wave")["summary"]["sheet_cards"][0]
assert card["collect_round_id"] == "round"
assert card["old_queue_sec"] == 10 and card["old_duration_sec"] == 5
assert card["start_skew_sec"] < 1
def test_unselected_route_does_not_block_regular_acceptance(db):
ct = db.get(BizCompareTemplate, "ct")
ct.metrics_json = [*ct.metrics_json, {"metric_id": "ip_route", "sheet_id": "route", "key_fields": ["prefix"]}]
db.commit()
assert svc.finish_batch(db, "wave")["accept_summary"]["passed"]
def test_incremental_evaluation_reuses_route_rows_and_paged_evidence(db):
from unittest.mock import patch
ct = db.get(BizCompareTemplate, "ct")
ct.metrics_json = [*ct.metrics_json, {"metric_id": "ip_route", "sheet_id": "route", "key_fields": ["prefix"], "compare_fields": ["nexthop"]}]
db.get(BizMonitorTemplate, "mt").collect_metric_ids_json = ["arp", "ip_route"]
# Legacy combined samples also exercise compatibility with existing projects.
for bid in ("ob", "nb", "oc", "nc"):
db.add(BizStateBatchCommand(id="route-cmd-" + bid, batch_id=bid, metric_id="ip_route", parse_status="ok"))
for i in range(2000):
db.add(BizStateMetricRow(id=f"route-{i}", batch_id="ob", metric_id="ip_route", data_json={"prefix": str(i), "nexthop": "old"}))
db.commit()
first = svc.run_evaluate(db, batch_id="wave", purpose="auto")
with patch.object(svc, "_load_metric_rows", wraps=svc._load_metric_rows) as load:
second = svc.run_evaluate(db, batch_id="wave", purpose="auto", _metric_ids={"arp"})
assert all(call.kwargs["metric_id"] == "arp" for call in load.call_args_list)
assert svc.list_run_diffs(db, second["id"], sheet_id="route", color="red", limit=100)["total"] == 2000
assert db.query(BizMigrationDiff).filter(BizMigrationDiff.run_id == second["id"]).count() == 1
assert second["summary"]["anomaly"] == first["summary"]["anomaly"]
assert svc.list_run_diffs(db, second["id"], only_expect=True)["total"] == 1
# Template correction invalidates cached verdicts and evidence.
ct.metrics_json = [ct.metrics_json[0], {**ct.metrics_json[1], "row_filters": [{"field": "prefix", "op": "eq", "value": "1"}]}]
db.commit()
third = svc.run_evaluate(db, batch_id="wave", purpose="auto", _metric_ids={"arp"})
assert svc.list_run_diffs(db, third["id"], sheet_id="route")["total"] == 1
def test_baseline_picker_never_loads_route_inventory(db):
from unittest.mock import patch
ct = db.get(BizCompareTemplate, "ct")
ct.metrics_json = [*ct.metrics_json, {"metric_id": "ip_route", "sheet_id": "route", "key_fields": ["prefix"]}]
db.get(BizMonitorTemplate, "mt").collect_metric_ids_json = ["arp", "ip_route"]
db.commit()
with patch.object(svc, "_load_metric_rows", wraps=svc._load_metric_rows) as load:
sheets = svc.list_baseline_expect_objects(db, "p")["sheets"]
assert {s["metric_id"] for s in sheets} == {"arp"}
assert all(c.kwargs["metric_id"] == "arp" for c in load.call_args_list)
def test_legacy_auto_monitor_refreshes_changed_samples_without_explicit_metric_bindings(db):
assert try_auto_monitor_for_task(db, "nh") == 1
db.get(BizStateBatch, "nc").status = "failed"
db.commit()
assert try_auto_monitor_for_task(db, "nh") == 1
last = db.query(BizMigrationRun).order_by(BizMigrationRun.created_at.desc()).first()
assert not last.summary_json["coverage_complete"]
@pytest.mark.parametrize("problem", ["missing_receipt", "failed_receipt", "stale", "skew", "duplicate", "skipped", "unknown_scope"])
def test_uncertain_coverage_cannot_pass_acceptance(db, problem):
if problem == "missing_receipt":

View file

@ -0,0 +1,180 @@
"""Paired admission/dispatch, actual cadence, and per-metric isolation."""
from concurrent.futures import ThreadPoolExecutor
from datetime import timedelta
from threading import Barrier, Event
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool
from netx_api.db import Base
from netx_api.models import BizStateTask, BizStateTaskItem, BizStateBatch
from netx_api.timeutil import utcnow_naive
from netx_api.biz_state import claim
from netx_api import biz_state_scheduler as scheduler
@pytest.fixture
def db(monkeypatch):
engine = create_engine("sqlite://", connect_args={"check_same_thread": False}, poolclass=StaticPool)
Base.metadata.create_all(engine)
factory = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
monkeypatch.setattr(claim, "SessionLocal", factory)
monkeypatch.setattr(scheduler, "SessionLocal", factory)
with factory() as session:
for tid in ("old", "new"):
session.add(BizStateTask(id=tid, source="managed", ne_id=tid,
status="running", interval_sec=60, collect_group_id="pair"))
session.add(BizStateTaskItem(id="item-" + tid, task_id=tid, source_profile_id="zte.arp", enabled=True))
session.commit()
yield session
scheduler._in_flight.clear()
def test_individual_trigger_enqueues_both_with_one_identity(db):
result = claim.enqueue_collect("new", manual=True)
assert result["queued"] and result["task_id"] == "new"
rows = db.query(BizStateBatch).all()
assert len(rows) == 2
assert {b.task_id for b in rows} == {"old", "new"}
assert len({b.collect_round_id for b in rows}) == 1
assert len({b.queued_at for b in rows}) == 1
assert not claim.enqueue_collect("old", manual=True)["queued"]
@pytest.mark.parametrize("problem", ["busy", "no_items", "paused", "same_device"])
def test_pair_admission_never_creates_half_round(db, problem):
peer = db.get(BizStateTask, "new")
if problem == "busy":
peer.collect_running = True
elif problem == "no_items":
db.delete(db.get(BizStateTaskItem, "item-new"))
elif problem == "paused":
peer.status = "paused"
else:
peer.ne_id = "old"
db.commit()
assert not claim.enqueue_collect("old")["queued"]
assert db.query(BizStateBatch).count() == 0
def test_pair_claim_requires_two_slots(db):
claim.enqueue_collect("old")
assert claim.claim_queued_batches(1) == []
assert len(claim.claim_queued_batches(2)) == 2
def test_busy_device_keeps_both_queued_and_no_ssh_overlap(db):
db.add(BizStateBatch(id="portrait", task_id="other", ne_id="old", status="running"))
db.commit()
claim.enqueue_collect("old")
assert claim.claim_queued_batches(8) == []
assert db.query(BizStateBatch).filter(BizStateBatch.status == "queued").count() == 2
def test_capacity_one_has_actionable_admission_error(db, monkeypatch):
monkeypatch.setattr(claim, "max_concurrent_tasks", lambda: 1)
assert claim.enqueue_collect("old")["reason"] == "collect_pair_capacity"
assert db.query(BizStateBatch).count() == 0
def test_manual_routes_are_excluded_from_periodic_scheduler(db):
for t in db.query(BizStateTask).all():
t.collect_manual_only = True
db.commit()
assert scheduler._enqueue_due_tasks() == 0
assert claim.enqueue_collect("old", manual=True)["queued"]
def test_pair_cadence_uses_later_start_and_slow_peer_blocks_restart(db):
now = utcnow_naive()
db.get(BizStateTask, "old").last_collect_started_at = now - timedelta(seconds=100)
db.get(BizStateTask, "new").last_collect_started_at = now - timedelta(seconds=10)
db.commit()
assert scheduler._enqueue_due_tasks() == 0
db.get(BizStateTask, "new").last_collect_started_at = now - timedelta(seconds=100)
db.get(BizStateTask, "new").collect_running = True
db.commit()
assert scheduler._enqueue_due_tasks() == 0
db.get(BizStateTask, "new").collect_running = False
db.commit()
assert scheduler._enqueue_due_tasks() == 1
assert db.query(BizStateBatch).count() == 2
def test_inline_pair_uses_two_workers_without_waiting_for_one_side(db, monkeypatch):
from netx_api.biz_state import collect_runner
both_started, release = Barrier(2), Event()
started, synchronized = [], []
def fake_collect(**job):
started.append(job["task_id"])
both_started.wait(timeout=3)
synchronized.append(job["task_id"])
release.wait(timeout=3)
monkeypatch.setattr(scheduler, "execute_claimed_batch", fake_collect)
monkeypatch.setattr(collect_runner, "_should_execute_inline", lambda: True)
with ThreadPoolExecutor(max_workers=2) as pool:
monkeypatch.setattr(scheduler, "_dispatch_pool_get", lambda: pool)
result = collect_runner.dispatch_collect("old", manual=True)
assert result["queued"]
# Both worker entries must reach a barrier; serial old→new would fail it.
release.set()
assert set(started) == {"old", "new"}
assert set(synchronized) == {"old", "new"}
def test_faster_metrics_claim_before_manual_route_round(db):
for tid in ("ro", "rn"):
db.add(BizStateTask(id=tid, ne_id="old" if tid == "ro" else "new",
status="paused", collect_group_id="routes", collect_manual_only=True))
db.add(BizStateTaskItem(id="i-" + tid, task_id=tid, source_profile_id="zte.ip_route", enabled=True))
db.commit()
claim.enqueue_collect("ro", manual=True)
claim.enqueue_collect("old")
jobs = claim.claim_queued_batches(4)
assert {job["task_id"] for job in jobs} == {"old", "new"}
def test_manual_route_round_cannot_starve_behind_continuous_fast_metrics(db):
for tid in ("ro", "rn"):
db.add(BizStateTask(id=tid, ne_id="old" if tid == "ro" else "new",
status="paused", collect_group_id="routes", collect_manual_only=True))
db.add(BizStateTaskItem(id="i-" + tid, task_id=tid, source_profile_id="zte.ip_route", enabled=True))
db.commit()
claim.enqueue_collect("ro", manual=True)
claim.enqueue_collect("old")
for batch in db.query(BizStateBatch).filter(BizStateBatch.collect_group_id == "routes").all():
batch.queued_at = utcnow_naive() - timedelta(seconds=301)
db.commit()
assert {job["task_id"] for job in claim.claim_queued_batches(4)} == {"ro", "rn"}
@pytest.mark.parametrize("running", [False, True])
def test_stopping_one_side_stops_the_whole_round(db, monkeypatch, running):
from netx_api.biz_state import collect_stop
monkeypatch.setattr(collect_stop, "SessionLocal", claim.SessionLocal)
claim.enqueue_collect("old")
if running:
claim.claim_queued_batches(2)
result = collect_stop.request_stop_collect("new")
assert result["ok"] and result["stopped"]
db.expire_all()
batches = db.query(BizStateBatch).all()
try:
if running:
assert len(result["signaled_batches"]) == 2
assert all(b.message == collect_stop.STOP_REQUEST_TOKEN for b in batches)
assert all(t.collect_running for t in db.query(BizStateTask).all())
else:
assert len(result["cancelled_batches"]) == 2
assert all(b.status == "cancelled" for b in batches)
assert not any(t.collect_running for t in db.query(BizStateTask).all())
assert claim.claim_queued_batches(2) == []
assert claim.enqueue_collect("new")["queued"]
finally:
for batch in batches:
collect_stop.clear_stop_requested(batch.id)