feat(B-WS): WebSocket Helpers + Redis Pub/Sub + Error-Handling
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
This commit is contained in:
Agent Zero
2026-08-13 16:43:54 +02:00
parent a3a26d1f66
commit 7a81a5f072
7 changed files with 969 additions and 76 deletions
+14 -28
View File
@@ -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()
@@ -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."""