Add ZTE-first port traffic monitoring with task wizard and ops wall.

Collect interface bit/s via CLI on a schedule, store samples in Postgres, and chart trends with uPlot.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-07-31 14:56:50 +08:00
parent a95c6598bb
commit efe8f58cd2
25 changed files with 2880 additions and 36 deletions

View file

@ -77,6 +77,9 @@ class Settings(BaseSettings):
config_sync_scheduler_tick_sec: int = 60
# After process start / unexpected restart, wait before any scheduled sync.
config_sync_startup_grace_sec: int = 3600
# Port traffic monitoring (CLI rate bit/s samples)
port_traffic_scheduler_enabled: bool = True
port_traffic_scheduler_tick_sec: int = 15
# Managed NE exec: max CLI commands per request (lab can raise; hard-capped in ne_exec).
ne_exec_max_commands: int = 5
# WebCRT interactive terminal sessions

View file

@ -27,6 +27,7 @@ from .db import Base, SessionLocal, engine, get_db
from .collection_router import router as collection_router
from .cli_router import router as cli_router
from .config_sync_router import router as config_sync_router
from .port_traffic_router import router as port_traffic_router
from .managed_ne_router import router as managed_ne_router
from .webcrt_router import router as webcrt_router
from .topology_router import router as topology_router
@ -135,6 +136,7 @@ app.include_router(managed_ne_router)
app.include_router(cli_router)
app.include_router(collection_router)
app.include_router(config_sync_router)
app.include_router(port_traffic_router)
app.include_router(webcrt_router)
app.include_router(topology_router)
parser_cfg = load_parser_config()
@ -836,11 +838,15 @@ def on_startup() -> None:
_schedule_log.info("startup: resumed %s pending ne collection runs", resumed)
from .config_sync_recovery import recover_config_sync_on_startup
from .config_sync_service import ensure_policy
from .port_traffic_recovery import recover_port_traffic_on_startup
ensure_policy(db)
cfg_resumed = recover_config_sync_on_startup(db)
if cfg_resumed:
_schedule_log.info("startup: resumed %s config_sync task(s) from interrupted cycle", cfg_resumed)
pt_cleared = recover_port_traffic_on_startup(db)
if pt_cleared:
_schedule_log.info("startup: cleared %s port_traffic stuck collect_running flag(s)", pt_cleared)
except Exception:
_schedule_log.exception("startup: ne collection / config_sync recovery failed")
finally:
@ -851,6 +857,12 @@ def on_startup() -> None:
start_config_sync_scheduler()
except Exception:
_schedule_log.exception("startup: config_sync scheduler init failed")
try:
from .port_traffic_scheduler import start_port_traffic_scheduler
start_port_traffic_scheduler()
except Exception:
_schedule_log.exception("startup: port_traffic scheduler init failed")
# Best-effort schema evolution for new columns (no migrations framework).
# Safe for Postgres (IF NOT EXISTS); ignored on failure.
try:

View file

@ -3,7 +3,7 @@ from __future__ import annotations
from datetime import datetime
from uuid import uuid4
from sqlalchemy import Boolean, DateTime, Float, ForeignKey, Integer, LargeBinary, String, Text, UniqueConstraint
from sqlalchemy import BigInteger, Boolean, DateTime, Float, ForeignKey, Integer, LargeBinary, String, Text, UniqueConstraint
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.types import JSON
@ -588,3 +588,65 @@ class NeConfigHistory(Base):
collected_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True)
cycle_id: Mapped[str] = mapped_column(String(64), default="")
task_id: Mapped[str] = mapped_column(String(64), default="")
class PortTrafficTask(Base):
"""Port traffic monitoring job definition."""
__tablename__ = "port_traffic_task"
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: uuid4().hex)
title: Mapped[str] = mapped_column(String(256), default="")
status: Mapped[str] = mapped_column(String(32), default="draft", index=True) # draft|running|paused|stopped
interval_sec: Mapped[int] = mapped_column(Integer, default=60)
retention_days: Mapped[int] = mapped_column(Integer, default=7)
concurrency: Mapped[int] = mapped_column(Integer, default=5)
collect_running: Mapped[bool] = mapped_column(Boolean, default=False)
last_collect_started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
last_collect_ended_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
last_error: Mapped[str] = mapped_column(String(1024), default="")
created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True)
updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow)
class PortTrafficTarget(Base):
"""Monitored interface under a port traffic task."""
__tablename__ = "port_traffic_target"
__table_args__ = (
UniqueConstraint("task_id", "source", "target_id", "ifname", name="uq_port_traffic_target_if"),
)
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: uuid4().hex)
task_id: Mapped[str] = mapped_column(String(64), index=True)
source: Mapped[str] = mapped_column(String(32), default="managed", index=True)
target_id: Mapped[str] = mapped_column(String(128), index=True)
ne_name: Mapped[str] = mapped_column(String(256), default="")
ne_ip: Mapped[str] = mapped_column(String(128), default="")
vendor: Mapped[str] = mapped_column(String(64), default="")
ifname: Mapped[str] = mapped_column(String(128), default="")
if_description: Mapped[str] = mapped_column(String(512), default="")
bw_bps: Mapped[int] = mapped_column(BigInteger, default=0)
status: Mapped[str] = mapped_column(String(32), default="active", index=True) # active|disabled
last_error: Mapped[str] = mapped_column(String(1024), default="")
last_sample_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow)
class PortTrafficSample(Base):
"""Time-series sample for a monitored interface."""
__tablename__ = "port_traffic_sample"
__table_args__ = (UniqueConstraint("target_row_id", "ts", name="uq_port_traffic_sample_ts"),)
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: uuid4().hex)
target_row_id: Mapped[str] = mapped_column(String(64), index=True)
ts: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True)
in_bps: Mapped[float] = mapped_column(Float, default=0.0)
out_bps: Mapped[float] = mapped_column(Float, default=0.0)
in_util_pct: Mapped[float] = mapped_column(Float, default=0.0)
out_util_pct: Mapped[float] = mapped_column(Float, default=0.0)
bw_bps: Mapped[int] = mapped_column(BigInteger, default=0)
rate_period_sec: Mapped[int] = mapped_column(Integer, default=0)
raw_ok: Mapped[bool] = mapped_column(Boolean, default=True)
message: Mapped[str] = mapped_column(String(512), default="")

View file

@ -0,0 +1,29 @@
"""Vendor → port traffic CLI command matrix (ZTE first)."""
from __future__ import annotations
from dataclasses import dataclass
from .config_sync_commands import normalize_vendor_key
@dataclass(frozen=True)
class PortTrafficCommands:
brief: str
detail_template: str # format with ifname=
vendor_key: str = "other"
def commands_for_vendor(vendor: str, device_type: str = "") -> PortTrafficCommands | None:
key = normalize_vendor_key(vendor, device_type)
if key == "zte":
return PortTrafficCommands(
brief="show interface brief",
detail_template="show interface {ifname}",
vendor_key=key,
)
return None
def detail_command(cmds: PortTrafficCommands, ifname: str) -> str:
return cmds.detail_template.format(ifname=ifname)

View file

@ -0,0 +1,237 @@
"""Parsers for ZTE show interface brief / detail (rate bit/s)."""
from __future__ import annotations
import re
from dataclasses import dataclass
from typing import Any
_BW_UNIT = {
"k": 1_000,
"m": 1_000_000,
"g": 1_000_000_000,
"t": 1_000_000_000_000,
}
_RE_BW_COMPACT = re.compile(r"^(\d+(?:\.\d+)?)\s*([kKmMgGtT])(?:bit)?s?$", re.I)
_RE_BW_DETAIL = re.compile(
r"\bBW\s+(\d+(?:\.\d+)?)\s*([kKmMgGtT])?\s*(?:G?bit|bit)/s\b",
re.I,
)
_RE_RATE_PERIOD = re.compile(r"Rate\s+period\s*:\s*(\d+)\s*s", re.I)
_RE_INPUT_BPS = re.compile(r"^\s*Input\s*:\s*([\d.]+)\s*bit/s", re.I | re.M)
_RE_OUTPUT_BPS = re.compile(r"^\s*Output\s*:\s*([\d.]+)\s*bit/s", re.I | re.M)
_RE_UTIL = re.compile(
r"Intf\s+utilization\s*:\s*input\s*([\d.]+)%\s*output\s*([\d.]+)%",
re.I,
)
_RE_IF_UP = re.compile(r"^(\S+)\s+is\s+(up|down)\b", re.I | re.M)
_RE_DESC = re.compile(r"^\s*Description:\s*(.+?)\s*$", re.I | re.M)
# Fixed-width columns from ZTE `show interface brief` header.
_COL_IF = (0, 24)
_COL_ATTR = (24, 35)
_COL_MODE = (35, 48)
_COL_BW = (48, 54)
_COL_ADMIN = (54, 60)
_COL_PHY = (60, 66)
_COL_PROT = (66, 72)
_COL_DESC = (72, None)
@dataclass(frozen=True)
class BriefPort:
ifname: str
attribute: str = ""
mode: str = ""
bw_raw: str = ""
bw_bps: int = 0
admin: str = ""
phy: str = ""
prot: str = ""
description: str = ""
@dataclass(frozen=True)
class DetailRates:
ifname: str = ""
admin_oper: str = ""
description: str = ""
bw_bps: int = 0
rate_period_sec: int = 0
in_bps: float = 0.0
out_bps: float = 0.0
in_util_pct: float = 0.0
out_util_pct: float = 0.0
def parse_bw_to_bps(raw: str) -> int:
text = (raw or "").strip()
if not text or text.upper() == "N/A":
return 0
m = _RE_BW_COMPACT.match(text.replace(" ", ""))
if m:
value = float(m.group(1))
mult = _BW_UNIT[m.group(2).lower()]
return int(value * mult)
m2 = _RE_BW_DETAIL.search(text)
if m2:
value = float(m2.group(1))
unit = (m2.group(2) or "g").lower()
# "BW 1 Gbit/s" — unit letter may be in "Gbit" when group2 empty after "1 "
if m2.group(2) is None and "gbit" in text.lower():
unit = "g"
elif m2.group(2) is None and "mbit" in text.lower():
unit = "m"
elif m2.group(2) is None and "kbit" in text.lower():
unit = "k"
mult = _BW_UNIT.get(unit, 1_000_000_000)
return int(value * mult)
# Plain "BW 1000000000" unlikely; try digits only
digits = re.sub(r"[^\d.]", "", text)
if digits:
try:
return int(float(digits))
except ValueError:
return 0
return 0
def _slice(line: str, start: int, end: int | None) -> str:
if end is None:
return line[start:].rstrip() if len(line) > start else ""
if len(line) <= start:
return ""
return line[start:end].strip()
def parse_zte_interface_brief(text: str) -> list[BriefPort]:
"""Parse ZTE `show interface brief` into port rows."""
lines = (text or "").replace("\r\n", "\n").replace("\r", "\n").split("\n")
out: list[BriefPort] = []
started = False
for raw in lines:
line = raw.rstrip()
if not line.strip():
continue
if line.lstrip().startswith("Interface") and "Admin" in line and "Description" in line:
started = True
continue
if not started:
continue
if " is " in line and "ifindex" in line.lower():
break
ifname = _slice(line, *_COL_IF).split()[0] if _slice(line, *_COL_IF) else ""
if not ifname or ifname.lower() == "interface":
continue
bw_raw = _slice(line, *_COL_BW)
out.append(
BriefPort(
ifname=ifname,
attribute=_slice(line, *_COL_ATTR),
mode=_slice(line, *_COL_MODE),
bw_raw=bw_raw,
bw_bps=parse_bw_to_bps(bw_raw),
admin=_slice(line, *_COL_ADMIN).lower(),
phy=_slice(line, *_COL_PHY).lower(),
prot=_slice(line, *_COL_PROT).lower(),
description=_slice(line, *_COL_DESC).strip(),
)
)
return out
def parse_zte_interface_detail(text: str) -> DetailRates:
"""Parse ZTE `show interface {ifname}` rate / util / BW."""
blob = text or ""
ifname = ""
admin_oper = ""
m_if = _RE_IF_UP.search(blob)
if m_if:
ifname = m_if.group(1)
admin_oper = m_if.group(2).lower()
desc = ""
m_desc = _RE_DESC.search(blob)
if m_desc:
desc = m_desc.group(1).strip()
bw_bps = 0
m_bw = _RE_BW_DETAIL.search(blob)
if m_bw:
bw_bps = parse_bw_to_bps(m_bw.group(0))
else:
# Fallback line scan
for line in blob.splitlines():
if re.search(r"\bBW\b", line, re.I):
bw_bps = parse_bw_to_bps(line)
if bw_bps:
break
rate_period = 0
m_rp = _RE_RATE_PERIOD.search(blob)
if m_rp:
rate_period = int(m_rp.group(1))
# Prefer Rate period block: first Input/Output after "Rate period"
in_bps = 0.0
out_bps = 0.0
rp_idx = blob.lower().find("rate period")
rate_blob = blob[rp_idx:] if rp_idx >= 0 else blob
# Stop before Peak rate to avoid peak Input/Output
peak_idx = rate_blob.lower().find("peak rate")
if peak_idx >= 0:
rate_blob = rate_blob[:peak_idx]
m_in = _RE_INPUT_BPS.search(rate_blob)
m_out = _RE_OUTPUT_BPS.search(rate_blob)
if m_in:
in_bps = float(m_in.group(1))
if m_out:
out_bps = float(m_out.group(1))
in_util = 0.0
out_util = 0.0
m_util = _RE_UTIL.search(blob)
if m_util:
in_util = float(m_util.group(1))
out_util = float(m_util.group(2))
return DetailRates(
ifname=ifname,
admin_oper=admin_oper,
description=desc,
bw_bps=bw_bps,
rate_period_sec=rate_period,
in_bps=in_bps,
out_bps=out_bps,
in_util_pct=in_util,
out_util_pct=out_util,
)
def brief_port_to_dict(row: BriefPort) -> dict[str, Any]:
return {
"ifname": row.ifname,
"attribute": row.attribute,
"mode": row.mode,
"bw_raw": row.bw_raw,
"bw_bps": row.bw_bps,
"admin": row.admin,
"phy": row.phy,
"prot": row.prot,
"description": row.description,
}
def detail_to_dict(row: DetailRates) -> dict[str, Any]:
return {
"ifname": row.ifname,
"admin_oper": row.admin_oper,
"description": row.description,
"bw_bps": row.bw_bps,
"rate_period_sec": row.rate_period_sec,
"in_bps": row.in_bps,
"out_bps": row.out_bps,
"in_util_pct": row.in_util_pct,
"out_util_pct": row.out_util_pct,
}

View file

@ -0,0 +1,31 @@
"""Startup recovery for interrupted port traffic collect rounds."""
from __future__ import annotations
import logging
from datetime import datetime
from sqlalchemy.orm import Session
from .models import PortTrafficTask
_log = logging.getLogger("netx.port_traffic.recovery")
def recover_port_traffic_on_startup(db: Session) -> int:
"""Clear stuck collect_running flags so scheduler can resume."""
now = datetime.utcnow()
rows = db.query(PortTrafficTask).filter(PortTrafficTask.collect_running.is_(True)).all()
n = 0
for task in rows:
task.collect_running = False
if not task.last_collect_ended_at:
task.last_collect_ended_at = now
if not task.last_error:
task.last_error = "requeued_after_restart"
task.updated_at = now
n += 1
if n:
db.commit()
_log.info("port_traffic recovery cleared collect_running on %s task(s)", n)
return n

View file

@ -0,0 +1,211 @@
"""HTTP API for port traffic monitoring."""
from __future__ import annotations
from datetime import datetime
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy.orm import Session
from .auth_service import write_audit
from .db import get_db
from .port_traffic_schemas import (
DiscoverPortsRequest,
PortTrafficTaskCreate,
PortTrafficTaskUpdate,
PortTrafficTargetsPut,
)
from .port_traffic_service import (
create_task,
dashboard,
delete_task,
discover_ports,
get_samples,
get_task,
list_targets,
list_tasks,
put_targets,
set_task_status,
update_task,
)
router = APIRouter(prefix="/v1/port-traffic", tags=["port-traffic"])
def _actor(request: Request) -> tuple[str, str]:
user = getattr(request.state, "auth_user", None)
if not user:
return "", ""
return str(getattr(user, "id", "") or ""), str(getattr(user, "username", "") or "")
@router.get("/dashboard")
def api_dashboard(db: Session = Depends(get_db)):
return dashboard(db).model_dump()
@router.get("/tasks")
def api_list_tasks(
page: int = Query(default=1, ge=1),
page_size: int = Query(default=20, ge=1, le=100),
db: Session = Depends(get_db),
):
return list_tasks(db, page=page, page_size=page_size)
@router.post("/tasks")
def api_create_task(
body: PortTrafficTaskCreate,
request: Request,
db: Session = Depends(get_db),
):
out = create_task(db, body)
uid, uname = _actor(request)
write_audit(
db,
action="port_traffic.task.create",
actor_user_id=uid,
actor_username=uname,
method="POST",
path="/v1/port-traffic/tasks",
status_code=200,
detail={"id": out.id, "title": out.title, "start_now": body.start_now},
)
return out.model_dump()
@router.get("/tasks/{task_id}")
def api_get_task(task_id: str, db: Session = Depends(get_db)):
return get_task(db, task_id).model_dump()
@router.patch("/tasks/{task_id}")
def api_patch_task(
task_id: str,
body: PortTrafficTaskUpdate,
request: Request,
db: Session = Depends(get_db),
):
out = update_task(db, task_id, body)
uid, uname = _actor(request)
write_audit(
db,
action="port_traffic.task.update",
actor_user_id=uid,
actor_username=uname,
method="PATCH",
path=f"/v1/port-traffic/tasks/{task_id}",
status_code=200,
detail=body.model_dump(exclude_unset=True),
)
return out.model_dump()
@router.delete("/tasks/{task_id}")
def api_delete_task(task_id: str, request: Request, db: Session = Depends(get_db)):
out = delete_task(db, task_id)
uid, uname = _actor(request)
write_audit(
db,
action="port_traffic.task.delete",
actor_user_id=uid,
actor_username=uname,
method="DELETE",
path=f"/v1/port-traffic/tasks/{task_id}",
status_code=200,
detail={"id": task_id},
)
return out
@router.post("/tasks/{task_id}/start")
def api_start_task(task_id: str, request: Request, db: Session = Depends(get_db)):
out = set_task_status(db, task_id, "running")
uid, uname = _actor(request)
write_audit(
db,
action="port_traffic.task.start",
actor_user_id=uid,
actor_username=uname,
method="POST",
path=f"/v1/port-traffic/tasks/{task_id}/start",
status_code=200,
detail={"id": task_id},
)
return out.model_dump()
@router.post("/tasks/{task_id}/pause")
def api_pause_task(task_id: str, request: Request, db: Session = Depends(get_db)):
out = set_task_status(db, task_id, "paused")
uid, uname = _actor(request)
write_audit(
db,
action="port_traffic.task.pause",
actor_user_id=uid,
actor_username=uname,
method="POST",
path=f"/v1/port-traffic/tasks/{task_id}/pause",
status_code=200,
detail={"id": task_id},
)
return out.model_dump()
@router.post("/tasks/{task_id}/stop")
def api_stop_task(task_id: str, request: Request, db: Session = Depends(get_db)):
out = set_task_status(db, task_id, "stopped")
uid, uname = _actor(request)
write_audit(
db,
action="port_traffic.task.stop",
actor_user_id=uid,
actor_username=uname,
method="POST",
path=f"/v1/port-traffic/tasks/{task_id}/stop",
status_code=200,
detail={"id": task_id},
)
return out.model_dump()
@router.get("/tasks/{task_id}/targets")
def api_list_targets(task_id: str, db: Session = Depends(get_db)):
return {"items": [t.model_dump() for t in list_targets(db, task_id)]}
@router.put("/tasks/{task_id}/targets")
def api_put_targets(
task_id: str,
body: PortTrafficTargetsPut,
request: Request,
db: Session = Depends(get_db),
):
items = put_targets(db, task_id, body)
uid, uname = _actor(request)
write_audit(
db,
action="port_traffic.targets.put",
actor_user_id=uid,
actor_username=uname,
method="PUT",
path=f"/v1/port-traffic/tasks/{task_id}/targets",
status_code=200,
detail={"id": task_id, "count": len(items)},
)
return {"items": [t.model_dump() for t in items]}
@router.post("/discover/ports")
def api_discover_ports(body: DiscoverPortsRequest, db: Session = Depends(get_db)):
return discover_ports(db, body).model_dump()
@router.get("/samples")
def api_samples(
target_id: str = Query(..., description="port_traffic_target row id"),
from_ts: datetime | None = Query(default=None, alias="from"),
to_ts: datetime | None = Query(default=None, alias="to"),
db: Session = Depends(get_db),
):
return get_samples(db, target_row_id=target_id, from_ts=from_ts, to_ts=to_ts).model_dump()

View file

@ -0,0 +1,258 @@
"""Port traffic collection worker: claim task round, sample interfaces via CLI."""
from __future__ import annotations
import logging
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from datetime import datetime
from threading import Lock
from typing import Any
from uuid import uuid4
from fastapi import HTTPException
from .cli_resolve import resolve_cli_target
from .config import settings
from .db import SessionLocal
from .models import PortTrafficSample, PortTrafficTarget, PortTrafficTask
from .ne_session_factory import close_netmiko_connection, open_netmiko_connection
from .port_traffic_commands import commands_for_vendor, detail_command
from .port_traffic_parsers import parse_zte_interface_detail
_log = logging.getLogger("netx.port_traffic.runner")
_pools: dict[str, ThreadPoolExecutor] = {}
_pools_lock = Lock()
def _utcnow() -> datetime:
return datetime.utcnow()
def _format_error(exc: BaseException) -> str:
return f"{type(exc).__name__}: {exc}"[:1020]
def _pool_for_task(task_id: str, concurrency: int) -> ThreadPoolExecutor:
with _pools_lock:
pool = _pools.get(task_id)
if pool is None:
workers = max(1, min(20, int(concurrency or 5)))
pool = ThreadPoolExecutor(max_workers=workers, thread_name_prefix=f"pt-{task_id[:8]}")
_pools[task_id] = pool
return pool
def _release_pool(task_id: str) -> None:
with _pools_lock:
pool = _pools.pop(task_id, None)
if pool is not None:
try:
pool.shutdown(wait=False, cancel_futures=False)
except TypeError:
pool.shutdown(wait=False)
except Exception:
_log.exception("port_traffic pool shutdown failed task=%s", task_id)
def _set_target_error(target_row_id: str, message: str) -> None:
db = SessionLocal()
try:
row = db.get(PortTrafficTarget, target_row_id)
if row:
row.last_error = message[:1020]
db.commit()
finally:
db.close()
def _claim_collect_round(task_id: str) -> list[str] | None:
"""Mark task collect_running and return active target row ids, or None if skip."""
for attempt in range(8):
db = SessionLocal()
try:
task = db.get(PortTrafficTask, task_id)
if not task:
return None
if str(task.status or "") != "running":
return None
if bool(task.collect_running):
return None
ended = task.last_collect_ended_at
interval = max(15, int(task.interval_sec or 60))
if ended is not None:
elapsed = (_utcnow() - ended).total_seconds()
if elapsed < interval:
return None
targets = (
db.query(PortTrafficTarget)
.filter(
PortTrafficTarget.task_id == task_id,
PortTrafficTarget.status == "active",
)
.all()
)
if not targets:
return None
task.collect_running = True
task.last_collect_started_at = _utcnow()
task.last_error = ""
task.updated_at = _utcnow()
db.commit()
return [str(t.id) for t in targets]
except Exception:
db.rollback()
_log.exception("port_traffic claim failed task=%s attempt=%s", task_id, attempt)
time.sleep(0.05 * (attempt + 1))
finally:
db.close()
return None
def _finish_collect_round(task_id: str, *, error: str = "") -> None:
db = SessionLocal()
try:
task = db.get(PortTrafficTask, task_id)
if not task:
return
task.collect_running = False
task.last_collect_ended_at = _utcnow()
if error:
task.last_error = error[:1020]
task.updated_at = _utcnow()
db.commit()
finally:
db.close()
_release_pool(task_id)
def _run_show(creds: dict[str, Any], command: str, read_timeout: int) -> str:
conn = open_netmiko_connection(creds, session_timeout=read_timeout + 60)
try:
return str(conn.send_command(command_string=command, read_timeout=read_timeout) or "")
finally:
close_netmiko_connection(conn)
def _sample_one_target(target_row_id: str) -> None:
creds: dict[str, Any] | None = None
cmd = ""
per_cmd = int(settings.ne_collect_read_timeout_sec or 120)
cap = int(settings.ne_collect_run_timeout_cap_sec or 600)
db = SessionLocal()
try:
row = db.get(PortTrafficTarget, target_row_id)
if not row or str(row.status or "") != "active":
return
source = str(row.source or "").strip().lower()
target_id = str(row.target_id or "").strip()
ifname = str(row.ifname or "").strip()
vendor_hint = str(row.vendor or "")
try:
if source == "managed":
creds, device = resolve_cli_target(db, managed_ne_id=target_id)
elif source == "ume":
creds, device = resolve_cli_target(db, ume_ne_id=target_id)
else:
row.last_error = "invalid_source"
db.commit()
return
except HTTPException as exc:
row.last_error = str(exc.detail or "resolve_failed")[:1020]
db.commit()
return
except Exception as exc:
row.last_error = _format_error(exc)
db.commit()
return
vendor = str(device.get("vendor") or vendor_hint or "")
device_type = str(device.get("device_type") or "")
cmds = commands_for_vendor(vendor, device_type)
if cmds is None:
row.last_error = "unsupported_vendor"
db.commit()
return
cmd = detail_command(cmds, ifname)
finally:
db.close()
if not creds or not cmd:
return
budget = min(cap, per_cmd + 90)
try:
with ThreadPoolExecutor(max_workers=1) as pool:
fut = pool.submit(_run_show, creds, cmd, per_cmd)
raw = fut.result(timeout=budget)
except Exception as exc:
_set_target_error(target_row_id, _format_error(exc))
return
parsed = parse_zte_interface_detail(raw)
if parsed.in_bps == 0 and parsed.out_bps == 0 and parsed.bw_bps == 0 and not parsed.ifname:
_set_target_error(target_row_id, "parse_empty")
return
now = _utcnow()
db = SessionLocal()
try:
row = db.get(PortTrafficTarget, target_row_id)
if not row:
return
bw = int(parsed.bw_bps or row.bw_bps or 0)
if bw and not row.bw_bps:
row.bw_bps = bw
db.add(
PortTrafficSample(
id=uuid4().hex,
target_row_id=target_row_id,
ts=now,
in_bps=float(parsed.in_bps),
out_bps=float(parsed.out_bps),
in_util_pct=float(parsed.in_util_pct),
out_util_pct=float(parsed.out_util_pct),
bw_bps=bw,
rate_period_sec=int(parsed.rate_period_sec or 0),
raw_ok=True,
message="",
)
)
row.last_error = ""
row.last_sample_at = now
db.commit()
except Exception:
db.rollback()
_log.exception("port_traffic sample save failed target=%s", target_row_id)
finally:
db.close()
def dispatch_collect(task_id: str) -> int:
"""Claim and sample all active targets for a running task. Returns target count."""
target_ids = _claim_collect_round(task_id)
if not target_ids:
return 0
db = SessionLocal()
try:
task = db.get(PortTrafficTask, task_id)
concurrency = int(task.concurrency or 5) if task else 5
finally:
db.close()
pool = _pool_for_task(task_id, concurrency)
futures = [pool.submit(_sample_one_target, tid) for tid in target_ids]
errors = 0
try:
for fut in as_completed(futures):
try:
fut.result()
except Exception:
errors += 1
_log.exception("port_traffic target worker failed task=%s", task_id)
finally:
err_msg = f"{errors}_target_errors" if errors else ""
_finish_collect_round(task_id, error=err_msg)
return len(target_ids)

View file

@ -0,0 +1,96 @@
"""Background scheduler for port traffic collection + retention purge."""
from __future__ import annotations
import logging
import threading
from datetime import datetime
from .config import settings
from .db import SessionLocal
from .models import PortTrafficTask
from .port_traffic_runner import dispatch_collect
from .port_traffic_service import purge_expired_samples
_log = logging.getLogger("netx.port_traffic.scheduler")
_stop = threading.Event()
_thread: threading.Thread | None = None
_purge_counter = 0
def _utcnow() -> datetime:
return datetime.utcnow()
def try_dispatch_due_tasks() -> int:
"""Dispatch collect rounds for due running tasks. Returns number started."""
db = SessionLocal()
try:
tasks = (
db.query(PortTrafficTask)
.filter(PortTrafficTask.status == "running", PortTrafficTask.collect_running.is_(False))
.all()
)
due_ids: list[str] = []
now = _utcnow()
for task in tasks:
interval = max(15, int(task.interval_sec or 60))
ended = task.last_collect_ended_at
if ended is None:
due_ids.append(str(task.id))
continue
if (now - ended).total_seconds() >= interval:
due_ids.append(str(task.id))
finally:
db.close()
started = 0
for tid in due_ids:
try:
n = dispatch_collect(tid)
if n:
started += 1
_log.info("port_traffic collect started task=%s targets=%s", tid, n)
except Exception:
_log.exception("port_traffic dispatch failed task=%s", tid)
return started
def _loop() -> None:
global _purge_counter
tick = max(5, int(settings.port_traffic_scheduler_tick_sec or 15))
_log.info("port_traffic scheduler started tick=%ss", tick)
while not _stop.is_set():
try:
if bool(settings.port_traffic_scheduler_enabled):
try_dispatch_due_tasks()
_purge_counter += 1
# Retention purge roughly every ~20 ticks
if _purge_counter >= 20:
_purge_counter = 0
db = SessionLocal()
try:
purge_expired_samples(db)
finally:
db.close()
except Exception:
_log.exception("port_traffic scheduler tick failed")
_stop.wait(tick)
_log.info("port_traffic scheduler stopped")
def start_port_traffic_scheduler() -> None:
global _thread
if not bool(settings.port_traffic_scheduler_enabled):
_log.info("port_traffic scheduler disabled by settings")
return
if _thread and _thread.is_alive():
return
_stop.clear()
_thread = threading.Thread(target=_loop, name="port-traffic-scheduler", daemon=True)
_thread.start()
_log.info("started thread %s alive=%s", _thread.name, _thread.is_alive())
def stop_port_traffic_scheduler() -> None:
_stop.set()

View file

@ -0,0 +1,131 @@
"""Pydantic schemas for port traffic monitoring API."""
from __future__ import annotations
from datetime import datetime
from typing import Literal
from pydantic import BaseModel, Field
class PortTrafficNeRef(BaseModel):
source: Literal["managed", "ume"]
id: str
ne_name: str = ""
ne_ip: str = ""
vendor: str = ""
class PortTrafficTargetIn(BaseModel):
source: Literal["managed", "ume"]
target_id: str
ne_name: str = ""
ne_ip: str = ""
vendor: str = ""
ifname: str
if_description: str = ""
bw_bps: int = 0
class PortTrafficTaskCreate(BaseModel):
title: str = Field(min_length=1, max_length=256)
interval_sec: int = Field(default=60, ge=15, le=3600)
retention_days: int = Field(default=7, ge=1, le=90)
concurrency: int = Field(default=5, ge=1, le=20)
targets: list[PortTrafficTargetIn] = Field(default_factory=list)
start_now: bool = False
class PortTrafficTaskUpdate(BaseModel):
title: str | None = Field(default=None, min_length=1, max_length=256)
interval_sec: int | None = Field(default=None, ge=15, le=3600)
retention_days: int | None = Field(default=None, ge=1, le=90)
concurrency: int | None = Field(default=None, ge=1, le=20)
class PortTrafficTaskOut(BaseModel):
id: str
title: str
status: str
interval_sec: int
retention_days: int
concurrency: int
collect_running: bool = False
target_count: int = 0
active_target_count: int = 0
last_collect_started_at: datetime | None = None
last_collect_ended_at: datetime | None = None
last_error: str = ""
created_at: datetime | None = None
updated_at: datetime | None = None
class PortTrafficTargetOut(BaseModel):
id: str
task_id: str
source: str
target_id: str
ne_name: str
ne_ip: str
vendor: str
ifname: str
if_description: str
bw_bps: int
status: str
last_error: str = ""
last_sample_at: datetime | None = None
created_at: datetime | None = None
class PortTrafficTargetsPut(BaseModel):
targets: list[PortTrafficTargetIn]
class DiscoverPortsRequest(BaseModel):
source: Literal["managed", "ume"]
id: str
class DiscoverPortItem(BaseModel):
ifname: str
attribute: str = ""
mode: str = ""
bw_raw: str = ""
bw_bps: int = 0
admin: str = ""
phy: str = ""
prot: str = ""
description: str = ""
class DiscoverPortsResponse(BaseModel):
source: str
id: str
ne_name: str = ""
ne_ip: str = ""
vendor: str = ""
vendor_key: str = ""
ports: list[DiscoverPortItem] = Field(default_factory=list)
class PortTrafficSamplePoint(BaseModel):
ts: datetime
in_bps: float
out_bps: float
in_util_pct: float
out_util_pct: float
bw_bps: int
rate_period_sec: int = 0
class PortTrafficSamplesOut(BaseModel):
target: PortTrafficTargetOut
points: list[PortTrafficSamplePoint] = Field(default_factory=list)
class PortTrafficDashboardOut(BaseModel):
task_count: int = 0
running_task_count: int = 0
active_target_count: int = 0
sample_count_24h: int = 0
last_sample_at: datetime | None = None

View file

@ -0,0 +1,398 @@
"""Port traffic monitoring service: CRUD, discover, samples, dashboard."""
from __future__ import annotations
import logging
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 resolve_cli_target
from .config import settings
from .models import PortTrafficSample, PortTrafficTarget, PortTrafficTask
from .ne_session_factory import close_netmiko_connection, open_netmiko_connection
from .port_traffic_commands import commands_for_vendor
from .port_traffic_parsers import brief_port_to_dict, parse_zte_interface_brief
from .port_traffic_schemas import (
DiscoverPortItem,
DiscoverPortsRequest,
DiscoverPortsResponse,
PortTrafficDashboardOut,
PortTrafficSamplePoint,
PortTrafficSamplesOut,
PortTrafficTargetIn,
PortTrafficTargetOut,
PortTrafficTargetsPut,
PortTrafficTaskCreate,
PortTrafficTaskOut,
PortTrafficTaskUpdate,
)
_log = logging.getLogger("netx.port_traffic.service")
def _utcnow() -> datetime:
return datetime.utcnow()
def _task_out(db: Session, task: PortTrafficTask) -> PortTrafficTaskOut:
tid = str(task.id)
total = db.query(PortTrafficTarget).filter(PortTrafficTarget.task_id == tid).count()
active = (
db.query(PortTrafficTarget)
.filter(PortTrafficTarget.task_id == tid, PortTrafficTarget.status == "active")
.count()
)
return PortTrafficTaskOut(
id=tid,
title=str(task.title or ""),
status=str(task.status or ""),
interval_sec=int(task.interval_sec or 60),
retention_days=int(task.retention_days or 7),
concurrency=int(task.concurrency or 5),
collect_running=bool(task.collect_running),
target_count=int(total),
active_target_count=int(active),
last_collect_started_at=task.last_collect_started_at,
last_collect_ended_at=task.last_collect_ended_at,
last_error=str(task.last_error or ""),
created_at=task.created_at,
updated_at=task.updated_at,
)
def _target_out(row: PortTrafficTarget) -> PortTrafficTargetOut:
return PortTrafficTargetOut(
id=str(row.id),
task_id=str(row.task_id),
source=str(row.source or ""),
target_id=str(row.target_id or ""),
ne_name=str(row.ne_name or ""),
ne_ip=str(row.ne_ip or ""),
vendor=str(row.vendor or ""),
ifname=str(row.ifname or ""),
if_description=str(row.if_description or ""),
bw_bps=int(row.bw_bps or 0),
status=str(row.status or ""),
last_error=str(row.last_error or ""),
last_sample_at=row.last_sample_at,
created_at=row.created_at,
)
def _assert_zte_targets(targets: list[PortTrafficTargetIn]) -> None:
for t in targets:
cmds = commands_for_vendor(t.vendor or "", "")
if cmds is None:
raise HTTPException(
status_code=400,
detail=f"vendor_not_supported_for_port_traffic: {t.vendor or 'unknown'} ({t.ne_name or t.target_id})",
)
def list_tasks(db: Session, *, page: int = 1, page_size: int = 20) -> dict[str, Any]:
q = db.query(PortTrafficTask).order_by(PortTrafficTask.created_at.desc())
total = q.count()
rows = q.offset((page - 1) * page_size).limit(page_size).all()
return {
"total": total,
"page": page,
"page_size": page_size,
"items": [_task_out(db, r).model_dump() for r in rows],
}
def get_task(db: Session, task_id: str) -> PortTrafficTaskOut:
task = db.get(PortTrafficTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
return _task_out(db, task)
def create_task(db: Session, body: PortTrafficTaskCreate) -> PortTrafficTaskOut:
_assert_zte_targets(body.targets)
now = _utcnow()
status = "running" if body.start_now and body.targets else "draft"
if body.start_now and not body.targets:
status = "draft"
task = PortTrafficTask(
id=uuid4().hex,
title=body.title.strip(),
status=status,
interval_sec=int(body.interval_sec),
retention_days=int(body.retention_days),
concurrency=int(body.concurrency),
created_at=now,
updated_at=now,
)
db.add(task)
db.flush()
for t in body.targets:
db.add(
PortTrafficTarget(
id=uuid4().hex,
task_id=task.id,
source=t.source,
target_id=t.target_id,
ne_name=t.ne_name or "",
ne_ip=t.ne_ip or "",
vendor=t.vendor or "",
ifname=t.ifname.strip(),
if_description=t.if_description or "",
bw_bps=int(t.bw_bps or 0),
status="active",
created_at=now,
)
)
db.commit()
db.refresh(task)
return _task_out(db, task)
def update_task(db: Session, task_id: str, body: PortTrafficTaskUpdate) -> PortTrafficTaskOut:
task = db.get(PortTrafficTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
data = body.model_dump(exclude_unset=True)
if "title" in data and data["title"] is not None:
task.title = str(data["title"]).strip()
if "interval_sec" in data and data["interval_sec"] is not None:
task.interval_sec = int(data["interval_sec"])
if "retention_days" in data and data["retention_days"] is not None:
task.retention_days = int(data["retention_days"])
if "concurrency" in data and data["concurrency"] is not None:
task.concurrency = int(data["concurrency"])
task.updated_at = _utcnow()
db.commit()
db.refresh(task)
return _task_out(db, task)
def delete_task(db: Session, task_id: str) -> dict[str, Any]:
task = db.get(PortTrafficTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
if bool(task.collect_running):
raise HTTPException(status_code=409, detail="collect_running")
targets = db.query(PortTrafficTarget).filter(PortTrafficTarget.task_id == task_id).all()
ids = [str(t.id) for t in targets]
if ids:
db.query(PortTrafficSample).filter(PortTrafficSample.target_row_id.in_(ids)).delete(
synchronize_session=False
)
db.query(PortTrafficTarget).filter(PortTrafficTarget.task_id == task_id).delete(
synchronize_session=False
)
db.delete(task)
db.commit()
return {"ok": True, "id": task_id}
def set_task_status(db: Session, task_id: str, status: str) -> PortTrafficTaskOut:
task = db.get(PortTrafficTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
if status == "running":
active = (
db.query(PortTrafficTarget)
.filter(PortTrafficTarget.task_id == task_id, PortTrafficTarget.status == "active")
.count()
)
if active <= 0:
raise HTTPException(status_code=400, detail="no_active_targets")
task.status = status
task.updated_at = _utcnow()
if status in ("stopped", "paused"):
# leave collect_running for runner to finish; recovery clears stuck flags
pass
db.commit()
db.refresh(task)
return _task_out(db, task)
def put_targets(db: Session, task_id: str, body: PortTrafficTargetsPut) -> list[PortTrafficTargetOut]:
task = db.get(PortTrafficTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
if bool(task.collect_running):
raise HTTPException(status_code=409, detail="collect_running")
_assert_zte_targets(body.targets)
old = db.query(PortTrafficTarget).filter(PortTrafficTarget.task_id == task_id).all()
old_ids = [str(t.id) for t in old]
if old_ids:
db.query(PortTrafficSample).filter(PortTrafficSample.target_row_id.in_(old_ids)).delete(
synchronize_session=False
)
db.query(PortTrafficTarget).filter(PortTrafficTarget.task_id == task_id).delete(
synchronize_session=False
)
now = _utcnow()
rows: list[PortTrafficTarget] = []
for t in body.targets:
row = PortTrafficTarget(
id=uuid4().hex,
task_id=task_id,
source=t.source,
target_id=t.target_id,
ne_name=t.ne_name or "",
ne_ip=t.ne_ip or "",
vendor=t.vendor or "",
ifname=t.ifname.strip(),
if_description=t.if_description or "",
bw_bps=int(t.bw_bps or 0),
status="active",
created_at=now,
)
db.add(row)
rows.append(row)
task.updated_at = now
db.commit()
return [_target_out(r) for r in rows]
def list_targets(db: Session, task_id: str) -> list[PortTrafficTargetOut]:
task = db.get(PortTrafficTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
rows = (
db.query(PortTrafficTarget)
.filter(PortTrafficTarget.task_id == task_id)
.order_by(PortTrafficTarget.ne_name, PortTrafficTarget.ifname)
.all()
)
return [_target_out(r) for r in rows]
def discover_ports(db: Session, body: DiscoverPortsRequest) -> DiscoverPortsResponse:
try:
if body.source == "managed":
creds, device = resolve_cli_target(db, managed_ne_id=body.id)
else:
creds, device = resolve_cli_target(db, ume_ne_id=body.id)
except HTTPException:
raise
except Exception as exc:
raise HTTPException(status_code=400, detail=f"resolve_failed: {exc}") from exc
vendor = str(device.get("vendor") or creds.get("vendor") or "")
device_type = str(device.get("device_type") or creds.get("device_type") or "")
ne_name = str(device.get("name") or creds.get("host") or "")
ne_ip = str(device.get("ip_address") or creds.get("host") or "")
cmds = commands_for_vendor(vendor, device_type)
if cmds is None:
raise HTTPException(
status_code=400,
detail=f"vendor_not_supported_for_port_traffic: {vendor or 'unknown'}",
)
per_cmd = int(settings.ne_collect_read_timeout_sec or 120)
conn = open_netmiko_connection(creds, session_timeout=per_cmd + 60)
try:
raw = str(conn.send_command(command_string=cmds.brief, read_timeout=per_cmd) or "")
finally:
close_netmiko_connection(conn)
ports = [DiscoverPortItem(**brief_port_to_dict(p)) for p in parse_zte_interface_brief(raw)]
return DiscoverPortsResponse(
source=body.source,
id=body.id,
ne_name=ne_name,
ne_ip=ne_ip,
vendor=vendor,
vendor_key=cmds.vendor_key,
ports=ports,
)
def get_samples(
db: Session,
*,
target_row_id: str,
from_ts: datetime | None = None,
to_ts: datetime | None = None,
) -> PortTrafficSamplesOut:
target = db.get(PortTrafficTarget, target_row_id)
if not target:
raise HTTPException(status_code=404, detail="target_not_found")
now = _utcnow()
if to_ts is None:
to_ts = now
if from_ts is None:
from_ts = to_ts - timedelta(hours=1)
rows = (
db.query(PortTrafficSample)
.filter(
PortTrafficSample.target_row_id == target_row_id,
PortTrafficSample.ts >= from_ts,
PortTrafficSample.ts <= to_ts,
PortTrafficSample.raw_ok.is_(True),
)
.order_by(PortTrafficSample.ts.asc())
.all()
)
points = [
PortTrafficSamplePoint(
ts=r.ts,
in_bps=float(r.in_bps or 0),
out_bps=float(r.out_bps or 0),
in_util_pct=float(r.in_util_pct or 0),
out_util_pct=float(r.out_util_pct or 0),
bw_bps=int(r.bw_bps or 0),
rate_period_sec=int(r.rate_period_sec or 0),
)
for r in rows
]
return PortTrafficSamplesOut(target=_target_out(target), points=points)
def dashboard(db: Session) -> PortTrafficDashboardOut:
task_count = db.query(PortTrafficTask).count()
running = db.query(PortTrafficTask).filter(PortTrafficTask.status == "running").count()
active_targets = db.query(PortTrafficTarget).filter(PortTrafficTarget.status == "active").count()
since = _utcnow() - timedelta(hours=24)
sample_count = (
db.query(PortTrafficSample)
.filter(PortTrafficSample.ts >= since, PortTrafficSample.raw_ok.is_(True))
.count()
)
last = db.query(func.max(PortTrafficSample.ts)).scalar()
return PortTrafficDashboardOut(
task_count=int(task_count),
running_task_count=int(running),
active_target_count=int(active_targets),
sample_count_24h=int(sample_count),
last_sample_at=last,
)
def purge_expired_samples(db: Session) -> int:
"""Delete samples older than each task's retention_days."""
tasks = db.query(PortTrafficTask).all()
deleted = 0
now = _utcnow()
for task in tasks:
days = max(1, int(task.retention_days or 7))
cutoff = now - timedelta(days=days)
target_ids = [
str(t.id)
for t in db.query(PortTrafficTarget.id).filter(PortTrafficTarget.task_id == task.id).all()
]
if not target_ids:
continue
n = (
db.query(PortTrafficSample)
.filter(
PortTrafficSample.target_row_id.in_(target_ids),
PortTrafficSample.ts < cutoff,
)
.delete(synchronize_session=False)
)
deleted += int(n or 0)
if deleted:
db.commit()
_log.info("port_traffic retention purged samples=%s", deleted)
return deleted