feat(B-WS): WebSocket Helpers + Redis Pub/Sub + Error-Handling
Check Cross-Plugin Imports / check (push) Has been cancelled
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:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user