diff --git a/netx_api/config_sync_common.py b/netx_api/config_sync_common.py new file mode 100644 index 0000000..fa5cba9 --- /dev/null +++ b/netx_api/config_sync_common.py @@ -0,0 +1,129 @@ +"""Config sync shared helpers and policy ensure/prune.""" +from __future__ import annotations + +import logging +from datetime import datetime +from typing import Any + +from sqlalchemy.orm import Session + +from .config_sync_schemas import ( + ConfigSyncCycleOut, + ConfigSyncPolicyOut, + ConfigSyncTargetRef, + ConfigSyncTaskOut, +) +from .models import ConfigSyncCycle, ConfigSyncPolicy, ConfigSyncTask +from .timeutil import utcnow_naive + +_log = logging.getLogger("netx.config_sync") + +POLICY_ID = 1 +DEFAULT_CYCLE_KEEP = 30 + + +def _utcnow() -> datetime: + return utcnow_naive() + + +def ensure_policy(db: Session) -> ConfigSyncPolicy: + row = db.get(ConfigSyncPolicy, POLICY_ID) + if row is None: + row = ConfigSyncPolicy(id=POLICY_ID, enabled=False) + db.add(row) + db.commit() + db.refresh(row) + return row + + +def prune_config_sync_cycles(db: Session, *, keep: int = DEFAULT_CYCLE_KEEP) -> int: + """Delete finished cycles beyond ``keep`` (newest kept). Active cycles always retained.""" + keep = max(0, min(200, int(keep))) + finished = ( + db.query(ConfigSyncCycle) + .filter(ConfigSyncCycle.status.in_(("success", "fail", "cancelled"))) + .order_by(ConfigSyncCycle.created_at.desc()) + .all() + ) + to_drop = finished if keep == 0 else finished[keep:] + if not to_drop: + return 0 + dropped = 0 + for cycle in to_drop: + cid = str(cycle.id) + db.query(ConfigSyncTask).filter(ConfigSyncTask.cycle_id == cid).delete( + synchronize_session=False + ) + db.delete(cycle) + dropped += 1 + if dropped: + db.commit() + return dropped + + +def _cycle_keep_value(row: ConfigSyncPolicy) -> int: + return max(0, min(200, int(getattr(row, "cycle_keep", None) or DEFAULT_CYCLE_KEEP))) + + +def _targets_from_json(raw: Any) -> list[ConfigSyncTargetRef]: + items: list[ConfigSyncTargetRef] = [] + if not isinstance(raw, list): + return items + for x in raw: + if not isinstance(x, dict): + continue + src = str(x.get("source") or "").strip().lower() + tid = str(x.get("id") or "").strip() + if src not in ("managed", "ume") or not tid: + continue + items.append(ConfigSyncTargetRef(source=src, id=tid)) # type: ignore[arg-type] + return items + + +def policy_to_out(row: ConfigSyncPolicy) -> ConfigSyncPolicyOut: + return ConfigSyncPolicyOut( + enabled=bool(row.enabled), + interval_days=max(1, int(row.interval_days or 3)), + concurrency=max(1, min(30, int(row.concurrency or 5))), + scope_mode=str(row.scope_mode or "all"), + selected_targets=_targets_from_json(row.selected_targets), + history_keep=max(0, min(30, int(row.history_keep if row.history_keep is not None else 3))), + cycle_keep=_cycle_keep_value(row), + updated_at=row.updated_at, + ) + + + +def cycle_to_out(row: ConfigSyncCycle) -> ConfigSyncCycleOut: + return ConfigSyncCycleOut( + id=str(row.id), + trigger_mode=str(row.trigger_mode or ""), + status=str(row.status or ""), + concurrency=int(row.concurrency or 0), + planned_count=int(row.planned_count or 0), + success_count=int(row.success_count or 0), + fail_count=int(row.fail_count or 0), + skip_count=int(row.skip_count or 0), + error_message=str(row.error_message or ""), + started_at=row.started_at, + ended_at=row.ended_at, + created_at=row.created_at, + ) + + +def task_to_out(row: ConfigSyncTask) -> ConfigSyncTaskOut: + return ConfigSyncTaskOut( + id=str(row.id), + cycle_id=str(row.cycle_id), + source=str(row.source), + target_id=str(row.target_id), + ne_name=str(row.ne_name or ""), + ne_ip=str(row.ne_ip or ""), + vendor=str(row.vendor or ""), + status=str(row.status or ""), + message=str(row.message or ""), + started_at=row.started_at, + ended_at=row.ended_at, + ) + + diff --git a/netx_api/config_sync_cycles.py b/netx_api/config_sync_cycles.py new file mode 100644 index 0000000..e9a703b --- /dev/null +++ b/netx_api/config_sync_cycles.py @@ -0,0 +1,468 @@ +"""Config sync policy updates, cycles, and dashboard.""" +from __future__ import annotations + +from datetime import datetime, timedelta +from typing import Any +from uuid import uuid4 + +from fastapi import HTTPException +from sqlalchemy import func +from sqlalchemy.orm import Session + +from .cli_resolve import cli_profile_ready +from .config_sync_common import ( + DEFAULT_CYCLE_KEEP, + _cycle_keep_value, + _targets_from_json, + _utcnow, + cycle_to_out, + ensure_policy, + policy_to_out, + prune_config_sync_cycles, + task_to_out, +) +from .config_sync_schemas import ( + ConfigSyncCycleCreate, + ConfigSyncCycleOut, + ConfigSyncDashboardOut, + ConfigSyncPolicyOut, + ConfigSyncPolicyUpdate, + ConfigSyncTaskOut, +) +from .models import ( + ConfigSyncCycle, + ConfigSyncPolicy, + ConfigSyncTask, + ManagedNE, + NeConfigSnapshot, + UmeInventoryNE, +) + +def get_policy(db: Session) -> ConfigSyncPolicyOut: + return policy_to_out(ensure_policy(db)) + + +def update_policy(db: Session, body: ConfigSyncPolicyUpdate) -> ConfigSyncPolicyOut: + row = ensure_policy(db) + data = body.model_dump(exclude_unset=True) + if "enabled" in data and data["enabled"] is not None: + row.enabled = bool(data["enabled"]) + if "interval_days" in data and data["interval_days"] is not None: + row.interval_days = int(data["interval_days"]) + if "concurrency" in data and data["concurrency"] is not None: + row.concurrency = max(1, min(30, int(data["concurrency"]))) + if "scope_mode" in data and data["scope_mode"] is not None: + row.scope_mode = str(data["scope_mode"]) + if "selected_targets" in data and data["selected_targets"] is not None: + refs = data["selected_targets"] + row.selected_targets = [ + {"source": r.source if hasattr(r, "source") else r["source"], "id": r.id if hasattr(r, "id") else r["id"]} + for r in refs + ] + if "history_keep" in data and data["history_keep"] is not None: + row.history_keep = max(0, min(30, int(data["history_keep"]))) + if "cycle_keep" in data and data["cycle_keep"] is not None: + row.cycle_keep = max(0, min(200, int(data["cycle_keep"]))) + row.updated_at = _utcnow() + db.commit() + db.refresh(row) + prune_config_sync_cycles(db, keep=_cycle_keep_value(row)) + return policy_to_out(row) + + +def expand_targets(db: Session, policy: ConfigSyncPolicy) -> list[dict[str, str]]: + """Return list of {source, id, ne_name, ne_ip, vendor, device_type}.""" + mode = str(policy.scope_mode or "all").strip().lower() + out: list[dict[str, str]] = [] + seen: set[tuple[str, str]] = set() + + def _add(source: str, tid: str, name: str, ip: str, vendor: str, device_type: str) -> None: + key = (source, tid) + if key in seen: + return + seen.add(key) + out.append( + { + "source": source, + "id": tid, + "ne_name": name, + "ne_ip": ip, + "vendor": vendor, + "device_type": device_type, + } + ) + + if mode == "selected": + for ref in _targets_from_json(policy.selected_targets): + if ref.source == "managed": + ne = db.get(ManagedNE, ref.id) + if ne: + _add("managed", str(ne.id), str(ne.name or ""), str(ne.ip_address or ""), str(ne.vendor or ""), str(ne.device_type or "")) + else: + inv = db.get(UmeInventoryNE, ref.id) + if inv: + _add( + "ume", + str(inv.ne_id), + str(inv.host_name or inv.ne_name or inv.user_label or inv.ne_id or ""), + str(inv.ip_address or ""), + str(inv.vendor or ""), + str(inv.ne_type or ""), + ) + return out + + for ne in db.query(ManagedNE).order_by(ManagedNE.updated_at.desc()).all(): + _add("managed", str(ne.id), str(ne.name or ""), str(ne.ip_address or ""), str(ne.vendor or ""), str(ne.device_type or "")) + + if cli_profile_ready(db): + for inv in db.query(UmeInventoryNE).order_by(UmeInventoryNE.ne_id.asc()).all(): + if not str(inv.ip_address or "").strip(): + continue + _add( + "ume", + str(inv.ne_id), + str(inv.host_name or inv.ne_name or inv.user_label or inv.ne_id or ""), + str(inv.ip_address or ""), + str(inv.vendor or ""), + str(inv.ne_type or ""), + ) + return out + + +def has_active_cycle(db: Session) -> ConfigSyncCycle | None: + """Any non-terminal cycle occupies the single-flight slot (incl. paused).""" + return ( + db.query(ConfigSyncCycle) + .filter(ConfigSyncCycle.status.in_(("running", "pending", "paused"))) + .order_by(ConfigSyncCycle.created_at.desc()) + .first() + ) + + +def has_running_cycle(db: Session) -> ConfigSyncCycle | None: + """Backward-compatible alias: treat paused as active so a new cycle cannot start.""" + return has_active_cycle(db) + + +def last_finished_cycle(db: Session) -> ConfigSyncCycle | None: + return ( + db.query(ConfigSyncCycle) + .filter(ConfigSyncCycle.status.in_(("success", "fail", "cancelled"))) + .order_by(ConfigSyncCycle.ended_at.desc().nullslast(), ConfigSyncCycle.created_at.desc()) + .first() + ) + + +def next_due_at(db: Session, policy: ConfigSyncPolicy | None = None) -> datetime | None: + pol = policy or ensure_policy(db) + if not pol.enabled: + return None + from .config_sync_scheduler import startup_grace_until + + last = ( + db.query(ConfigSyncCycle) + .filter(ConfigSyncCycle.status == "success", ConfigSyncCycle.ended_at.isnot(None)) + .order_by(ConfigSyncCycle.ended_at.desc()) + .first() + ) + days = max(1, int(pol.interval_days or 3)) + if last and last.ended_at: + due = last.ended_at + timedelta(days=days) + else: + # Never synced successfully: do not fire immediately on enable / first boot. + due = _utcnow() + timedelta(days=days) + grace_until = startup_grace_until() + if grace_until is not None and due < grace_until: + return grace_until + return due + + +def create_cycle(db: Session, body: ConfigSyncCycleCreate) -> ConfigSyncCycleOut: + if has_running_cycle(db): + raise HTTPException(status_code=409, detail="config_sync_cycle_already_running") + policy = ensure_policy(db) + mode = str(body.mode or "full").strip().lower() + trigger = "retry_failed" if mode == "retry_failed" else "manual" + concurrency = max(1, min(30, int(policy.concurrency or 5))) + + targets: list[dict[str, str]] = [] + if mode == "retry_failed": + src_cycle_id = str(body.cycle_id or "").strip() + src = None + if src_cycle_id: + src = db.get(ConfigSyncCycle, src_cycle_id) + if src is None: + src = ( + db.query(ConfigSyncCycle) + .filter(ConfigSyncCycle.fail_count > 0) + .order_by(ConfigSyncCycle.created_at.desc()) + .first() + ) + if src is None: + raise HTTPException(status_code=404, detail="no_failed_cycle") + fails = ( + db.query(ConfigSyncTask) + .filter(ConfigSyncTask.cycle_id == src.id, ConfigSyncTask.status == "fail") + .all() + ) + for t in fails: + targets.append( + { + "source": str(t.source), + "id": str(t.target_id), + "ne_name": str(t.ne_name or ""), + "ne_ip": str(t.ne_ip or ""), + "vendor": str(t.vendor or ""), + "device_type": "", + } + ) + if not targets: + raise HTTPException(status_code=400, detail="no_failed_tasks") + else: + targets = expand_targets(db, policy) + if not targets: + raise HTTPException(status_code=400, detail="no_targets") + + cycle = ConfigSyncCycle( + id=uuid4().hex, + trigger_mode=trigger, + status="running", + concurrency=concurrency, + planned_count=len(targets), + started_at=_utcnow(), + created_at=_utcnow(), + ) + db.add(cycle) + db.flush() + for t in targets: + db.add( + ConfigSyncTask( + id=uuid4().hex, + cycle_id=cycle.id, + source=t["source"], + target_id=t["id"], + ne_name=t.get("ne_name") or "", + ne_ip=t.get("ne_ip") or "", + vendor=t.get("vendor") or "", + status="pending", + ) + ) + db.commit() + db.refresh(cycle) + return cycle_to_out(cycle) + + +def list_cycles(db: Session, *, page: int, page_size: int) -> dict[str, Any]: + q = db.query(ConfigSyncCycle).order_by(ConfigSyncCycle.created_at.desc()) + total = int(q.count()) + rows = q.offset((page - 1) * page_size).limit(page_size).all() + return {"total": total, "page": page, "page_size": page_size, "items": [cycle_to_out(r) for r in rows]} + + +def get_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: + row = db.get(ConfigSyncCycle, cycle_id) + if not row: + raise HTTPException(status_code=404, detail="cycle_not_found") + return cycle_to_out(row) + + +def list_cycle_tasks( + db: Session, + cycle_id: str, + *, + page: int, + page_size: int, + status: str = "", + keyword: str = "", +) -> dict[str, Any]: + if not db.get(ConfigSyncCycle, cycle_id): + raise HTTPException(status_code=404, detail="cycle_not_found") + q = db.query(ConfigSyncTask).filter(ConfigSyncTask.cycle_id == cycle_id) + st = str(status or "").strip() + if st: + q = q.filter(ConfigSyncTask.status == st) + kw = str(keyword or "").strip() + if kw: + like = f"%{kw}%" + q = q.filter( + or_( + ConfigSyncTask.ne_name.ilike(like), + ConfigSyncTask.ne_ip.ilike(like), + ConfigSyncTask.target_id.ilike(like), + ConfigSyncTask.message.ilike(like), + ) + ) + total = int(q.count()) + rows = q.order_by(ConfigSyncTask.ne_name.asc()).offset((page - 1) * page_size).limit(page_size).all() + return {"total": total, "page": page, "page_size": page_size, "items": [task_to_out(r) for r in rows]} + + +def pause_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: + row = db.get(ConfigSyncCycle, cycle_id) + if not row: + raise HTTPException(status_code=404, detail="cycle_not_found") + if str(row.status) not in ("running", "pending"): + raise HTTPException(status_code=400, detail="cycle_not_running") + row.status = "paused" + db.commit() + db.refresh(row) + return cycle_to_out(row) + + +def resume_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: + row = db.get(ConfigSyncCycle, cycle_id) + if not row: + raise HTTPException(status_code=404, detail="cycle_not_found") + if str(row.status) != "paused": + raise HTTPException(status_code=400, detail="cycle_not_paused") + pending = ( + db.query(ConfigSyncTask) + .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status == "pending") + .count() + ) + if pending <= 0: + raise HTTPException(status_code=400, detail="no_pending_tasks") + other = has_running_cycle(db) + if other and str(other.id) != cycle_id: + raise HTTPException(status_code=409, detail="config_sync_cycle_already_running") + row.status = "running" + db.commit() + db.refresh(row) + return cycle_to_out(row) + + +def stop_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: + """Cancel remaining work and close the cycle (running/paused/pending).""" + row = db.get(ConfigSyncCycle, cycle_id) + if not row: + raise HTTPException(status_code=404, detail="cycle_not_found") + if str(row.status) not in ("running", "paused", "pending"): + raise HTTPException(status_code=400, detail="cycle_not_active") + now = _utcnow() + pending = ( + db.query(ConfigSyncTask) + .filter( + ConfigSyncTask.cycle_id == cycle_id, + ConfigSyncTask.status.in_(("pending", "running")), + ) + .all() + ) + for task in pending: + # In-flight workers may still finish and overwrite; pending must not start. + if str(task.status) == "pending": + task.status = "cancelled" + task.message = "stopped_by_user" + task.ended_at = now + else: + task.message = (str(task.message or "") + " · stop_requested")[:1020] + row.status = "cancelled" + row.error_message = "stopped_by_user" + row.ended_at = now + db.commit() + sync_cycle_progress(db, cycle_id) + db.refresh(row) + try: + from .config_sync_runner import _release_pool + + _release_pool(cycle_id) + except Exception: + pass + out = cycle_to_out(row) + try: + prune_config_sync_cycles(db, keep=_cycle_keep_value(ensure_policy(db))) + except Exception: + _log.exception("prune_config_sync_cycles after stop failed") + return out + + +def dashboard(db: Session) -> ConfigSyncDashboardOut: + policy = ensure_policy(db) + snap_count = int(db.query(func.count()).select_from(NeConfigSnapshot).scalar() or 0) + running = ( + db.query(ConfigSyncCycle) + .filter(ConfigSyncCycle.status.in_(("running", "paused", "pending"))) + .order_by(ConfigSyncCycle.created_at.desc()) + .first() + ) + last = last_finished_cycle(db) + fail_by_vendor: dict[str, int] = {} + if last: + rows = ( + db.query(ConfigSyncTask.vendor, func.count()) + .filter(ConfigSyncTask.cycle_id == last.id, ConfigSyncTask.status == "fail") + .group_by(ConfigSyncTask.vendor) + .all() + ) + for vendor, cnt in rows: + fail_by_vendor[str(vendor or "unknown") or "unknown"] = int(cnt) + return ConfigSyncDashboardOut( + policy=policy_to_out(policy), + snapshot_count=snap_count, + last_cycle=cycle_to_out(last) if last else None, + running_cycle=cycle_to_out(running) if running else None, + next_due_at=next_due_at(db, policy), + fail_by_vendor=fail_by_vendor, + ) + +def sync_cycle_progress(db: Session, cycle_id: str) -> None: + cycle = db.get(ConfigSyncCycle, cycle_id) + if not cycle: + return + success = ( + db.query(func.count()) + .select_from(ConfigSyncTask) + .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status == "success") + .scalar() + or 0 + ) + fail = ( + db.query(func.count()) + .select_from(ConfigSyncTask) + .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status == "fail") + .scalar() + or 0 + ) + skip = ( + db.query(func.count()) + .select_from(ConfigSyncTask) + .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status.in_(("skipped", "cancelled"))) + .scalar() + or 0 + ) + cycle.success_count = int(success) + cycle.fail_count = int(fail) + cycle.skip_count = int(skip) + db.commit() + + +def finalize_cycle(db: Session, cycle_id: str) -> None: + cycle = db.get(ConfigSyncCycle, cycle_id) + if not cycle: + return + if str(cycle.status) in ("paused", "cancelled"): + return + pending = ( + db.query(func.count()) + .select_from(ConfigSyncTask) + .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status.in_(("pending", "running"))) + .scalar() + or 0 + ) + if int(pending) > 0: + return + sync_cycle_progress(db, cycle_id) + db.refresh(cycle) + if str(cycle.status) in ("paused", "cancelled"): + return + # Cycle outcome is about finishing the run, not per-NE results. + # Individual task failures stay in fail_count for retry/dashboard. + cycle.status = "success" + if cycle.error_message == "completed_with_failures": + cycle.error_message = "" + cycle.ended_at = _utcnow() + db.commit() + try: + prune_config_sync_cycles(db, keep=_cycle_keep_value(ensure_policy(db))) + except Exception: + _log.exception("prune_config_sync_cycles after finish failed") + diff --git a/netx_api/config_sync_service.py b/netx_api/config_sync_service.py index 8db370d..8bcc106 100644 --- a/netx_api/config_sync_service.py +++ b/netx_api/config_sync_service.py @@ -1,737 +1,66 @@ """Config sync policy, cycles, dashboard, and snapshot queries.""" - from __future__ import annotations -import io -import logging -import re -import zipfile -from datetime import datetime, timedelta -from typing import Any -from uuid import uuid4 - -from fastapi import HTTPException -from sqlalchemy import func, or_ -from sqlalchemy.orm import Session - -from .config_sync_codec import decompress_text -from .config_sync_schemas import ( - ConfigSyncCycleCreate, - ConfigSyncCycleOut, - ConfigSyncDashboardOut, - ConfigSyncPolicyOut, - ConfigSyncPolicyUpdate, - ConfigSyncTargetRef, - ConfigSyncTaskOut, - NeConfigHistoryOut, - NeConfigSnapshotDetailOut, - NeConfigSnapshotMetaOut, +from .config_sync_common import ( + DEFAULT_CYCLE_KEEP, + POLICY_ID, + _cycle_keep_value, + ensure_policy, + policy_to_out, + prune_config_sync_cycles, ) -from .models import ( - ConfigSyncCycle, - ConfigSyncPolicy, - ConfigSyncTask, - ManagedNE, - NeConfigHistory, - NeConfigSnapshot, - UmeInventoryNE, +from .config_sync_cycles import ( + create_cycle, + dashboard, + expand_targets, + finalize_cycle, + get_cycle, + get_policy, + has_active_cycle, + has_running_cycle, + last_finished_cycle, + list_cycle_tasks, + list_cycles, + next_due_at, + pause_cycle, + resume_cycle, + stop_cycle, + sync_cycle_progress, + update_policy, ) -from .cli_resolve import cli_profile_ready - -_log = logging.getLogger("netx.config_sync") - -POLICY_ID = 1 - - -def _utcnow() -> datetime: - return datetime.utcnow() - - -DEFAULT_CYCLE_KEEP = 30 - - -def ensure_policy(db: Session) -> ConfigSyncPolicy: - row = db.get(ConfigSyncPolicy, POLICY_ID) - if row is None: - row = ConfigSyncPolicy(id=POLICY_ID, enabled=False) - db.add(row) - db.commit() - db.refresh(row) - return row - - -def prune_config_sync_cycles(db: Session, *, keep: int = DEFAULT_CYCLE_KEEP) -> int: - """Delete finished cycles beyond ``keep`` (newest kept). Active cycles always retained.""" - keep = max(0, min(200, int(keep))) - finished = ( - db.query(ConfigSyncCycle) - .filter(ConfigSyncCycle.status.in_(("success", "fail", "cancelled"))) - .order_by(ConfigSyncCycle.created_at.desc()) - .all() - ) - to_drop = finished if keep == 0 else finished[keep:] - if not to_drop: - return 0 - dropped = 0 - for cycle in to_drop: - cid = str(cycle.id) - db.query(ConfigSyncTask).filter(ConfigSyncTask.cycle_id == cid).delete( - synchronize_session=False - ) - db.delete(cycle) - dropped += 1 - if dropped: - db.commit() - return dropped - - -def _cycle_keep_value(row: ConfigSyncPolicy) -> int: - return max(0, min(200, int(getattr(row, "cycle_keep", None) or DEFAULT_CYCLE_KEEP))) - - -def _targets_from_json(raw: Any) -> list[ConfigSyncTargetRef]: - items: list[ConfigSyncTargetRef] = [] - if not isinstance(raw, list): - return items - for x in raw: - if not isinstance(x, dict): - continue - src = str(x.get("source") or "").strip().lower() - tid = str(x.get("id") or "").strip() - if src not in ("managed", "ume") or not tid: - continue - items.append(ConfigSyncTargetRef(source=src, id=tid)) # type: ignore[arg-type] - return items - - -def policy_to_out(row: ConfigSyncPolicy) -> ConfigSyncPolicyOut: - return ConfigSyncPolicyOut( - enabled=bool(row.enabled), - interval_days=max(1, int(row.interval_days or 3)), - concurrency=max(1, min(30, int(row.concurrency or 5))), - scope_mode=str(row.scope_mode or "all"), - selected_targets=_targets_from_json(row.selected_targets), - history_keep=max(0, min(30, int(row.history_keep if row.history_keep is not None else 3))), - cycle_keep=_cycle_keep_value(row), - updated_at=row.updated_at, - ) - - -def get_policy(db: Session) -> ConfigSyncPolicyOut: - return policy_to_out(ensure_policy(db)) - - -def update_policy(db: Session, body: ConfigSyncPolicyUpdate) -> ConfigSyncPolicyOut: - row = ensure_policy(db) - data = body.model_dump(exclude_unset=True) - if "enabled" in data and data["enabled"] is not None: - row.enabled = bool(data["enabled"]) - if "interval_days" in data and data["interval_days"] is not None: - row.interval_days = int(data["interval_days"]) - if "concurrency" in data and data["concurrency"] is not None: - row.concurrency = max(1, min(30, int(data["concurrency"]))) - if "scope_mode" in data and data["scope_mode"] is not None: - row.scope_mode = str(data["scope_mode"]) - if "selected_targets" in data and data["selected_targets"] is not None: - refs = data["selected_targets"] - row.selected_targets = [ - {"source": r.source if hasattr(r, "source") else r["source"], "id": r.id if hasattr(r, "id") else r["id"]} - for r in refs - ] - if "history_keep" in data and data["history_keep"] is not None: - row.history_keep = max(0, min(30, int(data["history_keep"]))) - if "cycle_keep" in data and data["cycle_keep"] is not None: - row.cycle_keep = max(0, min(200, int(data["cycle_keep"]))) - row.updated_at = _utcnow() - db.commit() - db.refresh(row) - prune_config_sync_cycles(db, keep=_cycle_keep_value(row)) - return policy_to_out(row) - - -def cycle_to_out(row: ConfigSyncCycle) -> ConfigSyncCycleOut: - return ConfigSyncCycleOut( - id=str(row.id), - trigger_mode=str(row.trigger_mode or ""), - status=str(row.status or ""), - concurrency=int(row.concurrency or 0), - planned_count=int(row.planned_count or 0), - success_count=int(row.success_count or 0), - fail_count=int(row.fail_count or 0), - skip_count=int(row.skip_count or 0), - error_message=str(row.error_message or ""), - started_at=row.started_at, - ended_at=row.ended_at, - created_at=row.created_at, - ) - - -def task_to_out(row: ConfigSyncTask) -> ConfigSyncTaskOut: - return ConfigSyncTaskOut( - id=str(row.id), - cycle_id=str(row.cycle_id), - source=str(row.source), - target_id=str(row.target_id), - ne_name=str(row.ne_name or ""), - ne_ip=str(row.ne_ip or ""), - vendor=str(row.vendor or ""), - status=str(row.status or ""), - message=str(row.message or ""), - started_at=row.started_at, - ended_at=row.ended_at, - ) - - -def expand_targets(db: Session, policy: ConfigSyncPolicy) -> list[dict[str, str]]: - """Return list of {source, id, ne_name, ne_ip, vendor, device_type}.""" - mode = str(policy.scope_mode or "all").strip().lower() - out: list[dict[str, str]] = [] - seen: set[tuple[str, str]] = set() - - def _add(source: str, tid: str, name: str, ip: str, vendor: str, device_type: str) -> None: - key = (source, tid) - if key in seen: - return - seen.add(key) - out.append( - { - "source": source, - "id": tid, - "ne_name": name, - "ne_ip": ip, - "vendor": vendor, - "device_type": device_type, - } - ) - - if mode == "selected": - for ref in _targets_from_json(policy.selected_targets): - if ref.source == "managed": - ne = db.get(ManagedNE, ref.id) - if ne: - _add("managed", str(ne.id), str(ne.name or ""), str(ne.ip_address or ""), str(ne.vendor or ""), str(ne.device_type or "")) - else: - inv = db.get(UmeInventoryNE, ref.id) - if inv: - _add( - "ume", - str(inv.ne_id), - str(inv.host_name or inv.ne_name or inv.user_label or inv.ne_id or ""), - str(inv.ip_address or ""), - str(inv.vendor or ""), - str(inv.ne_type or ""), - ) - return out - - for ne in db.query(ManagedNE).order_by(ManagedNE.updated_at.desc()).all(): - _add("managed", str(ne.id), str(ne.name or ""), str(ne.ip_address or ""), str(ne.vendor or ""), str(ne.device_type or "")) - - if cli_profile_ready(db): - for inv in db.query(UmeInventoryNE).order_by(UmeInventoryNE.ne_id.asc()).all(): - if not str(inv.ip_address or "").strip(): - continue - _add( - "ume", - str(inv.ne_id), - str(inv.host_name or inv.ne_name or inv.user_label or inv.ne_id or ""), - str(inv.ip_address or ""), - str(inv.vendor or ""), - str(inv.ne_type or ""), - ) - return out - - -def has_active_cycle(db: Session) -> ConfigSyncCycle | None: - """Any non-terminal cycle occupies the single-flight slot (incl. paused).""" - return ( - db.query(ConfigSyncCycle) - .filter(ConfigSyncCycle.status.in_(("running", "pending", "paused"))) - .order_by(ConfigSyncCycle.created_at.desc()) - .first() - ) - - -def has_running_cycle(db: Session) -> ConfigSyncCycle | None: - """Backward-compatible alias: treat paused as active so a new cycle cannot start.""" - return has_active_cycle(db) - - -def last_finished_cycle(db: Session) -> ConfigSyncCycle | None: - return ( - db.query(ConfigSyncCycle) - .filter(ConfigSyncCycle.status.in_(("success", "fail", "cancelled"))) - .order_by(ConfigSyncCycle.ended_at.desc().nullslast(), ConfigSyncCycle.created_at.desc()) - .first() - ) - - -def next_due_at(db: Session, policy: ConfigSyncPolicy | None = None) -> datetime | None: - pol = policy or ensure_policy(db) - if not pol.enabled: - return None - from .config_sync_scheduler import startup_grace_until - - last = ( - db.query(ConfigSyncCycle) - .filter(ConfigSyncCycle.status == "success", ConfigSyncCycle.ended_at.isnot(None)) - .order_by(ConfigSyncCycle.ended_at.desc()) - .first() - ) - days = max(1, int(pol.interval_days or 3)) - if last and last.ended_at: - due = last.ended_at + timedelta(days=days) - else: - # Never synced successfully: do not fire immediately on enable / first boot. - due = _utcnow() + timedelta(days=days) - grace_until = startup_grace_until() - if grace_until is not None and due < grace_until: - return grace_until - return due - - -def create_cycle(db: Session, body: ConfigSyncCycleCreate) -> ConfigSyncCycleOut: - if has_running_cycle(db): - raise HTTPException(status_code=409, detail="config_sync_cycle_already_running") - policy = ensure_policy(db) - mode = str(body.mode or "full").strip().lower() - trigger = "retry_failed" if mode == "retry_failed" else "manual" - concurrency = max(1, min(30, int(policy.concurrency or 5))) - - targets: list[dict[str, str]] = [] - if mode == "retry_failed": - src_cycle_id = str(body.cycle_id or "").strip() - src = None - if src_cycle_id: - src = db.get(ConfigSyncCycle, src_cycle_id) - if src is None: - src = ( - db.query(ConfigSyncCycle) - .filter(ConfigSyncCycle.fail_count > 0) - .order_by(ConfigSyncCycle.created_at.desc()) - .first() - ) - if src is None: - raise HTTPException(status_code=404, detail="no_failed_cycle") - fails = ( - db.query(ConfigSyncTask) - .filter(ConfigSyncTask.cycle_id == src.id, ConfigSyncTask.status == "fail") - .all() - ) - for t in fails: - targets.append( - { - "source": str(t.source), - "id": str(t.target_id), - "ne_name": str(t.ne_name or ""), - "ne_ip": str(t.ne_ip or ""), - "vendor": str(t.vendor or ""), - "device_type": "", - } - ) - if not targets: - raise HTTPException(status_code=400, detail="no_failed_tasks") - else: - targets = expand_targets(db, policy) - if not targets: - raise HTTPException(status_code=400, detail="no_targets") - - cycle = ConfigSyncCycle( - id=uuid4().hex, - trigger_mode=trigger, - status="running", - concurrency=concurrency, - planned_count=len(targets), - started_at=_utcnow(), - created_at=_utcnow(), - ) - db.add(cycle) - db.flush() - for t in targets: - db.add( - ConfigSyncTask( - id=uuid4().hex, - cycle_id=cycle.id, - source=t["source"], - target_id=t["id"], - ne_name=t.get("ne_name") or "", - ne_ip=t.get("ne_ip") or "", - vendor=t.get("vendor") or "", - status="pending", - ) - ) - db.commit() - db.refresh(cycle) - return cycle_to_out(cycle) - - -def list_cycles(db: Session, *, page: int, page_size: int) -> dict[str, Any]: - q = db.query(ConfigSyncCycle).order_by(ConfigSyncCycle.created_at.desc()) - total = int(q.count()) - rows = q.offset((page - 1) * page_size).limit(page_size).all() - return {"total": total, "page": page, "page_size": page_size, "items": [cycle_to_out(r) for r in rows]} - - -def get_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: - row = db.get(ConfigSyncCycle, cycle_id) - if not row: - raise HTTPException(status_code=404, detail="cycle_not_found") - return cycle_to_out(row) - - -def list_cycle_tasks( - db: Session, - cycle_id: str, - *, - page: int, - page_size: int, - status: str = "", - keyword: str = "", -) -> dict[str, Any]: - if not db.get(ConfigSyncCycle, cycle_id): - raise HTTPException(status_code=404, detail="cycle_not_found") - q = db.query(ConfigSyncTask).filter(ConfigSyncTask.cycle_id == cycle_id) - st = str(status or "").strip() - if st: - q = q.filter(ConfigSyncTask.status == st) - kw = str(keyword or "").strip() - if kw: - like = f"%{kw}%" - q = q.filter( - or_( - ConfigSyncTask.ne_name.ilike(like), - ConfigSyncTask.ne_ip.ilike(like), - ConfigSyncTask.target_id.ilike(like), - ConfigSyncTask.message.ilike(like), - ) - ) - total = int(q.count()) - rows = q.order_by(ConfigSyncTask.ne_name.asc()).offset((page - 1) * page_size).limit(page_size).all() - return {"total": total, "page": page, "page_size": page_size, "items": [task_to_out(r) for r in rows]} - - -def pause_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: - row = db.get(ConfigSyncCycle, cycle_id) - if not row: - raise HTTPException(status_code=404, detail="cycle_not_found") - if str(row.status) not in ("running", "pending"): - raise HTTPException(status_code=400, detail="cycle_not_running") - row.status = "paused" - db.commit() - db.refresh(row) - return cycle_to_out(row) - - -def resume_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: - row = db.get(ConfigSyncCycle, cycle_id) - if not row: - raise HTTPException(status_code=404, detail="cycle_not_found") - if str(row.status) != "paused": - raise HTTPException(status_code=400, detail="cycle_not_paused") - pending = ( - db.query(ConfigSyncTask) - .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status == "pending") - .count() - ) - if pending <= 0: - raise HTTPException(status_code=400, detail="no_pending_tasks") - other = has_running_cycle(db) - if other and str(other.id) != cycle_id: - raise HTTPException(status_code=409, detail="config_sync_cycle_already_running") - row.status = "running" - db.commit() - db.refresh(row) - return cycle_to_out(row) - - -def stop_cycle(db: Session, cycle_id: str) -> ConfigSyncCycleOut: - """Cancel remaining work and close the cycle (running/paused/pending).""" - row = db.get(ConfigSyncCycle, cycle_id) - if not row: - raise HTTPException(status_code=404, detail="cycle_not_found") - if str(row.status) not in ("running", "paused", "pending"): - raise HTTPException(status_code=400, detail="cycle_not_active") - now = _utcnow() - pending = ( - db.query(ConfigSyncTask) - .filter( - ConfigSyncTask.cycle_id == cycle_id, - ConfigSyncTask.status.in_(("pending", "running")), - ) - .all() - ) - for task in pending: - # In-flight workers may still finish and overwrite; pending must not start. - if str(task.status) == "pending": - task.status = "cancelled" - task.message = "stopped_by_user" - task.ended_at = now - else: - task.message = (str(task.message or "") + " · stop_requested")[:1020] - row.status = "cancelled" - row.error_message = "stopped_by_user" - row.ended_at = now - db.commit() - sync_cycle_progress(db, cycle_id) - db.refresh(row) - try: - from .config_sync_runner import _release_pool - - _release_pool(cycle_id) - except Exception: - pass - out = cycle_to_out(row) - try: - prune_config_sync_cycles(db, keep=_cycle_keep_value(ensure_policy(db))) - except Exception: - _log.exception("prune_config_sync_cycles after stop failed") - return out - - -def dashboard(db: Session) -> ConfigSyncDashboardOut: - policy = ensure_policy(db) - snap_count = int(db.query(func.count()).select_from(NeConfigSnapshot).scalar() or 0) - running = ( - db.query(ConfigSyncCycle) - .filter(ConfigSyncCycle.status.in_(("running", "paused", "pending"))) - .order_by(ConfigSyncCycle.created_at.desc()) - .first() - ) - last = last_finished_cycle(db) - fail_by_vendor: dict[str, int] = {} - if last: - rows = ( - db.query(ConfigSyncTask.vendor, func.count()) - .filter(ConfigSyncTask.cycle_id == last.id, ConfigSyncTask.status == "fail") - .group_by(ConfigSyncTask.vendor) - .all() - ) - for vendor, cnt in rows: - fail_by_vendor[str(vendor or "unknown") or "unknown"] = int(cnt) - return ConfigSyncDashboardOut( - policy=policy_to_out(policy), - snapshot_count=snap_count, - last_cycle=cycle_to_out(last) if last else None, - running_cycle=cycle_to_out(running) if running else None, - next_due_at=next_due_at(db, policy), - fail_by_vendor=fail_by_vendor, - ) - - -def _snap_meta(row: NeConfigSnapshot) -> NeConfigSnapshotMetaOut: - cmds = row.commands_json if isinstance(row.commands_json, list) else [] - return NeConfigSnapshotMetaOut( - source=str(row.source), - target_id=str(row.target_id), - vendor=str(row.vendor or ""), - device_type=str(row.device_type or ""), - ne_name=str(row.ne_name or ""), - ne_ip=str(row.ne_ip or ""), - config_sha256=str(row.config_sha256 or ""), - config_alt_sha256=str(row.config_alt_sha256 or ""), - plain_size=int(row.plain_size or 0), - plain_alt_size=int(row.plain_alt_size or 0), - zlib_size=int(row.zlib_size or 0), - zlib_alt_size=int(row.zlib_alt_size or 0), - has_alt=bool(row.config_alt_zlib), - commands=[str(c) for c in cmds], - collected_at=row.collected_at, - last_cycle_id=str(row.last_cycle_id or ""), - ) - - -def list_snapshots( - db: Session, - *, - page: int, - page_size: int, - keyword: str = "", - source: str = "", - vendor: str = "", -) -> dict[str, Any]: - q = db.query(NeConfigSnapshot) - src = str(source or "").strip().lower() - if src in ("managed", "ume"): - q = q.filter(NeConfigSnapshot.source == src) - vend = str(vendor or "").strip() - if vend: - q = q.filter(NeConfigSnapshot.vendor.ilike(f"%{vend}%")) - kw = str(keyword or "").strip() - if kw: - like = f"%{kw}%" - q = q.filter( - or_( - NeConfigSnapshot.ne_name.ilike(like), - NeConfigSnapshot.ne_ip.ilike(like), - NeConfigSnapshot.target_id.ilike(like), - ) - ) - total = int(q.count()) - rows = q.order_by(NeConfigSnapshot.collected_at.desc()).offset((page - 1) * page_size).limit(page_size).all() - return {"total": total, "page": page, "page_size": page_size, "items": [_snap_meta(r) for r in rows]} - - -def get_snapshot_detail( - db: Session, - source: str, - target_id: str, - *, - field: str = "both", -) -> NeConfigSnapshotDetailOut: - src = str(source or "").strip().lower() - tid = str(target_id or "").strip() - row = db.get(NeConfigSnapshot, {"source": src, "target_id": tid}) - if not row: - raise HTTPException(status_code=404, detail="snapshot_not_found") - meta = _snap_meta(row) - primary = "" - alt = "" - f = str(field or "both").strip().lower() - if f in ("primary", "both", ""): - primary = decompress_text(row.config_zlib) - if f in ("alt", "both") and row.config_alt_zlib: - alt = decompress_text(row.config_alt_zlib) - return NeConfigSnapshotDetailOut(**meta.model_dump(), config_text=primary, config_alt_text=alt) - - -def _safe_export_part(text: str) -> str: - s = re.sub(r'[<>:"/\\|?*\s]+', "_", str(text or "").strip()) - return (s[:80] or "ne").strip("._") or "ne" - - -def build_snapshot_export( - db: Session, - source: str, - target_id: str, - *, - field: str = "primary", -) -> tuple[str, bytes, str]: - """Return (filename, payload, media_type) for download.""" - detail = get_snapshot_detail(db, source, target_id, field="both") - name = _safe_export_part(detail.ne_name or detail.target_id) - ip = _safe_export_part(detail.ne_ip or "ip") - base = f"{name}-{ip}-{detail.source}" - f = str(field or "primary").strip().lower() - - if f == "alt": - if not detail.has_alt or not detail.config_alt_text: - raise HTTPException(status_code=404, detail="alt_config_not_found") - filename = f"{base}-hierarchical.txt" - return filename, detail.config_alt_text.encode("utf-8"), "text/plain; charset=utf-8" - - if f == "both" and detail.has_alt and detail.config_alt_text: - buf = io.BytesIO() - with zipfile.ZipFile(buf, "w", compression=zipfile.ZIP_DEFLATED) as zf: - zf.writestr(f"{base}-set.txt", detail.config_text or "") - zf.writestr(f"{base}-hierarchical.txt", detail.config_alt_text or "") - return f"{base}-configs.zip", buf.getvalue(), "application/zip" - - filename = f"{base}-config.txt" - return filename, (detail.config_text or "").encode("utf-8"), "text/plain; charset=utf-8" - - -def list_snapshot_history( - db: Session, - source: str, - target_id: str, - *, - page: int, - page_size: int, -) -> dict[str, Any]: - src = str(source or "").strip().lower() - tid = str(target_id or "").strip() - q = ( - db.query(NeConfigHistory) - .filter(NeConfigHistory.source == src, NeConfigHistory.target_id == tid) - .order_by(NeConfigHistory.collected_at.desc()) - ) - total = int(q.count()) - rows = q.offset((page - 1) * page_size).limit(page_size).all() - items: list[NeConfigHistoryOut] = [] - for row in rows: - cmds = row.commands_json if isinstance(row.commands_json, list) else [] - items.append( - NeConfigHistoryOut( - id=str(row.id), - source=str(row.source), - target_id=str(row.target_id), - vendor=str(row.vendor or ""), - device_type=str(row.device_type or ""), - ne_name=str(row.ne_name or ""), - ne_ip=str(row.ne_ip or ""), - config_sha256=str(row.config_sha256 or ""), - config_alt_sha256=str(row.config_alt_sha256 or ""), - plain_size=int(row.plain_size or 0), - plain_alt_size=int(row.plain_alt_size or 0), - zlib_size=int(row.zlib_size or 0), - zlib_alt_size=int(row.zlib_alt_size or 0), - has_alt=bool(row.config_alt_zlib), - commands=[str(c) for c in cmds], - collected_at=row.collected_at, - last_cycle_id=str(row.cycle_id or ""), - cycle_id=str(row.cycle_id or ""), - ) - ) - return {"total": total, "page": page, "page_size": page_size, "items": items} - - -def sync_cycle_progress(db: Session, cycle_id: str) -> None: - cycle = db.get(ConfigSyncCycle, cycle_id) - if not cycle: - return - success = ( - db.query(func.count()) - .select_from(ConfigSyncTask) - .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status == "success") - .scalar() - or 0 - ) - fail = ( - db.query(func.count()) - .select_from(ConfigSyncTask) - .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status == "fail") - .scalar() - or 0 - ) - skip = ( - db.query(func.count()) - .select_from(ConfigSyncTask) - .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status.in_(("skipped", "cancelled"))) - .scalar() - or 0 - ) - cycle.success_count = int(success) - cycle.fail_count = int(fail) - cycle.skip_count = int(skip) - db.commit() - - -def finalize_cycle(db: Session, cycle_id: str) -> None: - cycle = db.get(ConfigSyncCycle, cycle_id) - if not cycle: - return - if str(cycle.status) in ("paused", "cancelled"): - return - pending = ( - db.query(func.count()) - .select_from(ConfigSyncTask) - .filter(ConfigSyncTask.cycle_id == cycle_id, ConfigSyncTask.status.in_(("pending", "running"))) - .scalar() - or 0 - ) - if int(pending) > 0: - return - sync_cycle_progress(db, cycle_id) - db.refresh(cycle) - if str(cycle.status) in ("paused", "cancelled"): - return - # Cycle outcome is about finishing the run, not per-NE results. - # Individual task failures stay in fail_count for retry/dashboard. - cycle.status = "success" - if cycle.error_message == "completed_with_failures": - cycle.error_message = "" - cycle.ended_at = _utcnow() - db.commit() - try: - prune_config_sync_cycles(db, keep=_cycle_keep_value(ensure_policy(db))) - except Exception: - _log.exception("prune_config_sync_cycles after finish failed") +from .config_sync_snapshots import ( + build_snapshot_export, + get_snapshot_detail, + list_snapshot_history, + list_snapshots, +) + +__all__ = [ + "DEFAULT_CYCLE_KEEP", + "POLICY_ID", + "_cycle_keep_value", + "build_snapshot_export", + "create_cycle", + "dashboard", + "ensure_policy", + "expand_targets", + "finalize_cycle", + "get_cycle", + "get_policy", + "get_snapshot_detail", + "has_active_cycle", + "has_running_cycle", + "last_finished_cycle", + "list_cycle_tasks", + "list_cycles", + "list_snapshot_history", + "list_snapshots", + "next_due_at", + "pause_cycle", + "policy_to_out", + "prune_config_sync_cycles", + "resume_cycle", + "stop_cycle", + "sync_cycle_progress", + "update_policy", +] diff --git a/netx_api/config_sync_snapshots.py b/netx_api/config_sync_snapshots.py new file mode 100644 index 0000000..97299af --- /dev/null +++ b/netx_api/config_sync_snapshots.py @@ -0,0 +1,177 @@ +"""Config sync snapshot list/detail/export/history.""" +from __future__ import annotations + +import io +import re +import zipfile +from typing import Any + +from fastapi import HTTPException +from sqlalchemy import or_ +from sqlalchemy.orm import Session + +from .config_sync_codec import decompress_text +from .config_sync_schemas import ( + NeConfigHistoryOut, + NeConfigSnapshotDetailOut, + NeConfigSnapshotMetaOut, +) +from .models import NeConfigHistory, NeConfigSnapshot + +def _snap_meta(row: NeConfigSnapshot) -> NeConfigSnapshotMetaOut: + cmds = row.commands_json if isinstance(row.commands_json, list) else [] + return NeConfigSnapshotMetaOut( + source=str(row.source), + target_id=str(row.target_id), + vendor=str(row.vendor or ""), + device_type=str(row.device_type or ""), + ne_name=str(row.ne_name or ""), + ne_ip=str(row.ne_ip or ""), + config_sha256=str(row.config_sha256 or ""), + config_alt_sha256=str(row.config_alt_sha256 or ""), + plain_size=int(row.plain_size or 0), + plain_alt_size=int(row.plain_alt_size or 0), + zlib_size=int(row.zlib_size or 0), + zlib_alt_size=int(row.zlib_alt_size or 0), + has_alt=bool(row.config_alt_zlib), + commands=[str(c) for c in cmds], + collected_at=row.collected_at, + last_cycle_id=str(row.last_cycle_id or ""), + ) + + +def list_snapshots( + db: Session, + *, + page: int, + page_size: int, + keyword: str = "", + source: str = "", + vendor: str = "", +) -> dict[str, Any]: + q = db.query(NeConfigSnapshot) + src = str(source or "").strip().lower() + if src in ("managed", "ume"): + q = q.filter(NeConfigSnapshot.source == src) + vend = str(vendor or "").strip() + if vend: + q = q.filter(NeConfigSnapshot.vendor.ilike(f"%{vend}%")) + kw = str(keyword or "").strip() + if kw: + like = f"%{kw}%" + q = q.filter( + or_( + NeConfigSnapshot.ne_name.ilike(like), + NeConfigSnapshot.ne_ip.ilike(like), + NeConfigSnapshot.target_id.ilike(like), + ) + ) + total = int(q.count()) + rows = q.order_by(NeConfigSnapshot.collected_at.desc()).offset((page - 1) * page_size).limit(page_size).all() + return {"total": total, "page": page, "page_size": page_size, "items": [_snap_meta(r) for r in rows]} + + +def get_snapshot_detail( + db: Session, + source: str, + target_id: str, + *, + field: str = "both", +) -> NeConfigSnapshotDetailOut: + src = str(source or "").strip().lower() + tid = str(target_id or "").strip() + row = db.get(NeConfigSnapshot, {"source": src, "target_id": tid}) + if not row: + raise HTTPException(status_code=404, detail="snapshot_not_found") + meta = _snap_meta(row) + primary = "" + alt = "" + f = str(field or "both").strip().lower() + if f in ("primary", "both", ""): + primary = decompress_text(row.config_zlib) + if f in ("alt", "both") and row.config_alt_zlib: + alt = decompress_text(row.config_alt_zlib) + return NeConfigSnapshotDetailOut(**meta.model_dump(), config_text=primary, config_alt_text=alt) + + +def _safe_export_part(text: str) -> str: + s = re.sub(r'[<>:"/\\|?*\s]+', "_", str(text or "").strip()) + return (s[:80] or "ne").strip("._") or "ne" + + +def build_snapshot_export( + db: Session, + source: str, + target_id: str, + *, + field: str = "primary", +) -> tuple[str, bytes, str]: + """Return (filename, payload, media_type) for download.""" + detail = get_snapshot_detail(db, source, target_id, field="both") + name = _safe_export_part(detail.ne_name or detail.target_id) + ip = _safe_export_part(detail.ne_ip or "ip") + base = f"{name}-{ip}-{detail.source}" + f = str(field or "primary").strip().lower() + + if f == "alt": + if not detail.has_alt or not detail.config_alt_text: + raise HTTPException(status_code=404, detail="alt_config_not_found") + filename = f"{base}-hierarchical.txt" + return filename, detail.config_alt_text.encode("utf-8"), "text/plain; charset=utf-8" + + if f == "both" and detail.has_alt and detail.config_alt_text: + buf = io.BytesIO() + with zipfile.ZipFile(buf, "w", compression=zipfile.ZIP_DEFLATED) as zf: + zf.writestr(f"{base}-set.txt", detail.config_text or "") + zf.writestr(f"{base}-hierarchical.txt", detail.config_alt_text or "") + return f"{base}-configs.zip", buf.getvalue(), "application/zip" + + filename = f"{base}-config.txt" + return filename, (detail.config_text or "").encode("utf-8"), "text/plain; charset=utf-8" + + +def list_snapshot_history( + db: Session, + source: str, + target_id: str, + *, + page: int, + page_size: int, +) -> dict[str, Any]: + src = str(source or "").strip().lower() + tid = str(target_id or "").strip() + q = ( + db.query(NeConfigHistory) + .filter(NeConfigHistory.source == src, NeConfigHistory.target_id == tid) + .order_by(NeConfigHistory.collected_at.desc()) + ) + total = int(q.count()) + rows = q.offset((page - 1) * page_size).limit(page_size).all() + items: list[NeConfigHistoryOut] = [] + for row in rows: + cmds = row.commands_json if isinstance(row.commands_json, list) else [] + items.append( + NeConfigHistoryOut( + id=str(row.id), + source=str(row.source), + target_id=str(row.target_id), + vendor=str(row.vendor or ""), + device_type=str(row.device_type or ""), + ne_name=str(row.ne_name or ""), + ne_ip=str(row.ne_ip or ""), + config_sha256=str(row.config_sha256 or ""), + config_alt_sha256=str(row.config_alt_sha256 or ""), + plain_size=int(row.plain_size or 0), + plain_alt_size=int(row.plain_alt_size or 0), + zlib_size=int(row.zlib_size or 0), + zlib_alt_size=int(row.zlib_alt_size or 0), + has_alt=bool(row.config_alt_zlib), + commands=[str(c) for c in cmds], + collected_at=row.collected_at, + last_cycle_id=str(row.cycle_id or ""), + cycle_id=str(row.cycle_id or ""), + ) + ) + return {"total": total, "page": page, "page_size": page_size, "items": items} + + diff --git a/netx_api/topology_classify.py b/netx_api/topology_classify.py index 62f389a..51c9d96 100644 --- a/netx_api/topology_classify.py +++ b/netx_api/topology_classify.py @@ -1,927 +1,35 @@ """Regex-based fabric role/region classification + slice map generation.""" - from __future__ import annotations -import re -from datetime import datetime -from typing import Any -from uuid import uuid4 - -from fastapi import HTTPException -from sqlalchemy.orm import Session - -from .models import TopoClassifyRule, TopoFabricEdge, TopoFabricNode, TopoFolder, TopoView, TopoViewNode -from .topology_membership import ( - VIEW_KIND_CUSTOM, - VIEW_ROLE_ACCESS, - VIEW_ROLE_AGGREGATION, - VIEW_ROLE_CORE, - VIEW_ROLES, - merge_filter_with_membership, - normalize_view_role, +from .topology_classify_apply import ( + apply_classify, + apply_classify_empty_only, + bulk_tag_fabric_nodes, + list_unmatched, + match_fabric_nodes, + patch_fabric_node_tags, + preview_classify, ) -from .topology_schemas import ( - ClassifyApplyOut, - ClassifyPreviewOut, - ClassifyRuleCreate, - ClassifyRuleOut, - ClassifyRuleUpdate, - FabricNodeOut, - FabricNodesBulkTagOut, - FabricNodesBulkTagRequest, - FabricNodesMatchOut, - FabricNodesMatchRequest, - FabricNodeTagPatch, - SliceGenerateOut, - SliceGenerateRequest, - SliceMapPlan, - TopologyFolderCreate, - TopologyViewCreate, +from .topology_classify_rules import create_rule, delete_rule, list_rules, update_rule +from .topology_classify_slices import ( + generate_slices, + preview_slices, + search_fabric_nodes_with_views, ) -_MAX_PATTERN_LEN = 512 -_ROLE_VALUES = VIEW_ROLES | {"unknown"} -_MATCH_FIELDS = frozenset({"name", "ip", "name_ip"}) -_SCOPES = frozenset({"role", "region"}) -_SLICE_TEMPLATES = frozenset({"core_only", "core_agg", "agg_access"}) - - -def _utcnow() -> datetime: - return datetime.utcnow() - - -def _compile_pattern(pattern: str) -> re.Pattern[str]: - p = str(pattern or "").strip() - if not p: - raise HTTPException(status_code=400, detail="pattern_required") - if len(p) > _MAX_PATTERN_LEN: - raise HTTPException(status_code=400, detail="pattern_too_long") - try: - return re.compile(p, re.IGNORECASE) - except re.error as exc: - raise HTTPException(status_code=400, detail=f"invalid_pattern:{exc}") from exc - - -def _match_text(node: TopoFabricNode, match_field: str) -> str: - name = str(node.name or "") - ip = str(node.ip or "") - mf = str(match_field or "name").strip().lower() - if mf == "ip": - return ip - if mf == "name_ip": - return f"{name} {ip}".strip() - return name - - -def _rule_out(row: TopoClassifyRule) -> ClassifyRuleOut: - return ClassifyRuleOut( - id=row.id, - scope=str(row.scope or "role"), - name=str(row.name or ""), - pattern=str(row.pattern or ""), - match_field=str(row.match_field or "name"), - priority=int(row.priority or 100), - enabled=bool(row.enabled), - payload=dict(row.payload or {}), - remark=str(row.remark or ""), - created_at=row.created_at, - updated_at=row.updated_at, - ) - - -def _validate_payload(scope: str, payload: dict[str, Any]) -> dict[str, Any]: - out = dict(payload or {}) - if scope == "role": - role = normalize_view_role(str(out.get("role") or "")) - if str(out.get("role") or "").strip().lower() not in VIEW_ROLES: - raise HTTPException(status_code=400, detail="role_payload_invalid") - return {"role": role} - if "folder_id" in out and str(out.get("folder_id") or "").strip(): - return {"folder_id": str(out["folder_id"]).strip()} - if "region_name_from_group" in out: - try: - g = int(out.get("region_name_from_group")) - except (TypeError, ValueError) as exc: - raise HTTPException(status_code=400, detail="region_group_invalid") from exc - if g < 1: - raise HTTPException(status_code=400, detail="region_group_invalid") - return {"region_name_from_group": g} - raise HTTPException(status_code=400, detail="region_payload_invalid") - - -def list_rules(db: Session, *, scope: str = "") -> list[ClassifyRuleOut]: - q = db.query(TopoClassifyRule) - if scope.strip(): - q = q.filter(TopoClassifyRule.scope == scope.strip().lower()) - rows = q.order_by( - TopoClassifyRule.scope.asc(), - TopoClassifyRule.priority.asc(), - TopoClassifyRule.name.asc(), - ).all() - return [_rule_out(r) for r in rows] - - -def create_rule(db: Session, body: ClassifyRuleCreate) -> ClassifyRuleOut: - scope = str(body.scope or "").strip().lower() - if scope not in _SCOPES: - raise HTTPException(status_code=400, detail="scope_invalid") - match_field = str(body.match_field or "name").strip().lower() - if match_field not in _MATCH_FIELDS: - raise HTTPException(status_code=400, detail="match_field_invalid") - _compile_pattern(body.pattern) - payload = _validate_payload(scope, dict(body.payload or {})) - if scope == "region" and payload.get("folder_id"): - folder = db.get(TopoFolder, payload["folder_id"]) - if folder is None or str(folder.kind or "") != "region": - raise HTTPException(status_code=400, detail="folder_not_found") - now = _utcnow() - row = TopoClassifyRule( - id=uuid4().hex, - scope=scope, - name=str(body.name or "").strip()[:256] or f"{scope}-rule", - pattern=str(body.pattern or "").strip()[:_MAX_PATTERN_LEN], - match_field=match_field, - priority=int(body.priority if body.priority is not None else 100), - enabled=bool(body.enabled if body.enabled is not None else True), - payload=payload, - remark=str(body.remark or "")[:512], - created_at=now, - updated_at=now, - ) - db.add(row) - db.commit() - db.refresh(row) - return _rule_out(row) - - -def update_rule(db: Session, rule_id: str, body: ClassifyRuleUpdate) -> ClassifyRuleOut: - row = db.get(TopoClassifyRule, rule_id) - if row is None: - raise HTTPException(status_code=404, detail="rule_not_found") - if body.name is not None: - row.name = str(body.name or "").strip()[:256] - if body.pattern is not None: - _compile_pattern(body.pattern) - row.pattern = str(body.pattern or "").strip()[:_MAX_PATTERN_LEN] - if body.match_field is not None: - mf = str(body.match_field or "name").strip().lower() - if mf not in _MATCH_FIELDS: - raise HTTPException(status_code=400, detail="match_field_invalid") - row.match_field = mf - if body.priority is not None: - row.priority = int(body.priority) - if body.enabled is not None: - row.enabled = bool(body.enabled) - if body.remark is not None: - row.remark = str(body.remark or "")[:512] - if body.payload is not None: - row.payload = _validate_payload(str(row.scope or "role"), dict(body.payload or {})) - if str(row.scope) == "region" and row.payload.get("folder_id"): - folder = db.get(TopoFolder, row.payload["folder_id"]) - if folder is None or str(folder.kind or "") != "region": - raise HTTPException(status_code=400, detail="folder_not_found") - row.updated_at = _utcnow() - db.commit() - db.refresh(row) - return _rule_out(row) - - -def delete_rule(db: Session, rule_id: str) -> dict[str, Any]: - row = db.get(TopoClassifyRule, rule_id) - if row is None: - raise HTTPException(status_code=404, detail="rule_not_found") - db.delete(row) - db.commit() - return {"ok": True, "id": rule_id, "deleted": True} - - -def _enabled_rules(db: Session, scope: str) -> list[tuple[TopoClassifyRule, re.Pattern[str]]]: - rows = ( - db.query(TopoClassifyRule) - .filter(TopoClassifyRule.scope == scope, TopoClassifyRule.enabled.is_(True)) - .order_by(TopoClassifyRule.priority.asc(), TopoClassifyRule.name.asc()) - .all() - ) - out: list[tuple[TopoClassifyRule, re.Pattern[str]]] = [] - for r in rows: - try: - out.append((r, _compile_pattern(r.pattern))) - except HTTPException: - continue - return out - - -def _ensure_region_by_name(db: Session, name: str) -> TopoFolder: - from .topology_service import bootstrap_topology_tree, create_folder - - name = str(name or "").strip()[:256] - if not name: - raise HTTPException(status_code=400, detail="region_name_empty") - existing = ( - db.query(TopoFolder) - .filter(TopoFolder.kind == "region", TopoFolder.name == name) - .first() - ) - if existing is not None: - return existing - bootstrap_topology_tree(db) - created = create_folder(db, TopologyFolderCreate(name=name, kind="region")) - folder = db.get(TopoFolder, created.id) - assert folder is not None - return folder - - -def _resolve_role_hit( - node: TopoFabricNode, rules: list[tuple[TopoClassifyRule, re.Pattern[str]]] -) -> tuple[str | None, str | None, bool]: - """Return (role, rule_id, multi_hit).""" - hits: list[tuple[str, str]] = [] - for rule, cre in rules: - text = _match_text(node, rule.match_field) - if not text: - continue - if cre.search(text): - role = normalize_view_role(str((rule.payload or {}).get("role") or "")) - hits.append((role, rule.id)) - if not hits: - return None, None, False - return hits[0][0], hits[0][1], len(hits) > 1 - - -def _resolve_region_hit( - db: Session, - node: TopoFabricNode, - rules: list[tuple[TopoClassifyRule, re.Pattern[str]]], - *, - create_missing: bool, -) -> tuple[str | None, str | None, bool]: - hits: list[tuple[str, str]] = [] - for rule, cre in rules: - text = _match_text(node, rule.match_field) - if not text: - continue - m = cre.search(text) - if not m: - continue - payload = dict(rule.payload or {}) - folder_id = str(payload.get("folder_id") or "").strip() - if folder_id: - hits.append((folder_id, rule.id)) - continue - g = int(payload.get("region_name_from_group") or 0) - try: - region_name = m.group(g) - except IndexError: - continue - region_name = str(region_name or "").strip() - if not region_name: - continue - if create_missing: - folder = _ensure_region_by_name(db, region_name) - hits.append((folder.id, rule.id)) - else: - existing = ( - db.query(TopoFolder) - .filter(TopoFolder.kind == "region", TopoFolder.name == region_name) - .first() - ) - hits.append((existing.id if existing else f"new:{region_name}", rule.id)) - if not hits: - return None, None, False - return hits[0][0], hits[0][1], len(hits) > 1 - - -def preview_classify(db: Session, *, sample_limit: int = 20) -> ClassifyPreviewOut: - role_rules = _enabled_rules(db, "role") - region_rules = _enabled_rules(db, "region") - nodes = db.query(TopoFabricNode).order_by(TopoFabricNode.name.asc()).all() - role_matched = role_unmatched = role_conflict = 0 - region_matched = region_unmatched = region_conflict = 0 - role_samples: list[dict[str, Any]] = [] - region_samples: list[dict[str, Any]] = [] - unmatched_samples: list[dict[str, Any]] = [] - - for n in nodes: - role, _rid, multi_r = _resolve_role_hit(n, role_rules) - if role is None: - role_unmatched += 1 - else: - role_matched += 1 - if multi_r: - role_conflict += 1 - if len(role_samples) < sample_limit: - role_samples.append( - { - "fabric_node_id": n.id, - "name": n.name, - "ip": n.ip, - "role": role, - "multi_hit": multi_r, - } - ) - - region_id, _rrid, multi_reg = _resolve_region_hit( - db, n, region_rules, create_missing=False - ) - if region_id is None: - region_unmatched += 1 - else: - region_matched += 1 - if multi_reg: - region_conflict += 1 - if len(region_samples) < sample_limit: - region_samples.append( - { - "fabric_node_id": n.id, - "name": n.name, - "ip": n.ip, - "region": region_id, - "multi_hit": multi_reg, - } - ) - - if role is None and region_id is None and len(unmatched_samples) < sample_limit: - unmatched_samples.append( - {"fabric_node_id": n.id, "name": n.name, "ip": n.ip, "vendor": n.vendor} - ) - - return ClassifyPreviewOut( - total_nodes=len(nodes), - role_matched=role_matched, - role_unmatched=role_unmatched, - role_conflicts=role_conflict, - region_matched=region_matched, - region_unmatched=region_unmatched, - region_conflicts=region_conflict, - role_samples=role_samples, - region_samples=region_samples, - unmatched_samples=unmatched_samples, - ) - - -def apply_classify( - db: Session, - *, - skip_manual: bool = True, - fill_empty_only: bool = False, -) -> ClassifyApplyOut: - role_rules = _enabled_rules(db, "role") - region_rules = _enabled_rules(db, "region") - nodes = db.query(TopoFabricNode).all() - role_updated = region_updated = skipped_manual = 0 - - for n in nodes: - role, _, _ = _resolve_role_hit(n, role_rules) - if role is not None: - if skip_manual and str(n.role_source or "") == "manual": - skipped_manual += 1 - elif fill_empty_only and str(n.role or "").strip(): - pass - else: - n.role = role - n.role_source = "rule" - n.updated_at = _utcnow() - role_updated += 1 - elif not str(n.role or "").strip() and str(n.role_source or "") != "manual": - n.role = "unknown" - n.role_source = "rule" - n.updated_at = _utcnow() - - region_id, _, _ = _resolve_region_hit(db, n, region_rules, create_missing=True) - if region_id is not None and not str(region_id).startswith("new:"): - if skip_manual and str(n.region_source or "") == "manual": - skipped_manual += 1 - elif fill_empty_only and str(n.region_folder_id or "").strip(): - pass - else: - n.region_folder_id = region_id - n.region_source = "rule" - n.updated_at = _utcnow() - region_updated += 1 - - db.commit() - return ClassifyApplyOut( - role_updated=role_updated, - region_updated=region_updated, - skipped_manual=skipped_manual, - total_nodes=len(nodes), - ) - - -def list_unmatched( - db: Session, - *, - kind: str = "any", - page: int = 1, - page_size: int = 50, -) -> dict[str, Any]: - from sqlalchemy import or_ - - q = db.query(TopoFabricNode) - k = str(kind or "any").strip().lower() - role_miss = or_(TopoFabricNode.role == "", TopoFabricNode.role == "unknown") - region_miss = or_( - TopoFabricNode.region_folder_id.is_(None), - TopoFabricNode.region_folder_id == "", - ) - if k == "role": - q = q.filter(role_miss) - elif k == "region": - q = q.filter(region_miss) - else: - q = q.filter(or_(role_miss, region_miss)) - total = q.count() - rows = ( - q.order_by(TopoFabricNode.name.asc()) - .offset(max(0, (page - 1) * page_size)) - .limit(page_size) - .all() - ) - from .topology_service import _node_out - - return { - "total": total, - "page": page, - "page_size": page_size, - "items": [_node_out(n).model_dump() for n in rows], - } - - -def patch_fabric_node_tags( - db: Session, fabric_node_id: str, body: FabricNodeTagPatch -) -> FabricNodeOut: - from .topology_service import _node_out - - n = db.get(TopoFabricNode, fabric_node_id) - if n is None: - raise HTTPException(status_code=404, detail="fabric_node_not_found") - if body.role is not None: - role = str(body.role or "").strip().lower() - if role and role not in _ROLE_VALUES: - raise HTTPException(status_code=400, detail="role_invalid") - n.role = role - n.role_source = "manual" - if body.region_folder_id is not None: - fid = str(body.region_folder_id or "").strip() - if fid: - folder = db.get(TopoFolder, fid) - if folder is None or str(folder.kind or "") != "region": - raise HTTPException(status_code=400, detail="folder_not_found") - n.region_folder_id = fid - else: - n.region_folder_id = None - n.region_source = "manual" - n.updated_at = _utcnow() - db.commit() - db.refresh(n) - return _node_out(n) - - -def _iter_regex_matches( - db: Session, *, pattern: str, match_field: str -) -> list[TopoFabricNode]: - cre = _compile_pattern(pattern) - mf = str(match_field or "name").strip().lower() - if mf not in _MATCH_FIELDS: - raise HTTPException(status_code=400, detail="match_field_invalid") - out: list[TopoFabricNode] = [] - for n in db.query(TopoFabricNode).order_by(TopoFabricNode.name.asc()).all(): - text = _match_text(n, mf) - if text and cre.search(text): - out.append(n) - return out - - -def match_fabric_nodes(db: Session, body: FabricNodesMatchRequest) -> FabricNodesMatchOut: - """Ephemeral regex lookup — does not persist rules.""" - from .topology_inventory_lifecycle import fabric_link_status - - matched = _iter_regex_matches( - db, pattern=body.pattern, match_field=body.match_field - ) - limit = max(1, min(200, int(body.sample_limit or 50))) - samples = [ - { - "fabric_node_id": n.id, - "name": n.name, - "ip": n.ip, - "role": n.role or "", - "region_folder_id": n.region_folder_id or "", - "link_status": fabric_link_status(n), - } - for n in matched[:limit] - ] - return FabricNodesMatchOut( - pattern=str(body.pattern or "").strip(), - match_field=str(body.match_field or "name"), - total_matched=len(matched), - samples=samples, - fabric_node_ids=[n.id for n in matched], - ) - - -def bulk_tag_fabric_nodes( - db: Session, body: FabricNodesBulkTagRequest -) -> FabricNodesBulkTagOut: - """Assign role/region after user confirms a regex or explicit selection.""" - if body.role is None and body.region_folder_id is None: - raise HTTPException(status_code=400, detail="role_or_region_required") - - role_v: str | None = None - if body.role is not None: - role_v = str(body.role or "").strip().lower() - if role_v and role_v not in _ROLE_VALUES: - raise HTTPException(status_code=400, detail="role_invalid") - - region_v: str | None = None - if body.region_folder_id is not None: - region_v = str(body.region_folder_id or "").strip() - if region_v: - folder = db.get(TopoFolder, region_v) - if folder is None or str(folder.kind or "") != "region": - raise HTTPException(status_code=400, detail="folder_not_found") - else: - region_v = "" - - ids = [str(x).strip() for x in (body.fabric_node_ids or []) if str(x).strip()] - if str(body.pattern or "").strip(): - matched_nodes = _iter_regex_matches( - db, pattern=body.pattern, match_field=body.match_field - ) - elif ids: - matched_nodes = ( - db.query(TopoFabricNode) - .filter(TopoFabricNode.id.in_(ids)) - .order_by(TopoFabricNode.name.asc()) - .all() - ) - else: - raise HTTPException(status_code=400, detail="ids_or_pattern_required") - - samples = [ - { - "fabric_node_id": n.id, - "name": n.name, - "ip": n.ip, - "role": n.role or "", - "region_folder_id": n.region_folder_id or "", - } - for n in matched_nodes[:50] - ] - if body.dry_run: - return FabricNodesBulkTagOut( - dry_run=True, - matched=len(matched_nodes), - updated=0, - role=role_v, - region_folder_id=region_v, - samples=samples, - ) - - now = _utcnow() - updated = 0 - for n in matched_nodes: - if role_v is not None: - n.role = role_v - n.role_source = "manual" - if region_v is not None: - n.region_folder_id = region_v or None - n.region_source = "manual" - n.updated_at = now - updated += 1 - db.commit() - return FabricNodesBulkTagOut( - dry_run=False, - matched=len(matched_nodes), - updated=updated, - role=role_v, - region_folder_id=region_v, - samples=samples, - ) - - -def apply_classify_empty_only(db: Session) -> ClassifyApplyOut: - """Incremental classify for newly synced nodes (fill empty tags only).""" - return apply_classify(db, skip_manual=True, fill_empty_only=True) - - -# --- Slice generation ------------------------------------------------------- - - -def _active_neighbors(db: Session, seed_ids: set[str], *, hops: int = 1) -> set[str]: - if not seed_ids or hops <= 0: - return set() - frontier = set(seed_ids) - found: set[str] = set() - for _ in range(hops): - if not frontier: - break - rows = ( - db.query(TopoFabricEdge) - .filter( - TopoFabricEdge.layer == "physical", - TopoFabricEdge.status == "active", - (TopoFabricEdge.a_node_id.in_(frontier) | TopoFabricEdge.b_node_id.in_(frontier)), - ) - .all() - ) - nxt: set[str] = set() - for e in rows: - for a, b in ((e.a_node_id, e.b_node_id), (e.b_node_id, e.a_node_id)): - if a in frontier and b not in seed_ids and b not in found: - nxt.add(str(b)) - found |= nxt - frontier = nxt - return found - - -def _connected_components(db: Session, node_ids: list[str]) -> list[list[str]]: - ids = [str(x) for x in node_ids if str(x)] - if not ids: - return [] - id_set = set(ids) - adj: dict[str, set[str]] = {i: set() for i in ids} - rows = ( - db.query(TopoFabricEdge) - .filter( - TopoFabricEdge.layer == "physical", - TopoFabricEdge.status == "active", - TopoFabricEdge.a_node_id.in_(ids), - TopoFabricEdge.b_node_id.in_(ids), - ) - .all() - ) - for e in rows: - a, b = str(e.a_node_id), str(e.b_node_id) - if a in id_set and b in id_set: - adj[a].add(b) - adj[b].add(a) - seen: set[str] = set() - comps: list[list[str]] = [] - for nid in ids: - if nid in seen: - continue - stack = [nid] - seen.add(nid) - comp: list[str] = [] - while stack: - cur = stack.pop() - comp.append(cur) - for nb in adj.get(cur, ()): - if nb not in seen: - seen.add(nb) - stack.append(nb) - comps.append(sorted(comp)) - return comps - - -def _nodes_in_region(db: Session, folder_id: str, *, role: str = "") -> list[TopoFabricNode]: - q = db.query(TopoFabricNode).filter(TopoFabricNode.region_folder_id == folder_id) - if role: - q = q.filter(TopoFabricNode.role == role) - return q.order_by(TopoFabricNode.name.asc()).all() - - -def preview_slices(db: Session, body: SliceGenerateRequest) -> SliceGenerateOut: - folder = db.get(TopoFolder, body.folder_id) - if folder is None or str(folder.kind or "") != "region": - raise HTTPException(status_code=400, detail="folder_not_found") - template = str(body.template or "").strip().lower() - if template not in _SLICE_TEMPLATES: - raise HTTPException(status_code=400, detail="template_invalid") - max_nodes = max(1, min(2000, int(body.max_nodes or 300))) - plans: list[SliceMapPlan] = [] - overlap_ids: set[str] = set() - seen_in_maps: dict[str, int] = {} - - def _track(ids: list[str]) -> None: - for i in ids: - seen_in_maps[i] = seen_in_maps.get(i, 0) + 1 - - if template == "core_only": - cores = _nodes_in_region(db, folder.id, role=VIEW_ROLE_CORE) - comps = _connected_components(db, [n.id for n in cores]) or [ - [n.id] for n in cores - ] - for idx, comp in enumerate(comps, start=1): - if len(comp) > max_nodes: - raise HTTPException( - status_code=400, - detail=f"slice_exceeds_max_nodes:{len(comp)}>{max_nodes}", - ) - name = f"Core-{idx}" if len(comps) > 1 else "Core" - plans.append( - SliceMapPlan( - name=name, - role=VIEW_ROLE_CORE, - seed_fabric_node_ids=comp, - member_fabric_node_ids=comp, - node_count=len(comp), - ) - ) - _track(comp) - elif template == "core_agg": - cores = _nodes_in_region(db, folder.id, role=VIEW_ROLE_CORE) - comps = _connected_components(db, [n.id for n in cores]) or [ - [n.id] for n in cores - ] - for idx, comp in enumerate(comps, start=1): - peers = _active_neighbors(db, set(comp), hops=1) - agg_ids = [ - p - for p in peers - if (fn := db.get(TopoFabricNode, p)) is not None - and str(fn.role or "") == VIEW_ROLE_AGGREGATION - and str(fn.region_folder_id or "") == folder.id - ] - members = sorted(set(comp) | set(agg_ids)) - if len(members) > max_nodes: - raise HTTPException( - status_code=400, - detail=f"slice_exceeds_max_nodes:{len(members)}>{max_nodes}", - ) - name = f"CoreAgg-{idx}" if len(comps) > 1 else "Core+Agg" - plans.append( - SliceMapPlan( - name=name, - role=VIEW_ROLE_CORE, - seed_fabric_node_ids=comp, - member_fabric_node_ids=members, - node_count=len(members), - ) - ) - _track(members) - else: # agg_access - aggs = _nodes_in_region(db, folder.id, role=VIEW_ROLE_AGGREGATION) - comps = _connected_components(db, [n.id for n in aggs]) or [ - [n.id] for n in aggs - ] - for idx, comp in enumerate(comps, start=1): - peers = _active_neighbors(db, set(comp), hops=1) - acc_ids = [ - p - for p in peers - if (fn := db.get(TopoFabricNode, p)) is not None - and str(fn.role or "") == VIEW_ROLE_ACCESS - and str(fn.region_folder_id or "") == folder.id - ] - members = sorted(set(comp) | set(acc_ids)) - if len(members) > max_nodes: - raise HTTPException( - status_code=400, - detail=f"slice_exceeds_max_nodes:{len(members)}>{max_nodes}", - ) - name = f"AggAccess-{idx}" if len(comps) > 1 else "Agg+Access" - plans.append( - SliceMapPlan( - name=name, - role=VIEW_ROLE_AGGREGATION, - seed_fabric_node_ids=comp, - member_fabric_node_ids=members, - node_count=len(members), - ) - ) - _track(members) - - overlap_ids = {nid for nid, cnt in seen_in_maps.items() if cnt > 1} - return SliceGenerateOut( - folder_id=folder.id, - template=template, - dry_run=True, - maps=plans, - map_count=len(plans), - overlap_node_count=len(overlap_ids), - created_view_ids=[], - ) - - -def generate_slices(db: Session, body: SliceGenerateRequest) -> SliceGenerateOut: - from .topology_service import create_view, _place_fabric_ids_on_view - - preview = preview_slices(db, body) - if body.dry_run: - return preview - - created: list[str] = [] - for plan in preview.maps: - view = create_view( - db, - TopologyViewCreate( - name=plan.name, - folder_id=body.folder_id, - kind=VIEW_KIND_CUSTOM, - role=plan.role, - remark=f"slice:{body.template}", - ), - ) - mem = { - "mode": "hybrid", - "seed_fabric_node_ids": list(plan.member_fabric_node_ids), - "expand_hops": 0, - "max_nodes": int(body.max_nodes or 300), - "frozen": True, - "managed_ne_ids": [], - "tags_any": [], - "vendors": [], - "device_types": [], - "keyword": "", - } - row = db.get(TopoView, view.id) - assert row is not None - row.filter = merge_filter_with_membership( - dict(row.filter or {}), role=normalize_view_role(plan.role), membership=mem - ) - _place_fabric_ids_on_view(db, row, list(plan.member_fabric_node_ids), existing=set()) - row.updated_at = _utcnow() - db.commit() - created.append(view.id) - - # Optionally seed physical overview with cores only - if body.seed_physical_cores: - from .topology_service import ensure_region_physical_view - - phys = ensure_region_physical_view(db, body.folder_id, commit=True) - cores = [n.id for n in _nodes_in_region(db, body.folder_id, role=VIEW_ROLE_CORE)] - existing = { - vn.fabric_node_id - for vn in db.query(TopoViewNode).filter(TopoViewNode.view_id == phys.id).all() - } - to_add = [c for c in cores if c not in existing][: int(body.max_nodes or 300)] - if to_add: - _place_fabric_ids_on_view(db, phys, to_add, existing=existing) - mem = merge_filter_with_membership( - dict(phys.filter or {}), - role=VIEW_ROLE_CORE, - kind="physical", - membership={ - **dict((phys.filter or {}).get("membership") or {}), - "frozen": True, - "max_nodes": int(body.max_nodes or 500), - }, - ) - phys.filter = mem - phys.updated_at = _utcnow() - db.commit() - - return SliceGenerateOut( - folder_id=body.folder_id, - template=str(body.template), - dry_run=False, - maps=preview.maps, - map_count=len(preview.maps), - overlap_node_count=preview.overlap_node_count, - created_view_ids=created, - ) - - -def search_fabric_nodes_with_views( - db: Session, - *, - keyword: str = "", - page: int = 1, - page_size: int = 50, -) -> dict[str, Any]: - from .topology_service import _node_out - - q = db.query(TopoFabricNode) - kw = str(keyword or "").strip() - if kw: - like = f"%{kw}%" - q = q.filter( - (TopoFabricNode.name.ilike(like)) - | (TopoFabricNode.ip.ilike(like)) - | (TopoFabricNode.vendor.ilike(like)) - ) - total = q.count() - rows = ( - q.order_by(TopoFabricNode.name.asc()) - .offset(max(0, (page - 1) * page_size)) - .limit(page_size) - .all() - ) - node_ids = [n.id for n in rows] - placements: dict[str, list[dict[str, Any]]] = {nid: [] for nid in node_ids} - if node_ids: - vnodes = ( - db.query(TopoViewNode, TopoView, TopoFolder) - .join(TopoView, TopoView.id == TopoViewNode.view_id) - .outerjoin(TopoFolder, TopoFolder.id == TopoView.folder_id) - .filter(TopoViewNode.fabric_node_id.in_(node_ids)) - .all() - ) - for vn, view, folder in vnodes: - placements.setdefault(vn.fabric_node_id, []).append( - { - "view_id": view.id, - "view_name": view.name, - "folder_id": view.folder_id or "", - "folder_name": (folder.name if folder else "") or "", - "kind": view.kind or "custom", - } - ) - items = [] - for n in rows: - d = _node_out(n).model_dump() - d["views"] = placements.get(n.id, []) - items.append(d) - return {"total": total, "page": page, "page_size": page_size, "items": items} +__all__ = [ + "apply_classify", + "apply_classify_empty_only", + "bulk_tag_fabric_nodes", + "create_rule", + "delete_rule", + "generate_slices", + "list_rules", + "list_unmatched", + "match_fabric_nodes", + "patch_fabric_node_tags", + "preview_classify", + "preview_slices", + "search_fabric_nodes_with_views", + "update_rule", +] diff --git a/netx_api/topology_classify_apply.py b/netx_api/topology_classify_apply.py new file mode 100644 index 0000000..1570cb4 --- /dev/null +++ b/netx_api/topology_classify_apply.py @@ -0,0 +1,347 @@ +"""Topology classify preview/apply and fabric node tagging.""" +from __future__ import annotations + +import re +from typing import Any + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .models import TopoClassifyRule, TopoFabricNode, TopoFolder +from .topology_classify_common import ( + _MATCH_FIELDS, + _ROLE_VALUES, + _compile_pattern, + _enabled_rules, + _ensure_region_by_name, + _match_text, + _resolve_region_hit, + _resolve_role_hit, + _utcnow, +) +from .topology_membership import normalize_view_role +from .topology_schemas import ( + ClassifyApplyOut, + ClassifyPreviewOut, + FabricNodeOut, + FabricNodesBulkTagOut, + FabricNodesBulkTagRequest, + FabricNodesMatchOut, + FabricNodesMatchRequest, + FabricNodeTagPatch, +) + +def preview_classify(db: Session, *, sample_limit: int = 20) -> ClassifyPreviewOut: + role_rules = _enabled_rules(db, "role") + region_rules = _enabled_rules(db, "region") + nodes = db.query(TopoFabricNode).order_by(TopoFabricNode.name.asc()).all() + role_matched = role_unmatched = role_conflict = 0 + region_matched = region_unmatched = region_conflict = 0 + role_samples: list[dict[str, Any]] = [] + region_samples: list[dict[str, Any]] = [] + unmatched_samples: list[dict[str, Any]] = [] + + for n in nodes: + role, _rid, multi_r = _resolve_role_hit(n, role_rules) + if role is None: + role_unmatched += 1 + else: + role_matched += 1 + if multi_r: + role_conflict += 1 + if len(role_samples) < sample_limit: + role_samples.append( + { + "fabric_node_id": n.id, + "name": n.name, + "ip": n.ip, + "role": role, + "multi_hit": multi_r, + } + ) + + region_id, _rrid, multi_reg = _resolve_region_hit( + db, n, region_rules, create_missing=False + ) + if region_id is None: + region_unmatched += 1 + else: + region_matched += 1 + if multi_reg: + region_conflict += 1 + if len(region_samples) < sample_limit: + region_samples.append( + { + "fabric_node_id": n.id, + "name": n.name, + "ip": n.ip, + "region": region_id, + "multi_hit": multi_reg, + } + ) + + if role is None and region_id is None and len(unmatched_samples) < sample_limit: + unmatched_samples.append( + {"fabric_node_id": n.id, "name": n.name, "ip": n.ip, "vendor": n.vendor} + ) + + return ClassifyPreviewOut( + total_nodes=len(nodes), + role_matched=role_matched, + role_unmatched=role_unmatched, + role_conflicts=role_conflict, + region_matched=region_matched, + region_unmatched=region_unmatched, + region_conflicts=region_conflict, + role_samples=role_samples, + region_samples=region_samples, + unmatched_samples=unmatched_samples, + ) + + +def apply_classify( + db: Session, + *, + skip_manual: bool = True, + fill_empty_only: bool = False, +) -> ClassifyApplyOut: + role_rules = _enabled_rules(db, "role") + region_rules = _enabled_rules(db, "region") + nodes = db.query(TopoFabricNode).all() + role_updated = region_updated = skipped_manual = 0 + + for n in nodes: + role, _, _ = _resolve_role_hit(n, role_rules) + if role is not None: + if skip_manual and str(n.role_source or "") == "manual": + skipped_manual += 1 + elif fill_empty_only and str(n.role or "").strip(): + pass + else: + n.role = role + n.role_source = "rule" + n.updated_at = _utcnow() + role_updated += 1 + elif not str(n.role or "").strip() and str(n.role_source or "") != "manual": + n.role = "unknown" + n.role_source = "rule" + n.updated_at = _utcnow() + + region_id, _, _ = _resolve_region_hit(db, n, region_rules, create_missing=True) + if region_id is not None and not str(region_id).startswith("new:"): + if skip_manual and str(n.region_source or "") == "manual": + skipped_manual += 1 + elif fill_empty_only and str(n.region_folder_id or "").strip(): + pass + else: + n.region_folder_id = region_id + n.region_source = "rule" + n.updated_at = _utcnow() + region_updated += 1 + + db.commit() + return ClassifyApplyOut( + role_updated=role_updated, + region_updated=region_updated, + skipped_manual=skipped_manual, + total_nodes=len(nodes), + ) + + +def list_unmatched( + db: Session, + *, + kind: str = "any", + page: int = 1, + page_size: int = 50, +) -> dict[str, Any]: + from sqlalchemy import or_ + + q = db.query(TopoFabricNode) + k = str(kind or "any").strip().lower() + role_miss = or_(TopoFabricNode.role == "", TopoFabricNode.role == "unknown") + region_miss = or_( + TopoFabricNode.region_folder_id.is_(None), + TopoFabricNode.region_folder_id == "", + ) + if k == "role": + q = q.filter(role_miss) + elif k == "region": + q = q.filter(region_miss) + else: + q = q.filter(or_(role_miss, region_miss)) + total = q.count() + rows = ( + q.order_by(TopoFabricNode.name.asc()) + .offset(max(0, (page - 1) * page_size)) + .limit(page_size) + .all() + ) + from .topology_service import _node_out + + return { + "total": total, + "page": page, + "page_size": page_size, + "items": [_node_out(n).model_dump() for n in rows], + } + + +def patch_fabric_node_tags( + db: Session, fabric_node_id: str, body: FabricNodeTagPatch +) -> FabricNodeOut: + from .topology_service import _node_out + + n = db.get(TopoFabricNode, fabric_node_id) + if n is None: + raise HTTPException(status_code=404, detail="fabric_node_not_found") + if body.role is not None: + role = str(body.role or "").strip().lower() + if role and role not in _ROLE_VALUES: + raise HTTPException(status_code=400, detail="role_invalid") + n.role = role + n.role_source = "manual" + if body.region_folder_id is not None: + fid = str(body.region_folder_id or "").strip() + if fid: + folder = db.get(TopoFolder, fid) + if folder is None or str(folder.kind or "") != "region": + raise HTTPException(status_code=400, detail="folder_not_found") + n.region_folder_id = fid + else: + n.region_folder_id = None + n.region_source = "manual" + n.updated_at = _utcnow() + db.commit() + db.refresh(n) + return _node_out(n) + + +def _iter_regex_matches( + db: Session, *, pattern: str, match_field: str +) -> list[TopoFabricNode]: + cre = _compile_pattern(pattern) + mf = str(match_field or "name").strip().lower() + if mf not in _MATCH_FIELDS: + raise HTTPException(status_code=400, detail="match_field_invalid") + out: list[TopoFabricNode] = [] + for n in db.query(TopoFabricNode).order_by(TopoFabricNode.name.asc()).all(): + text = _match_text(n, mf) + if text and cre.search(text): + out.append(n) + return out + + +def match_fabric_nodes(db: Session, body: FabricNodesMatchRequest) -> FabricNodesMatchOut: + """Ephemeral regex lookup — does not persist rules.""" + from .topology_inventory_lifecycle import fabric_link_status + + matched = _iter_regex_matches( + db, pattern=body.pattern, match_field=body.match_field + ) + limit = max(1, min(200, int(body.sample_limit or 50))) + samples = [ + { + "fabric_node_id": n.id, + "name": n.name, + "ip": n.ip, + "role": n.role or "", + "region_folder_id": n.region_folder_id or "", + "link_status": fabric_link_status(n), + } + for n in matched[:limit] + ] + return FabricNodesMatchOut( + pattern=str(body.pattern or "").strip(), + match_field=str(body.match_field or "name"), + total_matched=len(matched), + samples=samples, + fabric_node_ids=[n.id for n in matched], + ) + + +def bulk_tag_fabric_nodes( + db: Session, body: FabricNodesBulkTagRequest +) -> FabricNodesBulkTagOut: + """Assign role/region after user confirms a regex or explicit selection.""" + if body.role is None and body.region_folder_id is None: + raise HTTPException(status_code=400, detail="role_or_region_required") + + role_v: str | None = None + if body.role is not None: + role_v = str(body.role or "").strip().lower() + if role_v and role_v not in _ROLE_VALUES: + raise HTTPException(status_code=400, detail="role_invalid") + + region_v: str | None = None + if body.region_folder_id is not None: + region_v = str(body.region_folder_id or "").strip() + if region_v: + folder = db.get(TopoFolder, region_v) + if folder is None or str(folder.kind or "") != "region": + raise HTTPException(status_code=400, detail="folder_not_found") + else: + region_v = "" + + ids = [str(x).strip() for x in (body.fabric_node_ids or []) if str(x).strip()] + if str(body.pattern or "").strip(): + matched_nodes = _iter_regex_matches( + db, pattern=body.pattern, match_field=body.match_field + ) + elif ids: + matched_nodes = ( + db.query(TopoFabricNode) + .filter(TopoFabricNode.id.in_(ids)) + .order_by(TopoFabricNode.name.asc()) + .all() + ) + else: + raise HTTPException(status_code=400, detail="ids_or_pattern_required") + + samples = [ + { + "fabric_node_id": n.id, + "name": n.name, + "ip": n.ip, + "role": n.role or "", + "region_folder_id": n.region_folder_id or "", + } + for n in matched_nodes[:50] + ] + if body.dry_run: + return FabricNodesBulkTagOut( + dry_run=True, + matched=len(matched_nodes), + updated=0, + role=role_v, + region_folder_id=region_v, + samples=samples, + ) + + now = _utcnow() + updated = 0 + for n in matched_nodes: + if role_v is not None: + n.role = role_v + n.role_source = "manual" + if region_v is not None: + n.region_folder_id = region_v or None + n.region_source = "manual" + n.updated_at = now + updated += 1 + db.commit() + return FabricNodesBulkTagOut( + dry_run=False, + matched=len(matched_nodes), + updated=updated, + role=role_v, + region_folder_id=region_v, + samples=samples, + ) + + +def apply_classify_empty_only(db: Session) -> ClassifyApplyOut: + """Incremental classify for newly synced nodes (fill empty tags only).""" + return apply_classify(db, skip_manual=True, fill_empty_only=True) + + diff --git a/netx_api/topology_classify_common.py b/netx_api/topology_classify_common.py new file mode 100644 index 0000000..bae8dd7 --- /dev/null +++ b/netx_api/topology_classify_common.py @@ -0,0 +1,184 @@ +"""Shared helpers for topology classify rules and apply.""" +from __future__ import annotations + +import re +from datetime import datetime +from typing import Any + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .models import TopoClassifyRule, TopoFabricNode, TopoFolder +from .timeutil import utcnow_naive +from .topology_membership import VIEW_ROLES, normalize_view_role +from .topology_schemas import ( + ClassifyRuleOut, + TopologyFolderCreate, +) + +_MAX_PATTERN_LEN = 512 +_ROLE_VALUES = VIEW_ROLES | {"unknown"} +_MATCH_FIELDS = frozenset({"name", "ip", "name_ip"}) +_SCOPES = frozenset({"role", "region"}) +_SLICE_TEMPLATES = frozenset({"core_only", "core_agg", "agg_access"}) + + +def _utcnow() -> datetime: + return utcnow_naive() + + +def _compile_pattern(pattern: str) -> re.Pattern[str]: + p = str(pattern or "").strip() + if not p: + raise HTTPException(status_code=400, detail="pattern_required") + if len(p) > _MAX_PATTERN_LEN: + raise HTTPException(status_code=400, detail="pattern_too_long") + try: + return re.compile(p, re.IGNORECASE) + except re.error as exc: + raise HTTPException(status_code=400, detail=f"invalid_pattern:{exc}") from exc + + +def _match_text(node: TopoFabricNode, match_field: str) -> str: + name = str(node.name or "") + ip = str(node.ip or "") + mf = str(match_field or "name").strip().lower() + if mf == "ip": + return ip + if mf == "name_ip": + return f"{name} {ip}".strip() + return name + + +def _rule_out(row: TopoClassifyRule) -> ClassifyRuleOut: + return ClassifyRuleOut( + id=row.id, + scope=str(row.scope or "role"), + name=str(row.name or ""), + pattern=str(row.pattern or ""), + match_field=str(row.match_field or "name"), + priority=int(row.priority or 100), + enabled=bool(row.enabled), + payload=dict(row.payload or {}), + remark=str(row.remark or ""), + created_at=row.created_at, + updated_at=row.updated_at, + ) + + +def _validate_payload(scope: str, payload: dict[str, Any]) -> dict[str, Any]: + out = dict(payload or {}) + if scope == "role": + role = normalize_view_role(str(out.get("role") or "")) + if str(out.get("role") or "").strip().lower() not in VIEW_ROLES: + raise HTTPException(status_code=400, detail="role_payload_invalid") + return {"role": role} + if "folder_id" in out and str(out.get("folder_id") or "").strip(): + return {"folder_id": str(out["folder_id"]).strip()} + if "region_name_from_group" in out: + try: + g = int(out.get("region_name_from_group")) + except (TypeError, ValueError) as exc: + raise HTTPException(status_code=400, detail="region_group_invalid") from exc + if g < 1: + raise HTTPException(status_code=400, detail="region_group_invalid") + return {"region_name_from_group": g} + raise HTTPException(status_code=400, detail="region_payload_invalid") + + +def _enabled_rules(db: Session, scope: str) -> list[tuple[TopoClassifyRule, re.Pattern[str]]]: + rows = ( + db.query(TopoClassifyRule) + .filter(TopoClassifyRule.scope == scope, TopoClassifyRule.enabled.is_(True)) + .order_by(TopoClassifyRule.priority.asc(), TopoClassifyRule.name.asc()) + .all() + ) + out: list[tuple[TopoClassifyRule, re.Pattern[str]]] = [] + for r in rows: + try: + out.append((r, _compile_pattern(r.pattern))) + except HTTPException: + continue + return out + + +def _ensure_region_by_name(db: Session, name: str) -> TopoFolder: + from .topology_service import bootstrap_topology_tree, create_folder + + name = str(name or "").strip()[:256] + if not name: + raise HTTPException(status_code=400, detail="region_name_empty") + existing = ( + db.query(TopoFolder) + .filter(TopoFolder.kind == "region", TopoFolder.name == name) + .first() + ) + if existing is not None: + return existing + bootstrap_topology_tree(db) + created = create_folder(db, TopologyFolderCreate(name=name, kind="region")) + folder = db.get(TopoFolder, created.id) + assert folder is not None + return folder + + +def _resolve_role_hit( + node: TopoFabricNode, rules: list[tuple[TopoClassifyRule, re.Pattern[str]]] +) -> tuple[str | None, str | None, bool]: + """Return (role, rule_id, multi_hit).""" + hits: list[tuple[str, str]] = [] + for rule, cre in rules: + text = _match_text(node, rule.match_field) + if not text: + continue + if cre.search(text): + role = normalize_view_role(str((rule.payload or {}).get("role") or "")) + hits.append((role, rule.id)) + if not hits: + return None, None, False + return hits[0][0], hits[0][1], len(hits) > 1 + + +def _resolve_region_hit( + db: Session, + node: TopoFabricNode, + rules: list[tuple[TopoClassifyRule, re.Pattern[str]]], + *, + create_missing: bool, +) -> tuple[str | None, str | None, bool]: + hits: list[tuple[str, str]] = [] + for rule, cre in rules: + text = _match_text(node, rule.match_field) + if not text: + continue + m = cre.search(text) + if not m: + continue + payload = dict(rule.payload or {}) + folder_id = str(payload.get("folder_id") or "").strip() + if folder_id: + hits.append((folder_id, rule.id)) + continue + g = int(payload.get("region_name_from_group") or 0) + try: + region_name = m.group(g) + except IndexError: + continue + region_name = str(region_name or "").strip() + if not region_name: + continue + if create_missing: + folder = _ensure_region_by_name(db, region_name) + hits.append((folder.id, rule.id)) + else: + existing = ( + db.query(TopoFolder) + .filter(TopoFolder.kind == "region", TopoFolder.name == region_name) + .first() + ) + hits.append((existing.id if existing else f"new:{region_name}", rule.id)) + if not hits: + return None, None, False + return hits[0][0], hits[0][1], len(hits) > 1 + + diff --git a/netx_api/topology_classify_rules.py b/netx_api/topology_classify_rules.py new file mode 100644 index 0000000..f892722 --- /dev/null +++ b/netx_api/topology_classify_rules.py @@ -0,0 +1,107 @@ +"""Topology classify rule CRUD.""" +from __future__ import annotations + +from typing import Any +from uuid import uuid4 + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .models import TopoClassifyRule, TopoFolder +from .topology_classify_common import ( + _MATCH_FIELDS, + _MAX_PATTERN_LEN, + _SCOPES, + _compile_pattern, + _rule_out, + _utcnow, + _validate_payload, +) +from .topology_schemas import ClassifyRuleCreate, ClassifyRuleOut, ClassifyRuleUpdate + +def list_rules(db: Session, *, scope: str = "") -> list[ClassifyRuleOut]: + q = db.query(TopoClassifyRule) + if scope.strip(): + q = q.filter(TopoClassifyRule.scope == scope.strip().lower()) + rows = q.order_by( + TopoClassifyRule.scope.asc(), + TopoClassifyRule.priority.asc(), + TopoClassifyRule.name.asc(), + ).all() + return [_rule_out(r) for r in rows] + + +def create_rule(db: Session, body: ClassifyRuleCreate) -> ClassifyRuleOut: + scope = str(body.scope or "").strip().lower() + if scope not in _SCOPES: + raise HTTPException(status_code=400, detail="scope_invalid") + match_field = str(body.match_field or "name").strip().lower() + if match_field not in _MATCH_FIELDS: + raise HTTPException(status_code=400, detail="match_field_invalid") + _compile_pattern(body.pattern) + payload = _validate_payload(scope, dict(body.payload or {})) + if scope == "region" and payload.get("folder_id"): + folder = db.get(TopoFolder, payload["folder_id"]) + if folder is None or str(folder.kind or "") != "region": + raise HTTPException(status_code=400, detail="folder_not_found") + now = _utcnow() + row = TopoClassifyRule( + id=uuid4().hex, + scope=scope, + name=str(body.name or "").strip()[:256] or f"{scope}-rule", + pattern=str(body.pattern or "").strip()[:_MAX_PATTERN_LEN], + match_field=match_field, + priority=int(body.priority if body.priority is not None else 100), + enabled=bool(body.enabled if body.enabled is not None else True), + payload=payload, + remark=str(body.remark or "")[:512], + created_at=now, + updated_at=now, + ) + db.add(row) + db.commit() + db.refresh(row) + return _rule_out(row) + + +def update_rule(db: Session, rule_id: str, body: ClassifyRuleUpdate) -> ClassifyRuleOut: + row = db.get(TopoClassifyRule, rule_id) + if row is None: + raise HTTPException(status_code=404, detail="rule_not_found") + if body.name is not None: + row.name = str(body.name or "").strip()[:256] + if body.pattern is not None: + _compile_pattern(body.pattern) + row.pattern = str(body.pattern or "").strip()[:_MAX_PATTERN_LEN] + if body.match_field is not None: + mf = str(body.match_field or "name").strip().lower() + if mf not in _MATCH_FIELDS: + raise HTTPException(status_code=400, detail="match_field_invalid") + row.match_field = mf + if body.priority is not None: + row.priority = int(body.priority) + if body.enabled is not None: + row.enabled = bool(body.enabled) + if body.remark is not None: + row.remark = str(body.remark or "")[:512] + if body.payload is not None: + row.payload = _validate_payload(str(row.scope or "role"), dict(body.payload or {})) + if str(row.scope) == "region" and row.payload.get("folder_id"): + folder = db.get(TopoFolder, row.payload["folder_id"]) + if folder is None or str(folder.kind or "") != "region": + raise HTTPException(status_code=400, detail="folder_not_found") + row.updated_at = _utcnow() + db.commit() + db.refresh(row) + return _rule_out(row) + + +def delete_rule(db: Session, rule_id: str) -> dict[str, Any]: + row = db.get(TopoClassifyRule, rule_id) + if row is None: + raise HTTPException(status_code=404, detail="rule_not_found") + db.delete(row) + db.commit() + return {"ok": True, "id": rule_id, "deleted": True} + + diff --git a/netx_api/topology_classify_slices.py b/netx_api/topology_classify_slices.py new file mode 100644 index 0000000..0b0a63b --- /dev/null +++ b/netx_api/topology_classify_slices.py @@ -0,0 +1,344 @@ +"""Topology slice map preview/generation and fabric search.""" +from __future__ import annotations + +from typing import Any + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .models import TopoFabricEdge, TopoFabricNode, TopoFolder, TopoView, TopoViewNode +from .topology_classify_common import _SLICE_TEMPLATES, _utcnow +from .topology_membership import ( + VIEW_KIND_CUSTOM, + VIEW_ROLE_ACCESS, + VIEW_ROLE_AGGREGATION, + VIEW_ROLE_CORE, + merge_filter_with_membership, + normalize_view_role, +) +from .topology_schemas import ( + FabricNodeOut, + SliceGenerateOut, + SliceGenerateRequest, + SliceMapPlan, + TopologyViewCreate, +) + +def _active_neighbors(db: Session, seed_ids: set[str], *, hops: int = 1) -> set[str]: + if not seed_ids or hops <= 0: + return set() + frontier = set(seed_ids) + found: set[str] = set() + for _ in range(hops): + if not frontier: + break + rows = ( + db.query(TopoFabricEdge) + .filter( + TopoFabricEdge.layer == "physical", + TopoFabricEdge.status == "active", + (TopoFabricEdge.a_node_id.in_(frontier) | TopoFabricEdge.b_node_id.in_(frontier)), + ) + .all() + ) + nxt: set[str] = set() + for e in rows: + for a, b in ((e.a_node_id, e.b_node_id), (e.b_node_id, e.a_node_id)): + if a in frontier and b not in seed_ids and b not in found: + nxt.add(str(b)) + found |= nxt + frontier = nxt + return found + + +def _connected_components(db: Session, node_ids: list[str]) -> list[list[str]]: + ids = [str(x) for x in node_ids if str(x)] + if not ids: + return [] + id_set = set(ids) + adj: dict[str, set[str]] = {i: set() for i in ids} + rows = ( + db.query(TopoFabricEdge) + .filter( + TopoFabricEdge.layer == "physical", + TopoFabricEdge.status == "active", + TopoFabricEdge.a_node_id.in_(ids), + TopoFabricEdge.b_node_id.in_(ids), + ) + .all() + ) + for e in rows: + a, b = str(e.a_node_id), str(e.b_node_id) + if a in id_set and b in id_set: + adj[a].add(b) + adj[b].add(a) + seen: set[str] = set() + comps: list[list[str]] = [] + for nid in ids: + if nid in seen: + continue + stack = [nid] + seen.add(nid) + comp: list[str] = [] + while stack: + cur = stack.pop() + comp.append(cur) + for nb in adj.get(cur, ()): + if nb not in seen: + seen.add(nb) + stack.append(nb) + comps.append(sorted(comp)) + return comps + + +def _nodes_in_region(db: Session, folder_id: str, *, role: str = "") -> list[TopoFabricNode]: + q = db.query(TopoFabricNode).filter(TopoFabricNode.region_folder_id == folder_id) + if role: + q = q.filter(TopoFabricNode.role == role) + return q.order_by(TopoFabricNode.name.asc()).all() + + +def preview_slices(db: Session, body: SliceGenerateRequest) -> SliceGenerateOut: + folder = db.get(TopoFolder, body.folder_id) + if folder is None or str(folder.kind or "") != "region": + raise HTTPException(status_code=400, detail="folder_not_found") + template = str(body.template or "").strip().lower() + if template not in _SLICE_TEMPLATES: + raise HTTPException(status_code=400, detail="template_invalid") + max_nodes = max(1, min(2000, int(body.max_nodes or 300))) + plans: list[SliceMapPlan] = [] + overlap_ids: set[str] = set() + seen_in_maps: dict[str, int] = {} + + def _track(ids: list[str]) -> None: + for i in ids: + seen_in_maps[i] = seen_in_maps.get(i, 0) + 1 + + if template == "core_only": + cores = _nodes_in_region(db, folder.id, role=VIEW_ROLE_CORE) + comps = _connected_components(db, [n.id for n in cores]) or [ + [n.id] for n in cores + ] + for idx, comp in enumerate(comps, start=1): + if len(comp) > max_nodes: + raise HTTPException( + status_code=400, + detail=f"slice_exceeds_max_nodes:{len(comp)}>{max_nodes}", + ) + name = f"Core-{idx}" if len(comps) > 1 else "Core" + plans.append( + SliceMapPlan( + name=name, + role=VIEW_ROLE_CORE, + seed_fabric_node_ids=comp, + member_fabric_node_ids=comp, + node_count=len(comp), + ) + ) + _track(comp) + elif template == "core_agg": + cores = _nodes_in_region(db, folder.id, role=VIEW_ROLE_CORE) + comps = _connected_components(db, [n.id for n in cores]) or [ + [n.id] for n in cores + ] + for idx, comp in enumerate(comps, start=1): + peers = _active_neighbors(db, set(comp), hops=1) + agg_ids = [ + p + for p in peers + if (fn := db.get(TopoFabricNode, p)) is not None + and str(fn.role or "") == VIEW_ROLE_AGGREGATION + and str(fn.region_folder_id or "") == folder.id + ] + members = sorted(set(comp) | set(agg_ids)) + if len(members) > max_nodes: + raise HTTPException( + status_code=400, + detail=f"slice_exceeds_max_nodes:{len(members)}>{max_nodes}", + ) + name = f"CoreAgg-{idx}" if len(comps) > 1 else "Core+Agg" + plans.append( + SliceMapPlan( + name=name, + role=VIEW_ROLE_CORE, + seed_fabric_node_ids=comp, + member_fabric_node_ids=members, + node_count=len(members), + ) + ) + _track(members) + else: # agg_access + aggs = _nodes_in_region(db, folder.id, role=VIEW_ROLE_AGGREGATION) + comps = _connected_components(db, [n.id for n in aggs]) or [ + [n.id] for n in aggs + ] + for idx, comp in enumerate(comps, start=1): + peers = _active_neighbors(db, set(comp), hops=1) + acc_ids = [ + p + for p in peers + if (fn := db.get(TopoFabricNode, p)) is not None + and str(fn.role or "") == VIEW_ROLE_ACCESS + and str(fn.region_folder_id or "") == folder.id + ] + members = sorted(set(comp) | set(acc_ids)) + if len(members) > max_nodes: + raise HTTPException( + status_code=400, + detail=f"slice_exceeds_max_nodes:{len(members)}>{max_nodes}", + ) + name = f"AggAccess-{idx}" if len(comps) > 1 else "Agg+Access" + plans.append( + SliceMapPlan( + name=name, + role=VIEW_ROLE_AGGREGATION, + seed_fabric_node_ids=comp, + member_fabric_node_ids=members, + node_count=len(members), + ) + ) + _track(members) + + overlap_ids = {nid for nid, cnt in seen_in_maps.items() if cnt > 1} + return SliceGenerateOut( + folder_id=folder.id, + template=template, + dry_run=True, + maps=plans, + map_count=len(plans), + overlap_node_count=len(overlap_ids), + created_view_ids=[], + ) + + +def generate_slices(db: Session, body: SliceGenerateRequest) -> SliceGenerateOut: + from .topology_service import create_view, _place_fabric_ids_on_view + + preview = preview_slices(db, body) + if body.dry_run: + return preview + + created: list[str] = [] + for plan in preview.maps: + view = create_view( + db, + TopologyViewCreate( + name=plan.name, + folder_id=body.folder_id, + kind=VIEW_KIND_CUSTOM, + role=plan.role, + remark=f"slice:{body.template}", + ), + ) + mem = { + "mode": "hybrid", + "seed_fabric_node_ids": list(plan.member_fabric_node_ids), + "expand_hops": 0, + "max_nodes": int(body.max_nodes or 300), + "frozen": True, + "managed_ne_ids": [], + "tags_any": [], + "vendors": [], + "device_types": [], + "keyword": "", + } + row = db.get(TopoView, view.id) + assert row is not None + row.filter = merge_filter_with_membership( + dict(row.filter or {}), role=normalize_view_role(plan.role), membership=mem + ) + _place_fabric_ids_on_view(db, row, list(plan.member_fabric_node_ids), existing=set()) + row.updated_at = _utcnow() + db.commit() + created.append(view.id) + + # Optionally seed physical overview with cores only + if body.seed_physical_cores: + from .topology_service import ensure_region_physical_view + + phys = ensure_region_physical_view(db, body.folder_id, commit=True) + cores = [n.id for n in _nodes_in_region(db, body.folder_id, role=VIEW_ROLE_CORE)] + existing = { + vn.fabric_node_id + for vn in db.query(TopoViewNode).filter(TopoViewNode.view_id == phys.id).all() + } + to_add = [c for c in cores if c not in existing][: int(body.max_nodes or 300)] + if to_add: + _place_fabric_ids_on_view(db, phys, to_add, existing=existing) + mem = merge_filter_with_membership( + dict(phys.filter or {}), + role=VIEW_ROLE_CORE, + kind="physical", + membership={ + **dict((phys.filter or {}).get("membership") or {}), + "frozen": True, + "max_nodes": int(body.max_nodes or 500), + }, + ) + phys.filter = mem + phys.updated_at = _utcnow() + db.commit() + + return SliceGenerateOut( + folder_id=body.folder_id, + template=str(body.template), + dry_run=False, + maps=preview.maps, + map_count=len(preview.maps), + overlap_node_count=preview.overlap_node_count, + created_view_ids=created, + ) + + +def search_fabric_nodes_with_views( + db: Session, + *, + keyword: str = "", + page: int = 1, + page_size: int = 50, +) -> dict[str, Any]: + from .topology_service import _node_out + + q = db.query(TopoFabricNode) + kw = str(keyword or "").strip() + if kw: + like = f"%{kw}%" + q = q.filter( + (TopoFabricNode.name.ilike(like)) + | (TopoFabricNode.ip.ilike(like)) + | (TopoFabricNode.vendor.ilike(like)) + ) + total = q.count() + rows = ( + q.order_by(TopoFabricNode.name.asc()) + .offset(max(0, (page - 1) * page_size)) + .limit(page_size) + .all() + ) + node_ids = [n.id for n in rows] + placements: dict[str, list[dict[str, Any]]] = {nid: [] for nid in node_ids} + if node_ids: + vnodes = ( + db.query(TopoViewNode, TopoView, TopoFolder) + .join(TopoView, TopoView.id == TopoViewNode.view_id) + .outerjoin(TopoFolder, TopoFolder.id == TopoView.folder_id) + .filter(TopoViewNode.fabric_node_id.in_(node_ids)) + .all() + ) + for vn, view, folder in vnodes: + placements.setdefault(vn.fabric_node_id, []).append( + { + "view_id": view.id, + "view_name": view.name, + "folder_id": view.folder_id or "", + "folder_name": (folder.name if folder else "") or "", + "kind": view.kind or "custom", + } + ) + items = [] + for n in rows: + d = _node_out(n).model_dump() + d["views"] = placements.get(n.id, []) + items.append(d) + return {"total": total, "page": page, "page_size": page_size, "items": items} +