| 1 | import asyncio |
| 2 | import sys |
| 3 | import threading |
| 4 | import time |
| 5 | from pathlib import Path |
| 6 | from typing import Any |
| 7 | from unittest.mock import AsyncMock, patch |
| 8 | |
| 9 | import pytest |
| 10 | |
| 11 | PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| 12 | if str(PROJECT_ROOT) not in sys.path: |
| 13 | sys.path.insert(0, str(PROJECT_ROOT)) |
| 14 | |
| 15 | from helpers.ws import ConnectionNotFoundError, WsHandler |
| 16 | from helpers.ws_manager import ( |
| 17 | WsManager, |
| 18 | WsResult, |
| 19 | BUFFER_TTL, |
| 20 | DIAGNOSTIC_EVENT, |
| 21 | LIFECYCLE_CONNECT_EVENT, |
| 22 | LIFECYCLE_DISCONNECT_EVENT, |
| 23 | ) |
| 24 | |
| 25 | NAMESPACE = "/test" |
| 26 | |
| 27 | |
| 28 | class FakeSocketIOServer: |
| 29 | def __init__(self): |
| 30 | self.emit = AsyncMock() |
| 31 | self.disconnect = AsyncMock() |
| 32 | |
| 33 | |
| 34 | class DummyHandler(WsHandler): |
| 35 | def __init__(self, socketio, lock, results=None): |
| 36 | super().__init__(socketio, lock) |
| 37 | self.results = results if results is not None else [] |
| 38 | |
| 39 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 40 | response = {"sid": sid, "data": data} |
| 41 | self.results.append(response) |
| 42 | return response |
| 43 | |
| 44 | |
| 45 | @pytest.mark.asyncio |
| 46 | async def test_connect_disconnect_updates_registry(): |
| 47 | socketio = FakeSocketIOServer() |
| 48 | manager = WsManager(socketio, threading.RLock()) |
| 49 | |
| 50 | await manager.handle_connect(NAMESPACE, "abc") |
| 51 | assert (NAMESPACE, "abc") in manager.connections |
| 52 | |
| 53 | await manager.handle_disconnect(NAMESPACE, "abc") |
| 54 | assert (NAMESPACE, "abc") not in manager.connections |
| 55 | |
| 56 | |
| 57 | @pytest.mark.asyncio |
| 58 | async def test_server_restart_broadcast_emitted_when_enabled(): |
| 59 | socketio = FakeSocketIOServer() |
| 60 | manager = WsManager(socketio, threading.RLock()) |
| 61 | manager.set_server_restart_broadcast(True) |
| 62 | |
| 63 | await manager.handle_connect(NAMESPACE, "sid-restart") |
| 64 | |
| 65 | socketio.emit.assert_awaited() |
| 66 | args, kwargs = socketio.emit.await_args_list[0] |
| 67 | assert args[0] == "server_restart" |
| 68 | envelope = args[1] |
| 69 | assert envelope["handlerId"] == manager._identifier # noqa: SLF001 |
| 70 | assert envelope["data"]["runtimeId"] |
| 71 | assert kwargs == {"to": "sid-restart", "namespace": NAMESPACE} |
| 72 | |
| 73 | |
| 74 | @pytest.mark.asyncio |
| 75 | async def test_server_restart_broadcast_skipped_when_disabled(): |
| 76 | socketio = FakeSocketIOServer() |
| 77 | manager = WsManager(socketio, threading.RLock()) |
| 78 | manager.set_server_restart_broadcast(False) |
| 79 | |
| 80 | await manager.handle_connect(NAMESPACE, "sid-no-restart") |
| 81 | |
| 82 | assert socketio.emit.await_count == 0 |
| 83 | |
| 84 | |
| 85 | @pytest.mark.asyncio |
| 86 | async def test_broadcast_performance_smoke(monkeypatch): |
| 87 | socketio = FakeSocketIOServer() |
| 88 | manager = WsManager(socketio, threading.RLock()) |
| 89 | |
| 90 | for idx in range(50): |
| 91 | await manager.handle_connect(NAMESPACE, f"sid-{idx}") |
| 92 | |
| 93 | # Drain lifecycle broadcast tasks queued by handle_connect and reset the |
| 94 | # mock so we only measure the explicit broadcast below. |
| 95 | for _ in range(55): |
| 96 | await asyncio.sleep(0) |
| 97 | socketio.emit.reset_mock() |
| 98 | |
| 99 | import time |
| 100 | |
| 101 | start = time.perf_counter() |
| 102 | await manager.broadcast(NAMESPACE, "perf_event", {"ok": True}) |
| 103 | duration_ms = (time.perf_counter() - start) * 1000 |
| 104 | |
| 105 | assert socketio.emit.await_count == 50 |
| 106 | assert duration_ms < 300 |
| 107 | |
| 108 | |
| 109 | @pytest.mark.asyncio |
| 110 | async def test_route_event_invokes_handler_and_ack(): |
| 111 | socketio = FakeSocketIOServer() |
| 112 | manager = WsManager(socketio, threading.RLock()) |
| 113 | |
| 114 | results = [] |
| 115 | handler = DummyHandler(socketio, threading.RLock(), results) |
| 116 | manager.register_handlers({NAMESPACE: [handler]}) |
| 117 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 118 | |
| 119 | response = await manager.route_event(NAMESPACE, "dummy", {"foo": "bar"}, "sid-1") |
| 120 | |
| 121 | assert results[0]["sid"] == "sid-1" |
| 122 | assert results[0]["data"]["foo"] == "bar" |
| 123 | assert "correlationId" in results[0]["data"] |
| 124 | |
| 125 | assert isinstance(response, dict) |
| 126 | assert "correlationId" in response |
| 127 | assert isinstance(response["results"], list) |
| 128 | entry = response["results"][0] |
| 129 | assert entry["ok"] is True |
| 130 | assert entry["data"]["sid"] == "sid-1" |
| 131 | assert entry["data"]["data"]["foo"] == "bar" |
| 132 | |
| 133 | |
| 134 | @pytest.mark.asyncio |
| 135 | async def test_route_event_no_handler_returns_standard_error(): |
| 136 | socketio = FakeSocketIOServer() |
| 137 | manager = WsManager(socketio, threading.RLock()) |
| 138 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 139 | |
| 140 | response = await manager.route_event(NAMESPACE, "missing", {}, "sid-1") |
| 141 | |
| 142 | assert len(response["results"]) == 1 |
| 143 | result = response["results"][0] |
| 144 | assert result["handlerId"].endswith("WsManager") |
| 145 | assert result["ok"] is False |
| 146 | assert result["error"]["code"] == "NO_HANDLERS" |
| 147 | assert ( |
| 148 | result["error"]["error"] |
| 149 | == f"No handler for namespace '{NAMESPACE}'" |
| 150 | ) |
| 151 | |
| 152 | |
| 153 | @pytest.mark.asyncio |
| 154 | async def test_route_event_all_returns_empty_when_no_connections(): |
| 155 | socketio = FakeSocketIOServer() |
| 156 | manager = WsManager(socketio, threading.RLock()) |
| 157 | |
| 158 | results = await manager.route_event_all(NAMESPACE, "event", {}, timeout_ms=1000) |
| 159 | |
| 160 | assert results == [] |
| 161 | |
| 162 | |
| 163 | @pytest.mark.asyncio |
| 164 | async def test_route_event_all_aggregates_results(): |
| 165 | socketio = FakeSocketIOServer() |
| 166 | manager = WsManager(socketio, threading.RLock()) |
| 167 | |
| 168 | class EchoHandler(WsHandler): |
| 169 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 170 | return {"sid": sid, "echo": data} |
| 171 | |
| 172 | handler = EchoHandler(socketio, threading.RLock()) |
| 173 | manager.register_handlers({NAMESPACE: [handler]}) |
| 174 | |
| 175 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 176 | await manager.handle_connect(NAMESPACE, "sid-2") |
| 177 | |
| 178 | aggregated = await manager.route_event_all( |
| 179 | NAMESPACE, "multi", {"value": 42}, timeout_ms=1000 |
| 180 | ) |
| 181 | |
| 182 | assert len(aggregated) == 2 |
| 183 | by_sid = {entry["sid"]: entry for entry in aggregated} |
| 184 | assert by_sid["sid-1"]["results"][0]["ok"] is True |
| 185 | payload_sid1 = by_sid["sid-1"]["results"][0]["data"] |
| 186 | assert payload_sid1["sid"] == "sid-1" |
| 187 | assert payload_sid1["echo"]["value"] == 42 |
| 188 | assert "correlationId" in payload_sid1["echo"] |
| 189 | assert by_sid["sid-2"]["results"][0]["ok"] is True |
| 190 | payload_sid2 = by_sid["sid-2"]["results"][0]["data"] |
| 191 | assert payload_sid2["sid"] == "sid-2" |
| 192 | assert payload_sid2["echo"]["value"] == 42 |
| 193 | assert by_sid["sid-1"]["correlationId"] |
| 194 | |
| 195 | |
| 196 | @pytest.mark.asyncio |
| 197 | async def test_route_event_all_timeout_marks_error(): |
| 198 | socketio = FakeSocketIOServer() |
| 199 | manager = WsManager(socketio, threading.RLock()) |
| 200 | |
| 201 | class SlowHandler(WsHandler): |
| 202 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 203 | await asyncio.sleep(0.2) |
| 204 | return {"status": "done"} |
| 205 | |
| 206 | handler = SlowHandler(socketio, threading.RLock()) |
| 207 | manager.register_handlers({NAMESPACE: [handler]}) |
| 208 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 209 | |
| 210 | aggregated = await manager.route_event_all(NAMESPACE, "slow", {}, timeout_ms=50) |
| 211 | |
| 212 | assert len(aggregated) == 1 |
| 213 | first_entry = aggregated[0] |
| 214 | result = first_entry["results"][0] |
| 215 | assert result["ok"] is False |
| 216 | assert result["error"] == {"code": "TIMEOUT", "error": "Request timeout"} |
| 217 | assert first_entry["correlationId"] |
| 218 | |
| 219 | |
| 220 | @pytest.mark.asyncio |
| 221 | async def test_route_event_exception_standardizes_error_payload(): |
| 222 | socketio = FakeSocketIOServer() |
| 223 | manager = WsManager(socketio, threading.RLock()) |
| 224 | |
| 225 | class FailingHandler(WsHandler): |
| 226 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 227 | raise RuntimeError("kaboom") |
| 228 | |
| 229 | handler = FailingHandler(socketio, threading.RLock()) |
| 230 | manager.register_handlers({NAMESPACE: [handler]}) |
| 231 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 232 | |
| 233 | response = await manager.route_event(NAMESPACE, "boom", {}, "sid-1") |
| 234 | |
| 235 | assert len(response["results"]) == 1 |
| 236 | result = response["results"][0] |
| 237 | assert result["handlerId"].endswith("FailingHandler") |
| 238 | assert result["ok"] is False |
| 239 | assert result["error"]["code"] == "HANDLER_ERROR" |
| 240 | assert result["error"]["error"] == "Internal server error" |
| 241 | assert "details" in result["error"] |
| 242 | |
| 243 | |
| 244 | @pytest.mark.asyncio |
| 245 | async def test_route_event_offloads_blocking_handlers(): |
| 246 | socketio = FakeSocketIOServer() |
| 247 | manager = WsManager(socketio, threading.RLock()) |
| 248 | |
| 249 | class BlockingHandler(WsHandler): |
| 250 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 251 | time.sleep(0.2) |
| 252 | return {"status": "done"} |
| 253 | |
| 254 | handler = BlockingHandler(socketio, threading.RLock()) |
| 255 | manager.register_handlers({NAMESPACE: [handler]}) |
| 256 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 257 | |
| 258 | route_task = asyncio.create_task( |
| 259 | manager.route_event(NAMESPACE, "block", {}, "sid-1") |
| 260 | ) |
| 261 | await asyncio.sleep(0) |
| 262 | |
| 263 | t0 = time.perf_counter() |
| 264 | await asyncio.sleep(0.05) |
| 265 | elapsed = time.perf_counter() - t0 |
| 266 | assert elapsed < 0.15 |
| 267 | |
| 268 | response = await route_task |
| 269 | assert response["results"] |
| 270 | |
| 271 | |
| 272 | @pytest.mark.asyncio |
| 273 | async def test_route_event_unwraps_ts_data_envelope_and_preserves_correlation_id(): |
| 274 | socketio = FakeSocketIOServer() |
| 275 | manager = WsManager(socketio, threading.RLock()) |
| 276 | |
| 277 | results: list[dict[str, Any]] = [] |
| 278 | handler = DummyHandler(socketio, threading.RLock(), results) |
| 279 | manager.register_handlers({NAMESPACE: [handler]}) |
| 280 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 281 | |
| 282 | response = await manager.route_event( |
| 283 | NAMESPACE, |
| 284 | "dummy", |
| 285 | { |
| 286 | "correlationId": "client-1", |
| 287 | "ts": "2025-10-29T12:00:00.000Z", |
| 288 | "data": {"value": 123}, |
| 289 | }, |
| 290 | "sid-1", |
| 291 | ) |
| 292 | |
| 293 | assert response["correlationId"] == "client-1" |
| 294 | assert len(results) == 1 |
| 295 | handler_payload = results[0]["data"] |
| 296 | assert handler_payload["value"] == 123 |
| 297 | assert handler_payload["correlationId"] == "client-1" |
| 298 | assert "ts" not in handler_payload |
| 299 | assert "data" not in handler_payload |
| 300 | |
| 301 | |
| 302 | @pytest.mark.asyncio |
| 303 | async def test_emit_to_unknown_sid_raises_error(): |
| 304 | socketio = FakeSocketIOServer() |
| 305 | manager = WsManager(socketio, threading.RLock()) |
| 306 | |
| 307 | with pytest.raises(ConnectionNotFoundError): |
| 308 | await manager.emit_to(NAMESPACE, "unknown", "event", {}) |
| 309 | |
| 310 | |
| 311 | @pytest.mark.asyncio |
| 312 | async def test_emit_to_known_disconnected_sid_buffers(): |
| 313 | socketio = FakeSocketIOServer() |
| 314 | manager = WsManager(socketio, threading.RLock()) |
| 315 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 316 | await manager.handle_disconnect(NAMESPACE, "sid-1") |
| 317 | |
| 318 | await manager.emit_to( |
| 319 | NAMESPACE, "sid-1", "event", {"a": 1}, correlation_id="corr-1" |
| 320 | ) |
| 321 | |
| 322 | assert (NAMESPACE, "sid-1") in manager.buffers |
| 323 | buffered = list(manager.buffers[(NAMESPACE, "sid-1")]) |
| 324 | assert len(buffered) == 1 |
| 325 | assert buffered[0].event_type == "event" |
| 326 | assert buffered[0].data == {"a": 1} |
| 327 | assert buffered[0].correlation_id == "corr-1" |
| 328 | |
| 329 | |
| 330 | @pytest.mark.asyncio |
| 331 | async def test_buffer_overflow_drops_oldest(monkeypatch): |
| 332 | socketio = FakeSocketIOServer() |
| 333 | manager = WsManager(socketio, threading.RLock()) |
| 334 | |
| 335 | await manager.handle_connect(NAMESPACE, "offline") |
| 336 | await manager.handle_disconnect(NAMESPACE, "offline") |
| 337 | |
| 338 | monkeypatch.setattr("helpers.ws_manager.BUFFER_MAX_SIZE", 2) |
| 339 | |
| 340 | await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 0}) |
| 341 | await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 1}) |
| 342 | await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 2}) |
| 343 | |
| 344 | buffer = manager.buffers[(NAMESPACE, "offline")] |
| 345 | assert len(buffer) == 2 |
| 346 | assert buffer[0].data["idx"] == 1 |
| 347 | assert buffer[1].data["idx"] == 2 |
| 348 | |
| 349 | |
| 350 | @pytest.mark.asyncio |
| 351 | async def test_expired_buffer_entries_are_discarded(monkeypatch): |
| 352 | socketio = FakeSocketIOServer() |
| 353 | manager = WsManager(socketio, threading.RLock()) |
| 354 | |
| 355 | await manager.handle_connect(NAMESPACE, "sid-expired") |
| 356 | await manager.handle_disconnect(NAMESPACE, "sid-expired") |
| 357 | |
| 358 | from datetime import timedelta, timezone, datetime |
| 359 | |
| 360 | past = datetime.now(timezone.utc) - (BUFFER_TTL + timedelta(seconds=5)) |
| 361 | future = past + BUFFER_TTL + timedelta(seconds=10) |
| 362 | |
| 363 | await manager.emit_to(NAMESPACE, "sid-expired", "event", {"a": 1}) |
| 364 | manager.buffers[(NAMESPACE, "sid-expired")][0].timestamp = past |
| 365 | |
| 366 | socketio.emit.reset_mock() |
| 367 | |
| 368 | monkeypatch.setattr( |
| 369 | "helpers.ws_manager._utcnow", |
| 370 | lambda: future, |
| 371 | ) |
| 372 | await manager.handle_connect(NAMESPACE, "sid-expired") |
| 373 | |
| 374 | assert socketio.emit.await_count == 0 |
| 375 | assert (NAMESPACE, "sid-expired") not in manager.buffers |
| 376 | |
| 377 | |
| 378 | @pytest.mark.asyncio |
| 379 | async def test_flush_buffer_delivers_and_logs(monkeypatch): |
| 380 | socketio = FakeSocketIOServer() |
| 381 | manager = WsManager(socketio, threading.RLock()) |
| 382 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 383 | await manager.handle_disconnect(NAMESPACE, "sid-1") |
| 384 | |
| 385 | await manager.emit_to(NAMESPACE, "sid-1", "event", {"a": 1}) |
| 386 | |
| 387 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 388 | |
| 389 | assert len(socketio.emit.await_args_list) == 1 |
| 390 | awaited_call = socketio.emit.await_args_list[0] |
| 391 | assert awaited_call.args[0] == "event" |
| 392 | envelope = awaited_call.args[1] |
| 393 | assert envelope["data"] == {"a": 1} |
| 394 | assert "eventId" in envelope and "handlerId" in envelope and "ts" in envelope |
| 395 | assert awaited_call.kwargs == {"to": "sid-1", "namespace": NAMESPACE} |
| 396 | assert (NAMESPACE, "sid-1") not in manager.buffers |
| 397 | |
| 398 | |
| 399 | @pytest.mark.asyncio |
| 400 | async def test_known_sid_expires_after_buffer_ttl(monkeypatch): |
| 401 | """After BUFFER_TTL, a disconnected sid is swept from _known_sids and emit_to raises.""" |
| 402 | socketio = FakeSocketIOServer() |
| 403 | manager = WsManager(socketio, threading.RLock()) |
| 404 | |
| 405 | await manager.handle_connect(NAMESPACE, "sid-stale") |
| 406 | await manager.handle_disconnect(NAMESPACE, "sid-stale") |
| 407 | |
| 408 | # Immediately after disconnect, buffering still works |
| 409 | await manager.emit_to(NAMESPACE, "sid-stale", "event", {"x": 1}) |
| 410 | assert (NAMESPACE, "sid-stale") in manager.buffers |
| 411 | |
| 412 | from datetime import timedelta, timezone, datetime |
| 413 | |
| 414 | future = datetime.now(timezone.utc) + BUFFER_TTL + timedelta(seconds=10) |
| 415 | monkeypatch.setattr("helpers.ws_manager._utcnow", lambda: future) |
| 416 | |
| 417 | # After TTL, emit_to should raise because the sid is no longer known |
| 418 | with pytest.raises(ConnectionNotFoundError): |
| 419 | await manager.emit_to(NAMESPACE, "sid-stale", "event", {"x": 2}) |
| 420 | |
| 421 | # _known_sids and buffers should be cleaned |
| 422 | assert (NAMESPACE, "sid-stale") not in manager._known_sids |
| 423 | assert (NAMESPACE, "sid-stale") not in manager.buffers |
| 424 | assert (NAMESPACE, "sid-stale") not in manager._disconnect_times |
| 425 | |
| 426 | |
| 427 | @pytest.mark.asyncio |
| 428 | async def test_sweep_cleans_stale_sids_on_disconnect(monkeypatch): |
| 429 | """_sweep_stale_sids runs during handle_disconnect and cleans expired entries.""" |
| 430 | socketio = FakeSocketIOServer() |
| 431 | manager = WsManager(socketio, threading.RLock()) |
| 432 | |
| 433 | await manager.handle_connect(NAMESPACE, "old-sid") |
| 434 | await manager.handle_disconnect(NAMESPACE, "old-sid") |
| 435 | |
| 436 | from datetime import timedelta, timezone, datetime |
| 437 | |
| 438 | future = datetime.now(timezone.utc) + BUFFER_TTL + timedelta(seconds=10) |
| 439 | monkeypatch.setattr("helpers.ws_manager._utcnow", lambda: future) |
| 440 | |
| 441 | # A new connect/disconnect triggers sweep which cleans old-sid |
| 442 | await manager.handle_connect(NAMESPACE, "new-sid") |
| 443 | await manager.handle_disconnect(NAMESPACE, "new-sid") |
| 444 | |
| 445 | assert (NAMESPACE, "old-sid") not in manager._known_sids |
| 446 | assert (NAMESPACE, "old-sid") not in manager._disconnect_times |
| 447 | |
| 448 | |
| 449 | @pytest.mark.asyncio |
| 450 | async def test_broadcast_excludes_multiple_sids(): |
| 451 | socketio = FakeSocketIOServer() |
| 452 | manager = WsManager(socketio, threading.RLock()) |
| 453 | |
| 454 | for sid in ("sid-1", "sid-2", "sid-3"): |
| 455 | await manager.handle_connect(NAMESPACE, sid) |
| 456 | |
| 457 | # Drain lifecycle broadcast tasks from handle_connect |
| 458 | for _ in range(10): |
| 459 | await asyncio.sleep(0) |
| 460 | socketio.emit.reset_mock() |
| 461 | |
| 462 | await manager.broadcast( |
| 463 | NAMESPACE, |
| 464 | "event", |
| 465 | {"foo": "bar"}, |
| 466 | exclude_sids={"sid-1", "sid-3"}, |
| 467 | handler_id="custom.broadcast", |
| 468 | correlation_id="corr-b", |
| 469 | ) |
| 470 | |
| 471 | assert len(socketio.emit.await_args_list) == 1 |
| 472 | awaited_call = socketio.emit.await_args_list[0] |
| 473 | assert awaited_call.args[0] == "event" |
| 474 | envelope = awaited_call.args[1] |
| 475 | assert envelope["data"] == {"foo": "bar"} |
| 476 | assert envelope["handlerId"] == "custom.broadcast" |
| 477 | assert envelope["correlationId"] == "corr-b" |
| 478 | assert "eventId" in envelope and "ts" in envelope |
| 479 | assert awaited_call.kwargs == {"to": "sid-2", "namespace": NAMESPACE} |
| 480 | |
| 481 | |
| 482 | @pytest.mark.asyncio |
| 483 | async def test_emit_to_wraps_envelope_with_metadata(): |
| 484 | socketio = FakeSocketIOServer() |
| 485 | manager = WsManager(socketio, threading.RLock()) |
| 486 | await manager.handle_connect(NAMESPACE, "sid-meta") |
| 487 | |
| 488 | await manager.emit_to( |
| 489 | NAMESPACE, |
| 490 | "sid-meta", |
| 491 | "meta_event", |
| 492 | {"payload": True}, |
| 493 | handler_id="custom.handler", |
| 494 | correlation_id="corr-meta", |
| 495 | ) |
| 496 | |
| 497 | socketio.emit.assert_awaited_once() |
| 498 | args, kwargs = socketio.emit.await_args_list[0] |
| 499 | assert args[0] == "meta_event" |
| 500 | envelope = args[1] |
| 501 | assert envelope["handlerId"] == "custom.handler" |
| 502 | assert envelope["correlationId"] == "corr-meta" |
| 503 | assert envelope["data"] == {"payload": True} |
| 504 | assert kwargs == {"to": "sid-meta", "namespace": NAMESPACE} |
| 505 | |
| 506 | |
| 507 | @pytest.mark.asyncio |
| 508 | async def test_timestamps_are_timezone_aware(): |
| 509 | socketio = FakeSocketIOServer() |
| 510 | manager = WsManager(socketio, threading.RLock()) |
| 511 | |
| 512 | await manager.handle_connect(NAMESPACE, "sid-utc") |
| 513 | info = manager.connections[(NAMESPACE, "sid-utc")] |
| 514 | |
| 515 | assert info.connected_at.tzinfo is not None |
| 516 | assert info.last_activity.tzinfo is not None |
| 517 | |
| 518 | with patch("helpers.ws_manager._utcnow") as mocked_now: |
| 519 | mocked_now.return_value = info.last_activity |
| 520 | await manager.route_event(NAMESPACE, "unknown", {}, "sid-utc") |
| 521 | assert info.last_activity.tzinfo is not None |
| 522 | |
| 523 | class DuplicateHandler(WsHandler): |
| 524 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 525 | return {"handledBy": self.identifier} |
| 526 | |
| 527 | |
| 528 | class AnotherDuplicateHandler(WsHandler): |
| 529 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 530 | return {"handledBy": self.identifier} |
| 531 | |
| 532 | |
| 533 | def test_register_handlers_warns_on_duplicates(monkeypatch): |
| 534 | socketio = FakeSocketIOServer() |
| 535 | manager = WsManager(socketio, threading.RLock()) |
| 536 | |
| 537 | warnings: list[str] = [] |
| 538 | |
| 539 | def capture_warning(message: str) -> None: |
| 540 | warnings.append(message) |
| 541 | |
| 542 | monkeypatch.setattr( |
| 543 | "helpers.print_style.PrintStyle.warning", staticmethod(capture_warning) |
| 544 | ) |
| 545 | |
| 546 | handler_a = DuplicateHandler(socketio, threading.RLock()) |
| 547 | |
| 548 | manager.register_handlers({NAMESPACE: [handler_a, handler_a]}) |
| 549 | |
| 550 | assert any("Duplicate handler registration" in msg for msg in warnings) |
| 551 | |
| 552 | |
| 553 | class NonDictHandler(WsHandler): |
| 554 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 555 | return "raw-value" |
| 556 | |
| 557 | |
| 558 | @pytest.mark.asyncio |
| 559 | async def test_route_event_standardizes_success_payload(): |
| 560 | socketio = FakeSocketIOServer() |
| 561 | manager = WsManager(socketio, threading.RLock()) |
| 562 | |
| 563 | handler = NonDictHandler(socketio, threading.RLock()) |
| 564 | manager.register_handlers({NAMESPACE: [handler]}) |
| 565 | |
| 566 | response = await manager.route_event(NAMESPACE, "non_dict", {}, "sid-123") |
| 567 | |
| 568 | assert len(response["results"]) == 1 |
| 569 | assert response["results"][0]["ok"] is True |
| 570 | assert response["results"][0]["data"] == {"result": "raw-value"} |
| 571 | |
| 572 | |
| 573 | class ErrorHandler(WsHandler): |
| 574 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 575 | raise RuntimeError("BOOM") |
| 576 | |
| 577 | |
| 578 | class ResultHandler(WsHandler): |
| 579 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 580 | if event == "result_event": |
| 581 | return WsResult.ok({"sid": sid}, correlation_id="explicit", duration_ms=1.234) |
| 582 | return WsResult.error( |
| 583 | code="E_RESULT", |
| 584 | message="boom", |
| 585 | details="test", |
| 586 | ) |
| 587 | |
| 588 | |
| 589 | @pytest.mark.asyncio |
| 590 | async def test_route_event_standardizes_error_payload(): |
| 591 | socketio = FakeSocketIOServer() |
| 592 | manager = WsManager(socketio, threading.RLock()) |
| 593 | |
| 594 | handler = ErrorHandler(socketio, threading.RLock()) |
| 595 | manager.register_handlers({NAMESPACE: [handler]}) |
| 596 | |
| 597 | response = await manager.route_event(NAMESPACE, "boom", {}, "sid-123") |
| 598 | |
| 599 | assert len(response["results"]) == 1 |
| 600 | payload = response["results"][0] |
| 601 | assert payload["ok"] is False |
| 602 | assert payload["error"]["code"] == "HANDLER_ERROR" |
| 603 | assert payload["error"]["error"] == "Internal server error" |
| 604 | |
| 605 | |
| 606 | @pytest.mark.asyncio |
| 607 | async def test_route_event_accepts_websocket_result_instances(): |
| 608 | socketio = FakeSocketIOServer() |
| 609 | manager = WsManager(socketio, threading.RLock()) |
| 610 | |
| 611 | handler = ResultHandler(socketio, threading.RLock()) |
| 612 | manager.register_handlers({NAMESPACE: [handler]}) |
| 613 | |
| 614 | response = await manager.route_event(NAMESPACE, "result_event", {}, "sid-123") |
| 615 | |
| 616 | assert response["results"] |
| 617 | payload = response["results"][0] |
| 618 | assert payload["ok"] is True |
| 619 | assert payload["data"] == {"sid": "sid-123"} |
| 620 | assert payload["correlationId"] == "explicit" |
| 621 | assert payload["durationMs"] == pytest.approx(1.234, rel=1e-3) |
| 622 | |
| 623 | |
| 624 | @pytest.mark.asyncio |
| 625 | async def test_route_event_preserves_websocket_result_errors(): |
| 626 | socketio = FakeSocketIOServer() |
| 627 | manager = WsManager(socketio, threading.RLock()) |
| 628 | |
| 629 | handler = ResultHandler(socketio, threading.RLock()) |
| 630 | manager.register_handlers({NAMESPACE: [handler]}) |
| 631 | |
| 632 | response = await manager.route_event(NAMESPACE, "result_error", {}, "sid-123") |
| 633 | |
| 634 | payload = response["results"][0] |
| 635 | assert payload["ok"] is False |
| 636 | assert payload["error"] == {"code": "E_RESULT", "error": "boom", "details": "test"} |
| 637 | |
| 638 | |
| 639 | class AlphaFilterHandler(WsHandler): |
| 640 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 641 | return {"handledBy": self.identifier, "sid": sid} |
| 642 | |
| 643 | |
| 644 | class BetaFilterHandler(WsHandler): |
| 645 | async def process(self, event: str, data: dict[str, Any], sid: str): |
| 646 | return {"handledBy": self.identifier, "sid": sid} |
| 647 | |
| 648 | |
| 649 | @pytest.mark.asyncio |
| 650 | async def test_route_event_include_handlers_filters_results(): |
| 651 | socketio = FakeSocketIOServer() |
| 652 | manager = WsManager(socketio, threading.RLock()) |
| 653 | |
| 654 | alpha = AlphaFilterHandler(socketio, threading.RLock()) |
| 655 | beta = BetaFilterHandler(socketio, threading.RLock()) |
| 656 | manager.register_handlers({NAMESPACE: [alpha, beta]}) |
| 657 | await manager.handle_connect(NAMESPACE, "sid-filter") |
| 658 | |
| 659 | response = await manager.route_event( |
| 660 | NAMESPACE, |
| 661 | "filter_event", |
| 662 | { |
| 663 | "includeHandlers": [alpha.identifier], |
| 664 | "payload": True, |
| 665 | }, |
| 666 | "sid-filter", |
| 667 | ) |
| 668 | |
| 669 | assert response["correlationId"] |
| 670 | results = response["results"] |
| 671 | assert len(results) == 1 |
| 672 | assert results[0]["handlerId"] == alpha.identifier |
| 673 | assert results[0]["data"]["handledBy"] == alpha.identifier |
| 674 | |
| 675 | |
| 676 | @pytest.mark.asyncio |
| 677 | async def test_route_event_rejects_exclude_handlers_without_permission(): |
| 678 | socketio = FakeSocketIOServer() |
| 679 | manager = WsManager(socketio, threading.RLock()) |
| 680 | |
| 681 | handler = AlphaFilterHandler(socketio, threading.RLock()) |
| 682 | manager.register_handlers({NAMESPACE: [handler]}) |
| 683 | await manager.handle_connect(NAMESPACE, "sid-exclude") |
| 684 | |
| 685 | response = await manager.route_event( |
| 686 | NAMESPACE, |
| 687 | "filter_event", |
| 688 | {"excludeHandlers": [handler.identifier]}, |
| 689 | "sid-exclude", |
| 690 | ) |
| 691 | |
| 692 | result = response["results"][0] |
| 693 | assert result["error"]["code"] == "INVALID_FILTER" |
| 694 | assert "excludeHandlers" in result["error"]["error"] |
| 695 | |
| 696 | |
| 697 | @pytest.mark.asyncio |
| 698 | async def test_route_event_all_respects_exclude_handlers(): |
| 699 | socketio = FakeSocketIOServer() |
| 700 | manager = WsManager(socketio, threading.RLock()) |
| 701 | |
| 702 | alpha = AlphaFilterHandler(socketio, threading.RLock()) |
| 703 | beta = BetaFilterHandler(socketio, threading.RLock()) |
| 704 | manager.register_handlers({NAMESPACE: [alpha, beta]}) |
| 705 | |
| 706 | await manager.handle_connect(NAMESPACE, "sid-a") |
| 707 | await manager.handle_connect(NAMESPACE, "sid-b") |
| 708 | |
| 709 | aggregated = await manager.route_event_all( |
| 710 | NAMESPACE, |
| 711 | "filter_event", |
| 712 | {"excludeHandlers": [beta.identifier]}, |
| 713 | handler_id="test.manager", |
| 714 | ) |
| 715 | |
| 716 | assert aggregated |
| 717 | for entry in aggregated: |
| 718 | assert entry["correlationId"] |
| 719 | assert entry["results"] |
| 720 | assert all(result["handlerId"] == alpha.identifier for result in entry["results"]) |
| 721 | |
| 722 | |
| 723 | @pytest.mark.asyncio |
| 724 | async def test_route_event_preserves_correlation_id(): |
| 725 | socketio = FakeSocketIOServer() |
| 726 | manager = WsManager(socketio, threading.RLock()) |
| 727 | |
| 728 | results = [] |
| 729 | handler = DummyHandler(socketio, threading.RLock(), results) |
| 730 | manager.register_handlers({NAMESPACE: [handler]}) |
| 731 | await manager.handle_connect(NAMESPACE, "sid-correlation") |
| 732 | |
| 733 | response = await manager.route_event( |
| 734 | NAMESPACE, |
| 735 | "dummy", |
| 736 | {"foo": "bar", "correlationId": "manual-correlation"}, |
| 737 | "sid-correlation", |
| 738 | ) |
| 739 | |
| 740 | assert response["correlationId"] == "manual-correlation" |
| 741 | result = response["results"][0] |
| 742 | assert result["correlationId"] == "manual-correlation" |
| 743 | |
| 744 | |
| 745 | @pytest.mark.asyncio |
| 746 | async def test_request_preserves_explicit_correlation_id(): |
| 747 | socketio = FakeSocketIOServer() |
| 748 | manager = WsManager(socketio, threading.RLock()) |
| 749 | |
| 750 | handler = DummyHandler(socketio, threading.RLock()) |
| 751 | manager.register_handlers({NAMESPACE: [handler]}) |
| 752 | await manager.handle_connect(NAMESPACE, "sid-request") |
| 753 | |
| 754 | response = await manager.request_for_sid( |
| 755 | namespace=NAMESPACE, |
| 756 | sid="sid-request", |
| 757 | event_type="dummy", |
| 758 | data={"payload": True, "correlationId": "req-correlation"}, |
| 759 | handler_id="tester", |
| 760 | ) |
| 761 | |
| 762 | assert response["correlationId"] == "req-correlation" |
| 763 | result = response["results"][0] |
| 764 | assert result["correlationId"] == "req-correlation" |
| 765 | |
| 766 | |
| 767 | @pytest.mark.asyncio |
| 768 | async def test_request_all_entries_include_correlation_id(): |
| 769 | socketio = FakeSocketIOServer() |
| 770 | manager = WsManager(socketio, threading.RLock()) |
| 771 | |
| 772 | handler = DummyHandler(socketio, threading.RLock()) |
| 773 | manager.register_handlers({NAMESPACE: [handler]}) |
| 774 | |
| 775 | await manager.handle_connect(NAMESPACE, "sid-1") |
| 776 | await manager.handle_connect(NAMESPACE, "sid-2") |
| 777 | |
| 778 | aggregated = await manager.route_event_all( |
| 779 | NAMESPACE, |
| 780 | "dummy", |
| 781 | {"value": 1, "correlationId": "agg-correlation"}, |
| 782 | ) |
| 783 | |
| 784 | assert aggregated |
| 785 | for entry in aggregated: |
| 786 | assert entry["correlationId"] == "agg-correlation" |
| 787 | assert entry["results"] |
| 788 | assert entry["results"][0]["correlationId"] == "agg-correlation" |
| 789 | |
| 790 | |
| 791 | def test_debug_logging_respects_runtime_flag(monkeypatch): |
| 792 | socketio = FakeSocketIOServer() |
| 793 | manager = WsManager(socketio, threading.RLock()) |
| 794 | |
| 795 | logs: list[str] = [] |
| 796 | |
| 797 | def capture(message: str) -> None: |
| 798 | logs.append(message) |
| 799 | |
| 800 | monkeypatch.setattr("helpers.print_style.PrintStyle.debug", staticmethod(capture)) |
| 801 | monkeypatch.setenv("A0_WS_DEBUG", "") |
| 802 | |
| 803 | manager._debug("should-not-log") # noqa: SLF001 |
| 804 | assert logs == [] |
| 805 | |
| 806 | monkeypatch.setenv("A0_WS_DEBUG", "1") |
| 807 | manager._debug("should-log") # noqa: SLF001 |
| 808 | assert logs == ["should-log"] |
| 809 | |
| 810 | |
| 811 | @pytest.mark.asyncio |
| 812 | async def test_diagnostic_event_emitted_for_inbound(): |
| 813 | socketio = FakeSocketIOServer() |
| 814 | manager = WsManager(socketio, threading.RLock()) |
| 815 | |
| 816 | results: list[dict[str, Any]] = [] |
| 817 | handler = DummyHandler(socketio, threading.RLock(), results) |
| 818 | manager.register_handlers({NAMESPACE: [handler]}) |
| 819 | |
| 820 | await manager.handle_connect(NAMESPACE, "observer") |
| 821 | assert manager.register_diagnostic_watcher(NAMESPACE, "observer") is True |
| 822 | await manager.handle_connect(NAMESPACE, "sid-client") |
| 823 | |
| 824 | await manager.route_event(NAMESPACE, "dummy", {"payload": "value"}, "sid-client") |
| 825 | |
| 826 | emitted_events = [call.args[0] for call in socketio.emit.await_args_list] |
| 827 | assert DIAGNOSTIC_EVENT in emitted_events |
| 828 | |
| 829 | |
| 830 | @pytest.mark.asyncio |
| 831 | async def test_lifecycle_events_broadcast(monkeypatch): |
| 832 | socketio = FakeSocketIOServer() |
| 833 | manager = WsManager(socketio, threading.RLock()) |
| 834 | |
| 835 | broadcast_mock = AsyncMock() |
| 836 | monkeypatch.setattr(manager, "broadcast", broadcast_mock) |
| 837 | |
| 838 | await manager.handle_connect(NAMESPACE, "sid-life") |
| 839 | await asyncio.sleep(0) |
| 840 | await manager.handle_disconnect(NAMESPACE, "sid-life") |
| 841 | await asyncio.sleep(0) |
| 842 | |
| 843 | events = [call.args[1] for call in broadcast_mock.await_args_list] |
| 844 | assert LIFECYCLE_CONNECT_EVENT in events |
| 845 | assert LIFECYCLE_DISCONNECT_EVENT in events |