from fastapi import APIRouter, Depends, Query, WebSocket, WebSocketDisconnect from typing import Optional, List from auth.models import TokenPayload from auth.dependencies import require_permission from mqtt.models import ( MqttCommandRequest, CommandSendResponse, MqttStatusResponse, DeviceMqttStatus, LogListResponse, HeartbeatListResponse, CommandListResponse, HeartbeatEntry, AlertEventEntry, AlertEventListResponse, BootEventListResponse, PingSampleListResponse, DiagnosticsReportListResponse, LatestMetricsResponse, LatestDiagnosticsEntry, LatestPingEntry, DeviceReportListResponse, ) from mqtt.client import mqtt_manager import database as db from datetime import datetime, timezone router = APIRouter(prefix="/api/mqtt", tags=["mqtt"]) @router.get("/status", response_model=MqttStatusResponse) async def get_all_device_status( _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): heartbeats = await db.get_latest_heartbeats() alert_events = await db.get_latest_alert_events() alert_by_serial = {a["device_serial"]: a for a in alert_events} now = datetime.now(timezone.utc) devices = [] for hb in heartbeats: received_str = hb["received_at"] try: received = datetime.fromisoformat(received_str) if received.tzinfo is None: received = received.replace(tzinfo=timezone.utc) seconds_ago = int((now - received).total_seconds()) except (ValueError, TypeError): seconds_ago = 9999 alert_event = alert_by_serial.get(hb["device_serial"]) devices.append(DeviceMqttStatus( device_serial=hb["device_serial"], online=seconds_ago < 90, last_heartbeat=HeartbeatEntry(**hb), seconds_since_heartbeat=seconds_ago, last_alert_event=AlertEventEntry(**alert_event) if alert_event else None, )) return MqttStatusResponse( devices=devices, broker_connected=mqtt_manager.connected, ) @router.get("/latest-metrics", response_model=LatestMetricsResponse) async def get_latest_metrics( _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): # Fleet-wide "last known" CPU temp + ping RTT, one query each — used by # DeviceList to render optional columns without polling any device. # Uptime/firmware/RSSI don't need this: they're already in /mqtt/status. diag_reports = await db.get_latest_diagnostics_reports() ping_samples = await db.get_latest_ping_samples() return LatestMetricsResponse( diagnostics=[LatestDiagnosticsEntry(**d) for d in diag_reports], pings=[LatestPingEntry(**p) for p in ping_samples], ) @router.post("/command/{device_serial}", response_model=CommandSendResponse) async def send_command( device_serial: str, body: MqttCommandRequest, _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): command_id = await db.insert_command( device_serial=device_serial, command_name=body.cmd, command_payload={"cmd": body.cmd, "contents": body.contents}, ) success = mqtt_manager.publish_command( device_serial=device_serial, cmd=body.cmd, contents=body.contents, ) if not success: await db.update_command_response( command_id, "error", {"error": "MQTT broker not connected"}, ) return CommandSendResponse( success=False, command_id=command_id, message="MQTT broker not connected", ) return CommandSendResponse( success=True, command_id=command_id, message=f"Command '{body.cmd}' sent to {device_serial}", ) @router.get("/logs/{device_serial}", response_model=LogListResponse) async def get_device_logs( device_serial: str, level: Optional[str] = Query(None, description="Filter: INFO, WARN, ERROR"), min_level: bool = Query(False, description="If true, level is a floor — also includes higher-severity levels"), search: Optional[str] = Query(None), source: Optional[List[str]] = Query(None, description="Filter by source: log, info. Repeat param to include several."), limit: int = Query(100, ge=1, le=1000), offset: int = Query(0, ge=0), since: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/after this time"), until: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/before this time"), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): logs, total = await db.get_logs( device_serial, level=level, search=search, source=source, min_level=min_level, limit=limit, offset=offset, since=since, until=until, ) return LogListResponse(logs=logs, total=total) @router.get("/heartbeats/{device_serial}", response_model=HeartbeatListResponse) async def get_device_heartbeats( device_serial: str, limit: int = Query(100, ge=1, le=5000), offset: int = Query(0, ge=0), since: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/after this time"), until: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/before this time"), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): heartbeats, total = await db.get_heartbeats( device_serial, limit=limit, offset=offset, since=since, until=until, ) return HeartbeatListResponse(heartbeats=heartbeats, total=total) @router.get("/commands/{device_serial}", response_model=CommandListResponse) async def get_device_commands( device_serial: str, limit: int = Query(100, ge=1, le=1000), offset: int = Query(0, ge=0), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): commands, total = await db.get_commands( device_serial, limit=limit, offset=offset, ) return CommandListResponse(commands=commands, total=total) @router.get("/alert-events/{device_serial}", response_model=AlertEventListResponse) async def get_device_alert_events( device_serial: str, limit: int = Query(100, ge=1, le=1000), offset: int = Query(0, ge=0), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): events, total = await db.get_alert_events( device_serial, limit=limit, offset=offset, ) return AlertEventListResponse(events=events, total=total) @router.get("/boot-events/{device_serial}", response_model=BootEventListResponse) async def get_device_boot_events( device_serial: str, limit: int = Query(100, ge=1, le=2000), offset: int = Query(0, ge=0), since: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/after this time"), until: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/before this time"), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): events, total = await db.get_boot_events( device_serial, limit=limit, offset=offset, since=since, until=until, ) return BootEventListResponse(events=events, total=total) @router.get("/ping-samples/{device_serial}", response_model=PingSampleListResponse) async def get_device_ping_samples( device_serial: str, limit: int = Query(200, ge=1, le=5000), offset: int = Query(0, ge=0), since: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/after this time"), until: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/before this time"), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): samples, total = await db.get_ping_samples( device_serial, limit=limit, offset=offset, since=since, until=until, ) return PingSampleListResponse(samples=samples, total=total) @router.get("/diagnostics-reports/{device_serial}", response_model=DiagnosticsReportListResponse) async def get_device_diagnostics_reports( device_serial: str, limit: int = Query(200, ge=1, le=5000), offset: int = Query(0, ge=0), since: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/after this time"), until: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/before this time"), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): reports, total = await db.get_diagnostics_reports( device_serial, limit=limit, offset=offset, since=since, until=until, ) return DiagnosticsReportListResponse(reports=reports, total=total) @router.get("/reports/{device_serial}", response_model=DeviceReportListResponse) async def get_device_reports( device_serial: str, limit: int = Query(200, ge=1, le=2000), offset: int = Query(0, ge=0), since: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/after this time"), until: Optional[datetime] = Query(None, description="ISO timestamp — only rows at/before this time"), _user: TokenPayload = Depends(require_permission("mqtt", "view")), ): """Critical, unsolicited board-initiated events from control/reports (currently only bell_overload). History/audit only — the tablets are the real-time consumer of this data, not the console.""" reports, total = await db.get_reports( device_serial, limit=limit, offset=offset, since=since, until=until, ) return DeviceReportListResponse(reports=reports, total=total) @router.websocket("/ws") async def mqtt_websocket(websocket: WebSocket): """Live MQTT data stream. Auth via query param: ?token=JWT""" token = websocket.query_params.get("token") if not token: await websocket.close(code=4001, reason="Missing token") return try: from auth.utils import decode_access_token from sqlalchemy import select from database.postgres import AsyncSessionLocal from staff.orm import Staff payload = decode_access_token(token) role = payload.get("role", "") # sysadmin and admin always have MQTT access if role not in ("sysadmin", "admin"): user_sub = payload.get("sub", "") async with AsyncSessionLocal() as session: result = await session.execute( select(Staff).where(Staff.id == user_sub).limit(1) ) staff = result.scalar_one_or_none() if staff is None: await websocket.close(code=4003, reason="User not found") return perms = staff.permissions or {} if not perms.get("mqtt", {}).get("access", False): await websocket.close(code=4003, reason="MQTT access denied") return except Exception: await websocket.close(code=4001, reason="Invalid token") return await websocket.accept() mqtt_manager.add_ws_subscriber(websocket) try: while True: await websocket.receive_text() except WebSocketDisconnect: pass finally: mqtt_manager.remove_ws_subscriber(websocket)