main
py 845 lines 27.6 KB
Raw
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