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
+184
View File
@@ -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))
+65
View File
@@ -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)