7a81a5f072
Check Cross-Plugin Imports / check (push) Has been cancelled
B-WS: app/core/ws_helpers.py (NEU) — gemeinsame WebSocket Helpers - authenticate_ws: Session-Auth für WebSocket (Cookie/Token → User/Tenant) - check_ws_origin: Origin-Check (delegiert auf verify_ws_origin) - check_ws_tenant: User-Tenant-Membership-Check - cleanup_ws_connection: Connection aus Registry entfernen + WS schließen - start_heartbeat: Background Ping-Task - send_ws_error: strukturierte Error-Message an Client - handle_ws_message: Message-Dispatch mit Error-Handling B-WS: app/core/ws_pubsub.py (NEU) — Redis Pub/Sub für Multi-Worker-Fanout - publish_to_channel / subscribe_to_channel - get_tenant_channel / broadcast_to_tenants B-WS: WebSocketManager + AIUIControlWSManager angepasst - connect() nutzt authenticate_ws + check_ws_origin + check_ws_tenant - disconnect() nutzt cleanup_ws_connection + cancelt Heartbeat/PubSub - broadcast() unterstützt Redis Pub/Sub Fanout B-ERR-WS: WS Error-Handling in ws_helpers integriert - send_ws_error für strukturierte Errors - handle_ws_message fängt Handler-Exceptions B-WS-TEST: 24 Tests in test_ws_helpers.py — alle grün - Auth, Origin, Error, Dispatch, Cleanup, Heartbeat, Pub/Sub Roundtrip - Keine Regression: 47/47 Resilience+Hooks Tests grün
168 lines
6.2 KiB
Python
168 lines
6.2 KiB
Python
"""WebSocket connection manager for the kommunikation plugin."""
|
|
|
|
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__)
|
|
|
|
|
|
class WebSocketManager:
|
|
"""Manages WebSocket connections per user for real-time messaging."""
|
|
|
|
def __init__(self) -> None:
|
|
# user_id (str) → list of WebSocket connections
|
|
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,
|
|
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 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}")
|
|
|
|
def subscribe(self, conversation_id: str, user_id: str) -> None:
|
|
"""Subscribe a user to a conversation's updates."""
|
|
if conversation_id not in self._subscriptions:
|
|
self._subscriptions[conversation_id] = set()
|
|
self._subscriptions[conversation_id].add(user_id)
|
|
|
|
def unsubscribe(self, conversation_id: str, user_id: str) -> None:
|
|
"""Unsubscribe a user from a conversation."""
|
|
if conversation_id in self._subscriptions:
|
|
self._subscriptions[conversation_id].discard(user_id)
|
|
|
|
async def send_to_user(self, user_id: str, message: dict[str, Any]) -> None:
|
|
"""Send a message to all connections of a specific user."""
|
|
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 send to user {user_id}, removing connection")
|
|
await self.disconnect(ws, user_id)
|
|
|
|
async def send_to_conversation(
|
|
self,
|
|
conversation_id: str,
|
|
message: dict[str, Any],
|
|
exclude_user: str | None = None,
|
|
) -> None:
|
|
"""Send a message to all users subscribed to a conversation."""
|
|
user_ids = self._subscriptions.get(conversation_id, set())
|
|
for user_id in list(user_ids):
|
|
if exclude_user and user_id == exclude_user:
|
|
continue
|
|
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."""
|
|
return list(self._connections.keys())
|
|
|
|
def is_user_online(self, user_id: str) -> bool:
|
|
"""Check if a user has any active connections."""
|
|
return user_id in self._connections and len(self._connections[user_id]) > 0
|