From d3618deff664eb0b3aa8582f84ec8c78312599d2 Mon Sep 17 00:00:00 2001 From: oliver Date: Sun, 11 Oct 2026 08:46:46 +0800 Subject: [PATCH] feat(cutover): pair metric collection rounds and isolate manual routes --- netx_api/biz_migration/auto_monitor.py | 14 +- netx_api/biz_migration/monitor_templates.py | 3 + netx_api/biz_migration/service.py | 361 ++++++++++++++------ netx_api/biz_state/claim.py | 144 ++++++-- netx_api/biz_state/collect_runner.py | 11 + netx_api/biz_state/collect_stop.py | 47 ++- netx_api/biz_state/schema_ensure.py | 9 + netx_api/biz_state/service.py | 4 + netx_api/biz_state_scheduler.py | 25 +- netx_api/models/biz_state.py | 7 + tests/test_biz_migration_evidence.py | 3 + tests/test_biz_migration_hf_decouple.py | 64 +++- tests/test_biz_migration_integrity.py | 91 +++++ tests/test_biz_state_paired_sampling.py | 180 ++++++++++ 14 files changed, 809 insertions(+), 154 deletions(-) create mode 100644 tests/test_biz_state_paired_sampling.py diff --git a/netx_api/biz_migration/auto_monitor.py b/netx_api/biz_migration/auto_monitor.py index e037472..4cd8174 100644 --- a/netx_api/biz_migration/auto_monitor.py +++ b/netx_api/biz_migration/auto_monitor.py @@ -65,10 +65,20 @@ def try_auto_monitor_for_task(db: Session, task_id: str) -> int: if last and (last.summary_json or {}).get("auto_fingerprint") == fingerprint: db.rollback() continue + previous_samples = {s[0]: (s[1], s[2]) for s in (last.summary_json or {}).get("auto_samples", [])} if last else {} + if last and not previous_samples: + for sample in (last.summary_json or {}).get("metric_batches", {}).values(): + for side in ("old", "new"): + previous_samples[sample.get(f"{side}_task_id", "")] = ( + sample.get(f"{side}_batch_id", ""), sample.get(f"{side}_status", "success")) + changed_tasks = {tid for tid, bid, status in samples if previous_samples.get(tid) != (bid, status)} + changed_metrics = {mid for side in ("old", "new") for binding in service._hf_bindings(proj, side) + if binding["task_id"] in changed_tasks for mid in binding.get("metric_ids") or payload[-1]} try: - result = service.run_evaluate(db, batch_id=batch.id, purpose="auto", _commit=False) + result = service.run_evaluate(db, batch_id=batch.id, purpose="auto", _commit=False, + _metric_ids=changed_metrics) run = db.get(BizMigrationRun, result["id"]) - run.summary_json = {**run.summary_json, "auto_fingerprint": fingerprint} + run.summary_json = {**run.summary_json, "auto_fingerprint": fingerprint, "auto_samples": samples} db.commit() count += 1 except Exception: diff --git a/netx_api/biz_migration/monitor_templates.py b/netx_api/biz_migration/monitor_templates.py index cee7ccf..9ebc420 100644 --- a/netx_api/biz_migration/monitor_templates.py +++ b/netx_api/biz_migration/monitor_templates.py @@ -34,6 +34,8 @@ def _out(row: BizMonitorTemplate, compare_name: str = "", *, db: Session | None effective = seen if not effective: effective = [PORT_METRIC_ID] + ct = db.get(BizCompareTemplate, row.compare_template_id) if db and row.compare_template_id else None + available = list(dict.fromkeys(str(s.get("metric_id") or "") for s in cmp_svc.template_metrics(ct))) if ct else effective out: dict[str, Any] = { "id": row.id, "name": row.name, @@ -41,6 +43,7 @@ def _out(row: BizMonitorTemplate, compare_name: str = "", *, db: Session | None "compare_template_name": compare_name, "collect_metric_ids": collect, "collect_metric_ids_effective": effective, + "available_metric_ids": [m for m in available if m], "defaults": dict(row.defaults_json or {}), "sheet_overrides": list(row.sheet_overrides_json or []), "note": row.note or "", diff --git a/netx_api/biz_migration/service.py b/netx_api/biz_migration/service.py index 8138df8..d895809 100644 --- a/netx_api/biz_migration/service.py +++ b/netx_api/biz_migration/service.py @@ -2,12 +2,14 @@ from __future__ import annotations +import hashlib +from collections import Counter from datetime import datetime, timezone from typing import Any from uuid import uuid4 from fastapi import HTTPException -from sqlalchemy import insert +from sqlalchemy import and_, insert, or_ from sqlalchemy.orm import Session from ..biz_state.compare_rules import apply_row_filters, scalar_text @@ -42,6 +44,7 @@ from .evaluate import ( ) PURPOSE_CUTOVER_HF = "cutover_hf" +MANUAL_ROUTE_METRICS = frozenset({"ip_route", "ipv6_route", "bgp_route"}) def _parse_dt(raw: Any) -> datetime | None: @@ -217,7 +220,9 @@ def _hf_bindings(proj: BizMigrationProject, side: str) -> list[dict[str, Any]]: iv = max(60, int(row.get("interval_sec") or 60)) except (TypeError, ValueError): iv = 60 - out.append({"task_id": tid, "metric_ids": mids, "interval_sec": iv}) + out.append({"task_id": tid, "metric_ids": mids, "interval_sec": iv, + "collect_group_id": str(row.get("collect_group_id") or ""), + "manual_only": bool(row.get("manual_only"))}) if out: return out # Legacy single slot @@ -296,11 +301,45 @@ def _metric_sample_ok(db: Session, batch: BizStateBatch, metric_id: str) -> bool return bool(statuses) and all(s in ("ok", "ok_aux") for s in statuses) +def _current_pair_for_metric(db: Session, proj: BizMigrationProject, mid: str, + pin_old: str = "", pin_new: str = "") -> tuple[BizStateBatch | None, BizStateBatch | None, dict[str, Any]]: + old_tid, new_tid = (_hf_task_id_for_metric(proj, side, mid) for side in ("old", "new")) + old_task = db.get(BizStateTask, old_tid) if old_tid else None + new_task = db.get(BizStateTask, new_tid) if new_tid else None + gid = old_task.collect_group_id if old_task else "" + meta: dict[str, Any] = {"collect_group_id": gid, "manual_only": bool(old_task and old_task.collect_manual_only)} + if gid and not (pin_old or pin_new): + if not new_task or new_task.collect_group_id != gid: + return None, None, {**meta, "sample_issues": ["collect_pair_invalid"]} + latest = db.query(BizStateBatch).filter(BizStateBatch.collect_group_id == gid, + BizStateBatch.task_id.in_((old_tid, new_tid))).order_by( + BizStateBatch.queued_at.desc(), BizStateBatch.started_at.desc(), BizStateBatch.id.desc()).first() + peers = db.query(BizStateBatch).filter(BizStateBatch.collect_round_id == latest.collect_round_id, + BizStateBatch.task_id.in_((old_tid, new_tid))).all() if latest and latest.collect_round_id else [] + meta.update({"collect_round_id": latest.collect_round_id if latest else "", + "old_collect_status": next((b.status for b in peers if b.task_id == old_tid), "missing"), + "new_collect_status": next((b.status for b in peers if b.task_id == new_tid), "missing")}) + if len(peers) != 2: + return None, None, {**meta, "sample_issues": ["collect_pair_missing"]} + if any(b.status in ("running", "queued") for b in peers): + return None, None, {**meta, "sample_issues": ["collect_pair_pending"]} + by_task = {b.task_id: b for b in peers} + if not all(_metric_sample_ok(db, b, mid) for b in peers): + return None, None, {**meta, "sample_issues": ["collect_pair_failed"]} + return by_task[old_tid], by_task[new_tid], meta + old = _current_batch_for_metric(db, proj, "old", mid, pinned_batch_id=pin_old) + new = _current_batch_for_metric(db, proj, "new", mid, pinned_batch_id=pin_new) + if old and new and (old.collect_round_id or new.collect_round_id) and old.collect_round_id != new.collect_round_id: + return None, None, {**meta, "sample_issues": ["collect_round_mismatch"]} + return old, new, meta + + def _sample_timing(proj: BizMigrationProject, mid: str, old: BizStateBatch, new: BizStateBatch, *, check_age: bool) -> dict[str, Any]: interval = max(60, _metric_interval_map(proj).get(mid) or proj.hf_interval_sec or 60) old_at, new_at = old.ended_at or old.started_at, new.ended_at or new.started_at - age_limit, skew_limit = interval * 3, interval * 2 + durations = [max(0, ((b.ended_at or b.started_at) - b.started_at).total_seconds()) for b in (old, new)] + age_limit, skew_limit = interval * 3 + max(durations), interval * 2 reasons: list[str] = [] if not old_at or not new_at: reasons.append("sample_time_missing") @@ -312,9 +351,15 @@ def _sample_timing(proj: BizMigrationProject, mid: str, old: BizStateBatch, if check_age and max((utcnow_naive() - old_at).total_seconds(), (utcnow_naive() - new_at).total_seconds()) > age_limit: reasons.append("sample_stale") + start_skew = abs((old.started_at - new.started_at).total_seconds()) + if old.collect_round_id and start_skew > min(30, interval / 2): + reasons.append("sample_start_skew") return {"old_collected_at": _dt_iso(old_at), "new_collected_at": _dt_iso(new_at), "age_limit_sec": age_limit, "skew_limit_sec": skew_limit, - "skew_sec": skew, "sample_issues": reasons} + "skew_sec": skew, "start_skew_sec": start_skew, "sample_issues": reasons, + "old_duration_sec": durations[0], "new_duration_sec": durations[1], + "old_queue_sec": max(0, (old.started_at - (old.queued_at or old.started_at)).total_seconds()), + "new_queue_sec": max(0, (new.started_at - (new.queued_at or new.started_at)).total_seconds())} def find_portrait_task_for_ne(db: Session, *, source: str, ne_id: str) -> BizStateTask | None: @@ -762,6 +807,19 @@ def patch_project(db: Session, project_id: str, body: dict[str, Any]) -> dict[st p.hf_start_at = _parse_dt(body.get("hf_start_at")) if "hf_end_at" in body: p.hf_end_at = _parse_dt(body.get("hf_end_at")) + sampling_changed = any(k in body for k in ("collect_metric_ids", "metric_interval_sec", "hf_interval_sec", "monitor_template_id")) + has_hf = bool(_all_hf_task_ids(p, "old") and _all_hf_task_ids(p, "new")) + if sampling_changed and has_hf: + # Validate newly selected commands before saving the project selection. + for side in ("old", "new"): + portrait = db.get(BizStateTask, portrait_task_id(p, side)) + existing = db.get(BizStateTask, hf_task_id(p, side)) + fields = _template_from_task_or_ne(db, portrait=portrait, existing_hf=existing, ne=None) + for mid in resolve_collect_metric_ids(db, p): + task = db.get(BizStateTask, _hf_task_id_for_metric(p, side, mid)) + if (task and mid in _enabled_metric_ids(db, task.id)) or _portrait_metric_items(db, portrait, mid): + continue + _catalog_item_for_metric(vendor=fields["vendor"], device_type=fields["device_type"], metric_id=mid) p.updated_at = utcnow_naive() db.commit() db.refresh(p) @@ -769,6 +827,9 @@ def patch_project(db: Session, project_id: str, body: dict[str, Any]) -> dict[st if any(k in body for k in ("status", "hf_start_at", "hf_end_at")): _apply_hf_window_to_tasks(db, p) db.refresh(p) + if sampling_changed and has_hf: + ensure_highfreq(db, p.id, collect_now=False) + db.refresh(p) return project_to_dict(db, p) @@ -870,6 +931,7 @@ def run_evaluate( acceptance: bool = False, purpose: str = "manual", _commit: bool = True, + _metric_ids: set[str] | None = None, ) -> dict[str, Any]: mb = db.get(BizMigrationBatch, batch_id) if not mb: @@ -904,12 +966,14 @@ def run_evaluate( # Primary current batch ids (for run record / pin UX); per-metric may differ when multi-HF if not old_batch: for tid in _all_hf_task_ids(proj, "old") or [current_task_id(proj, "old")]: - old_batch = _latest_success_batch(db, tid) + old_batch = db.query(BizStateBatch).filter(BizStateBatch.task_id == tid).order_by( + BizStateBatch.started_at.desc(), BizStateBatch.id.desc()).first() if old_batch: break if not new_batch: for tid in _all_hf_task_ids(proj, "new") or [current_task_id(proj, "new")]: - new_batch = _latest_success_batch(db, tid) + new_batch = db.query(BizStateBatch).filter(BizStateBatch.task_id == tid).order_by( + BizStateBatch.started_at.desc(), BizStateBatch.id.desc()).first() if new_batch: break if not old_batch: @@ -971,12 +1035,26 @@ def run_evaluate( all_rows: list[dict[str, Any]] = [] metric_batches: dict[str, dict[str, str]] = {} seq = 0 - verdict_counts: dict[str, int] = {} rows_cache: dict[tuple[str, str], list[dict[str, Any]]] = {} command_cache: dict[str, dict[str, Any]] = {} - scope_ids = {str(s.get("metric_id") or "") for s in sheets} | {sheet_key(s) for s in sheets} | {"_ports"} + enabled_sheets = [s for s in sheets if str(s.get("metric_id") or "") in collect_ids and not + override_for_sheet(sheet_overrides, sheet_id=sheet_key(s), metric_id=s.get("metric_id", "")).get("skip_dual")] + scope_ids = {str(s.get("metric_id") or "") for s in enabled_sheets} | {sheet_key(s) for s in enabled_sheets} + if any(s.get("iface_fields") for s in enabled_sheets): + scope_ids.add("_ports") uncovered_scope = sorted({k for scope in (expect, previous_expect) for k, values in scope.items() if values and k not in scope_ids}) + config_snapshot = {"sheets": sheets, "sheet_overrides": sheet_overrides, + "defaults": defaults, "iface_normalize": iface_norm, "port_map": port_map, + "expect_set": mb.expect_set_json, "previous_expect": {k: sorted(v) for k, v in previous_expect.items()}, + "baseline_ids": [proj.old_baseline_batch_id, proj.new_baseline_batch_id], + "collect_metric_ids": sorted(collect_ids), "window_active": window_active, + "bindings": [_hf_bindings(proj, "old"), _hf_bindings(proj, "new")]} + prior_run = db.query(BizMigrationRun).filter(BizMigrationRun.batch_id == batch_id).order_by( + BizMigrationRun.created_at.desc()).first() if _metric_ids is not None else None + reusable = (prior_run.summary_json or {}).get("config_snapshot") == config_snapshot if prior_run else False + prior_cards = {sheet_key(c): c for c in (prior_run.summary_json or {}).get("sheet_cards", [])} if reusable else {} + pair_cache: dict[str, tuple[BizStateBatch | None, BizStateBatch | None, dict[str, Any]]] = {} def metric_rows(bid: str, mid: str) -> list[dict[str, Any]]: key = (bid, mid) @@ -1021,8 +1099,17 @@ def run_evaluate( ) continue - old_cur_batch = _current_batch_for_metric(db, proj, "old", mid, pinned_batch_id=pin_old) - new_cur_batch = _current_batch_for_metric(db, proj, "new", mid, pinned_batch_id=pin_new) + if _metric_ids is not None and mid not in _metric_ids and sid in prior_cards: + cached = {**prior_cards[sid], "source_run_id": prior_cards[sid].get("source_run_id") or prior_run.id} + times = [_parse_dt(cached.get(k)) for k in ("old_collected_at", "new_collected_at")] + if any(at and (utcnow_naive() - at).total_seconds() > cached.get("age_limit_sec", 180) for at in times): + cached.update({"current_missing": True, "sample_issues": ["sample_stale"]}) + sheet_cards.append(cached) + metric_batches[mid] = (prior_run.summary_json or {}).get("metric_batches", {}).get(mid, {}) + continue + if mid not in pair_cache: + pair_cache[mid] = _current_pair_for_metric(db, proj, mid, pin_old, pin_new) + old_cur_batch, new_cur_batch, pair_meta = pair_cache[mid] # Do NOT fall back to another HF task's batch — empty/wrong metric → false red if not old_cur_batch or not new_cur_batch: sheet_cards.append( @@ -1041,13 +1128,14 @@ def run_evaluate( "current_missing": True, "collect_incomplete": True, "collect_skipped": False, + **pair_meta, } ) continue old_cur_mid = old_cur_batch.id new_cur_mid = new_cur_batch.id - timing = _sample_timing(proj, mid, old_cur_batch, new_cur_batch, - check_age=acceptance or not (pin_old and pin_new)) + timing = {**pair_meta, **_sample_timing(proj, mid, old_cur_batch, new_cur_batch, + check_age=acceptance or not (pin_old and pin_new))} baseline_missing = [side for side in ("old", "new") if not _metric_sample_ok( db, db.get(BizStateBatch, getattr(proj, f"{side}_baseline_batch_id")), mid)] if timing["sample_issues"] or baseline_missing: @@ -1114,6 +1202,8 @@ def run_evaluate( "new_batch_id": new_cur_mid, "old_task_id": old_tid, "new_task_id": new_tid, + "old_status": old_cur_batch.status, + "new_status": new_cur_batch.status, **timing, } sheet_cards.append( @@ -1138,6 +1228,8 @@ def run_evaluate( "duplicate_keys_after": int(one.get("duplicate_keys_after") or 0), **timing, "awaiting_peer": awaiting, + "verdict_counts": dict(Counter(r["verdict"] for r in one["rows"] if r.get("verdict"))), + "expected_count": sum(bool(r.get("in_expect")) for r in one["rows"]), } ) for r in one["rows"]: @@ -1154,10 +1246,12 @@ def run_evaluate( r["sheet_id"] = sid seq += 1 all_rows.append(r) - v = str(r.get("verdict") or "") - if v: - verdict_counts[v] = verdict_counts.get(v, 0) + 1 + verdict_counts = {} + for card in sheet_cards: + if not card.get("current_missing"): + for verdict, count in card.get("verdict_counts", {}).items(): + verdict_counts[verdict] = verdict_counts.get(verdict, 0) + count active_cards = [ c for c in sheet_cards if not c.get("collect_skipped") and not c.get("current_missing") ] @@ -1202,11 +1296,8 @@ def run_evaluate( "steady_outside": sum(int(c.get("steady_outside") or 0) for c in active_cards), "duplicate_keys": sum(int(c.get("duplicate_keys_before") or 0) + int(c.get("duplicate_keys_after") or 0) for c in active_cards), "uncovered_scope": uncovered_scope, - "coverage_complete": bool(sheet_cards) and not uncovered_scope and not any(c.get("collect_skipped") or c.get("current_missing") or c.get("awaiting_peer") or c.get("duplicate_keys_before") or c.get("duplicate_keys_after") for c in sheet_cards), - "config_snapshot": {"sheets": sheets, "sheet_overrides": sheet_overrides, - "defaults": defaults, "iface_normalize": iface_norm, - "port_map": port_map, "expect_set": mb.expect_set_json, - "previous_expect": {k: sorted(v) for k, v in previous_expect.items()}}, + "coverage_complete": bool(enabled_sheets) and not uncovered_scope and not any(c.get("current_missing") or c.get("awaiting_peer") or c.get("duplicate_keys_before") or c.get("duplicate_keys_after") for c in sheet_cards if not c.get("collect_skipped")), + "config_snapshot": config_snapshot, "window_active": window_active, "expect_ports": sorted(expect.get("_ports") or ()), "old_current_task_id": current_task_id(proj, "old"), @@ -1216,6 +1307,9 @@ def run_evaluate( ) db.add(run) db.flush() + for card in sheet_cards: + card.setdefault("source_run_id", run.id) + run.summary_json = {**run.summary_json, "sheet_cards": sheet_cards} pending_diffs: list[dict[str, Any]] = [] for r in all_rows: key_list = r.get("key") or [] @@ -1468,6 +1562,9 @@ def list_baseline_expect_objects(db: Session, project_id: str) -> dict[str, Any] # Only offer expect keys for metrics that are actually HF-collected if collect_ids and mid not in collect_ids: continue + if mid in MANUAL_ROUTE_METRICS: + # Never enumerate million-row route tables into batch scope pickers. + continue sheet_ov = override_for_sheet(sheet_overrides, sheet_id=sid, metric_id=mid) if sheet_ov.get("skip_dual"): continue @@ -1526,20 +1623,18 @@ def _catalog_item_for_metric(*, vendor: str, device_type: str, metric_id: str) - mid = str(metric_id or "").strip() vkey = resolve_vendor_key(vendor, device_type) - for p in profiles_for_vendor(vkey): - if p.metric_id == mid and p.kind == "collect": - if getattr(p, "placeholders", None): - ph = ",".join(str(x.name) for x in (p.placeholders or []) if getattr(x, "name", None)) - raise HTTPException( - status_code=400, - detail=f"metric_needs_bindings:{mid}:{ph or 'required'}", - ) + candidates = [p for p in profiles_for_vendor(vkey) if p.metric_id == mid and p.kind == "collect"] + for p in sorted(candidates, key=lambda p: bool(p.placeholders)): + if not any(getattr(ph, "required", True) for ph in (p.placeholders or [])): return { "source_profile_id": p.profile_id, "kind": "catalog", "enabled": True, "title": p.title or mid, } + if candidates: + ph = ",".join(x.name for x in candidates[0].placeholders if getattr(x, "required", True)) + raise HTTPException(status_code=400, detail=f"metric_needs_bindings:{mid}:{ph}") raise HTTPException( status_code=400, detail=f"no_profile_for_metric:{mid}:{vkey or vendor or 'unknown'}", @@ -1643,6 +1738,26 @@ def _template_from_task_or_ne( raise HTTPException(status_code=400, detail="old_new_ne_or_task_required") +def _portrait_metric_items(db: Session, portrait: BizStateTask | None, mid: str) -> list[dict[str, Any]]: + """Preserve CLI variants and bindings already used by the device baseline.""" + from ..biz_state.profiles import get_profile + from ..models import BizStateTaskItem, BizStateTaskItemBinding + + if not portrait: + return [] + out = [] + for item in db.query(BizStateTaskItem).filter(BizStateTaskItem.task_id == portrait.id, + BizStateTaskItem.enabled.is_(True)).order_by(BizStateTaskItem.sort_order).all(): + profile = get_profile(item.source_profile_id) if item.kind == "catalog" else None + if not profile or profile.metric_id != mid: + continue + bindings = db.query(BizStateTaskItemBinding).filter(BizStateTaskItemBinding.item_id == item.id).all() + out.append({"source_profile_id": item.source_profile_id, "kind": item.kind, + "enabled": True, "title": item.title, "command_override": item.command_override, + "bindings": [{"placeholder": b.placeholder, "value": b.value} for b in bindings]}) + return out + + def _ensure_side_highfreq( db: Session, *, @@ -1653,6 +1768,7 @@ def _ensure_side_highfreq( retention_days: int, metric_ids: list[str], status: str = "running", + configured_items: list[dict[str, Any]] | None = None, ) -> tuple[BizStateTask, bool]: """Return (hf_task, created). Reuse existing HF if metrics match; never mutate portrait.""" from ..biz_state import service as biz_svc @@ -1660,6 +1776,9 @@ def _ensure_side_highfreq( want = {str(m).strip() for m in metric_ids if str(m).strip()} if not want: want = {PORT_METRIC_ID} + if existing_hf and existing_hf.collect_running and _enabled_metric_ids(db, existing_hf.id) != want: + # A running collection owns its item configuration until completion. + existing_hf = None if existing_hf and _is_highfreq_task(db, existing_hf, want): patch: dict[str, Any] = { @@ -1676,7 +1795,7 @@ def _ensure_side_highfreq( # Existing HF (purpose already cutover_hf) with wrong metrics → update items in place. # Never rewrite a task that only "looks like" HF by note — create a sibling instead. if existing_hf and str(getattr(existing_hf, "purpose", None) or "").strip() == PURPOSE_CUTOVER_HF: - items = [ + items = configured_items or [ _catalog_item_for_metric( vendor=str(template_fields.get("vendor") or existing_hf.vendor), device_type=str(template_fields.get("device_type") or existing_hf.device_type), @@ -1700,7 +1819,7 @@ def _ensure_side_highfreq( refreshed = db.get(BizStateTask, existing_hf.id) return refreshed or existing_hf, False - items = [ + items = configured_items or [ _catalog_item_for_metric( vendor=str(template_fields.get("vendor") or ""), device_type=str(template_fields.get("device_type") or ""), @@ -1754,11 +1873,12 @@ def _apply_hf_window_to_tasks(db: Session, proj: BizMigrationProject) -> None: from ..biz_state import service as biz_svc log = logging.getLogger("netx.biz_migration.hf_window") - want = _hf_window_status(proj) + window_status = _hf_window_status(proj) for tid in _all_hf_task_ids(proj, "old") + _all_hf_task_ids(proj, "new"): task = db.get(BizStateTask, tid) if not task: continue + want = "paused" if task.collect_manual_only else window_status purpose = str(getattr(task, "purpose", None) or "").strip() if purpose and purpose != PURPOSE_CUTOVER_HF: continue @@ -1806,9 +1926,7 @@ def ensure_highfreq( old_ne: dict[str, Any] | None = None, new_ne: dict[str, Any] | None = None, ) -> dict[str, Any]: - """Create/bind HF tasks (one per interval group) — never overwrites portrait ids.""" - from ..biz_state.collect_runner import dispatch_collect - + """One task per metric/device, paired by project; route tables are manual-only.""" proj = db.get(BizMigrationProject, project_id) if not proj: raise HTTPException(status_code=404, detail="project_not_found") @@ -1822,33 +1940,28 @@ def ensure_highfreq( if not metric_ids: metric_ids = [PORT_METRIC_ID] by_metric = _metric_interval_map(proj) - groups = _group_metrics_by_interval(metric_ids, default_interval=iv_default, by_metric=by_metric) + groups = [(max(60, by_metric.get(mid) or iv_default), [mid]) for mid in metric_ids] want_status = _hf_window_status(proj) old_portrait = db.get(BizStateTask, proj.old_task_id) if proj.old_task_id else None new_portrait = db.get(BizStateTask, proj.new_task_id) if proj.new_task_id else None - # Reuse existing bindings by interval when possible - def _existing_by_interval(side: str) -> dict[int, BizStateTask]: - out: dict[int, BizStateTask] = {} - for b in _hf_bindings(proj, side): - tid = str(b.get("task_id") or "").strip() - t = db.get(BizStateTask, tid) if tid else None - if not t: - continue - try: - iv = max(60, int(b.get("interval_sec") or t.interval_sec or 60)) - except (TypeError, ValueError): - iv = 60 - out[iv] = t - # legacy primary slot - primary = db.get(BizStateTask, hf_task_id(proj, side)) if hf_task_id(proj, side) else None - if primary and iv_default not in out: - out[iv_default] = primary + # Reuse a single-metric binding; a legacy combined task can be assigned + # once, while the remaining metrics receive independent sibling tasks. + def _existing_by_metric(side: str) -> dict[str, BizStateTask]: + out: dict[str, BizStateTask] = {} + for binding in _hf_bindings(proj, side): + task = db.get(BizStateTask, binding["task_id"]) + if task: + mids = binding.get("metric_ids") or metric_ids + for mid in mids: + if mid in metric_ids: + out.setdefault(mid, task) + break return out - old_existing = _existing_by_interval("old") - new_existing = _existing_by_interval("new") + old_existing = _existing_by_metric("old") + new_existing = _existing_by_metric("new") old_prev_ids = set(_all_hf_task_ids(proj, "old")) new_prev_ids = set(_all_hf_task_ids(proj, "new")) @@ -1871,33 +1984,51 @@ def ensure_highfreq( new_any_created = False for g_iv, g_metrics in groups: + mid = g_metrics[0] + gid = "cutover:" + proj.id + ":" + hashlib.sha256(mid.encode()).hexdigest()[:16] + old_candidate, new_candidate = old_existing.get(mid), new_existing.get(mid) + if old_candidate and old_candidate.collect_group_id not in ("", gid): + old_candidate = None + if new_candidate and new_candidate.collect_group_id not in ("", gid): + new_candidate = None old_task, old_created = _ensure_side_highfreq( db, - existing_hf=old_existing.get(g_iv), + existing_hf=old_candidate, template_fields=old_fields, project_name=proj.name, interval_sec=g_iv, retention_days=ret, metric_ids=g_metrics, - status=want_status, + status="paused", + configured_items=_portrait_metric_items(db, old_portrait, mid), ) new_task, new_created = _ensure_side_highfreq( db, - existing_hf=new_existing.get(g_iv), + existing_hf=new_candidate, template_fields=new_fields, project_name=proj.name, interval_sec=g_iv, retention_days=ret, metric_ids=g_metrics, - status=want_status, + status="paused", + configured_items=_portrait_metric_items(db, new_portrait, mid), ) + if old_task.ne_id == new_task.ne_id: + raise HTTPException(status_code=400, detail="collect_pair_same_device") + for task in (old_task, new_task): + task.collect_group_id = gid + task.collect_manual_only = mid in MANUAL_ROUTE_METRICS + task.status = "paused" if task.collect_manual_only else want_status + db.commit() old_any_created = old_any_created or old_created new_any_created = new_any_created or new_created old_bindings.append( - {"task_id": old_task.id, "metric_ids": list(g_metrics), "interval_sec": g_iv} + {"task_id": old_task.id, "metric_ids": list(g_metrics), "interval_sec": g_iv, + "collect_group_id": gid, "manual_only": mid in MANUAL_ROUTE_METRICS} ) new_bindings.append( - {"task_id": new_task.id, "metric_ids": list(g_metrics), "interval_sec": g_iv} + {"task_id": new_task.id, "metric_ids": list(g_metrics), "interval_sec": g_iv, + "collect_group_id": gid, "manual_only": mid in MANUAL_ROUTE_METRICS} ) proj = db.get(BizMigrationProject, project_id) @@ -1921,27 +2052,24 @@ def ensure_highfreq( continue if str(getattr(t, "purpose", None) or "").strip() not in ("", PURPOSE_CUTOVER_HF): continue + if t.collect_group_id and not t.collect_group_id.startswith("cutover:" + proj.id + ":"): + continue + t.collect_group_id = "" + t.collect_manual_only = False + db.commit() if t.status == "running": try: biz_svc.update_task(db, tid, {"status": "paused"}) except Exception: # noqa: BLE001 pass + if t.collect_running: + from ..biz_state.collect_stop import request_stop_collect + + request_stop_collect(tid) collect: dict[str, Any] = {"old": [], "new": []} - if collect_now and want_status == "running": - for side, binds in (("old", old_bindings), ("new", new_bindings)): - for b in binds: - tid = b["task_id"] - try: - dispatch_collect(tid) - collect[side].append({"ok": True, "task_id": tid}) - except Exception as exc: # noqa: BLE001 - collect[side].append({"ok": False, "task_id": tid, "error": str(exc)[:200]}) - elif collect_now and want_status != "running": - collect = { - "old": [{"ok": False, "error": "hf_window_inactive"}], - "new": [{"ok": False, "error": "hf_window_inactive"}], - } + if collect_now: + collect = collect_project_now(db, project_id) return { "project": project_to_dict(db, proj), @@ -1984,26 +2112,33 @@ def collect_project_now(db: Session, project_id: str) -> dict[str, Any]: } out: dict[str, Any] = {"old": [], "new": [], "hf_status": want_status} - for side in ("old", "new"): - tids = _all_hf_task_ids(proj, side) - if not tids: - out[side] = [{"ok": False, "error": "hf_task_missing"}] + old_bindings, new_bindings = _hf_bindings(proj, "old"), _hf_bindings(proj, "new") + handled: set[str] = set() + for old_binding in old_bindings: + gid = old_binding.get("collect_group_id") + if not gid or gid in handled: continue - for tid in tids: - task = db.get(BizStateTask, tid) - if not task: - out[side].append({"ok": False, "error": "task_not_found", "task_id": tid}) - continue - if bool(task.collect_running): - out[side].append( - {"ok": False, "error": "collect_already_running", "task_id": tid} - ) - continue - try: - dispatch_collect(tid) - out[side].append({"ok": True, "task_id": tid}) - except Exception as exc: # noqa: BLE001 - out[side].append({"ok": False, "task_id": tid, "error": str(exc)[:200]}) + handled.add(gid) + new_binding = next((b for b in new_bindings if b.get("collect_group_id") == gid), None) + if not new_binding: + result = {"ok": False, "error": "collect_pair_invalid"} + else: + result = dispatch_collect(old_binding["task_id"], manual=True) + result = {**result, "ok": bool(result.get("queued")), "error": result.get("reason", "")} + for side, binding in (("old", old_binding), ("new", new_binding)): + out[side].append({**result, "task_id": binding["task_id"] if binding else "", + "metric_ids": old_binding.get("metric_ids", [])}) + # Legacy projects remain operable until their bindings are upgraded. + if not handled: + for side in ("old", "new"): + tids = _all_hf_task_ids(proj, side) + if not tids: + out[side] = [{"ok": False, "error": "hf_task_missing"}] + for tid in tids: + result = dispatch_collect(tid, manual=True) + out[side].append({**result, "ok": bool(result.get("queued")), "task_id": tid, + "error": result.get("reason", "")}) + return out @@ -2352,9 +2487,17 @@ def list_run_diffs( offset: int = 0, limit: int = 100, ) -> dict[str, Any]: - if not db.get(BizMigrationRun, run_id): + run = db.get(BizMigrationRun, run_id) + if not run: raise HTTPException(status_code=404, detail="run_not_found") - q = db.query(BizMigrationDiff).filter(BizMigrationDiff.run_id == run_id) + sources = [and_(BizMigrationDiff.run_id == c["source_run_id"], + BizMigrationDiff.key_json["sheet_id"].as_string() == sheet_key(c)) + for c in (run.summary_json or {}).get("sheet_cards", []) + if c.get("source_run_id") and c["source_run_id"] != run_id + and not (only_expect and c.get("expected_count") == 0) + and (not metric_id or c.get("metric_id") == metric_id) + and (not sheet_id or sheet_key(c) == sheet_id)] + q = db.query(BizMigrationDiff).filter(or_(BizMigrationDiff.run_id == run_id, *sources)) if metric_id: q = q.filter(BizMigrationDiff.metric_id == metric_id) if sheet_id: @@ -2369,7 +2512,7 @@ def list_run_diffs( literal = kw.strip().replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") q = q.filter(BizMigrationDiff.search_text.ilike(f"%{literal}%", escape="\\")) total = q.count() - rows = q.order_by(BizMigrationDiff.seq.asc()).offset(max(0, offset)).limit(min(500, max(1, limit))).all() + rows = q.order_by(BizMigrationDiff.metric_id.asc(), BizMigrationDiff.seq.asc(), BizMigrationDiff.id.asc()).offset(max(0, offset)).limit(min(500, max(1, limit))).all() return { "total": total, "offset": max(0, offset), @@ -2411,9 +2554,35 @@ def board(db: Session, batch_id: str, run_id: str = "") -> dict[str, Any]: "project": proj, "batch": batch_to_dict(mb), "run": run_to_dict(db, run) if run else None, + "sampling": sampling_state(db, db.get(BizMigrationProject, mb.project_id)) if not run_id else [], } +def sampling_state(db: Session, proj: BizMigrationProject) -> list[dict[str, Any]]: + """Small live read independent of evaluation and route row counts.""" + out = [] + for binding in _hf_bindings(proj, "old"): + gid = binding.get("collect_group_id") + if not gid: + continue + latest = db.query(BizStateBatch).filter(BizStateBatch.collect_group_id == gid).order_by( + BizStateBatch.queued_at.desc(), BizStateBatch.started_at.desc()).first() + peers = db.query(BizStateBatch).filter(BizStateBatch.collect_round_id == latest.collect_round_id).all() if latest else [] + old = next((b for b in peers if b.task_id == binding["task_id"]), None) + new = next((b for b in peers if b.task_id != binding["task_id"]), None) + card: dict[str, Any] = {"metric_id": (binding.get("metric_ids") or [""])[0], + "collect_group_id": gid, "collect_round_id": latest.collect_round_id if latest else "", + "manual_only": binding.get("manual_only", False), "interval_sec": binding["interval_sec"]} + for side, batch in (("old", old), ("new", new)): + card[f"{side}_collect_status"] = batch.status if batch else "idle" + card[f"{side}_duration_sec"] = max(0, ((batch.ended_at or utcnow_naive()) - batch.started_at).total_seconds()) if batch and batch.status != "queued" else 0 + card[f"{side}_queue_sec"] = max(0, ((utcnow_naive() if batch.status == "queued" else batch.started_at) - (batch.queued_at or batch.started_at)).total_seconds()) if batch else 0 + if old and new and all(b.status != "queued" for b in (old, new)): + card["start_skew_sec"] = abs((old.started_at - new.started_at).total_seconds()) + out.append(card) + return out + + def get_monitor_context( db: Session, *, diff --git a/netx_api/biz_state/claim.py b/netx_api/biz_state/claim.py index a86674e..2a592c0 100644 --- a/netx_api/biz_state/claim.py +++ b/netx_api/biz_state/claim.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging -from datetime import datetime +from datetime import datetime, timedelta from typing import Any from uuid import uuid4 @@ -39,6 +39,10 @@ def enqueue_collect(task_id: str, *, manual: bool = False) -> dict[str, Any]: task = db.get(BizStateTask, tid) if not task: return {"ok": False, "queued": False, "reason": "task_not_found", "task_id": tid} + if task.collect_group_id: + group_id = task.collect_group_id + db.close() + return enqueue_collect_group(group_id, manual=manual, requested_task_id=tid) if str(task.source or "").strip().lower() == "import": return { "ok": False, @@ -99,6 +103,7 @@ def enqueue_collect(task_id: str, *, manual: bool = False) -> dict[str, Any]: vendor=task.vendor, status="queued", started_at=now, + queued_at=now, message="queued" if not manual else "queued_manual", ) db.add(batch) @@ -121,6 +126,61 @@ def enqueue_collect(task_id: str, *, manual: bool = False) -> dict[str, Any]: db.close() +def enqueue_collect_group(group_id: str, *, manual: bool = False, + requested_task_id: str = "") -> dict[str, Any]: + """Atomically enqueue both devices with one round identity, or neither.""" + db = SessionLocal() + try: + tasks = db.query(BizStateTask).filter(BizStateTask.collect_group_id == group_id).order_by( + BizStateTask.id).with_for_update().populate_existing().all() + result: dict[str, Any] = {"ok": False, "queued": False, "task_id": requested_task_id, + "collect_group_id": group_id} + if len(tasks) != 2 or len({t.ne_id for t in tasks}) != 2: + return {**result, "reason": "collect_pair_invalid"} + if max_concurrent_tasks() < 2: + return {**result, "reason": "collect_pair_capacity"} + from ..cli_budget import clamp_cli_workers + + worker_cap = clamp_cli_workers(int(getattr(settings, "biz_state_worker_collect_threads", None) + or getattr(settings, "biz_state_dispatch_workers", 8) or 8)) + if worker_cap < 2: + return {**result, "reason": "collect_pair_capacity"} + if any(t.collect_running for t in tasks): + return {**result, "reason": "collect_pair_busy"} + for task in tasks: + if task.source == "import" or (not manual and (task.status != "running" or task.collect_manual_only)): + return {**result, "reason": "not_scheduled"} + if task.status in ("", "deleted"): + return {**result, "reason": "bad_status"} + if not db.query(BizStateTaskItem.id).filter(BizStateTaskItem.task_id == task.id, + BizStateTaskItem.enabled.is_(True)).first(): + return {**result, "reason": "no_enabled_items"} + now, round_id = _utcnow(), uuid4().hex + batches = [] + for task in tasks: + task.collect_running = True + task.collect_queued_at = now + task.last_error = "" + task.updated_at = now + batch = BizStateBatch(id=uuid4().hex, task_id=task.id, source=task.source, + ne_id=task.ne_id, ne_name=task.ne_name, vendor=task.vendor, status="queued", + started_at=now, queued_at=now, collect_group_id=group_id, + collect_round_id=round_id, collect_priority=10 if task.collect_manual_only else 20, + message="queued_manual" if manual else "queued") + db.add(batch) + batches.append({"batch_id": batch.id, "task_id": task.id}) + db.commit() + selected = next((b for b in batches if b["task_id"] == requested_task_id), batches[0]) + return {**result, **selected, "ok": True, "queued": True, "collect_round_id": round_id, + "batches": batches, "manual": manual} + except Exception: + db.rollback() + _log.exception("enqueue collect pair failed group=%s", group_id) + return {"ok": False, "queued": False, "reason": "enqueue_error", "task_id": requested_task_id} + finally: + db.close() + + def _dialect_supports_skip_locked(db: Session) -> bool: try: bind = db.get_bind() @@ -151,11 +211,16 @@ def claim_queued_batches(limit: int | None = None) -> list[dict[str, Any]]: Global ceiling is always ``max_concurrent_tasks()``; ``limit`` only caps how many this caller wants in one call. """ - from sqlalchemy import text + from sqlalchemy import and_, case, text global_cap = max_concurrent_tasks() db = SessionLocal() try: + pg = _dialect_supports_skip_locked(db) + if pg: + # A short transaction lock makes capacity reservation and pair claims + # atomic across worker processes. Never held during device I/O. + db.execute(text("SELECT pg_advisory_xact_lock(hashtext('biz-state-claim-capacity'))")) running_n = ( db.query(BizStateBatch).filter(BizStateBatch.status == "running").count() ) @@ -173,50 +238,54 @@ def claim_queued_batches(limit: int | None = None) -> list[dict[str, Any]]: if str(r[0] or "").strip() } + # Manual route rounds yield to fresh monitoring, but must eventually + # run even when a busy device always has regular metrics queued. + priority = case((and_(BizStateBatch.collect_priority == 10, + BizStateBatch.queued_at <= _utcnow() - timedelta(seconds=300)), 30), + else_=BizStateBatch.collect_priority) q = ( db.query(BizStateBatch) .filter(BizStateBatch.status == "queued") - .order_by(BizStateBatch.started_at.asc()) + .order_by(priority.desc(), BizStateBatch.started_at.asc(), BizStateBatch.id.asc()) ) - pg = _dialect_supports_skip_locked(db) if pg: q = q.with_for_update(skip_locked=True) else: q = q.with_for_update() - candidates = q.limit(max(slots * 4, slots)).all() + candidates = q.limit(max(slots * 8, 64)).all() claimed_rows: list[BizStateBatch] = [] + visited: set[str] = set() for batch in candidates: + if batch.id in visited: + continue if len(claimed_rows) >= slots: break - ne = str(batch.ne_id or "").strip() - if ne and ne in busy_nes: - continue - # Serialize claims per NE across workers (Postgres). - if ne and pg: - try: - db.execute( - text("SELECT pg_advisory_xact_lock(hashtext(:ne))"), - {"ne": ne}, - ) - except Exception: - _log.exception("advisory lock failed ne=%s", ne) - conflict = ( - db.query(BizStateBatch.id) - .filter( - BizStateBatch.status == "running", - BizStateBatch.ne_id == ne, - BizStateBatch.id != batch.id, - ) - .first() - ) - if conflict: + group = [batch] + if batch.collect_round_id: + peers = db.query(BizStateBatch).filter( + BizStateBatch.collect_round_id == batch.collect_round_id, + ).with_for_update(skip_locked=pg).all() + if len(peers) != 2 or any(p.status != "queued" for p in peers): continue - batch.status = "running" - batch.message = "" - claimed_rows.append(batch) - if ne: - busy_nes.add(ne) + group = peers + visited.update(p.id for p in group) + if len(group) > slots - len(claimed_rows): + continue + nes = {str(p.ne_id or "").strip() for p in group} + if nes & busy_nes: + continue + now = _utcnow() + for member in group: + member.status = "running" + member.message = "" + if member.collect_round_id: + member.started_at = now + task = db.get(BizStateTask, member.task_id) + if task: + task.last_collect_started_at = now + claimed_rows.append(member) + busy_nes.update(nes - {""}) if not claimed_rows: db.rollback() @@ -231,10 +300,13 @@ def claim_queued_batches(limit: int | None = None) -> list[dict[str, Any]]: db.query(BizStateBatch).filter(BizStateBatch.status == "running").count() ) while n_running > global_cap and claimed_rows: - demote = claimed_rows.pop() - demote.status = "queued" - demote.message = "queued" - n_running -= 1 + last = claimed_rows[-1] + demoted = [b for b in claimed_rows if b.collect_round_id == last.collect_round_id] if last.collect_round_id else [last] + for demote in demoted: + claimed_rows.remove(demote) + demote.status = "queued" + demote.message = "queued" + n_running -= 1 if not claimed_rows: db.rollback() diff --git a/netx_api/biz_state/collect_runner.py b/netx_api/biz_state/collect_runner.py index cba6dd9..91797ce 100644 --- a/netx_api/biz_state/collect_runner.py +++ b/netx_api/biz_state/collect_runner.py @@ -784,6 +784,12 @@ def dispatch_collect(task_id: str, *, manual: bool = False) -> dict[str, Any]: result = enqueue_collect(task_id, manual=manual) if not result.get("queued"): return result + if result.get("collect_group_id"): + if _should_execute_inline(): + from ..biz_state_scheduler import _claim_and_dispatch + + _claim_and_dispatch() + return result batch_id = str(result.get("batch_id") or "") if not batch_id: return result @@ -870,6 +876,11 @@ def execute_claimed_batch( task = db.get(BizStateTask, task_id) if task_id else None if not batch: return + if batch.collect_round_id: + batch.started_at = _utcnow() + if task: + task.last_collect_started_at = batch.started_at + db.commit() if not task_id: task_id = str(batch.task_id or "") task = db.get(BizStateTask, task_id) if task_id else None diff --git a/netx_api/biz_state/collect_stop.py b/netx_api/biz_state/collect_stop.py index 8c92db0..1f27e31 100644 --- a/netx_api/biz_state/collect_stop.py +++ b/netx_api/biz_state/collect_stop.py @@ -131,7 +131,7 @@ def _force_close_holders(batch_id: str) -> int: def request_stop_collect(task_id: str) -> dict[str, Any]: - """Cancel queued batches and signal running collect for this task to abort.""" + """Stop a task, including its paired collector, without stranding a half round.""" tid = str(task_id or "").strip() if not tid: return {"ok": False, "reason": "missing_task_id"} @@ -141,26 +141,37 @@ def request_stop_collect(task_id: str) -> dict[str, Any]: signaled_ids: list[str] = [] closed = 0 try: + from sqlalchemy import text + from .claim import _dialect_supports_skip_locked + + if _dialect_supports_skip_locked(db): + db.execute(text("SELECT pg_advisory_xact_lock(hashtext('biz-state-claim-capacity'))")) task = db.get(BizStateTask, tid) if not task: return {"ok": False, "reason": "task_not_found", "task_id": tid} + task_ids = [tid] + if task.collect_group_id: + task_ids = [r[0] for r in db.query(BizStateTask.id).filter( + BizStateTask.collect_group_id == task.collect_group_id).all()] active = ( db.query(BizStateBatch) .filter( - BizStateBatch.task_id == tid, + BizStateBatch.task_id.in_(task_ids), BizStateBatch.status.in_(("queued", "running")), ) - .all() + .order_by(BizStateBatch.id).with_for_update().all() ) + tasks = db.query(BizStateTask).filter(BizStateTask.id.in_(task_ids)).order_by( + BizStateTask.id).with_for_update().populate_existing().all() if not active: # Nothing to stop — clear sticky collect_running if orphaned. - if bool(task.collect_running): - task.collect_running = False - if hasattr(task, "collect_queued_at"): - task.collect_queued_at = None - task.updated_at = utcnow_naive() - db.commit() + for member in tasks: + if member.collect_running: + member.collect_running = False + member.collect_queued_at = None + member.updated_at = utcnow_naive() + db.commit() return { "ok": True, "task_id": tid, @@ -171,7 +182,7 @@ def request_stop_collect(task_id: str) -> dict[str, Any]: } now = utcnow_naive() - still_running = False + running_tasks: set[str] = set() for batch in active: bid = str(batch.id) mark_stop_requested(bid) @@ -183,16 +194,16 @@ def request_stop_collect(task_id: str) -> dict[str, Any]: else: batch.message = STOP_REQUEST_TOKEN signaled_ids.append(bid) - still_running = True + running_tasks.add(batch.task_id) closed += _force_close_holders(bid) - if not still_running: - task.collect_running = False - if hasattr(task, "collect_queued_at"): - task.collect_queued_at = None - task.last_collect_ended_at = now - task.last_error = STOP_USER_MESSAGE - task.updated_at = now + for member in tasks: + if member.id not in running_tasks: + member.collect_running = False + member.collect_queued_at = None + member.last_collect_ended_at = now + member.last_error = STOP_USER_MESSAGE + member.updated_at = now db.commit() _log.info( diff --git a/netx_api/biz_state/schema_ensure.py b/netx_api/biz_state/schema_ensure.py index 2179310..1c11933 100644 --- a/netx_api/biz_state/schema_ensure.py +++ b/netx_api/biz_state/schema_ensure.py @@ -82,6 +82,15 @@ def apply_biz_state_schema(conn: Connection) -> None: "ALTER TABLE biz_migration_project ADD COLUMN IF NOT EXISTS new_hf_bindings_json JSON DEFAULT '[]'", "ALTER TABLE biz_state_task ADD COLUMN IF NOT EXISTS purpose VARCHAR(32) DEFAULT ''", "CREATE INDEX IF NOT EXISTS ix_biz_state_task_purpose ON biz_state_task (purpose)", + "ALTER TABLE biz_state_task ADD COLUMN IF NOT EXISTS collect_group_id VARCHAR(128) DEFAULT ''", + "ALTER TABLE biz_state_task ADD COLUMN IF NOT EXISTS collect_manual_only BOOLEAN DEFAULT FALSE", + "CREATE INDEX IF NOT EXISTS ix_biz_state_task_collect_group_id ON biz_state_task (collect_group_id)", + "ALTER TABLE biz_state_batch ADD COLUMN IF NOT EXISTS collect_group_id VARCHAR(128) DEFAULT ''", + "ALTER TABLE biz_state_batch ADD COLUMN IF NOT EXISTS collect_round_id VARCHAR(64) DEFAULT ''", + "ALTER TABLE biz_state_batch ADD COLUMN IF NOT EXISTS collect_priority INTEGER DEFAULT 0", + "ALTER TABLE biz_state_batch ADD COLUMN IF NOT EXISTS queued_at TIMESTAMP", + "CREATE INDEX IF NOT EXISTS ix_biz_state_batch_collect_group_id ON biz_state_batch (collect_group_id)", + "CREATE INDEX IF NOT EXISTS ix_biz_state_batch_collect_round_id ON biz_state_batch (collect_round_id)", "ALTER TABLE biz_state_batch_command ADD COLUMN IF NOT EXISTS raw_line_count INTEGER DEFAULT 0", "ALTER TABLE biz_state_batch_command ADD COLUMN IF NOT EXISTS declared_total INTEGER DEFAULT 0", "ALTER TABLE biz_migration_red_ticket ADD COLUMN IF NOT EXISTS match_key_str VARCHAR(256) DEFAULT ''", diff --git a/netx_api/biz_state/service.py b/netx_api/biz_state/service.py index da5797c..95e191f 100644 --- a/netx_api/biz_state/service.py +++ b/netx_api/biz_state/service.py @@ -396,6 +396,8 @@ def _task_summary(task: BizStateTask) -> dict[str, Any]: "daily_keep_enabled": bool(getattr(task, "daily_keep_enabled", False)), "daily_keep_count": int(getattr(task, "daily_keep_count", None) or 10), "collect_running": bool(task.collect_running), + "collect_group_id": task.collect_group_id or "", + "collect_manual_only": bool(task.collect_manual_only), "last_collect_started_at": task.last_collect_started_at.isoformat() + "Z" if task.last_collect_started_at else None, @@ -444,6 +446,8 @@ def list_tasks(db: Session, *, purpose: str | None = None) -> list[dict[str, Any "status": t.status, "interval_sec": t.interval_sec, "collect_running": bool(t.collect_running), + "collect_group_id": t.collect_group_id or "", + "collect_manual_only": bool(t.collect_manual_only), "last_error": t.last_error, "last_collect_started_at": t.last_collect_started_at.isoformat() + "Z" if t.last_collect_started_at diff --git a/netx_api/biz_state_scheduler.py b/netx_api/biz_state_scheduler.py index d44250f..49dcd07 100644 --- a/netx_api/biz_state_scheduler.py +++ b/netx_api/biz_state_scheduler.py @@ -110,12 +110,21 @@ def _sync_cutover_hf_windows() -> None: projects = db.query(BizMigrationProject).all() for proj in projects: - want = _hf_window_status(proj) + tids = _all_hf_task_ids(proj, "old") + _all_hf_task_ids(proj, "new") + if tids and any((t := db.get(BizStateTask, tid)) and not t.collect_group_id for tid in tids): + try: + from .biz_migration.service import ensure_highfreq + + ensure_highfreq(db, proj.id, collect_now=False) + except Exception: + _log.exception("upgrade cutover collection pairs failed project=%s", proj.id) + window_status = _hf_window_status(proj) tids = _all_hf_task_ids(proj, "old") + _all_hf_task_ids(proj, "new") for tid in tids: task = db.get(BizStateTask, tid) if not task: continue + want = "paused" if task.collect_manual_only else window_status purpose = str(getattr(task, "purpose", None) or "").strip() if purpose and purpose != PURPOSE_CUTOVER_HF: continue @@ -149,12 +158,26 @@ def _enqueue_due_tasks() -> int: .filter( BizStateTask.status == "running", BizStateTask.collect_running.is_(False), + BizStateTask.collect_manual_only.is_(False), ) .all() ) due_ids: list[str] = [] + seen_groups: set[str] = set() now = _utcnow() for task in tasks: + if task.collect_group_id: + if task.collect_group_id in seen_groups: + continue + seen_groups.add(task.collect_group_id) + peers = db.query(BizStateTask).filter(BizStateTask.collect_group_id == task.collect_group_id).all() + if len(peers) != 2 or any(p.collect_running or p.status != "running" or p.collect_manual_only for p in peers): + continue + anchors = [p.last_collect_started_at or p.last_collect_ended_at for p in peers] + anchor = max((a for a in anchors if a), default=None) + if anchor is None or (now - anchor).total_seconds() >= max(p.interval_sec or 60 for p in peers): + due_ids.append(task.id) + continue interval = max(60, int(task.interval_sec or 300)) # Prefer start-based interval to avoid drift when collect duration # approaches the interval (end-based: 40s collect + 60s → 100s cycle). diff --git a/netx_api/models/biz_state.py b/netx_api/models/biz_state.py index a434dc7..864f85d 100644 --- a/netx_api/models/biz_state.py +++ b/netx_api/models/biz_state.py @@ -32,6 +32,9 @@ class BizStateTask(Base): note: Mapped[str] = mapped_column(String(256), default="") # "" | portrait | cutover_hf — filter cutover HF in biz-state list purpose: Mapped[str] = mapped_column(String(32), default="", index=True) + # Optional paired sampling; ordinary per-device tasks remain independent. + collect_group_id: Mapped[str] = mapped_column(String(128), default="", index=True) + collect_manual_only: Mapped[bool] = mapped_column(Boolean, default=False) status: Mapped[str] = mapped_column(String(32), default="draft", index=True) # draft|running|paused|stopped interval_sec: Mapped[int] = mapped_column(Integer, default=300) # Keep snapshots for N calendar days (protected baselines/refs never auto-deleted) @@ -86,6 +89,10 @@ class BizStateBatch(Base): id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: uuid4().hex) task_id: Mapped[str] = mapped_column(String(64), default="", index=True) + collect_group_id: Mapped[str] = mapped_column(String(128), default="", index=True) + collect_round_id: Mapped[str] = mapped_column(String(64), default="", index=True) + collect_priority: Mapped[int] = mapped_column(Integer, default=0) + queued_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) source: Mapped[str] = mapped_column(String(32), default="managed", index=True) ne_id: Mapped[str] = mapped_column(String(128), default="", index=True) ne_name: Mapped[str] = mapped_column(String(256), default="") diff --git a/tests/test_biz_migration_evidence.py b/tests/test_biz_migration_evidence.py index 89a4adb..8038f4e 100644 --- a/tests/test_biz_migration_evidence.py +++ b/tests/test_biz_migration_evidence.py @@ -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, diff --git a/tests/test_biz_migration_hf_decouple.py b/tests/test_biz_migration_hf_decouple.py index ae63da9..cf9a575 100644 --- a/tests/test_biz_migration_hf_decouple.py +++ b/tests/test_biz_migration_hf_decouple.py @@ -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) diff --git a/tests/test_biz_migration_integrity.py b/tests/test_biz_migration_integrity.py index a1e9140..6f2070d 100644 --- a/tests/test_biz_migration_integrity.py +++ b/tests/test_biz_migration_integrity.py @@ -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": diff --git a/tests/test_biz_state_paired_sampling.py b/tests/test_biz_state_paired_sampling.py new file mode 100644 index 0000000..1594824 --- /dev/null +++ b/tests/test_biz_state_paired_sampling.py @@ -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) +