501 lines
17 KiB
Python
501 lines
17 KiB
Python
|
|
"""Tests for WebSocket helpers (ws_helpers) and Redis Pub/Sub (ws_pubsub).
|
||
|
|
|
||
|
|
Covers B-WS-TEST + B-ERR-WS-TEST requirements.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import asyncio
|
||
|
|
import json
|
||
|
|
import uuid
|
||
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
import pytest_asyncio
|
||
|
|
|
||
|
|
from app.core.ws_helpers import (
|
||
|
|
authenticate_ws,
|
||
|
|
check_ws_origin,
|
||
|
|
cleanup_ws_connection,
|
||
|
|
handle_ws_message,
|
||
|
|
send_ws_error,
|
||
|
|
start_heartbeat,
|
||
|
|
)
|
||
|
|
from app.core.ws_pubsub import (
|
||
|
|
broadcast_to_tenants,
|
||
|
|
get_tenant_channel,
|
||
|
|
publish_to_channel,
|
||
|
|
subscribe_to_channel,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
# ─── Mock WebSocket ───────────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
class MockWebSocket:
|
||
|
|
"""Minimal mock WebSocket for testing WS helpers."""
|
||
|
|
|
||
|
|
def __init__(
|
||
|
|
self,
|
||
|
|
cookies: dict[str, str] | None = None,
|
||
|
|
headers: dict[str, str] | None = None,
|
||
|
|
query_params: dict[str, str] | None = None,
|
||
|
|
) -> None:
|
||
|
|
self.cookies = cookies or {}
|
||
|
|
self.headers = headers or {}
|
||
|
|
self.query_params = query_params or {}
|
||
|
|
self._closed = False
|
||
|
|
self._close_code: int | None = None
|
||
|
|
self._close_reason: str | None = None
|
||
|
|
self._sent: list[str] = []
|
||
|
|
self._accepted = False
|
||
|
|
|
||
|
|
async def accept(self) -> None:
|
||
|
|
self._accepted = True
|
||
|
|
|
||
|
|
async def close(self, code: int = 1000, reason: str = "") -> None:
|
||
|
|
self._closed = True
|
||
|
|
self._close_code = code
|
||
|
|
self._close_reason = reason
|
||
|
|
|
||
|
|
async def send_text(self, text: str) -> None:
|
||
|
|
if self._closed:
|
||
|
|
raise RuntimeError("WebSocket is closed")
|
||
|
|
self._sent.append(text)
|
||
|
|
|
||
|
|
@property
|
||
|
|
def sent_messages(self) -> list[dict]:
|
||
|
|
return [json.loads(t) for t in self._sent]
|
||
|
|
|
||
|
|
|
||
|
|
# ─── authenticate_ws ──────────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
class TestAuthenticateWs:
|
||
|
|
"""Tests for authenticate_ws."""
|
||
|
|
|
||
|
|
async def test_authenticate_ws_valid_session(self, db_session, redis_client):
|
||
|
|
"""authenticate_ws returns user info when session is valid."""
|
||
|
|
from app.core.auth import create_session, hash_password
|
||
|
|
from app.models.tenant import Tenant
|
||
|
|
from app.models.user import User, UserTenant
|
||
|
|
from app.models.role import Role
|
||
|
|
|
||
|
|
# Create tenant, user, role, membership
|
||
|
|
tenant = Tenant(name="Test Tenant", slug="test-tenant")
|
||
|
|
db_session.add(tenant)
|
||
|
|
await db_session.flush()
|
||
|
|
|
||
|
|
user = User(
|
||
|
|
email="wsauth@test.com",
|
||
|
|
name="WS Auth Test",
|
||
|
|
password_hash=hash_password("TestPass123!"),
|
||
|
|
is_active=True,
|
||
|
|
preferences={},
|
||
|
|
)
|
||
|
|
db_session.add(user)
|
||
|
|
await db_session.flush()
|
||
|
|
|
||
|
|
role = Role(
|
||
|
|
tenant_id=tenant.id,
|
||
|
|
name="admin",
|
||
|
|
permissions={"*": {"*": True}},
|
||
|
|
denied_permissions=[],
|
||
|
|
field_permissions={},
|
||
|
|
)
|
||
|
|
db_session.add(role)
|
||
|
|
await db_session.flush()
|
||
|
|
|
||
|
|
ut = UserTenant(
|
||
|
|
user_id=user.id,
|
||
|
|
tenant_id=tenant.id,
|
||
|
|
is_default=True,
|
||
|
|
role="admin",
|
||
|
|
role_id=role.id,
|
||
|
|
)
|
||
|
|
db_session.add(ut)
|
||
|
|
await db_session.flush()
|
||
|
|
await db_session.commit()
|
||
|
|
|
||
|
|
# Create session
|
||
|
|
session_id, csrf_token = await create_session(
|
||
|
|
db_session, redis_client, user, tenant.id, role="admin"
|
||
|
|
)
|
||
|
|
|
||
|
|
# Create mock WebSocket with session cookie
|
||
|
|
ws = MockWebSocket(cookies={"leocrm_session": session_id})
|
||
|
|
|
||
|
|
auth = await authenticate_ws(ws, db_session)
|
||
|
|
assert auth is not None
|
||
|
|
assert auth["user_id"] == str(user.id)
|
||
|
|
assert auth["tenant_id"] == str(tenant.id)
|
||
|
|
assert auth["session_id"] == session_id
|
||
|
|
assert auth["role"] == "admin"
|
||
|
|
assert ws._closed is False
|
||
|
|
|
||
|
|
async def test_authenticate_ws_missing_cookie(self, db_session):
|
||
|
|
"""authenticate_ws closes WS and returns None when no session cookie."""
|
||
|
|
ws = MockWebSocket(cookies={})
|
||
|
|
|
||
|
|
auth = await authenticate_ws(ws, db_session)
|
||
|
|
assert auth is None
|
||
|
|
assert ws._closed is True
|
||
|
|
assert ws._close_code == 4401
|
||
|
|
|
||
|
|
async def test_authenticate_ws_invalid_session(self, db_session):
|
||
|
|
"""authenticate_ws closes WS and returns None when session is invalid."""
|
||
|
|
ws = MockWebSocket(cookies={"leocrm_session": "nonexistent-session-id"})
|
||
|
|
|
||
|
|
auth = await authenticate_ws(ws, db_session)
|
||
|
|
assert auth is None
|
||
|
|
assert ws._closed is True
|
||
|
|
assert ws._close_code == 4401
|
||
|
|
|
||
|
|
|
||
|
|
# ─── check_ws_origin ──────────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
class TestCheckWsOrigin:
|
||
|
|
"""Tests for check_ws_origin."""
|
||
|
|
|
||
|
|
async def test_check_ws_origin_valid(self):
|
||
|
|
"""check_ws_origin returns True for valid origin + CSRF."""
|
||
|
|
from app.core.auth import create_session, hash_password
|
||
|
|
from app.models.tenant import Tenant
|
||
|
|
from app.models.user import User, UserTenant
|
||
|
|
from app.models.role import Role
|
||
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||
|
|
|
||
|
|
# We need a full setup with session in Redis for CSRF validation
|
||
|
|
# Use patch to mock verify_ws_origin returning True
|
||
|
|
with patch("app.core.ws_helpers.verify_ws_origin", return_value=True):
|
||
|
|
ws = MockWebSocket()
|
||
|
|
result = await check_ws_origin(ws)
|
||
|
|
assert result is True
|
||
|
|
assert ws._closed is False
|
||
|
|
|
||
|
|
async def test_check_ws_origin_invalid(self):
|
||
|
|
"""check_ws_origin closes WS and returns False for invalid origin."""
|
||
|
|
with patch("app.core.ws_helpers.verify_ws_origin", return_value=False):
|
||
|
|
ws = MockWebSocket()
|
||
|
|
result = await check_ws_origin(ws)
|
||
|
|
assert result is False
|
||
|
|
assert ws._closed is True
|
||
|
|
assert ws._close_code == 4403
|
||
|
|
|
||
|
|
|
||
|
|
# ─── send_ws_error ────────────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
class TestSendWsError:
|
||
|
|
"""Tests for send_ws_error."""
|
||
|
|
|
||
|
|
async def test_send_ws_error_basic(self):
|
||
|
|
"""send_ws_error sends structured error message."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
await send_ws_error(ws, "test_error", "Something went wrong")
|
||
|
|
|
||
|
|
assert len(ws.sent_messages) == 1
|
||
|
|
msg = ws.sent_messages[0]
|
||
|
|
assert msg["type"] == "error"
|
||
|
|
assert msg["code"] == "test_error"
|
||
|
|
assert msg["detail"] == "Something went wrong"
|
||
|
|
assert "trace_id" not in msg
|
||
|
|
|
||
|
|
async def test_send_ws_error_with_trace_id(self):
|
||
|
|
"""send_ws_error includes trace_id when provided."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
await send_ws_error(ws, "test_error", "Failed", trace_id="trace-123")
|
||
|
|
|
||
|
|
msg = ws.sent_messages[0]
|
||
|
|
assert msg["trace_id"] == "trace-123"
|
||
|
|
|
||
|
|
async def test_send_ws_error_on_closed_ws(self):
|
||
|
|
"""send_ws_error does not raise when WS is closed."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
ws._closed = True
|
||
|
|
# Should not raise
|
||
|
|
await send_ws_error(ws, "closed_error", "WS already closed")
|
||
|
|
|
||
|
|
|
||
|
|
# ─── handle_ws_message ─────────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
class TestHandleWsMessage:
|
||
|
|
"""Tests for handle_ws_message."""
|
||
|
|
|
||
|
|
async def test_handle_ws_message_dispatch(self):
|
||
|
|
"""handle_ws_message dispatches to correct handler."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
called = False
|
||
|
|
|
||
|
|
async def ping_handler(websocket, msg):
|
||
|
|
nonlocal called
|
||
|
|
called = True
|
||
|
|
assert msg["type"] == "ping"
|
||
|
|
|
||
|
|
handlers = {"ping": ping_handler}
|
||
|
|
message = json.dumps({"type": "ping"})
|
||
|
|
await handle_ws_message(ws, message, handlers)
|
||
|
|
|
||
|
|
assert called is True
|
||
|
|
|
||
|
|
async def test_handle_ws_message_unknown_type(self):
|
||
|
|
"""handle_ws_message sends error for unknown type."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
handlers = {"ping": AsyncMock()}
|
||
|
|
message = json.dumps({"type": "unknown_type"})
|
||
|
|
await handle_ws_message(ws, message, handlers)
|
||
|
|
|
||
|
|
assert len(ws.sent_messages) == 1
|
||
|
|
msg = ws.sent_messages[0]
|
||
|
|
assert msg["type"] == "error"
|
||
|
|
assert msg["code"] == "unknown_type"
|
||
|
|
|
||
|
|
async def test_handle_ws_message_handler_exception(self):
|
||
|
|
"""handle_ws_message sends error when handler raises exception."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
|
||
|
|
async def bad_handler(websocket, msg):
|
||
|
|
raise ValueError("Handler crashed")
|
||
|
|
|
||
|
|
handlers = {"ping": bad_handler}
|
||
|
|
message = json.dumps({"type": "ping"})
|
||
|
|
await handle_ws_message(ws, message, handlers)
|
||
|
|
|
||
|
|
assert len(ws.sent_messages) == 1
|
||
|
|
msg = ws.sent_messages[0]
|
||
|
|
assert msg["type"] == "error"
|
||
|
|
assert msg["code"] == "handler_error"
|
||
|
|
assert "Handler crashed" in msg["detail"]
|
||
|
|
|
||
|
|
async def test_handle_ws_message_invalid_json(self):
|
||
|
|
"""handle_ws_message sends error for invalid JSON."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
handlers = {}
|
||
|
|
await handle_ws_message(ws, "not json at all", handlers)
|
||
|
|
|
||
|
|
assert len(ws.sent_messages) == 1
|
||
|
|
msg = ws.sent_messages[0]
|
||
|
|
assert msg["code"] == "invalid_json"
|
||
|
|
|
||
|
|
async def test_handle_ws_message_missing_type(self):
|
||
|
|
"""handle_ws_message sends error when type field is missing."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
handlers = {}
|
||
|
|
message = json.dumps({"data": "no type here"})
|
||
|
|
await handle_ws_message(ws, message, handlers)
|
||
|
|
|
||
|
|
assert len(ws.sent_messages) == 1
|
||
|
|
msg = ws.sent_messages[0]
|
||
|
|
assert msg["code"] == "missing_type"
|
||
|
|
|
||
|
|
|
||
|
|
# ─── cleanup_ws_connection ─────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
class TestCleanupWsConnection:
|
||
|
|
"""Tests for cleanup_ws_connection."""
|
||
|
|
|
||
|
|
async def test_cleanup_removes_from_registry(self):
|
||
|
|
"""cleanup_ws_connection removes WebSocket from registry."""
|
||
|
|
ws1 = MockWebSocket()
|
||
|
|
ws2 = MockWebSocket()
|
||
|
|
registry: dict[str, list] = {"user1": [ws1, ws2]}
|
||
|
|
|
||
|
|
await cleanup_ws_connection(ws1, "user1", registry)
|
||
|
|
|
||
|
|
assert "user1" in registry
|
||
|
|
assert ws1 not in registry["user1"]
|
||
|
|
assert ws2 in registry["user1"]
|
||
|
|
assert len(registry["user1"]) == 1
|
||
|
|
|
||
|
|
async def test_cleanup_removes_empty_user(self):
|
||
|
|
"""cleanup_ws_connection removes user entry when no connections left."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
registry: dict[str, list] = {"user1": [ws]}
|
||
|
|
|
||
|
|
await cleanup_ws_connection(ws, "user1", registry)
|
||
|
|
|
||
|
|
assert "user1" not in registry
|
||
|
|
|
||
|
|
async def test_cleanup_unknown_user(self):
|
||
|
|
"""cleanup_ws_connection handles unknown user gracefully."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
registry: dict[str, list] = {}
|
||
|
|
|
||
|
|
await cleanup_ws_connection(ws, "unknown_user", registry)
|
||
|
|
assert "unknown_user" not in registry
|
||
|
|
|
||
|
|
|
||
|
|
# ─── start_heartbeat ───────────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
class TestStartHeartbeat:
|
||
|
|
"""Tests for start_heartbeat."""
|
||
|
|
|
||
|
|
async def test_heartbeat_sends_ping(self):
|
||
|
|
"""start_heartbeat sends ping messages at interval."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
task = await start_heartbeat(ws, interval=0)
|
||
|
|
|
||
|
|
# Wait a tiny bit for the first ping
|
||
|
|
await asyncio.sleep(0.05)
|
||
|
|
task.cancel()
|
||
|
|
try:
|
||
|
|
await task
|
||
|
|
except asyncio.CancelledError:
|
||
|
|
pass
|
||
|
|
|
||
|
|
assert len(ws._sent) >= 1
|
||
|
|
msg = ws.sent_messages[0]
|
||
|
|
assert msg["type"] == "ping"
|
||
|
|
|
||
|
|
async def test_heartbeat_cancellable(self):
|
||
|
|
"""start_heartbeat task can be cancelled."""
|
||
|
|
ws = MockWebSocket()
|
||
|
|
task = await start_heartbeat(ws, interval=10)
|
||
|
|
|
||
|
|
task.cancel()
|
||
|
|
try:
|
||
|
|
await task
|
||
|
|
except asyncio.CancelledError:
|
||
|
|
pass
|
||
|
|
|
||
|
|
assert task.cancelled() or task.done()
|
||
|
|
|
||
|
|
|
||
|
|
# ─── get_tenant_channel ─────────────────────────────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
class TestGetTenantChannel:
|
||
|
|
"""Tests for get_tenant_channel."""
|
||
|
|
|
||
|
|
def test_channel_naming(self):
|
||
|
|
"""get_tenant_channel returns ws:{tenant_id}:{topic}."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
channel = get_tenant_channel(tenant_id, "kommunikation")
|
||
|
|
assert channel == f"ws:{tenant_id}:kommunikation"
|
||
|
|
|
||
|
|
def test_channel_naming_different_topic(self):
|
||
|
|
"""get_tenant_channel works with different topics."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
channel = get_tenant_channel(tenant_id, "ai_ui_control")
|
||
|
|
assert channel == f"ws:{tenant_id}:ai_ui_control"
|
||
|
|
|
||
|
|
|
||
|
|
# ─── publish_to_channel + subscribe_to_channel ──────────────────────────────────
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
class TestPubSub:
|
||
|
|
"""Tests for Redis Pub/Sub helpers."""
|
||
|
|
|
||
|
|
async def test_publish_subscribe_roundtrip(self, redis_client):
|
||
|
|
"""publish_to_channel + subscribe_to_channel roundtrip."""
|
||
|
|
received: list[dict] = []
|
||
|
|
channel = f"test-roundtrip-{uuid.uuid4()}"
|
||
|
|
event = asyncio.Event()
|
||
|
|
|
||
|
|
async def handler(msg: dict) -> None:
|
||
|
|
received.append(msg)
|
||
|
|
event.set()
|
||
|
|
|
||
|
|
# Subscribe
|
||
|
|
task = await subscribe_to_channel(channel, handler)
|
||
|
|
|
||
|
|
# Give subscriber a moment to connect
|
||
|
|
await asyncio.sleep(0.1)
|
||
|
|
|
||
|
|
# Publish
|
||
|
|
test_msg = {"type": "test", "data": "hello"}
|
||
|
|
await publish_to_channel(channel, test_msg)
|
||
|
|
|
||
|
|
# Wait for message
|
||
|
|
await asyncio.wait_for(event.wait(), timeout=2.0)
|
||
|
|
|
||
|
|
assert len(received) == 1
|
||
|
|
assert received[0]["type"] == "test"
|
||
|
|
assert received[0]["data"] == "hello"
|
||
|
|
|
||
|
|
# Cleanup
|
||
|
|
task.cancel()
|
||
|
|
try:
|
||
|
|
await task
|
||
|
|
except asyncio.CancelledError:
|
||
|
|
pass
|
||
|
|
|
||
|
|
async def test_broadcast_to_tenants(self, redis_client):
|
||
|
|
"""broadcast_to_tenants publishes to the correct tenant channel."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
received: list[dict] = []
|
||
|
|
event = asyncio.Event()
|
||
|
|
|
||
|
|
channel = get_tenant_channel(tenant_id, "test_topic")
|
||
|
|
|
||
|
|
async def handler(msg: dict) -> None:
|
||
|
|
received.append(msg)
|
||
|
|
event.set()
|
||
|
|
|
||
|
|
task = await subscribe_to_channel(channel, handler)
|
||
|
|
await asyncio.sleep(0.1)
|
||
|
|
|
||
|
|
test_msg = {"type": "broadcast", "content": "tenant message"}
|
||
|
|
await broadcast_to_tenants(tenant_id, "test_topic", test_msg)
|
||
|
|
|
||
|
|
await asyncio.wait_for(event.wait(), timeout=2.0)
|
||
|
|
|
||
|
|
assert len(received) == 1
|
||
|
|
assert received[0]["type"] == "broadcast"
|
||
|
|
assert received[0]["content"] == "tenant message"
|
||
|
|
|
||
|
|
task.cancel()
|
||
|
|
try:
|
||
|
|
await task
|
||
|
|
except asyncio.CancelledError:
|
||
|
|
pass
|
||
|
|
|
||
|
|
async def test_publish_to_channel_no_subscribers(self, redis_client):
|
||
|
|
"""publish_to_channel works even with no subscribers."""
|
||
|
|
channel = f"test-no-sub-{uuid.uuid4()}"
|
||
|
|
await publish_to_channel(channel, {"type": "noop"})
|
||
|
|
# Should not raise
|
||
|
|
|
||
|
|
async def test_subscribe_multiple_messages(self, redis_client):
|
||
|
|
"""subscribe_to_channel handler receives multiple messages."""
|
||
|
|
received: list[dict] = []
|
||
|
|
channel = f"test-multi-{uuid.uuid4()}"
|
||
|
|
count_event = asyncio.Event()
|
||
|
|
msg_count = 0
|
||
|
|
|
||
|
|
async def handler(msg: dict) -> None:
|
||
|
|
nonlocal msg_count
|
||
|
|
received.append(msg)
|
||
|
|
msg_count += 1
|
||
|
|
if msg_count >= 3:
|
||
|
|
count_event.set()
|
||
|
|
|
||
|
|
task = await subscribe_to_channel(channel, handler)
|
||
|
|
await asyncio.sleep(0.1)
|
||
|
|
|
||
|
|
for i in range(3):
|
||
|
|
await publish_to_channel(channel, {"type": "msg", "index": i})
|
||
|
|
|
||
|
|
await asyncio.wait_for(count_event.wait(), timeout=3.0)
|
||
|
|
|
||
|
|
assert len(received) == 3
|
||
|
|
assert received[0]["index"] == 0
|
||
|
|
assert received[2]["index"] == 2
|
||
|
|
|
||
|
|
task.cancel()
|
||
|
|
try:
|
||
|
|
await task
|
||
|
|
except asyncio.CancelledError:
|
||
|
|
pass
|