"""Shared WebSocket helpers: auth, origin check, tenant check, cleanup, heartbeat, error handling, message dispatch.""" from __future__ import annotations import asyncio import json import logging import uuid from typing import Any, Callable from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from starlette.websockets import WebSocket from app.core.auth import get_redis, get_session_data, verify_ws_origin from app.config import get_settings logger = logging.getLogger(__name__) async def authenticate_ws(websocket: WebSocket, db: AsyncSession) -> dict[str, Any] | None: """Authenticate a WebSocket connection via session cookie. Extracts the session cookie, validates the session in Redis, and returns user/tenant info. On failure, closes the WebSocket with code 4401 and returns ``None``. """ settings = get_settings() session_id = websocket.cookies.get(settings.session_cookie_name) if not session_id: await websocket.close(code=4401, reason="Unauthorized") return None redis = get_redis() session_data = await get_session_data(redis, session_id) if session_data is None: await websocket.close(code=4401, reason="Unauthorized") return None if not session_data.get("is_active", False): await websocket.close(code=4401, reason="Unauthorized") return None return { "user_id": session_data["user_id"], "tenant_id": session_data["tenant_id"], "session_id": session_id, "role": session_data.get("role"), "email": session_data.get("email"), "name": session_data.get("name"), "is_system_admin": session_data.get("is_system_admin", False), } async def check_ws_origin(websocket: WebSocket) -> bool: """Verify the WebSocket origin and CSRF token. Delegates to :func:`verify_ws_origin`. On failure, closes the WebSocket with code 4403 and returns ``False``. """ result = await verify_ws_origin(websocket) if not result: await websocket.close(code=4403, reason="Forbidden origin") return False return True async def check_ws_tenant( websocket: WebSocket, tenant_id: uuid.UUID, user_id: uuid.UUID, db: AsyncSession, ) -> bool: """Verify that *user_id* belongs to *tenant_id*. On failure, closes the WebSocket with code 4403 and returns ``False``. """ from app.models.user import UserTenant result = await db.execute( select(UserTenant).where( UserTenant.user_id == user_id, UserTenant.tenant_id == tenant_id, ) ) if result.scalar_one_or_none() is None: await websocket.close(code=4403, reason="Forbidden tenant") return False return True async def cleanup_ws_connection( websocket: WebSocket, user_id: str, connection_registry: dict[str, list[WebSocket]], ) -> None: """Remove a WebSocket from the connection registry and close it cleanly.""" conns = connection_registry.get(user_id, []) if websocket in conns: conns.remove(websocket) if not conns: connection_registry.pop(user_id, None) try: await websocket.close() except Exception: logger.debug("WebSocket already closed during cleanup for user %s", user_id) async def start_heartbeat(websocket: WebSocket, interval: int = 30) -> asyncio.Task: """Start a background heartbeat task that sends periodic pings. Returns the :class:`asyncio.Task` so the caller can cancel it on disconnect. """ async def _heartbeat() -> None: while True: try: await asyncio.sleep(interval) await websocket.send_text(json.dumps({"type": "ping"})) except asyncio.CancelledError: break except Exception: logger.debug("Heartbeat stopped — WebSocket likely closed") break return asyncio.create_task(_heartbeat()) async def send_ws_error( websocket: WebSocket, code: str, detail: str, trace_id: str | None = None, ) -> None: """Send a structured error message to the WebSocket client.""" payload: dict[str, Any] = { "type": "error", "code": code, "detail": detail, } if trace_id is not None: payload["trace_id"] = trace_id try: await websocket.send_text(json.dumps(payload, default=str)) except Exception: logger.debug("Failed to send WS error to client") # ── Global connection registry for drain_all_connections ──────────────────── # Plugin WS endpoints register their connection registries here so that # drain_all_connections() can close them all during graceful shutdown. _global_ws_registries: list[dict[str, list[WebSocket]]] = [] def register_ws_registry(registry: dict[str, list[WebSocket]]) -> None: """Register a WebSocket connection registry for graceful shutdown.""" if registry not in _global_ws_registries: _global_ws_registries.append(registry) async def drain_all_connections(grace_period_seconds: float = 5.0) -> None: """Notify all connected WS clients about reconnect and close connections. Sends a ``reconnect`` hint message to every connected client, waits for ``grace_period_seconds``, then forcefully closes all sockets. Called during application graceful shutdown. """ total_connections = 0 for registry in _global_ws_registries: for user_id, conns in list(registry.items()): for ws in list(conns): try: await ws.send_text(json.dumps({ "type": "reconnect", "reason": "server_shutdown", "message": "Server is shutting down. Please reconnect shortly.", })) total_connections += 1 except Exception: logger.debug("Failed to send reconnect hint to WS client") logger.info(f"WS drain: notified {total_connections} connections, waiting {grace_period_seconds}s") if grace_period_seconds > 0: await asyncio.sleep(grace_period_seconds) # Close all connections for registry in _global_ws_registries: for user_id, conns in list(registry.items()): for ws in list(conns): try: await ws.close(code=1001, reason="Server shutting down") except Exception: logger.debug("WS already closed during drain") registry.clear() _global_ws_registries.clear() logger.info("WS drain: all connections closed") async def handle_ws_message( websocket: WebSocket, message: str, handlers: dict[str, Callable[[WebSocket, dict[str, Any]], Any]], ) -> None: """Dispatch a WebSocket text message to the appropriate handler. *handlers* maps message ``type`` strings to async callables that accept ``(websocket, msg)``. Unknown types and handler exceptions are reported back to the client via :func:`send_ws_error`. """ try: msg = json.loads(message) except (json.JSONDecodeError, TypeError): await send_ws_error(websocket, "invalid_json", "Message is not valid JSON") return if not isinstance(msg, dict): await send_ws_error(websocket, "invalid_message", "Message must be a JSON object") return msg_type = msg.get("type") if not msg_type: await send_ws_error(websocket, "missing_type", "Message missing 'type' field") return handler = handlers.get(msg_type) if handler is None: await send_ws_error(websocket, "unknown_type", f"Unknown message type: {msg_type}") return try: await handler(websocket, msg) except Exception as exc: logger.exception("Handler error for message type '%s'", msg_type) await send_ws_error(websocket, "handler_error", str(exc))