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