Split topology classify and config sync services by domain.

Keep public facades stable while moving rules/apply/slices and policy/cycles/snapshots into focused modules.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-08-02 17:21:50 +08:00
parent 58cdbe6165
commit 4ec890c407
9 changed files with 1846 additions and 1653 deletions

View file

@ -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,
)

View file

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

View file

@ -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,
from .config_sync_snapshots import (
build_snapshot_export,
get_snapshot_detail,
list_snapshot_history,
list_snapshots,
)
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
__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",
]
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")

View file

@ -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}

View file

@ -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]
__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",
]
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}

View file

@ -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)

View file

@ -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

View file

@ -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}

View file

@ -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}