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:
@@ -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))
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user