Files
leocrm/app/plugins/builtins/kommunikation/websocket_manager.py
T
Agent Zero 7a81a5f072
Check Cross-Plugin Imports / check (push) Has been cancelled
feat(B-WS): WebSocket Helpers + Redis Pub/Sub + Error-Handling
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
2026-08-13 16:43:54 +02:00

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