diff --git a/app/core/ws_helpers.py b/app/core/ws_helpers.py new file mode 100644 index 0000000..b4cec27 --- /dev/null +++ b/app/core/ws_helpers.py @@ -0,0 +1,184 @@ +"""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") + + +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)) diff --git a/app/core/ws_pubsub.py b/app/core/ws_pubsub.py new file mode 100644 index 0000000..4036376 --- /dev/null +++ b/app/core/ws_pubsub.py @@ -0,0 +1,65 @@ +"""Redis Pub/Sub helpers for WebSocket multi-worker fanout.""" + +from __future__ import annotations + +import asyncio +import json +import logging +import uuid +from typing import Awaitable, Callable + +from app.core.auth import get_redis + +logger = logging.getLogger(__name__) + + +async def publish_to_channel(channel: str, message: dict) -> None: + """Publish a JSON message to a Redis channel.""" + redis = get_redis() + await redis.publish(channel, json.dumps(message, default=str)) + + +async def subscribe_to_channel( + channel: str, + handler: Callable[[dict], Awaitable[None]], +) -> asyncio.Task: + """Subscribe to a Redis channel and call *handler* for every message. + + Returns the :class:`asyncio.Task` so the caller can cancel it on disconnect. + """ + + async def _subscriber() -> None: + redis = get_redis() + pubsub = redis.pubsub() + await pubsub.subscribe(channel) + try: + async for raw in pubsub.listen(): + if raw["type"] == "message": + try: + msg = json.loads(raw["data"]) + await handler(msg) + except Exception: + logger.exception("Error in pubsub handler for channel %s", channel) + finally: + try: + await pubsub.unsubscribe(channel) + await pubsub.aclose() + except Exception: + logger.debug("PubSub cleanup error for channel %s", channel) + + return asyncio.create_task(_subscriber()) + + +def get_tenant_channel(tenant_id: uuid.UUID, topic: str) -> str: + """Return the Redis channel name for a tenant + topic.""" + return f"ws:{tenant_id}:{topic}" + + +async def broadcast_to_tenants( + tenant_id: uuid.UUID, + topic: str, + message: dict, +) -> None: + """Publish a message to a tenant-specific Redis channel.""" + channel = get_tenant_channel(tenant_id, topic) + await publish_to_channel(channel, message) diff --git a/app/plugins/builtins/ai_ui_control/routes.py b/app/plugins/builtins/ai_ui_control/routes.py index b8f7605..c97e8d0 100644 --- a/app/plugins/builtins/ai_ui_control/routes.py +++ b/app/plugins/builtins/ai_ui_control/routes.py @@ -205,33 +205,28 @@ async def ai_ui_control_ws(websocket: WebSocket): Authentication: via session cookie (same pattern as kommunikation plugin). """ - from app.config import get_settings - from app.core.auth import get_session_data, get_redis, verify_ws_origin + from app.core.db import async_session_maker from app.core.service_container import get_container - settings = get_settings() - if not await verify_ws_origin(websocket): - await websocket.close(code=4003, reason="Origin not allowed") + container = get_container() + if not container.has("ai_ui_control_ws"): + await websocket.close(code=4003, reason="AI UI Control not available") return - session_id = websocket.cookies.get(settings.session_cookie_name) - if not session_id: - await websocket.close(code=4001, reason="Not authenticated") - return + ws_manager = container.get("ai_ui_control_ws") - redis = get_redis() - session_data = await get_session_data(redis, session_id) - if session_data is None: - await websocket.close(code=4001, reason="Session expired") - return + # connect() performs origin check, session auth, tenant check + async with async_session_maker() as db: + auth = await ws_manager.connect(websocket, db) + if auth is None: + return # connection was rejected and closed by ws_helpers - user_id = session_data["user_id"] - tenant_id = session_data["tenant_id"] + user_id = auth["user_id"] + tenant_id = auth["tenant_id"] # Plugin-Gate: check if ai_ui_control plugin is active (global + tenant) from app.core.permission_registry import get_permission_registry from sqlalchemy import text as sa_text - from app.core.db import async_session_maker import uuid as _uuid try: registry = get_permission_registry() @@ -251,14 +246,6 @@ async def ai_ui_control_ws(websocket: WebSocket): await websocket.close(code=4003, reason="Plugin check failed") return - container = get_container() - if not container.has("ai_ui_control_ws"): - await websocket.close(code=4003, reason="AI UI Control not available") - return - - ws_manager = container.get("ai_ui_control_ws") - await ws_manager.connect(websocket, user_id) - try: while True: data = await websocket.receive_text() diff --git a/app/plugins/builtins/ai_ui_control/websocket_manager.py b/app/plugins/builtins/ai_ui_control/websocket_manager.py index efd1284..0f30224 100644 --- a/app/plugins/builtins/ai_ui_control/websocket_manager.py +++ b/app/plugins/builtins/ai_ui_control/websocket_manager.py @@ -16,6 +16,20 @@ from typing import Any from fastapi import WebSocket +from app.core.ws_helpers import ( + authenticate_ws, + check_ws_origin, + check_ws_tenant, + cleanup_ws_connection, + start_heartbeat, +) +from app.core.ws_pubsub import ( + broadcast_to_tenants, + get_tenant_channel, + subscribe_to_channel, +) +from sqlalchemy.ext.asyncio import AsyncSession + logger = logging.getLogger(__name__) @@ -36,23 +50,84 @@ class AIUIControlWSManager: self._delivered_at: dict[str, float] = {} # command_id → user_id (to route feedback) self._command_user: dict[str, str] = {} + # user_id → list of heartbeat tasks + self._heartbeat_tasks: dict[str, list] = {} + # user_id → list of pubsub subscriber tasks + self._pubsub_tasks: dict[str, list] = {} - async def connect(self, websocket: WebSocket, user_id: str) -> None: - """Accept and register a new WebSocket connection.""" + async def connect( + self, + websocket: WebSocket, + db: AsyncSession, + ) -> dict[str, Any] | None: + """Accept, authenticate and register a new WebSocket connection. + + Performs origin check, session authentication and tenant validation. + Returns the auth dict (user_id, tenant_id, …) on success, or ``None`` + if the connection was rejected (already closed). + """ + # Origin / CSRF check + if not await check_ws_origin(websocket): + return None + + # Session authentication + auth = await authenticate_ws(websocket, db) + if auth is None: + return None + + user_id = auth["user_id"] + tenant_id_str = auth["tenant_id"] + tenant_id = uuid.UUID(tenant_id_str) + user_uuid = uuid.UUID(user_id) + + # Tenant membership check + if not await check_ws_tenant(websocket, tenant_id, user_uuid, db): + return None + + # Accept the WebSocket await websocket.accept() + + # Register connection if user_id not in self._connections: self._connections[user_id] = [] self._connections[user_id].append(websocket) + + # Start heartbeat + hb_task = await start_heartbeat(websocket) + self._heartbeat_tasks.setdefault(user_id, []).append(hb_task) + + # Start Redis Pub/Sub subscriber for tenant-wide fanout + channel = get_tenant_channel(tenant_id, "ai_ui_control") + pubsub_task = await subscribe_to_channel(channel, lambda msg: self._on_pubsub_message(user_id, msg)) + self._pubsub_tasks.setdefault(user_id, []).append(pubsub_task) + logger.debug(f"AI UI Control WS connected: user={user_id}, total={len(self._connections[user_id])}") + return auth + + async def _on_pubsub_message(self, user_id: str, msg: dict[str, Any]) -> None: + """Handle a message received via Redis Pub/Sub.""" + # Forward pubsub messages to the user's connections + conns = self._connections.get(user_id, []) + text = json.dumps(msg, default=str) + for ws in conns: + try: + await ws.send_text(text) + except Exception: + logger.warning(f"Failed to send pubsub message to user {user_id}") async def disconnect(self, websocket: WebSocket, user_id: str) -> None: - """Remove a WebSocket connection.""" - conns = self._connections.get(user_id, []) - if websocket in conns: - conns.remove(websocket) - if not conns: - self._connections.pop(user_id, None) - logger.debug(f"AI UI Control WS disconnected: user={user_id}, remaining={len(conns)}") + """Remove a WebSocket connection and clean up resources.""" + # Cancel heartbeat tasks for this user + for task in self._heartbeat_tasks.pop(user_id, []): + task.cancel() + # Cancel pubsub tasks for this user + for task in self._pubsub_tasks.pop(user_id, []): + task.cancel() + + # Clean up connection registry + await cleanup_ws_connection(websocket, user_id, self._connections) + + logger.debug(f"AI UI Control WS disconnected: user={user_id}") async def send_command(self, user_id: str, command: dict[str, Any]) -> str | None: """Send a UI command to all frontend connections of a user. @@ -120,6 +195,24 @@ class AIUIControlWSManager: """Get list of currently connected user IDs.""" return list(self._connections.keys()) + async def broadcast(self, message: dict[str, Any], tenant_id: uuid.UUID | None = None) -> None: + """Broadcast a message to all connected users. + + If *tenant_id* is provided, publishes via Redis Pub/Sub for multi-worker fanout. + Otherwise, sends directly to all locally connected users. + """ + if tenant_id is not None: + await broadcast_to_tenants(tenant_id, "ai_ui_control", message) + else: + for user_id in list(self._connections.keys()): + conns = self._connections.get(user_id, []) + text = json.dumps(message, default=str) + for ws in conns: + try: + await ws.send_text(text) + except Exception: + logger.warning(f"Failed to broadcast to user {user_id}") + def cleanup_stale(self, timeout_seconds: int = 60) -> None: """Remove stale command tracking entries older than timeout.""" now = time.time() diff --git a/app/plugins/builtins/kommunikation/routes.py b/app/plugins/builtins/kommunikation/routes.py index c3c1a3a..0c1b9d5 100644 --- a/app/plugins/builtins/kommunikation/routes.py +++ b/app/plugins/builtins/kommunikation/routes.py @@ -472,33 +472,29 @@ async def websocket_endpoint( Authenticates via session cookie. On connect, subscribes user to all their conversations. """ - # Verify Origin header against allowed CORS origins - from app.config import get_settings - from app.core.auth import get_session_data, get_redis, verify_ws_origin + from app.core.db import async_session_maker + from app.core.service_container import get_container - settings = get_settings() - if not await verify_ws_origin(websocket): - await websocket.close(code=4003, reason="Origin not allowed") + # Get WebSocket manager from service container + container = get_container() + if not container.has("comm_websocket"): + await websocket.close(code=4003, reason="Messaging not available") return - session_id = websocket.cookies.get(settings.session_cookie_name) - if not session_id: - await websocket.close(code=4001, reason="Not authenticated") - return + ws_manager = container.get("comm_websocket") - redis = get_redis() - session_data = await get_session_data(redis, session_id) - if session_data is None: - await websocket.close(code=4001, reason="Session expired") - return + # connect() performs origin check, session auth, tenant check + async with async_session_maker() as db: + auth = await ws_manager.connect(websocket, db) + if auth is None: + return # connection was rejected and closed by ws_helpers - user_id = session_data["user_id"] - tenant_id = session_data["tenant_id"] + user_id = auth["user_id"] + tenant_id = auth["tenant_id"] # Plugin-Gate: check if kommunikation plugin is active (global + tenant) from app.core.permission_registry import get_permission_registry from sqlalchemy import text as sa_text - from app.core.db import async_session_maker import uuid as _uuid try: registry = get_permission_registry() @@ -518,16 +514,6 @@ async def websocket_endpoint( await websocket.close(code=4003, reason="Plugin check failed") return - # Get WebSocket manager from service container - from app.core.service_container import get_container - container = get_container() - if not container.has("comm_websocket"): - await websocket.close(code=4003, reason="Messaging not available") - return - - ws_manager = container.get("comm_websocket") - await ws_manager.connect(websocket, user_id) - try: while True: data = await websocket.receive_text() diff --git a/app/plugins/builtins/kommunikation/websocket_manager.py b/app/plugins/builtins/kommunikation/websocket_manager.py index 4f25f2d..0c652bc 100644 --- a/app/plugins/builtins/kommunikation/websocket_manager.py +++ b/app/plugins/builtins/kommunikation/websocket_manager.py @@ -4,10 +4,25 @@ from __future__ import annotations import json import logging +import uuid from typing import Any from fastapi import WebSocket +from app.core.ws_helpers import ( + authenticate_ws, + check_ws_origin, + check_ws_tenant, + cleanup_ws_connection, + start_heartbeat, +) +from app.core.ws_pubsub import ( + broadcast_to_tenants, + get_tenant_channel, + subscribe_to_channel, +) +from sqlalchemy.ext.asyncio import AsyncSession + logger = logging.getLogger(__name__) @@ -19,26 +34,82 @@ class WebSocketManager: self._connections: dict[str, list[WebSocket]] = {} # conversation_id (str) → set of user_ids subscribed self._subscriptions: dict[str, set[str]] = {} + # user_id → list of heartbeat tasks + self._heartbeat_tasks: dict[str, list] = {} + # user_id → list of pubsub subscriber tasks + self._pubsub_tasks: dict[str, list] = {} - async def connect(self, websocket: WebSocket, user_id: str) -> None: - """Accept and register a new WebSocket connection.""" + async def connect( + self, + websocket: WebSocket, + db: AsyncSession, + ) -> dict[str, Any] | None: + """Accept, authenticate and register a new WebSocket connection. + + Performs origin check, session authentication and tenant validation. + Returns the auth dict (user_id, tenant_id, …) on success, or ``None`` + if the connection was rejected (already closed). + """ + # Origin / CSRF check + if not await check_ws_origin(websocket): + return None + + # Session authentication + auth = await authenticate_ws(websocket, db) + if auth is None: + return None + + user_id = auth["user_id"] + tenant_id_str = auth["tenant_id"] + tenant_id = uuid.UUID(tenant_id_str) + user_uuid = uuid.UUID(user_id) + + # Tenant membership check + if not await check_ws_tenant(websocket, tenant_id, user_uuid, db): + return None + + # Accept the WebSocket await websocket.accept() + + # Register connection if user_id not in self._connections: self._connections[user_id] = [] self._connections[user_id].append(websocket) + + # Start heartbeat + hb_task = await start_heartbeat(websocket) + self._heartbeat_tasks.setdefault(user_id, []).append(hb_task) + + # Start Redis Pub/Sub subscriber for tenant-wide fanout + channel = get_tenant_channel(tenant_id, "kommunikation") + pubsub_task = await subscribe_to_channel(channel, lambda msg: self._on_pubsub_message(user_id, msg)) + self._pubsub_tasks.setdefault(user_id, []).append(pubsub_task) + logger.debug(f"WebSocket connected: user={user_id}, total={len(self._connections[user_id])}") + return auth + + async def _on_pubsub_message(self, user_id: str, msg: dict[str, Any]) -> None: + """Handle a message received via Redis Pub/Sub.""" + await self.send_to_user(user_id, msg) async def disconnect(self, websocket: WebSocket, user_id: str) -> None: - """Remove a WebSocket connection.""" - conns = self._connections.get(user_id, []) - if websocket in conns: - conns.remove(websocket) - if not conns: - self._connections.pop(user_id, None) - # Remove from all subscriptions + """Remove a WebSocket connection and clean up resources.""" + # Cancel heartbeat tasks for this user + for task in self._heartbeat_tasks.pop(user_id, []): + task.cancel() + # Cancel pubsub tasks for this user + for task in self._pubsub_tasks.pop(user_id, []): + task.cancel() + + # Clean up connection registry + await cleanup_ws_connection(websocket, user_id, self._connections) + + # Remove from subscriptions if no more connections + if user_id not in self._connections: for conv_id, users in self._subscriptions.items(): users.discard(user_id) - logger.debug(f"WebSocket disconnected: user={user_id}, remaining={len(conns)}") + + logger.debug(f"WebSocket disconnected: user={user_id}") def subscribe(self, conversation_id: str, user_id: str) -> None: """Subscribe a user to a conversation's updates.""" @@ -75,10 +146,17 @@ class WebSocketManager: continue await self.send_to_user(user_id, message) - async def broadcast(self, message: dict[str, Any]) -> None: - """Broadcast a message to all connected users.""" - for user_id in list(self._connections.keys()): - await self.send_to_user(user_id, message) + async def broadcast(self, message: dict[str, Any], tenant_id: uuid.UUID | None = None) -> None: + """Broadcast a message to all connected users. + + If *tenant_id* is provided, publishes via Redis Pub/Sub for multi-worker fanout. + Otherwise, sends directly to all locally connected users. + """ + if tenant_id is not None: + await broadcast_to_tenants(tenant_id, "kommunikation", message) + else: + for user_id in list(self._connections.keys()): + await self.send_to_user(user_id, message) def get_online_users(self) -> list[str]: """Get list of currently connected user IDs.""" diff --git a/tests/test_ws_helpers.py b/tests/test_ws_helpers.py new file mode 100644 index 0000000..02cee39 --- /dev/null +++ b/tests/test_ws_helpers.py @@ -0,0 +1,500 @@ +"""Tests for WebSocket helpers (ws_helpers) and Redis Pub/Sub (ws_pubsub). + +Covers B-WS-TEST + B-ERR-WS-TEST requirements. +""" + +from __future__ import annotations + +import asyncio +import json +import uuid +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +import pytest_asyncio + +from app.core.ws_helpers import ( + authenticate_ws, + check_ws_origin, + cleanup_ws_connection, + handle_ws_message, + send_ws_error, + start_heartbeat, +) +from app.core.ws_pubsub import ( + broadcast_to_tenants, + get_tenant_channel, + publish_to_channel, + subscribe_to_channel, +) + + +# ─── Mock WebSocket ─────────────────────────────────────────────────────────── + + +class MockWebSocket: + """Minimal mock WebSocket for testing WS helpers.""" + + def __init__( + self, + cookies: dict[str, str] | None = None, + headers: dict[str, str] | None = None, + query_params: dict[str, str] | None = None, + ) -> None: + self.cookies = cookies or {} + self.headers = headers or {} + self.query_params = query_params or {} + self._closed = False + self._close_code: int | None = None + self._close_reason: str | None = None + self._sent: list[str] = [] + self._accepted = False + + async def accept(self) -> None: + self._accepted = True + + async def close(self, code: int = 1000, reason: str = "") -> None: + self._closed = True + self._close_code = code + self._close_reason = reason + + async def send_text(self, text: str) -> None: + if self._closed: + raise RuntimeError("WebSocket is closed") + self._sent.append(text) + + @property + def sent_messages(self) -> list[dict]: + return [json.loads(t) for t in self._sent] + + +# ─── authenticate_ws ────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +class TestAuthenticateWs: + """Tests for authenticate_ws.""" + + async def test_authenticate_ws_valid_session(self, db_session, redis_client): + """authenticate_ws returns user info when session is valid.""" + from app.core.auth import create_session, hash_password + from app.models.tenant import Tenant + from app.models.user import User, UserTenant + from app.models.role import Role + + # Create tenant, user, role, membership + tenant = Tenant(name="Test Tenant", slug="test-tenant") + db_session.add(tenant) + await db_session.flush() + + user = User( + email="wsauth@test.com", + name="WS Auth Test", + password_hash=hash_password("TestPass123!"), + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + + role = Role( + tenant_id=tenant.id, + name="admin", + permissions={"*": {"*": True}}, + denied_permissions=[], + field_permissions={}, + ) + db_session.add(role) + await db_session.flush() + + ut = UserTenant( + user_id=user.id, + tenant_id=tenant.id, + is_default=True, + role="admin", + role_id=role.id, + ) + db_session.add(ut) + await db_session.flush() + await db_session.commit() + + # Create session + session_id, csrf_token = await create_session( + db_session, redis_client, user, tenant.id, role="admin" + ) + + # Create mock WebSocket with session cookie + ws = MockWebSocket(cookies={"leocrm_session": session_id}) + + auth = await authenticate_ws(ws, db_session) + assert auth is not None + assert auth["user_id"] == str(user.id) + assert auth["tenant_id"] == str(tenant.id) + assert auth["session_id"] == session_id + assert auth["role"] == "admin" + assert ws._closed is False + + async def test_authenticate_ws_missing_cookie(self, db_session): + """authenticate_ws closes WS and returns None when no session cookie.""" + ws = MockWebSocket(cookies={}) + + auth = await authenticate_ws(ws, db_session) + assert auth is None + assert ws._closed is True + assert ws._close_code == 4401 + + async def test_authenticate_ws_invalid_session(self, db_session): + """authenticate_ws closes WS and returns None when session is invalid.""" + ws = MockWebSocket(cookies={"leocrm_session": "nonexistent-session-id"}) + + auth = await authenticate_ws(ws, db_session) + assert auth is None + assert ws._closed is True + assert ws._close_code == 4401 + + +# ─── check_ws_origin ────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +class TestCheckWsOrigin: + """Tests for check_ws_origin.""" + + async def test_check_ws_origin_valid(self): + """check_ws_origin returns True for valid origin + CSRF.""" + from app.core.auth import create_session, hash_password + from app.models.tenant import Tenant + from app.models.user import User, UserTenant + from app.models.role import Role + from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + + # We need a full setup with session in Redis for CSRF validation + # Use patch to mock verify_ws_origin returning True + with patch("app.core.ws_helpers.verify_ws_origin", return_value=True): + ws = MockWebSocket() + result = await check_ws_origin(ws) + assert result is True + assert ws._closed is False + + async def test_check_ws_origin_invalid(self): + """check_ws_origin closes WS and returns False for invalid origin.""" + with patch("app.core.ws_helpers.verify_ws_origin", return_value=False): + ws = MockWebSocket() + result = await check_ws_origin(ws) + assert result is False + assert ws._closed is True + assert ws._close_code == 4403 + + +# ─── send_ws_error ──────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +class TestSendWsError: + """Tests for send_ws_error.""" + + async def test_send_ws_error_basic(self): + """send_ws_error sends structured error message.""" + ws = MockWebSocket() + await send_ws_error(ws, "test_error", "Something went wrong") + + assert len(ws.sent_messages) == 1 + msg = ws.sent_messages[0] + assert msg["type"] == "error" + assert msg["code"] == "test_error" + assert msg["detail"] == "Something went wrong" + assert "trace_id" not in msg + + async def test_send_ws_error_with_trace_id(self): + """send_ws_error includes trace_id when provided.""" + ws = MockWebSocket() + await send_ws_error(ws, "test_error", "Failed", trace_id="trace-123") + + msg = ws.sent_messages[0] + assert msg["trace_id"] == "trace-123" + + async def test_send_ws_error_on_closed_ws(self): + """send_ws_error does not raise when WS is closed.""" + ws = MockWebSocket() + ws._closed = True + # Should not raise + await send_ws_error(ws, "closed_error", "WS already closed") + + +# ─── handle_ws_message ───────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +class TestHandleWsMessage: + """Tests for handle_ws_message.""" + + async def test_handle_ws_message_dispatch(self): + """handle_ws_message dispatches to correct handler.""" + ws = MockWebSocket() + called = False + + async def ping_handler(websocket, msg): + nonlocal called + called = True + assert msg["type"] == "ping" + + handlers = {"ping": ping_handler} + message = json.dumps({"type": "ping"}) + await handle_ws_message(ws, message, handlers) + + assert called is True + + async def test_handle_ws_message_unknown_type(self): + """handle_ws_message sends error for unknown type.""" + ws = MockWebSocket() + handlers = {"ping": AsyncMock()} + message = json.dumps({"type": "unknown_type"}) + await handle_ws_message(ws, message, handlers) + + assert len(ws.sent_messages) == 1 + msg = ws.sent_messages[0] + assert msg["type"] == "error" + assert msg["code"] == "unknown_type" + + async def test_handle_ws_message_handler_exception(self): + """handle_ws_message sends error when handler raises exception.""" + ws = MockWebSocket() + + async def bad_handler(websocket, msg): + raise ValueError("Handler crashed") + + handlers = {"ping": bad_handler} + message = json.dumps({"type": "ping"}) + await handle_ws_message(ws, message, handlers) + + assert len(ws.sent_messages) == 1 + msg = ws.sent_messages[0] + assert msg["type"] == "error" + assert msg["code"] == "handler_error" + assert "Handler crashed" in msg["detail"] + + async def test_handle_ws_message_invalid_json(self): + """handle_ws_message sends error for invalid JSON.""" + ws = MockWebSocket() + handlers = {} + await handle_ws_message(ws, "not json at all", handlers) + + assert len(ws.sent_messages) == 1 + msg = ws.sent_messages[0] + assert msg["code"] == "invalid_json" + + async def test_handle_ws_message_missing_type(self): + """handle_ws_message sends error when type field is missing.""" + ws = MockWebSocket() + handlers = {} + message = json.dumps({"data": "no type here"}) + await handle_ws_message(ws, message, handlers) + + assert len(ws.sent_messages) == 1 + msg = ws.sent_messages[0] + assert msg["code"] == "missing_type" + + +# ─── cleanup_ws_connection ───────────────────────────────────────────────────── + + +@pytest.mark.asyncio +class TestCleanupWsConnection: + """Tests for cleanup_ws_connection.""" + + async def test_cleanup_removes_from_registry(self): + """cleanup_ws_connection removes WebSocket from registry.""" + ws1 = MockWebSocket() + ws2 = MockWebSocket() + registry: dict[str, list] = {"user1": [ws1, ws2]} + + await cleanup_ws_connection(ws1, "user1", registry) + + assert "user1" in registry + assert ws1 not in registry["user1"] + assert ws2 in registry["user1"] + assert len(registry["user1"]) == 1 + + async def test_cleanup_removes_empty_user(self): + """cleanup_ws_connection removes user entry when no connections left.""" + ws = MockWebSocket() + registry: dict[str, list] = {"user1": [ws]} + + await cleanup_ws_connection(ws, "user1", registry) + + assert "user1" not in registry + + async def test_cleanup_unknown_user(self): + """cleanup_ws_connection handles unknown user gracefully.""" + ws = MockWebSocket() + registry: dict[str, list] = {} + + await cleanup_ws_connection(ws, "unknown_user", registry) + assert "unknown_user" not in registry + + +# ─── start_heartbeat ─────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +class TestStartHeartbeat: + """Tests for start_heartbeat.""" + + async def test_heartbeat_sends_ping(self): + """start_heartbeat sends ping messages at interval.""" + ws = MockWebSocket() + task = await start_heartbeat(ws, interval=0) + + # Wait a tiny bit for the first ping + await asyncio.sleep(0.05) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert len(ws._sent) >= 1 + msg = ws.sent_messages[0] + assert msg["type"] == "ping" + + async def test_heartbeat_cancellable(self): + """start_heartbeat task can be cancelled.""" + ws = MockWebSocket() + task = await start_heartbeat(ws, interval=10) + + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert task.cancelled() or task.done() + + +# ─── get_tenant_channel ───────────────────────────────────────────────────────── + + +class TestGetTenantChannel: + """Tests for get_tenant_channel.""" + + def test_channel_naming(self): + """get_tenant_channel returns ws:{tenant_id}:{topic}.""" + tenant_id = uuid.uuid4() + channel = get_tenant_channel(tenant_id, "kommunikation") + assert channel == f"ws:{tenant_id}:kommunikation" + + def test_channel_naming_different_topic(self): + """get_tenant_channel works with different topics.""" + tenant_id = uuid.uuid4() + channel = get_tenant_channel(tenant_id, "ai_ui_control") + assert channel == f"ws:{tenant_id}:ai_ui_control" + + +# ─── publish_to_channel + subscribe_to_channel ────────────────────────────────── + + +@pytest.mark.asyncio +class TestPubSub: + """Tests for Redis Pub/Sub helpers.""" + + async def test_publish_subscribe_roundtrip(self, redis_client): + """publish_to_channel + subscribe_to_channel roundtrip.""" + received: list[dict] = [] + channel = f"test-roundtrip-{uuid.uuid4()}" + event = asyncio.Event() + + async def handler(msg: dict) -> None: + received.append(msg) + event.set() + + # Subscribe + task = await subscribe_to_channel(channel, handler) + + # Give subscriber a moment to connect + await asyncio.sleep(0.1) + + # Publish + test_msg = {"type": "test", "data": "hello"} + await publish_to_channel(channel, test_msg) + + # Wait for message + await asyncio.wait_for(event.wait(), timeout=2.0) + + assert len(received) == 1 + assert received[0]["type"] == "test" + assert received[0]["data"] == "hello" + + # Cleanup + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + async def test_broadcast_to_tenants(self, redis_client): + """broadcast_to_tenants publishes to the correct tenant channel.""" + tenant_id = uuid.uuid4() + received: list[dict] = [] + event = asyncio.Event() + + channel = get_tenant_channel(tenant_id, "test_topic") + + async def handler(msg: dict) -> None: + received.append(msg) + event.set() + + task = await subscribe_to_channel(channel, handler) + await asyncio.sleep(0.1) + + test_msg = {"type": "broadcast", "content": "tenant message"} + await broadcast_to_tenants(tenant_id, "test_topic", test_msg) + + await asyncio.wait_for(event.wait(), timeout=2.0) + + assert len(received) == 1 + assert received[0]["type"] == "broadcast" + assert received[0]["content"] == "tenant message" + + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + async def test_publish_to_channel_no_subscribers(self, redis_client): + """publish_to_channel works even with no subscribers.""" + channel = f"test-no-sub-{uuid.uuid4()}" + await publish_to_channel(channel, {"type": "noop"}) + # Should not raise + + async def test_subscribe_multiple_messages(self, redis_client): + """subscribe_to_channel handler receives multiple messages.""" + received: list[dict] = [] + channel = f"test-multi-{uuid.uuid4()}" + count_event = asyncio.Event() + msg_count = 0 + + async def handler(msg: dict) -> None: + nonlocal msg_count + received.append(msg) + msg_count += 1 + if msg_count >= 3: + count_event.set() + + task = await subscribe_to_channel(channel, handler) + await asyncio.sleep(0.1) + + for i in range(3): + await publish_to_channel(channel, {"type": "msg", "index": i}) + + await asyncio.wait_for(count_event.wait(), timeout=3.0) + + assert len(received) == 3 + assert received[0]["index"] == 0 + assert received[2]["index"] == 2 + + task.cancel() + try: + await task + except asyncio.CancelledError: + pass