From a2f91f6ee2aa15b073ea1ebf59c590f955b7e61f Mon Sep 17 00:00:00 2001 From: oliver Date: Sun, 2 Aug 2026 17:06:46 +0800 Subject: [PATCH] Split port traffic service and replace deprecated UTC/lifespan APIs. Keep the public facade stable while moving device CRUD and sample/compare logic into focused modules, and use naive UTC via timeutil plus FastAPI lifespan. Co-authored-by: Cursor --- netx_api/alarms_router.py | 3 +- netx_api/auth_service.py | 17 +- netx_api/main.py | 33 +- netx_api/models.py | 123 ++-- netx_api/ne_service_common.py | 3 +- netx_api/port_traffic_common.py | 159 +++++ netx_api/port_traffic_devices.py | 504 +++++++++++++++ netx_api/port_traffic_samples.py | 354 +++++++++++ netx_api/port_traffic_service.py | 1004 ++---------------------------- netx_api/timeutil.py | 9 + netx_api/topology_common.py | 3 +- netx_api/ume_support.py | 3 +- 12 files changed, 1171 insertions(+), 1044 deletions(-) create mode 100644 netx_api/port_traffic_common.py create mode 100644 netx_api/port_traffic_devices.py create mode 100644 netx_api/port_traffic_samples.py create mode 100644 netx_api/timeutil.py diff --git a/netx_api/alarms_router.py b/netx_api/alarms_router.py index 2e13ea7..a402dc3 100644 --- a/netx_api/alarms_router.py +++ b/netx_api/alarms_router.py @@ -17,6 +17,7 @@ from .db import get_db from .importer import aggregate_alarms, import_alarm_excel, query_alarms from .models import AiAnalyzeHistory, AlarmBatch, AlarmNorm, ImportErrorRow, ImportJob from .parser_config import load_parser_config +from .timeutil import utcnow_naive from .schemas import ( AiAnalyzeHistoryItem, AiAnalyzeHistoryResponse, @@ -276,7 +277,7 @@ def ap_analyze(payload: dict, db: Session = Depends(get_db)) -> dict: answer=answer, error=err, evidence_json=json.dumps(diag or {}, ensure_ascii=False), - created_at=datetime.utcnow(), + created_at=utcnow_naive(), ) db.add(row) db.commit() diff --git a/netx_api/auth_service.py b/netx_api/auth_service.py index 32fee55..9fb36a7 100644 --- a/netx_api/auth_service.py +++ b/netx_api/auth_service.py @@ -21,6 +21,7 @@ from .auth_scopes import ( from .auth_tokens import hash_api_token, issue_access_token, new_api_token_plaintext from .config import settings from .models import ApiToken, AppUser, AuditLog +from .timeutil import utcnow_naive _log = logging.getLogger("netx.auth") @@ -117,7 +118,7 @@ def flag_default_password_users(db: Session) -> None: continue if verify_password(default_pwd, user.password_hash): user.must_change_password = True - user.updated_at = datetime.utcnow() + user.updated_at = utcnow_naive() changed += 1 if changed: db.commit() @@ -336,7 +337,7 @@ def update_user( user.must_change_password = True if scopes is not None: user.scopes = normalize_scopes(scopes) - user.updated_at = datetime.utcnow() + user.updated_at = utcnow_naive() db.commit() db.refresh(user) return user @@ -354,7 +355,7 @@ def change_password(db: Session, *, user: AppUser, old_password: str, new_passwo raise HTTPException(status_code=400, detail="password_must_differ_from_default") row.password_hash = hash_password(pwd) row.must_change_password = False - row.updated_at = datetime.utcnow() + row.updated_at = utcnow_naive() db.commit() @@ -371,7 +372,7 @@ def create_api_token( raise HTTPException(status_code=400, detail="token_name_too_long") expires_at: datetime | None = None if expires_in_days is not None and int(expires_in_days) > 0: - expires_at = datetime.utcnow() + timedelta(days=int(expires_in_days)) + expires_at = utcnow_naive() + timedelta(days=int(expires_in_days)) plaintext = new_api_token_plaintext() scope_list = normalize_scopes(scopes) if scopes is not None else [] # Cap token scopes to owner's effective scopes. @@ -393,7 +394,7 @@ def create_api_token( def _token_public(db: Session, r: ApiToken) -> dict[str, Any]: owner = get_user_by_id(db, r.user_id) - now = datetime.utcnow() + now = utcnow_naive() expired = bool(r.expires_at and r.expires_at <= now) return { "id": r.id, @@ -426,7 +427,7 @@ def revoke_api_token(db: Session, *, token_id: str, actor: AppUser) -> ApiToken: if actor.role != "admin" and row.user_id != actor.id: raise HTTPException(status_code=403, detail="forbidden") if row.revoked_at is None: - row.revoked_at = datetime.utcnow() + row.revoked_at = utcnow_naive() db.commit() db.refresh(row) return row @@ -441,12 +442,12 @@ def resolve_api_token_row(db: Session, plaintext: str) -> ApiToken | None: ) if row is None: return None - if row.expires_at is not None and row.expires_at <= datetime.utcnow(): + if row.expires_at is not None and row.expires_at <= utcnow_naive(): return None user = get_user_by_id(db, row.user_id) if user is None or not user.is_active: return None - row.last_used_at = datetime.utcnow() + row.last_used_at = utcnow_naive() try: db.commit() except Exception: diff --git a/netx_api/main.py b/netx_api/main.py index b647233..563158c 100644 --- a/netx_api/main.py +++ b/netx_api/main.py @@ -4,6 +4,8 @@ from __future__ import annotations import logging import time +from contextlib import asynccontextmanager +from typing import AsyncIterator import uvicorn from fastapi import FastAPI @@ -42,12 +44,28 @@ from .webcrt_router import router as webcrt_router _schedule_log = logging.getLogger("netx.ume.schedule") _BOOT_MONO = time.monotonic() + +@asynccontextmanager +async def lifespan(_app: FastAPI) -> AsyncIterator[None]: + from .app_startup import run_api_startup + + run_api_startup() + try: + yield + finally: + shutdown_oclaw_alarm_forwarder() + if ume_support._UME_WS_STOP_EVENT is not None: + ume_support._UME_WS_STOP_EVENT.set() + shutdown_ws_consumer() + + app = FastAPI( title="netx ops tool", version="0.1.0", docs_url="/docs" if bool(settings.docs_enabled) else None, redoc_url="/redoc" if bool(settings.docs_enabled) else None, openapi_url="/openapi.json" if bool(settings.docs_enabled) else None, + lifespan=lifespan, ) app.add_middleware(AuthAuditMiddleware) app.include_router(auth_router) @@ -67,21 +85,6 @@ app.include_router(alarms_router) parser_cfg = load_parser_config() -@app.on_event("startup") -def on_startup() -> None: - from .app_startup import run_api_startup - - run_api_startup() - - -@app.on_event("shutdown") -def on_shutdown() -> None: - shutdown_oclaw_alarm_forwarder() - if ume_support._UME_WS_STOP_EVENT is not None: - ume_support._UME_WS_STOP_EVENT.set() - shutdown_ws_consumer() - - @app.get("/health", status_code=200) def health() -> dict[str, str]: return {"status": "ok"} diff --git a/netx_api/models.py b/netx_api/models.py index a3134c1..b473e17 100644 --- a/netx_api/models.py +++ b/netx_api/models.py @@ -9,6 +9,7 @@ from sqlalchemy.orm import Mapped, mapped_column, relationship from sqlalchemy.types import JSON from .db import Base +from .timeutil import utcnow_naive # JSONB on Postgres; plain JSON elsewhere (unit tests / sqlite). _JsonType = JSON().with_variant(JSONB(), "postgresql") @@ -25,7 +26,7 @@ class AlarmBatch(Base): success_rows: Mapped[int] = mapped_column(Integer, default=0) failed_rows: Mapped[int] = mapped_column(Integer, default=0) status: Mapped[str] = mapped_column(String(32), default="done") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) alarms: Mapped[list["AlarmNorm"]] = relationship(back_populates="batch") errors: Mapped[list["ImportErrorRow"]] = relationship(back_populates="batch") @@ -88,7 +89,7 @@ class ImportJob(Base): batch_id: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True) ok: Mapped[int] = mapped_column(Integer, default=1) # 1 ok, 0 error summary: Mapped[str] = mapped_column(String(1024), default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class AiAnalyzeHistory(Base): @@ -103,7 +104,7 @@ class AiAnalyzeHistory(Base): answer: Mapped[str] = mapped_column(Text, default="") error: Mapped[str] = mapped_column(Text, default="") evidence_json: Mapped[str] = mapped_column(Text, default="{}") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class UmeSyncJob(Base): @@ -113,7 +114,7 @@ class UmeSyncJob(Base): domain: Mapped[str] = mapped_column(String(32), index=True, default="inventory") status: Mapped[str] = mapped_column(String(32), index=True, default="running") trigger_mode: Mapped[str] = mapped_column(String(32), default="manual") - started_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + started_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) ended_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) pulled_count: Mapped[int] = mapped_column(Integer, default=0) inserted_count: Mapped[int] = mapped_column(Integer, default=0) @@ -131,7 +132,7 @@ class UmeAlarmBatch(Base): total_rows: Mapped[int] = mapped_column(Integer, default=0) success_rows: Mapped[int] = mapped_column(Integer, default=0) failed_rows: Mapped[int] = mapped_column(Integer, default=0) - started_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + started_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) ended_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) error_message: Mapped[str] = mapped_column(String(1024), default="") raw_json: Mapped[str] = mapped_column(Text, default="{}") @@ -163,8 +164,8 @@ class UmeInventoryNE(Base): creator: Mapped[str] = mapped_column(String(128), default="") vendor: Mapped[str] = mapped_column(String(64), default="ZTE", comment="网元提供商") source_type: Mapped[str] = mapped_column(String(64), default="ume_restconf") - first_seen_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - last_seen_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + first_seen_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + last_seen_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) raw_json: Mapped[str] = mapped_column(Text, default="{}") @@ -182,8 +183,8 @@ class UmeAlarmCurrent(Base): time_created: Mapped[str] = mapped_column(Text, default="", index=True) root_cause_alarm_indication: Mapped[str] = mapped_column(Text, default="") notification_id: Mapped[str] = mapped_column(String(128), default="", index=True) - first_seen_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - last_seen_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + first_seen_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + last_seen_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) raw_json: Mapped[str] = mapped_column(Text, default="{}") @@ -201,8 +202,8 @@ class UmeAlarmHistory(Base): time_created: Mapped[str] = mapped_column(Text, default="", index=True) root_cause_alarm_indication: Mapped[str] = mapped_column(Text, default="") notification_id: Mapped[str] = mapped_column(String(128), default="", index=True) - first_seen_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - last_seen_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + first_seen_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + last_seen_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) raw_json: Mapped[str] = mapped_column(Text, default="{}") @@ -213,7 +214,7 @@ class UmeKeyAlertMonitorConfig(Base): id: Mapped[int] = mapped_column(Integer, primary_key=True) forward_on_clear: Mapped[int] = mapped_column(Integer, default=0) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class UmeKeyAlertRule(Base): @@ -228,8 +229,8 @@ class UmeKeyAlertRule(Base): forward_on_clear: Mapped[int] = mapped_column(Integer, default=0) label: Mapped[str] = mapped_column(String(256), default="") ne_types: Mapped[str] = mapped_column(Text, default="[]") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class UmeKeyAlertForwardLog(Base): @@ -240,7 +241,7 @@ class UmeKeyAlertForwardLog(Base): action: Mapped[str] = mapped_column(String(32), default="", index=True) rule_key: Mapped[str] = mapped_column(String(128), default="", index=True) notification_id: Mapped[str] = mapped_column(String(128), default="", index=True) - forwarded_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + forwarded_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) oclaw_ok: Mapped[int] = mapped_column(Integer, default=0) error: Mapped[str] = mapped_column(String(512), default="") @@ -254,8 +255,8 @@ class UmeAlarmSubscription(Base): subscription_id: Mapped[str] = mapped_column(String(128), default="", index=True) wss_uri: Mapped[str] = mapped_column(Text, default="") topic: Mapped[str] = mapped_column(String(64), default="ALARM") - established_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + established_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class UmeTokenCache(Base): @@ -266,7 +267,7 @@ class UmeTokenCache(Base): expires_at_epoch_s: Mapped[int] = mapped_column(Integer, default=0) lock_owner: Mapped[str] = mapped_column(String(128), default="", index=True) lock_expires_at_epoch_s: Mapped[int] = mapped_column(Integer, default=0, index=True) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class ManagedNE(Base): @@ -305,8 +306,8 @@ class ManagedNE(Base): hop_command_template: Mapped[str] = mapped_column(Text, default="") hop_vrf: Mapped[str] = mapped_column(String(128), default="") hop_target_auth_mode: Mapped[str] = mapped_column(String(32), default="bastion_managed") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class CliConnectProfile(Base): @@ -334,8 +335,8 @@ class CliConnectProfile(Base): hop_command_template: Mapped[str] = mapped_column(Text, default="") hop_vrf: Mapped[str] = mapped_column(String(128), default="") hop_target_auth_mode: Mapped[str] = mapped_column(String(32), default="bastion_managed") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class UmeCliOverride(Base): @@ -354,7 +355,7 @@ class UmeCliOverride(Base): connect_message: Mapped[str] = mapped_column(String(512), default="") connect_detail: Mapped[str] = mapped_column(Text, default="") connect_tested_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class NeCollectionJob(Base): @@ -370,7 +371,7 @@ class NeCollectionJob(Base): success_count: Mapped[int] = mapped_column(Integer, default=0) fail_count: Mapped[int] = mapped_column(Integer, default=0) error_message: Mapped[str] = mapped_column(String(1024), default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) ended_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) last_run_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True) @@ -420,8 +421,8 @@ class TopoFabricNode(Base): region_source: Mapped[str] = mapped_column(String(16), default="") attrs: Mapped[dict] = mapped_column(_JsonType, default=dict) last_seen_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class TopoClassifyRule(Base): @@ -441,8 +442,8 @@ class TopoClassifyRule(Base): # role: {role}; region: {folder_id} or {region_name_from_group} payload: Mapped[dict] = mapped_column(_JsonType, default=dict) remark: Mapped[str] = mapped_column(String(512), default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class TopoFabricEdge(Base): @@ -474,8 +475,8 @@ class TopoFabricEdge(Base): attrs: Mapped[dict] = mapped_column(_JsonType, default=dict) discovered_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) last_seen_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class TopoFolder(Base): @@ -490,8 +491,8 @@ class TopoFolder(Base): name: Mapped[str] = mapped_column(String(256), default="", index=True) sort_order: Mapped[int] = mapped_column(Integer, default=0) is_system: Mapped[bool] = mapped_column(Boolean, default=False) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class TopoView(Base): @@ -513,8 +514,8 @@ class TopoView(Base): # { layer?, status?, membership?: {...} } filter: Mapped[dict] = mapped_column(_JsonType, default=dict) viewport: Mapped[dict] = mapped_column(_JsonType, default=dict) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class TopoViewNode(Base): @@ -530,8 +531,8 @@ class TopoViewNode(Base): y: Mapped[float] = mapped_column(Float, default=0.0) label: Mapped[str] = mapped_column(String(256), default="") locked: Mapped[bool] = mapped_column(Boolean, default=False) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class TopoViewEdgeStyle(Base): @@ -546,8 +547,8 @@ class TopoViewEdgeStyle(Base): stroke_color: Mapped[str] = mapped_column(String(32), default="") stroke_width: Mapped[int] = mapped_column(Integer, default=0) line_style: Mapped[str] = mapped_column(String(16), default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class LldpCollectPolicy(Base): @@ -565,7 +566,7 @@ class LldpCollectPolicy(Base): auto_add_unmatched: Mapped[bool] = mapped_column(Boolean, default=True) # Keep N finished discover jobs (items + raw_preview); 0 = keep none finished. history_keep: Mapped[int] = mapped_column(Integer, default=30) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class TopoDiscoverJob(Base): @@ -588,8 +589,8 @@ class TopoDiscoverJob(Base): error: Mapped[str] = mapped_column(String(1024), default="") started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) ended_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class TopoDiscoverJobItem(Base): @@ -615,7 +616,7 @@ class TopoDiscoverJobItem(Base): parser_stub: Mapped[bool] = mapped_column(Boolean, default=False) error: Mapped[str] = mapped_column(String(1024), default="") raw_preview: Mapped[str] = mapped_column(Text, default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class TopoFabricStats(Base): @@ -629,7 +630,7 @@ class TopoFabricStats(Base): edge_active: Mapped[int] = mapped_column(Integer, default=0) edge_stale: Mapped[int] = mapped_column(Integer, default=0) last_discover_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class AppUser(Base): @@ -646,8 +647,8 @@ class AppUser(Base): is_active: Mapped[bool] = mapped_column(Boolean, default=True, index=True) must_change_password: Mapped[bool] = mapped_column(Boolean, default=False) created_by: Mapped[str] = mapped_column(String(64), default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class AuditLog(Base): @@ -656,7 +657,7 @@ class AuditLog(Base): __tablename__ = "audit_log" id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: uuid4().hex) - ts: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + ts: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) actor_user_id: Mapped[str] = mapped_column(String(64), default="", index=True) actor_username: Mapped[str] = mapped_column(String(128), default="", index=True) action: Mapped[str] = mapped_column(String(128), default="", index=True) @@ -679,7 +680,7 @@ class ApiToken(Base): user_id: Mapped[str] = mapped_column(String(64), index=True) # Capability subset; empty inherits owner user scopes (then intersected). scopes: Mapped[list] = mapped_column(_JsonType, default=list) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True) last_used_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) revoked_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) @@ -699,7 +700,7 @@ class ConfigSyncPolicy(Base): history_keep: Mapped[int] = mapped_column(Integer, default=3) # Finished sync cycles to retain (newest kept); active cycles always kept. cycle_keep: Mapped[int] = mapped_column(Integer, default=30) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) class ConfigSyncCycle(Base): @@ -718,7 +719,7 @@ class ConfigSyncCycle(Base): error_message: Mapped[str] = mapped_column(String(1024), default="") started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) ended_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class ConfigSyncTask(Base): @@ -760,7 +761,7 @@ class NeConfigSnapshot(Base): zlib_size: Mapped[int] = mapped_column(Integer, default=0) zlib_alt_size: Mapped[int] = mapped_column(Integer, default=0) commands_json: Mapped[list] = mapped_column(_JsonType, default=list) - collected_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + collected_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) last_cycle_id: Mapped[str] = mapped_column(String(64), default="") last_task_id: Mapped[str] = mapped_column(String(64), default="") @@ -786,7 +787,7 @@ class NeConfigHistory(Base): zlib_size: Mapped[int] = mapped_column(Integer, default=0) zlib_alt_size: Mapped[int] = mapped_column(Integer, default=0) commands_json: Mapped[list] = mapped_column(_JsonType, default=list) - collected_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + collected_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) cycle_id: Mapped[str] = mapped_column(String(64), default="") task_id: Mapped[str] = mapped_column(String(64), default="") @@ -814,8 +815,8 @@ class PortTrafficDevice(Base): 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) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) # Back-compat alias while callers migrate. @@ -832,7 +833,7 @@ class PortTrafficSeries(Base): device_id: Mapped[str] = mapped_column(String(64), default="", index=True) title: Mapped[str] = mapped_column(String(256), default="") status: Mapped[str] = mapped_column(String(32), default="active", index=True) # active|disabled - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) @property def task_id(self) -> str: @@ -862,7 +863,7 @@ class PortTrafficTarget(Base): status: Mapped[str] = mapped_column(String(32), default="active", index=True) # active|disabled|retired 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) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) @property def task_id(self) -> str: @@ -882,7 +883,7 @@ class PortTrafficSample(Base): 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) series_id: Mapped[str] = mapped_column(String(64), default="", index=True) - ts: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + ts: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, 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) @@ -904,7 +905,7 @@ class PortTrafficEvent(Base): ifname: Mapped[str] = mapped_column(String(128), default="") level: Mapped[str] = mapped_column(String(16), default="error", index=True) # info|warn|error message: Mapped[str] = mapped_column(Text, default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class PortTrafficBoard(Base): @@ -918,8 +919,8 @@ class PortTrafficBoard(Base): cols: Mapped[int] = mapped_column(Integer, default=2) created_by: Mapped[str] = mapped_column(String(64), default="") updated_by: Mapped[str] = mapped_column(String(64), default="") - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive, index=True) class PortTrafficPanel(Base): @@ -940,5 +941,5 @@ class PortTrafficPanel(Base): ord: Mapped[int] = mapped_column(Integer, default=0) col_span: Mapped[int] = mapped_column(Integer, default=1) row_span: Mapped[int] = mapped_column(Integer, default=1) - created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) + updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive) diff --git a/netx_api/ne_service_common.py b/netx_api/ne_service_common.py index dfc7ebc..1e189e9 100644 --- a/netx_api/ne_service_common.py +++ b/netx_api/ne_service_common.py @@ -19,6 +19,7 @@ from .models import ManagedNE from .ne_crypto import CredentialCryptoError, credentials_configured, decrypt_secret, encrypt_secret from .ne_schemas import ManagedNeCreate, ManagedNeOut, ManagedNeUpdate from .ne_session_factory import default_bastion_username_template, default_hop_command_template +from .timeutil import utcnow_naive IMPORT_COLUMNS = ( "device_type", @@ -45,7 +46,7 @@ _BUILTIN_NE_TYPE_RULES: list[tuple[re.Pattern[str], str, str]] = [ ] def _now() -> datetime: - return datetime.utcnow() + return utcnow_naive() def _require_crypto() -> None: diff --git a/netx_api/port_traffic_common.py b/netx_api/port_traffic_common.py new file mode 100644 index 0000000..8279f66 --- /dev/null +++ b/netx_api/port_traffic_common.py @@ -0,0 +1,159 @@ +"""Shared helpers for port traffic device/sample services.""" + +from __future__ import annotations + +from datetime import datetime +from typing import Any +from uuid import uuid4 + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .models import ( + PortTrafficDevice, + PortTrafficSeries, + PortTrafficTarget, +) +from .port_traffic_commands import commands_for_vendor +from .port_traffic_migrate import default_series_title, unique_series_title +from .port_traffic_schemas import ( + PortTrafficDeviceOut, + PortTrafficIfaceIn, + PortTrafficTargetOut, +) +from .timeutil import utcnow_naive + + +def _utcnow() -> datetime: + return utcnow_naive() + + +def _target_out(row: PortTrafficTarget) -> PortTrafficTargetOut: + did = str(row.device_id or "") + return PortTrafficTargetOut( + id=str(row.id), + device_id=did, + task_id=did, + series_id=str(row.series_id or ""), + 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 _device_out(db: Session, device: PortTrafficDevice) -> PortTrafficDeviceOut: + did = str(device.id) + total = db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == did).count() + active = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.device_id == did, PortTrafficTarget.status == "active") + .count() + ) + return PortTrafficDeviceOut( + id=did, + source=str(device.source or ""), + ne_id=str(device.ne_id or ""), + ne_name=str(device.ne_name or ""), + ne_ip=str(device.ne_ip or ""), + vendor=str(device.vendor or ""), + note=str(device.note or ""), + status=str(device.status or ""), + interval_sec=int(device.interval_sec or 60), + retention_days=int(device.retention_days or 7), + concurrency=int(device.concurrency or 1), + collect_running=bool(device.collect_running), + target_count=int(total), + active_target_count=int(active), + last_collect_started_at=device.last_collect_started_at, + last_collect_ended_at=device.last_collect_ended_at, + last_error=str(device.last_error or ""), + created_at=device.created_at, + updated_at=device.updated_at, + ) + + +def _assert_vendor(vendor: str, label: str) -> None: + cmds = commands_for_vendor(vendor or "", "") + if cmds is None: + raise HTTPException( + status_code=400, + detail=f"vendor_not_supported_for_port_traffic: {vendor or 'unknown'} ({label})", + ) + + +def _assert_ifaces_free( + db: Session, + *, + source: str, + ne_id: str, + ifnames: list[str], + exclude_device_id: str | None = None, +) -> None: + names = [str(x).strip() for x in ifnames if str(x).strip()] + if not names: + return + q = db.query(PortTrafficTarget).filter( + PortTrafficTarget.source == source, + PortTrafficTarget.target_id == ne_id, + PortTrafficTarget.ifname.in_(names), + PortTrafficTarget.status == "active", + ) + if exclude_device_id: + q = q.filter(PortTrafficTarget.device_id != exclude_device_id) + hit = q.first() + if hit: + raise HTTPException( + status_code=409, + detail=f"interface_already_monitored: {hit.ifname} on device {hit.device_id}", + ) + + +def _create_iface( + db: Session, + *, + device: PortTrafficDevice, + iface: PortTrafficIfaceIn, + now: datetime, +) -> PortTrafficTarget: + ifname = iface.ifname.strip() + title = unique_series_title( + db, str(device.id), default_series_title(device.ne_name or "", ifname) + ) + sid = uuid4().hex + db.add( + PortTrafficSeries( + id=sid, + device_id=str(device.id), + title=title, + status="active", + created_at=now, + ) + ) + row = PortTrafficTarget( + id=uuid4().hex, + device_id=str(device.id), + series_id=sid, + source=str(device.source), + target_id=str(device.ne_id), + ne_name=str(device.ne_name or ""), + ne_ip=str(device.ne_ip or ""), + vendor=str(device.vendor or ""), + ifname=ifname, + if_description=iface.if_description or "", + bw_bps=int(iface.bw_bps or 0), + status="active", + created_at=now, + ) + db.add(row) + return row + + diff --git a/netx_api/port_traffic_devices.py b/netx_api/port_traffic_devices.py new file mode 100644 index 0000000..aee5970 --- /dev/null +++ b/netx_api/port_traffic_devices.py @@ -0,0 +1,504 @@ +"""Port traffic device CRUD, interfaces, series, and discover.""" + +from __future__ import annotations + +import logging +from typing import Any +from uuid import uuid4 + +from fastapi import HTTPException +from sqlalchemy.orm import Session + +from .cli_resolve import resolve_cli_target +from .models import ( + ManagedNE, + PortTrafficDevice, + PortTrafficEvent, + PortTrafficSample, + PortTrafficSeries, + PortTrafficTarget, + UmeInventoryNE, +) +from .ne_netmiko import send_show_command +from .ne_session_factory import close_netmiko_connection, open_netmiko_connection +from .port_traffic_commands import commands_for_vendor +from .port_traffic_common import ( + _assert_ifaces_free, + _assert_vendor, + _create_iface, + _device_out, + _target_out, + _utcnow, +) +from .port_traffic_migrate import default_series_title, unique_series_title +from .port_traffic_parsers import brief_port_to_dict, parse_interface_brief +from .port_traffic_schemas import ( + DiscoverPortItem, + DiscoverPortsRequest, + DiscoverPortsResponse, + PortTrafficDeviceCreate, + PortTrafficDeviceOut, + PortTrafficDeviceRebind, + PortTrafficDeviceUpdate, + PortTrafficInterfacesPut, + PortTrafficReplacePortRequest, + PortTrafficSeriesOut, + PortTrafficTargetOut, +) + +_log = logging.getLogger("netx.port_traffic.service") + + +def list_devices(db: Session, *, page: int = 1, page_size: int = 20) -> dict[str, Any]: + q = db.query(PortTrafficDevice).order_by(PortTrafficDevice.ne_name.asc(), PortTrafficDevice.ne_ip.asc()) + total = q.count() + rows = q.offset((page - 1) * page_size).limit(page_size).all() + return { + "total": total, + "page": page, + "page_size": page_size, + "items": [_device_out(db, r).model_dump() for r in rows], + } + + +def get_device(db: Session, device_id: str) -> PortTrafficDeviceOut: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + return _device_out(db, device) + + +def create_device(db: Session, body: PortTrafficDeviceCreate) -> PortTrafficDeviceOut: + source = body.source + ne_id = body.ne_id.strip() + _assert_vendor(body.vendor, body.ne_name or ne_id) + clash = ( + db.query(PortTrafficDevice) + .filter(PortTrafficDevice.source == source, PortTrafficDevice.ne_id == ne_id) + .first() + ) + if clash: + raise HTTPException(status_code=409, detail="device_already_monitored") + ifnames = [i.ifname for i in body.interfaces] + _assert_ifaces_free(db, source=source, ne_id=ne_id, ifnames=ifnames) + + now = _utcnow() + status = "running" if body.start_now and body.interfaces else "draft" + device = PortTrafficDevice( + id=uuid4().hex, + source=source, + ne_id=ne_id, + ne_name=(body.ne_name or "").strip(), + ne_ip=(body.ne_ip or "").strip(), + vendor=(body.vendor or "").strip(), + note=(body.note or "").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(device) + db.flush() + for iface in body.interfaces: + if not iface.ifname.strip(): + continue + _create_iface(db, device=device, iface=iface, now=now) + db.commit() + db.refresh(device) + return _device_out(db, device) + + +def update_device(db: Session, device_id: str, body: PortTrafficDeviceUpdate) -> PortTrafficDeviceOut: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + data = body.model_dump(exclude_unset=True) + for key in ("note", "ne_name", "ne_ip", "vendor"): + if key in data and data[key] is not None: + setattr(device, key, str(data[key]).strip()) + if "interval_sec" in data and data["interval_sec"] is not None: + device.interval_sec = int(data["interval_sec"]) + if "retention_days" in data and data["retention_days"] is not None: + device.retention_days = int(data["retention_days"]) + if "concurrency" in data and data["concurrency"] is not None: + device.concurrency = int(data["concurrency"]) + device.updated_at = _utcnow() + db.commit() + db.refresh(device) + return _device_out(db, device) + + +def rebind_device( + db: Session, + device_id: str, + body: PortTrafficDeviceRebind, +) -> PortTrafficDeviceOut: + """Point a monitor device at an explicitly chosen inventory NE; keeps samples.""" + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + if bool(device.collect_running): + raise HTTPException(status_code=409, detail="collect_running") + + source = str(device.source or "").strip().lower() or "managed" + want_id = str(body.ne_id or "").strip() + if not want_id: + raise HTTPException(status_code=400, detail="ne_id_required") + + new_id = "" + new_name = "" + new_ip = "" + new_vendor = "" + + if source == "managed": + row = db.get(ManagedNE, want_id) + if not row: + raise HTTPException(status_code=404, detail="managed_ne_not_found") + new_id = str(row.id) + new_name = str(row.name or "") + new_ip = str(row.ip_address or "") + new_vendor = str(row.vendor or "") + elif source == "ume": + inv = db.get(UmeInventoryNE, want_id) + if not inv: + raise HTTPException(status_code=404, detail="ume_ne_not_found") + new_id = str(inv.ne_id) + new_ip = str(inv.ip_address or "") + new_name = str(inv.user_label or inv.ne_name or inv.host_name or new_ip or "").strip() + new_vendor = str(inv.vendor or "") + else: + raise HTTPException(status_code=400, detail="invalid_source") + + if new_id == str(device.ne_id or ""): + # Already bound; clear stale errors so collect can resume. + device.last_error = "" + device.updated_at = _utcnow() + for tgt in db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).all(): + if "managed_ne_not_found" in str(tgt.last_error or "") or "ume_ne_not_found" in str( + tgt.last_error or "" + ): + tgt.last_error = "" + db.commit() + db.refresh(device) + return _device_out(db, device) + + clash = ( + db.query(PortTrafficDevice) + .filter( + PortTrafficDevice.source == source, + PortTrafficDevice.ne_id == new_id, + PortTrafficDevice.id != device_id, + ) + .first() + ) + if clash: + raise HTTPException(status_code=409, detail="device_already_monitored") + + device.ne_id = new_id + if new_name: + device.ne_name = new_name + if new_ip: + device.ne_ip = new_ip + if new_vendor: + device.vendor = new_vendor + device.last_error = "" + device.updated_at = _utcnow() + + for tgt in db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).all(): + tgt.target_id = new_id + if new_name: + tgt.ne_name = new_name + if new_ip: + tgt.ne_ip = new_ip + if new_vendor: + tgt.vendor = new_vendor + if "managed_ne_not_found" in str(tgt.last_error or "") or "ume_ne_not_found" in str( + tgt.last_error or "" + ): + tgt.last_error = "" + + db.commit() + db.refresh(device) + _log.info("port_traffic rebind device=%s -> ne_id=%s ip=%s", device_id, new_id, new_ip) + return _device_out(db, device) + + +def delete_device(db: Session, device_id: str) -> dict[str, Any]: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + if bool(device.collect_running): + raise HTTPException(status_code=409, detail="collect_running") + targets = db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).all() + ids = [str(t.id) for t in targets] + series_ids = [str(t.series_id) for t in targets if t.series_id] + if ids: + from .port_traffic_board_service import delete_panels_for_targets + + delete_panels_for_targets(db, ids) + db.query(PortTrafficSample).filter(PortTrafficSample.target_row_id.in_(ids)).delete( + synchronize_session=False + ) + db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).delete( + synchronize_session=False + ) + db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id).delete( + synchronize_session=False + ) + if series_ids: + db.query(PortTrafficSeries).filter(PortTrafficSeries.id.in_(series_ids)).delete( + synchronize_session=False + ) + else: + db.query(PortTrafficSeries).filter(PortTrafficSeries.device_id == device_id).delete( + synchronize_session=False + ) + db.delete(device) + db.commit() + return {"ok": True, "id": device_id} + + +def set_device_status(db: Session, device_id: str, status: str) -> PortTrafficDeviceOut: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + if status == "running": + active = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.device_id == device_id, PortTrafficTarget.status == "active") + .count() + ) + if active <= 0: + raise HTTPException(status_code=400, detail="no_active_targets") + device.last_collect_ended_at = None + device.status = status + device.updated_at = _utcnow() + db.commit() + db.refresh(device) + return _device_out(db, device) + + +def put_interfaces(db: Session, device_id: str, body: PortTrafficInterfacesPut) -> list[PortTrafficTargetOut]: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + if bool(device.collect_running): + raise HTTPException(status_code=409, detail="collect_running") + + wanted = {i.ifname.strip(): i for i in body.interfaces if i.ifname.strip()} + _assert_ifaces_free( + db, + source=str(device.source), + ne_id=str(device.ne_id), + ifnames=list(wanted.keys()), + exclude_device_id=device_id, + ) + + now = _utcnow() + existing = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.device_id == device_id) + .all() + ) + by_if = {str(t.ifname): t for t in existing if str(t.status) == "active"} + + for ifname, row in list(by_if.items()): + if ifname not in wanted: + row.status = "retired" + + for ifname, iface in wanted.items(): + if ifname in by_if: + row = by_if[ifname] + row.if_description = iface.if_description or row.if_description + if iface.bw_bps: + row.bw_bps = int(iface.bw_bps) + continue + # Reactivate retired same ifname if present + retired = next( + ( + t + for t in existing + if str(t.ifname) == ifname and str(t.status) in ("retired", "disabled") + ), + None, + ) + if retired: + retired.status = "active" + retired.if_description = iface.if_description or retired.if_description + if iface.bw_bps: + retired.bw_bps = int(iface.bw_bps) + continue + _create_iface(db, device=device, iface=iface, now=now) + + device.updated_at = now + db.commit() + return list_targets(db, device_id) + + +def list_targets(db: Session, device_id: str) -> list[PortTrafficTargetOut]: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + rows = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.device_id == device_id) + .order_by(PortTrafficTarget.ifname) + .all() + ) + return [_target_out(r) for r in rows] + + +def list_series(db: Session, device_id: str) -> list[PortTrafficSeriesOut]: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + rows = ( + db.query(PortTrafficSeries) + .filter(PortTrafficSeries.device_id == device_id) + .order_by(PortTrafficSeries.title.asc()) + .all() + ) + out: list[PortTrafficSeriesOut] = [] + for s in rows: + active = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.series_id == s.id, PortTrafficTarget.status == "active") + .order_by(PortTrafficTarget.created_at.desc()) + .first() + ) + retired = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.series_id == s.id, PortTrafficTarget.status == "retired") + .count() + ) + did = str(s.device_id or "") + out.append( + PortTrafficSeriesOut( + id=str(s.id), + device_id=did, + task_id=did, + title=str(s.title or ""), + status=str(s.status or ""), + active_target=_target_out(active) if active else None, + retired_target_count=int(retired), + created_at=s.created_at, + ) + ) + return out + + +def replace_series_port( + db: Session, + device_id: str, + series_id: str, + body: PortTrafficReplacePortRequest, +) -> PortTrafficSeriesOut: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + if bool(device.collect_running): + raise HTTPException(status_code=409, detail="collect_running") + series = db.get(PortTrafficSeries, series_id) + if not series or str(series.device_id) != device_id: + raise HTTPException(status_code=404, detail="series_not_found") + ifname = body.ifname.strip() + _assert_ifaces_free( + db, + source=str(device.source), + ne_id=str(device.ne_id), + ifnames=[ifname], + exclude_device_id=device_id, + ) + now = _utcnow() + actives = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.series_id == series_id, PortTrafficTarget.status == "active") + .all() + ) + for old in actives: + old.status = "retired" + if body.series_title is not None and str(body.series_title).strip(): + wanted = str(body.series_title).strip()[:256] + if wanted != str(series.title or ""): + clash = ( + db.query(PortTrafficSeries.id) + .filter( + PortTrafficSeries.device_id == device_id, + PortTrafficSeries.title == wanted, + PortTrafficSeries.id != series_id, + ) + .first() + ) + if clash: + raise HTTPException(status_code=400, detail="series_title_exists") + series.title = wanted + row = PortTrafficTarget( + id=uuid4().hex, + device_id=device_id, + series_id=series_id, + source=str(device.source), + target_id=str(device.ne_id), + ne_name=str(device.ne_name or ""), + ne_ip=str(device.ne_ip or ""), + vendor=str(device.vendor or ""), + ifname=ifname, + if_description=body.if_description or "", + bw_bps=int(body.bw_bps or 0), + status="active", + created_at=now, + ) + db.add(row) + device.updated_at = now + db.commit() + for item in list_series(db, device_id): + if item.id == series_id: + return item + raise HTTPException(status_code=500, detail="series_replace_failed") + + +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 = send_show_command(conn, cmds.brief, read_timeout=per_cmd) + finally: + close_netmiko_connection(conn) + + ports = [ + DiscoverPortItem(**brief_port_to_dict(p)) + for p in parse_interface_brief(raw, cmds.vendor_key) + ] + 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, + ) + + diff --git a/netx_api/port_traffic_samples.py b/netx_api/port_traffic_samples.py new file mode 100644 index 0000000..843a401 --- /dev/null +++ b/netx_api/port_traffic_samples.py @@ -0,0 +1,354 @@ +"""Port traffic samples, compare, dashboard, events, and retention.""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from typing import Any +from uuid import uuid4 + +from fastapi import HTTPException +from sqlalchemy import func +from sqlalchemy.orm import Session + +from .config import settings +from .models import ( + PortTrafficDevice, + PortTrafficEvent, + PortTrafficSample, + PortTrafficSeries, + PortTrafficTarget, +) +from .port_traffic_common import _device_out, _target_out, _utcnow +from .port_traffic_parsers import resolve_util_pct +from .port_traffic_schemas import ( + PortTrafficCompareMeta, + PortTrafficCompareOut, + PortTrafficDashboardOut, + PortTrafficEventOut, + PortTrafficEventsOut, + PortTrafficSamplePoint, + PortTrafficSamplesOut, +) + +def _as_naive_utc(value: datetime | None) -> datetime | None: + if value is None: + return None + if value.tzinfo is None: + return value + return value.astimezone(timezone.utc).replace(tzinfo=None) + + +def _sample_points( + rows: list[PortTrafficSample], + *, + align_offset: timedelta | None = None, +) -> list[PortTrafficSamplePoint]: + points: list[PortTrafficSamplePoint] = [] + for r in rows: + ts_raw = r.ts + ts = ts_raw + if align_offset is not None and ts_raw is not None: + ts = ts_raw + align_offset + in_bps = float(r.in_bps or 0) + out_bps = float(r.out_bps or 0) + bw = int(r.bw_bps or 0) + points.append( + PortTrafficSamplePoint( + ts=ts, + ts_raw=ts_raw if align_offset is not None else None, + in_bps=in_bps, + out_bps=out_bps, + in_util_pct=resolve_util_pct(float(r.in_util_pct or 0), in_bps, bw), + out_util_pct=resolve_util_pct(float(r.out_util_pct or 0), out_bps, bw), + bw_bps=bw, + rate_period_sec=int(r.rate_period_sec or 0), + ) + ) + return points + + +def _query_target_samples( + db: Session, + *, + target_row_id: str, + from_ts: datetime, + to_ts: datetime, +) -> list[PortTrafficSample]: + return ( + 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() + ) + + +def baseline_offset_hours(baseline: str, range_hours: float, offset_hours: float | None) -> float | None: + key = str(baseline or "off").strip().lower() + if key in ("", "off", "none"): + return None + if key == "shift": + return float(range_hours) + if key == "day": + return 24.0 + if key == "week": + return 24.0 * 7 + if key == "custom": + if offset_hours is None or float(offset_hours) <= 0: + raise HTTPException(status_code=400, detail="offset_hours_required") + return float(offset_hours) + raise HTTPException(status_code=400, detail=f"invalid_baseline: {baseline}") + + +def compare_targets( + db: Session, + *, + target_row_id: str, + range_hours: float = 24, + baseline: str = "off", + offset_hours: float | None = None, + baseline_target_id: str | None = None, + ahead_hours: float = 0, + to_ts: datetime | None = None, +) -> PortTrafficCompareOut: + target = db.get(PortTrafficTarget, target_row_id) + if not target: + raise HTTPException(status_code=404, detail="target_not_found") + + mapped: PortTrafficTarget | None = None + mapped_id = str(baseline_target_id or "").strip() + if mapped_id: + if mapped_id == str(target.id): + raise HTTPException(status_code=400, detail="baseline_target_same_as_current") + mapped = db.get(PortTrafficTarget, mapped_id) + if not mapped: + raise HTTPException(status_code=404, detail="baseline_target_not_found") + + now = _utcnow() + # Anchor = "now" (or explicit to). Lookback is from anchor; ahead extends past it + # so period compare can show baseline trend after the current clock time. + anchor = _as_naive_utc(to_ts) or now + range_h = max(0.25, float(range_hours or 24)) + ahead_h = max(0.0, min(24.0, float(ahead_hours or 0))) + from_ts = anchor - timedelta(hours=range_h) + to_end = anchor + timedelta(hours=ahead_h) + to_q = to_end + timedelta(seconds=5) + current = _sample_points( + _query_target_samples(db, target_row_id=str(target.id), from_ts=from_ts, to_ts=to_q) + ) + + off_h = baseline_offset_hours(baseline, range_h, offset_hours) + baseline_points: list[PortTrafficSamplePoint] = [] + base_src = mapped if mapped is not None else target + want_baseline = off_h is not None or mapped is not None + if want_baseline: + if off_h is not None: + delta = timedelta(hours=off_h) + base_rows = _query_target_samples( + db, + target_row_id=str(base_src.id), + from_ts=from_ts - delta, + to_ts=to_q - delta, + ) + baseline_points = _sample_points(base_rows, align_offset=delta) + else: + base_rows = _query_target_samples( + db, target_row_id=str(base_src.id), from_ts=from_ts, to_ts=to_q + ) + baseline_points = _sample_points(base_rows) + + return PortTrafficCompareOut( + meta=PortTrafficCompareMeta( + target_id=str(target.id), + baseline=str(baseline or "off"), + offset_hours=float(off_h or 0), + range_hours=range_h, + ahead_hours=ahead_h, + current_target=_target_out(target), + baseline_target=_target_out(mapped) if mapped is not None else None, + baseline_target_id=str(mapped.id) if mapped is not None else "", + ), + current=current, + baseline=baseline_points, + ) + + +def compare_series(db: Session, **kwargs: Any) -> PortTrafficCompareOut: + target_row_id = kwargs.pop("target_row_id", None) or kwargs.pop("target_id", None) + series_id = kwargs.pop("series_id", None) + if not target_row_id and series_id: + active = ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.series_id == series_id, PortTrafficTarget.status == "active") + .order_by(PortTrafficTarget.created_at.desc()) + .first() + ) + if not active: + raise HTTPException(status_code=404, detail="series_active_target_not_found") + target_row_id = str(active.id) + if not target_row_id: + raise HTTPException(status_code=400, detail="target_id_required") + return compare_targets(db, target_row_id=str(target_row_id), **kwargs) + + +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() + to_ts = _as_naive_utc(to_ts) or now + from_ts = _as_naive_utc(from_ts) or (to_ts - timedelta(hours=1)) + to_ts = to_ts + timedelta(seconds=5) + rows = _query_target_samples(db, target_row_id=target_row_id, from_ts=from_ts, to_ts=to_ts) + return PortTrafficSamplesOut(target=_target_out(target), points=_sample_points(rows)) + + +def dashboard(db: Session) -> PortTrafficDashboardOut: + device_count = db.query(PortTrafficDevice).count() + running = db.query(PortTrafficDevice).filter(PortTrafficDevice.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( + device_count=int(device_count), + running_device_count=int(running), + active_target_count=int(active_targets), + sample_count_24h=int(sample_count), + last_sample_at=last, + task_count=int(device_count), + running_task_count=int(running), + ) + + +def list_device_events( + db: Session, + device_id: str, + *, + limit: int = 100, +) -> PortTrafficEventsOut: + device = db.get(PortTrafficDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="device_not_found") + lim = max(1, min(500, int(limit or 100))) + q = db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id) + total = q.count() + # Seed once from current last_error snapshot so older failures still show in log UI. + if total == 0: + seeded = False + if str(device.last_error or "").strip(): + append_device_event( + db, + device_id=device_id, + message=str(device.last_error), + level="error", + ) + seeded = True + for t in ( + db.query(PortTrafficTarget) + .filter(PortTrafficTarget.device_id == device_id) + .order_by(PortTrafficTarget.ifname) + .all() + ): + err = str(t.last_error or "").strip() + if not err: + continue + append_device_event( + db, + device_id=device_id, + target_row_id=str(t.id), + ifname=str(t.ifname or ""), + message=err, + level="error", + ) + seeded = True + if seeded: + db.commit() + total = db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id).count() + q = db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id) + rows = q.order_by(PortTrafficEvent.created_at.desc()).limit(lim).all() + items = [ + PortTrafficEventOut( + id=str(r.id), + device_id=str(r.device_id or ""), + target_row_id=str(r.target_row_id or ""), + ifname=str(r.ifname or ""), + level=str(r.level or "error"), + message=str(r.message or ""), + created_at=r.created_at, + ) + for r in rows + ] + return PortTrafficEventsOut(items=items, total=int(total)) + + +def append_device_event( + db: Session, + *, + device_id: str, + message: str, + level: str = "error", + target_row_id: str = "", + ifname: str = "", +) -> None: + msg = str(message or "").strip() + if not msg or not device_id: + return + db.add( + PortTrafficEvent( + id=uuid4().hex, + device_id=device_id, + target_row_id=str(target_row_id or ""), + ifname=str(ifname or ""), + level=str(level or "error")[:16], + message=msg[:4000], + created_at=_utcnow(), + ) + ) + + +def purge_expired_samples(db: Session) -> int: + devices = db.query(PortTrafficDevice).all() + deleted = 0 + now = _utcnow() + for device in devices: + days = max(1, int(device.retention_days or 7)) + cutoff = now - timedelta(days=days) + target_ids = [ + str(t.id) + for t in db.query(PortTrafficTarget.id) + .filter(PortTrafficTarget.device_id == device.id) + .all() + ] + if target_ids: + n = ( + db.query(PortTrafficSample) + .filter( + PortTrafficSample.target_row_id.in_(target_ids), + PortTrafficSample.ts < cutoff, + ) + .delete(synchronize_session=False) + ) + deleted += int(n or 0) + db.query(PortTrafficEvent).filter( + PortTrafficEvent.device_id == device.id, + PortTrafficEvent.created_at < cutoff, + ).delete(synchronize_session=False) + db.commit() + if deleted: + _log.info("port_traffic retention purged samples=%s", deleted) + return deleted diff --git a/netx_api/port_traffic_service.py b/netx_api/port_traffic_service.py index 1f329a9..0aa3ee8 100644 --- a/netx_api/port_traffic_service.py +++ b/netx_api/port_traffic_service.py @@ -1,962 +1,54 @@ -"""Port traffic monitoring service: device-centric CRUD, discover, samples, dashboard.""" +"""Port traffic monitoring service facade (device CRUD, discover, samples, dashboard).""" from __future__ import annotations -import logging -from datetime import datetime, timedelta, timezone -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 ( - ManagedNE, - PortTrafficDevice, - PortTrafficEvent, - PortTrafficSample, - PortTrafficSeries, - PortTrafficTarget, - UmeInventoryNE, +from .port_traffic_common import _target_out, _utcnow +from .port_traffic_devices import ( + create_device, + delete_device, + discover_ports, + get_device, + list_devices, + list_series, + list_targets, + put_interfaces, + rebind_device, + replace_series_port, + set_device_status, + update_device, ) -from .ne_session_factory import close_netmiko_connection, open_netmiko_connection -from .ne_netmiko import send_show_command -from .port_traffic_commands import commands_for_vendor -from .port_traffic_migrate import default_series_title, unique_series_title -from .port_traffic_parsers import brief_port_to_dict, parse_interface_brief, resolve_util_pct -from .port_traffic_schemas import ( - DiscoverPortItem, - DiscoverPortsRequest, - DiscoverPortsResponse, - PortTrafficCompareMeta, - PortTrafficCompareOut, - PortTrafficDashboardOut, - PortTrafficDeviceCreate, - PortTrafficDeviceOut, - PortTrafficDeviceRebind, - PortTrafficDeviceUpdate, - PortTrafficEventOut, - PortTrafficEventsOut, - PortTrafficIfaceIn, - PortTrafficInterfacesPut, - PortTrafficReplacePortRequest, - PortTrafficSamplePoint, - PortTrafficSamplesOut, - PortTrafficSeriesOut, - PortTrafficTargetOut, +from .port_traffic_samples import ( + append_device_event, + baseline_offset_hours, + compare_series, + compare_targets, + dashboard, + get_samples, + list_device_events, + purge_expired_samples, ) -_log = logging.getLogger("netx.port_traffic.service") - - -def _utcnow() -> datetime: - return datetime.utcnow() - - -def _target_out(row: PortTrafficTarget) -> PortTrafficTargetOut: - did = str(row.device_id or "") - return PortTrafficTargetOut( - id=str(row.id), - device_id=did, - task_id=did, - series_id=str(row.series_id or ""), - 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 _device_out(db: Session, device: PortTrafficDevice) -> PortTrafficDeviceOut: - did = str(device.id) - total = db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == did).count() - active = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.device_id == did, PortTrafficTarget.status == "active") - .count() - ) - return PortTrafficDeviceOut( - id=did, - source=str(device.source or ""), - ne_id=str(device.ne_id or ""), - ne_name=str(device.ne_name or ""), - ne_ip=str(device.ne_ip or ""), - vendor=str(device.vendor or ""), - note=str(device.note or ""), - status=str(device.status or ""), - interval_sec=int(device.interval_sec or 60), - retention_days=int(device.retention_days or 7), - concurrency=int(device.concurrency or 1), - collect_running=bool(device.collect_running), - target_count=int(total), - active_target_count=int(active), - last_collect_started_at=device.last_collect_started_at, - last_collect_ended_at=device.last_collect_ended_at, - last_error=str(device.last_error or ""), - created_at=device.created_at, - updated_at=device.updated_at, - ) - - -def _assert_vendor(vendor: str, label: str) -> None: - cmds = commands_for_vendor(vendor or "", "") - if cmds is None: - raise HTTPException( - status_code=400, - detail=f"vendor_not_supported_for_port_traffic: {vendor or 'unknown'} ({label})", - ) - - -def _assert_ifaces_free( - db: Session, - *, - source: str, - ne_id: str, - ifnames: list[str], - exclude_device_id: str | None = None, -) -> None: - names = [str(x).strip() for x in ifnames if str(x).strip()] - if not names: - return - q = db.query(PortTrafficTarget).filter( - PortTrafficTarget.source == source, - PortTrafficTarget.target_id == ne_id, - PortTrafficTarget.ifname.in_(names), - PortTrafficTarget.status == "active", - ) - if exclude_device_id: - q = q.filter(PortTrafficTarget.device_id != exclude_device_id) - hit = q.first() - if hit: - raise HTTPException( - status_code=409, - detail=f"interface_already_monitored: {hit.ifname} on device {hit.device_id}", - ) - - -def _create_iface( - db: Session, - *, - device: PortTrafficDevice, - iface: PortTrafficIfaceIn, - now: datetime, -) -> PortTrafficTarget: - ifname = iface.ifname.strip() - title = unique_series_title( - db, str(device.id), default_series_title(device.ne_name or "", ifname) - ) - sid = uuid4().hex - db.add( - PortTrafficSeries( - id=sid, - device_id=str(device.id), - title=title, - status="active", - created_at=now, - ) - ) - row = PortTrafficTarget( - id=uuid4().hex, - device_id=str(device.id), - series_id=sid, - source=str(device.source), - target_id=str(device.ne_id), - ne_name=str(device.ne_name or ""), - ne_ip=str(device.ne_ip or ""), - vendor=str(device.vendor or ""), - ifname=ifname, - if_description=iface.if_description or "", - bw_bps=int(iface.bw_bps or 0), - status="active", - created_at=now, - ) - db.add(row) - return row - - -def list_devices(db: Session, *, page: int = 1, page_size: int = 20) -> dict[str, Any]: - q = db.query(PortTrafficDevice).order_by(PortTrafficDevice.ne_name.asc(), PortTrafficDevice.ne_ip.asc()) - total = q.count() - rows = q.offset((page - 1) * page_size).limit(page_size).all() - return { - "total": total, - "page": page, - "page_size": page_size, - "items": [_device_out(db, r).model_dump() for r in rows], - } - - -def get_device(db: Session, device_id: str) -> PortTrafficDeviceOut: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - return _device_out(db, device) - - -def create_device(db: Session, body: PortTrafficDeviceCreate) -> PortTrafficDeviceOut: - source = body.source - ne_id = body.ne_id.strip() - _assert_vendor(body.vendor, body.ne_name or ne_id) - clash = ( - db.query(PortTrafficDevice) - .filter(PortTrafficDevice.source == source, PortTrafficDevice.ne_id == ne_id) - .first() - ) - if clash: - raise HTTPException(status_code=409, detail="device_already_monitored") - ifnames = [i.ifname for i in body.interfaces] - _assert_ifaces_free(db, source=source, ne_id=ne_id, ifnames=ifnames) - - now = _utcnow() - status = "running" if body.start_now and body.interfaces else "draft" - device = PortTrafficDevice( - id=uuid4().hex, - source=source, - ne_id=ne_id, - ne_name=(body.ne_name or "").strip(), - ne_ip=(body.ne_ip or "").strip(), - vendor=(body.vendor or "").strip(), - note=(body.note or "").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(device) - db.flush() - for iface in body.interfaces: - if not iface.ifname.strip(): - continue - _create_iface(db, device=device, iface=iface, now=now) - db.commit() - db.refresh(device) - return _device_out(db, device) - - -def update_device(db: Session, device_id: str, body: PortTrafficDeviceUpdate) -> PortTrafficDeviceOut: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - data = body.model_dump(exclude_unset=True) - for key in ("note", "ne_name", "ne_ip", "vendor"): - if key in data and data[key] is not None: - setattr(device, key, str(data[key]).strip()) - if "interval_sec" in data and data["interval_sec"] is not None: - device.interval_sec = int(data["interval_sec"]) - if "retention_days" in data and data["retention_days"] is not None: - device.retention_days = int(data["retention_days"]) - if "concurrency" in data and data["concurrency"] is not None: - device.concurrency = int(data["concurrency"]) - device.updated_at = _utcnow() - db.commit() - db.refresh(device) - return _device_out(db, device) - - -def rebind_device( - db: Session, - device_id: str, - body: PortTrafficDeviceRebind, -) -> PortTrafficDeviceOut: - """Point a monitor device at an explicitly chosen inventory NE; keeps samples.""" - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - if bool(device.collect_running): - raise HTTPException(status_code=409, detail="collect_running") - - source = str(device.source or "").strip().lower() or "managed" - want_id = str(body.ne_id or "").strip() - if not want_id: - raise HTTPException(status_code=400, detail="ne_id_required") - - new_id = "" - new_name = "" - new_ip = "" - new_vendor = "" - - if source == "managed": - row = db.get(ManagedNE, want_id) - if not row: - raise HTTPException(status_code=404, detail="managed_ne_not_found") - new_id = str(row.id) - new_name = str(row.name or "") - new_ip = str(row.ip_address or "") - new_vendor = str(row.vendor or "") - elif source == "ume": - inv = db.get(UmeInventoryNE, want_id) - if not inv: - raise HTTPException(status_code=404, detail="ume_ne_not_found") - new_id = str(inv.ne_id) - new_ip = str(inv.ip_address or "") - new_name = str(inv.user_label or inv.ne_name or inv.host_name or new_ip or "").strip() - new_vendor = str(inv.vendor or "") - else: - raise HTTPException(status_code=400, detail="invalid_source") - - if new_id == str(device.ne_id or ""): - # Already bound; clear stale errors so collect can resume. - device.last_error = "" - device.updated_at = _utcnow() - for tgt in db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).all(): - if "managed_ne_not_found" in str(tgt.last_error or "") or "ume_ne_not_found" in str( - tgt.last_error or "" - ): - tgt.last_error = "" - db.commit() - db.refresh(device) - return _device_out(db, device) - - clash = ( - db.query(PortTrafficDevice) - .filter( - PortTrafficDevice.source == source, - PortTrafficDevice.ne_id == new_id, - PortTrafficDevice.id != device_id, - ) - .first() - ) - if clash: - raise HTTPException(status_code=409, detail="device_already_monitored") - - device.ne_id = new_id - if new_name: - device.ne_name = new_name - if new_ip: - device.ne_ip = new_ip - if new_vendor: - device.vendor = new_vendor - device.last_error = "" - device.updated_at = _utcnow() - - for tgt in db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).all(): - tgt.target_id = new_id - if new_name: - tgt.ne_name = new_name - if new_ip: - tgt.ne_ip = new_ip - if new_vendor: - tgt.vendor = new_vendor - if "managed_ne_not_found" in str(tgt.last_error or "") or "ume_ne_not_found" in str( - tgt.last_error or "" - ): - tgt.last_error = "" - - db.commit() - db.refresh(device) - _log.info("port_traffic rebind device=%s -> ne_id=%s ip=%s", device_id, new_id, new_ip) - return _device_out(db, device) - - -def delete_device(db: Session, device_id: str) -> dict[str, Any]: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - if bool(device.collect_running): - raise HTTPException(status_code=409, detail="collect_running") - targets = db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).all() - ids = [str(t.id) for t in targets] - series_ids = [str(t.series_id) for t in targets if t.series_id] - if ids: - from .port_traffic_board_service import delete_panels_for_targets - - delete_panels_for_targets(db, ids) - db.query(PortTrafficSample).filter(PortTrafficSample.target_row_id.in_(ids)).delete( - synchronize_session=False - ) - db.query(PortTrafficTarget).filter(PortTrafficTarget.device_id == device_id).delete( - synchronize_session=False - ) - db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id).delete( - synchronize_session=False - ) - if series_ids: - db.query(PortTrafficSeries).filter(PortTrafficSeries.id.in_(series_ids)).delete( - synchronize_session=False - ) - else: - db.query(PortTrafficSeries).filter(PortTrafficSeries.device_id == device_id).delete( - synchronize_session=False - ) - db.delete(device) - db.commit() - return {"ok": True, "id": device_id} - - -def set_device_status(db: Session, device_id: str, status: str) -> PortTrafficDeviceOut: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - if status == "running": - active = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.device_id == device_id, PortTrafficTarget.status == "active") - .count() - ) - if active <= 0: - raise HTTPException(status_code=400, detail="no_active_targets") - device.last_collect_ended_at = None - device.status = status - device.updated_at = _utcnow() - db.commit() - db.refresh(device) - return _device_out(db, device) - - -def put_interfaces(db: Session, device_id: str, body: PortTrafficInterfacesPut) -> list[PortTrafficTargetOut]: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - if bool(device.collect_running): - raise HTTPException(status_code=409, detail="collect_running") - - wanted = {i.ifname.strip(): i for i in body.interfaces if i.ifname.strip()} - _assert_ifaces_free( - db, - source=str(device.source), - ne_id=str(device.ne_id), - ifnames=list(wanted.keys()), - exclude_device_id=device_id, - ) - - now = _utcnow() - existing = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.device_id == device_id) - .all() - ) - by_if = {str(t.ifname): t for t in existing if str(t.status) == "active"} - - for ifname, row in list(by_if.items()): - if ifname not in wanted: - row.status = "retired" - - for ifname, iface in wanted.items(): - if ifname in by_if: - row = by_if[ifname] - row.if_description = iface.if_description or row.if_description - if iface.bw_bps: - row.bw_bps = int(iface.bw_bps) - continue - # Reactivate retired same ifname if present - retired = next( - ( - t - for t in existing - if str(t.ifname) == ifname and str(t.status) in ("retired", "disabled") - ), - None, - ) - if retired: - retired.status = "active" - retired.if_description = iface.if_description or retired.if_description - if iface.bw_bps: - retired.bw_bps = int(iface.bw_bps) - continue - _create_iface(db, device=device, iface=iface, now=now) - - device.updated_at = now - db.commit() - return list_targets(db, device_id) - - -def list_targets(db: Session, device_id: str) -> list[PortTrafficTargetOut]: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - rows = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.device_id == device_id) - .order_by(PortTrafficTarget.ifname) - .all() - ) - return [_target_out(r) for r in rows] - - -def list_series(db: Session, device_id: str) -> list[PortTrafficSeriesOut]: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - rows = ( - db.query(PortTrafficSeries) - .filter(PortTrafficSeries.device_id == device_id) - .order_by(PortTrafficSeries.title.asc()) - .all() - ) - out: list[PortTrafficSeriesOut] = [] - for s in rows: - active = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.series_id == s.id, PortTrafficTarget.status == "active") - .order_by(PortTrafficTarget.created_at.desc()) - .first() - ) - retired = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.series_id == s.id, PortTrafficTarget.status == "retired") - .count() - ) - did = str(s.device_id or "") - out.append( - PortTrafficSeriesOut( - id=str(s.id), - device_id=did, - task_id=did, - title=str(s.title or ""), - status=str(s.status or ""), - active_target=_target_out(active) if active else None, - retired_target_count=int(retired), - created_at=s.created_at, - ) - ) - return out - - -def replace_series_port( - db: Session, - device_id: str, - series_id: str, - body: PortTrafficReplacePortRequest, -) -> PortTrafficSeriesOut: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - if bool(device.collect_running): - raise HTTPException(status_code=409, detail="collect_running") - series = db.get(PortTrafficSeries, series_id) - if not series or str(series.device_id) != device_id: - raise HTTPException(status_code=404, detail="series_not_found") - ifname = body.ifname.strip() - _assert_ifaces_free( - db, - source=str(device.source), - ne_id=str(device.ne_id), - ifnames=[ifname], - exclude_device_id=device_id, - ) - now = _utcnow() - actives = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.series_id == series_id, PortTrafficTarget.status == "active") - .all() - ) - for old in actives: - old.status = "retired" - if body.series_title is not None and str(body.series_title).strip(): - wanted = str(body.series_title).strip()[:256] - if wanted != str(series.title or ""): - clash = ( - db.query(PortTrafficSeries.id) - .filter( - PortTrafficSeries.device_id == device_id, - PortTrafficSeries.title == wanted, - PortTrafficSeries.id != series_id, - ) - .first() - ) - if clash: - raise HTTPException(status_code=400, detail="series_title_exists") - series.title = wanted - row = PortTrafficTarget( - id=uuid4().hex, - device_id=device_id, - series_id=series_id, - source=str(device.source), - target_id=str(device.ne_id), - ne_name=str(device.ne_name or ""), - ne_ip=str(device.ne_ip or ""), - vendor=str(device.vendor or ""), - ifname=ifname, - if_description=body.if_description or "", - bw_bps=int(body.bw_bps or 0), - status="active", - created_at=now, - ) - db.add(row) - device.updated_at = now - db.commit() - for item in list_series(db, device_id): - if item.id == series_id: - return item - raise HTTPException(status_code=500, detail="series_replace_failed") - - -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 = send_show_command(conn, cmds.brief, read_timeout=per_cmd) - finally: - close_netmiko_connection(conn) - - ports = [ - DiscoverPortItem(**brief_port_to_dict(p)) - for p in parse_interface_brief(raw, cmds.vendor_key) - ] - 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 _as_naive_utc(value: datetime | None) -> datetime | None: - if value is None: - return None - if value.tzinfo is None: - return value - return value.astimezone(timezone.utc).replace(tzinfo=None) - - -def _sample_points( - rows: list[PortTrafficSample], - *, - align_offset: timedelta | None = None, -) -> list[PortTrafficSamplePoint]: - points: list[PortTrafficSamplePoint] = [] - for r in rows: - ts_raw = r.ts - ts = ts_raw - if align_offset is not None and ts_raw is not None: - ts = ts_raw + align_offset - in_bps = float(r.in_bps or 0) - out_bps = float(r.out_bps or 0) - bw = int(r.bw_bps or 0) - points.append( - PortTrafficSamplePoint( - ts=ts, - ts_raw=ts_raw if align_offset is not None else None, - in_bps=in_bps, - out_bps=out_bps, - in_util_pct=resolve_util_pct(float(r.in_util_pct or 0), in_bps, bw), - out_util_pct=resolve_util_pct(float(r.out_util_pct or 0), out_bps, bw), - bw_bps=bw, - rate_period_sec=int(r.rate_period_sec or 0), - ) - ) - return points - - -def _query_target_samples( - db: Session, - *, - target_row_id: str, - from_ts: datetime, - to_ts: datetime, -) -> list[PortTrafficSample]: - return ( - 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() - ) - - -def baseline_offset_hours(baseline: str, range_hours: float, offset_hours: float | None) -> float | None: - key = str(baseline or "off").strip().lower() - if key in ("", "off", "none"): - return None - if key == "shift": - return float(range_hours) - if key == "day": - return 24.0 - if key == "week": - return 24.0 * 7 - if key == "custom": - if offset_hours is None or float(offset_hours) <= 0: - raise HTTPException(status_code=400, detail="offset_hours_required") - return float(offset_hours) - raise HTTPException(status_code=400, detail=f"invalid_baseline: {baseline}") - - -def compare_targets( - db: Session, - *, - target_row_id: str, - range_hours: float = 24, - baseline: str = "off", - offset_hours: float | None = None, - baseline_target_id: str | None = None, - ahead_hours: float = 0, - to_ts: datetime | None = None, -) -> PortTrafficCompareOut: - target = db.get(PortTrafficTarget, target_row_id) - if not target: - raise HTTPException(status_code=404, detail="target_not_found") - - mapped: PortTrafficTarget | None = None - mapped_id = str(baseline_target_id or "").strip() - if mapped_id: - if mapped_id == str(target.id): - raise HTTPException(status_code=400, detail="baseline_target_same_as_current") - mapped = db.get(PortTrafficTarget, mapped_id) - if not mapped: - raise HTTPException(status_code=404, detail="baseline_target_not_found") - - now = _utcnow() - # Anchor = "now" (or explicit to). Lookback is from anchor; ahead extends past it - # so period compare can show baseline trend after the current clock time. - anchor = _as_naive_utc(to_ts) or now - range_h = max(0.25, float(range_hours or 24)) - ahead_h = max(0.0, min(24.0, float(ahead_hours or 0))) - from_ts = anchor - timedelta(hours=range_h) - to_end = anchor + timedelta(hours=ahead_h) - to_q = to_end + timedelta(seconds=5) - current = _sample_points( - _query_target_samples(db, target_row_id=str(target.id), from_ts=from_ts, to_ts=to_q) - ) - - off_h = baseline_offset_hours(baseline, range_h, offset_hours) - baseline_points: list[PortTrafficSamplePoint] = [] - base_src = mapped if mapped is not None else target - want_baseline = off_h is not None or mapped is not None - if want_baseline: - if off_h is not None: - delta = timedelta(hours=off_h) - base_rows = _query_target_samples( - db, - target_row_id=str(base_src.id), - from_ts=from_ts - delta, - to_ts=to_q - delta, - ) - baseline_points = _sample_points(base_rows, align_offset=delta) - else: - base_rows = _query_target_samples( - db, target_row_id=str(base_src.id), from_ts=from_ts, to_ts=to_q - ) - baseline_points = _sample_points(base_rows) - - return PortTrafficCompareOut( - meta=PortTrafficCompareMeta( - target_id=str(target.id), - baseline=str(baseline or "off"), - offset_hours=float(off_h or 0), - range_hours=range_h, - ahead_hours=ahead_h, - current_target=_target_out(target), - baseline_target=_target_out(mapped) if mapped is not None else None, - baseline_target_id=str(mapped.id) if mapped is not None else "", - ), - current=current, - baseline=baseline_points, - ) - - -def compare_series(db: Session, **kwargs: Any) -> PortTrafficCompareOut: - target_row_id = kwargs.pop("target_row_id", None) or kwargs.pop("target_id", None) - series_id = kwargs.pop("series_id", None) - if not target_row_id and series_id: - active = ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.series_id == series_id, PortTrafficTarget.status == "active") - .order_by(PortTrafficTarget.created_at.desc()) - .first() - ) - if not active: - raise HTTPException(status_code=404, detail="series_active_target_not_found") - target_row_id = str(active.id) - if not target_row_id: - raise HTTPException(status_code=400, detail="target_id_required") - return compare_targets(db, target_row_id=str(target_row_id), **kwargs) - - -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() - to_ts = _as_naive_utc(to_ts) or now - from_ts = _as_naive_utc(from_ts) or (to_ts - timedelta(hours=1)) - to_ts = to_ts + timedelta(seconds=5) - rows = _query_target_samples(db, target_row_id=target_row_id, from_ts=from_ts, to_ts=to_ts) - return PortTrafficSamplesOut(target=_target_out(target), points=_sample_points(rows)) - - -def dashboard(db: Session) -> PortTrafficDashboardOut: - device_count = db.query(PortTrafficDevice).count() - running = db.query(PortTrafficDevice).filter(PortTrafficDevice.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( - device_count=int(device_count), - running_device_count=int(running), - active_target_count=int(active_targets), - sample_count_24h=int(sample_count), - last_sample_at=last, - task_count=int(device_count), - running_task_count=int(running), - ) - - -def list_device_events( - db: Session, - device_id: str, - *, - limit: int = 100, -) -> PortTrafficEventsOut: - device = db.get(PortTrafficDevice, device_id) - if not device: - raise HTTPException(status_code=404, detail="device_not_found") - lim = max(1, min(500, int(limit or 100))) - q = db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id) - total = q.count() - # Seed once from current last_error snapshot so older failures still show in log UI. - if total == 0: - seeded = False - if str(device.last_error or "").strip(): - append_device_event( - db, - device_id=device_id, - message=str(device.last_error), - level="error", - ) - seeded = True - for t in ( - db.query(PortTrafficTarget) - .filter(PortTrafficTarget.device_id == device_id) - .order_by(PortTrafficTarget.ifname) - .all() - ): - err = str(t.last_error or "").strip() - if not err: - continue - append_device_event( - db, - device_id=device_id, - target_row_id=str(t.id), - ifname=str(t.ifname or ""), - message=err, - level="error", - ) - seeded = True - if seeded: - db.commit() - total = db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id).count() - q = db.query(PortTrafficEvent).filter(PortTrafficEvent.device_id == device_id) - rows = q.order_by(PortTrafficEvent.created_at.desc()).limit(lim).all() - items = [ - PortTrafficEventOut( - id=str(r.id), - device_id=str(r.device_id or ""), - target_row_id=str(r.target_row_id or ""), - ifname=str(r.ifname or ""), - level=str(r.level or "error"), - message=str(r.message or ""), - created_at=r.created_at, - ) - for r in rows - ] - return PortTrafficEventsOut(items=items, total=int(total)) - - -def append_device_event( - db: Session, - *, - device_id: str, - message: str, - level: str = "error", - target_row_id: str = "", - ifname: str = "", -) -> None: - msg = str(message or "").strip() - if not msg or not device_id: - return - db.add( - PortTrafficEvent( - id=uuid4().hex, - device_id=device_id, - target_row_id=str(target_row_id or ""), - ifname=str(ifname or ""), - level=str(level or "error")[:16], - message=msg[:4000], - created_at=_utcnow(), - ) - ) - - -def purge_expired_samples(db: Session) -> int: - devices = db.query(PortTrafficDevice).all() - deleted = 0 - now = _utcnow() - for device in devices: - days = max(1, int(device.retention_days or 7)) - cutoff = now - timedelta(days=days) - target_ids = [ - str(t.id) - for t in db.query(PortTrafficTarget.id) - .filter(PortTrafficTarget.device_id == device.id) - .all() - ] - if target_ids: - n = ( - db.query(PortTrafficSample) - .filter( - PortTrafficSample.target_row_id.in_(target_ids), - PortTrafficSample.ts < cutoff, - ) - .delete(synchronize_session=False) - ) - deleted += int(n or 0) - db.query(PortTrafficEvent).filter( - PortTrafficEvent.device_id == device.id, - PortTrafficEvent.created_at < cutoff, - ).delete(synchronize_session=False) - db.commit() - if deleted: - _log.info("port_traffic retention purged samples=%s", deleted) - return deleted +__all__ = [ + "_target_out", + "_utcnow", + "append_device_event", + "baseline_offset_hours", + "compare_series", + "compare_targets", + "create_device", + "dashboard", + "delete_device", + "discover_ports", + "get_device", + "get_samples", + "list_device_events", + "list_devices", + "list_series", + "list_targets", + "purge_expired_samples", + "put_interfaces", + "rebind_device", + "replace_series_port", + "set_device_status", + "update_device", +] diff --git a/netx_api/timeutil.py b/netx_api/timeutil.py new file mode 100644 index 0000000..5dff643 --- /dev/null +++ b/netx_api/timeutil.py @@ -0,0 +1,9 @@ +"""Timezone helpers. DB columns store naive UTC; keep values naive for comparisons.""" + +from __future__ import annotations + +from datetime import datetime, timezone + + +def utcnow_naive() -> datetime: + return datetime.now(timezone.utc).replace(tzinfo=None) diff --git a/netx_api/topology_common.py b/netx_api/topology_common.py index aedafd6..df8af9f 100644 --- a/netx_api/topology_common.py +++ b/netx_api/topology_common.py @@ -12,6 +12,7 @@ from sqlalchemy import text from sqlalchemy.orm import Session from .models import TopoFabricEdge, TopoViewEdgeStyle +from .timeutil import utcnow_naive ROOT_FOLDER_NAME = "Network" PHYSICAL_VIEW_NAME = "Physical topology" @@ -92,7 +93,7 @@ def _purge_edge_if_due(db: Session, e: TopoFabricEdge) -> bool: def _utcnow() -> datetime: - return datetime.utcnow() + return utcnow_naive() def _norm_host(s: str) -> str: diff --git a/netx_api/ume_support.py b/netx_api/ume_support.py index c9b13af..040506a 100644 --- a/netx_api/ume_support.py +++ b/netx_api/ume_support.py @@ -14,6 +14,7 @@ from sqlalchemy.orm import Session from .config import settings from .db import SessionLocal from .models import UmeAlarmCurrent, UmeInventoryNE, UmeSyncJob +from .timeutil import utcnow_naive from .runtime_task_messages import ( RT_ALARMS_SYNC_IN_PROGRESS_SKIP, RT_KEEPALIVE_FAILED, @@ -240,7 +241,7 @@ def _fail_stale_running_sync_jobs_on_startup() -> None: ) if not rows: return - now_naive = datetime.utcnow() + now_naive = utcnow_naive() for row in rows: row.status = "failed" row.ended_at = now_naive