chore: delete all legacy WebSocket system files
keyboardstaff committed
Mar 26, 2026 at 01:07 UTC
07b5e056e00ded562fba5b3def9e02128bd52021
18 files changed
-4949
helpers/websocket.py
deleted
-550
@@ -1,550 +0,0 @@
1
-from __future__ import annotations
2
-
3
-import re
4
-import threading
5
-from abc import ABC, abstractmethod
6
-from urllib.parse import urlparse
7
-from typing import Any, Iterable, Optional, TYPE_CHECKING
8
-
9
-import socketio
10
-
11
-if TYPE_CHECKING: # pragma: no cover - hints only
12
- from helpers.websocket_manager import WebSocketManager
13
-
14
-_EVENT_NAME_PATTERN = re.compile(r"^[a-z][a-z0-9_]*$")
15
-_RESERVED_EVENT_NAMES: set[str] = {
16
- "connect",
17
- "disconnect",
18
- "error",
19
- "ping",
20
- "pong",
21
- "connect_error",
22
- "reconnect",
23
- "reconnect_attempt",
24
- "reconnect_error",
25
- "reconnect_failed",
26
-}
27
-
28
-
29
-def _default_port_for_scheme(scheme: str) -> int | None:
30
- if scheme == "http":
31
- return 80
32
- if scheme == "https":
33
- return 443
34
- return None
35
-
36
-
37
-def normalize_origin(value: Any) -> str | None:
38
- """Normalize an Origin/Referer header value to scheme://host[:port]."""
39
- if not isinstance(value, str) or not value.strip():
40
- return None
41
- parsed = urlparse(value.strip())
42
- if not parsed.scheme or not parsed.hostname:
43
- return None
44
- origin = f"{parsed.scheme}://{parsed.hostname}"
45
- if parsed.port:
46
- origin += f":{parsed.port}"
47
- return origin
48
-
49
-
50
-def _parse_host_header(value: Any) -> tuple[str | None, int | None]:
51
- if not isinstance(value, str) or not value.strip():
52
- return None, None
53
- parsed = urlparse(f"http://{value.strip()}")
54
- return parsed.hostname, parsed.port
55
-
56
-
57
-def validate_ws_origin(environ: dict[str, Any]) -> tuple[bool, str | None]:
58
- """Validate the browser Origin during the Socket.IO handshake.
59
-
60
- This is the minimum baseline recommended by RFC 6455 (Origin considerations)
61
- and OWASP (CSWSH mitigation): reject cross-origin WebSocket handshakes when
62
- the server is intended for a specific web UI origin.
63
- """
64
-
65
- raw_origin = environ.get("HTTP_ORIGIN") or environ.get("HTTP_REFERER")
66
- origin = normalize_origin(raw_origin)
67
- if origin is None:
68
- return False, "missing_origin"
69
-
70
- origin_parsed = urlparse(origin)
71
- origin_host = origin_parsed.hostname.lower() if origin_parsed.hostname else None
72
- origin_port = origin_parsed.port or _default_port_for_scheme(origin_parsed.scheme)
73
- if origin_host is None or origin_port is None:
74
- return False, "invalid_origin"
75
-
76
- # Build candidate request host/port pairs. Prefer explicit Host header, fall back to
77
- # forwarded headers (reverse proxies) and finally SERVER_NAME.
78
- raw_host = environ.get("HTTP_HOST")
79
- req_host, req_port = _parse_host_header(raw_host)
80
- if not req_host:
81
- req_host = environ.get("SERVER_NAME")
82
-
83
- if req_port is None:
84
- server_port_raw = environ.get("SERVER_PORT")
85
- try:
86
- server_port = int(server_port_raw) if server_port_raw is not None else None
87
- except (TypeError, ValueError):
88
- server_port = None
89
- if server_port is not None and server_port > 0:
90
- req_port = server_port
91
-
92
- if req_host:
93
- req_host = req_host.lower()
94
- if req_port is None:
95
- req_port = origin_port
96
-
97
- forwarded_host_raw = environ.get("HTTP_X_FORWARDED_HOST")
98
- forwarded_host = None
99
- forwarded_port = None
100
- if isinstance(forwarded_host_raw, str) and forwarded_host_raw.strip():
101
- first = forwarded_host_raw.split(",")[0].strip()
102
- forwarded_host, forwarded_port = _parse_host_header(first)
103
- if forwarded_host:
104
- forwarded_host = forwarded_host.lower()
105
-
106
- forwarded_proto_raw = environ.get("HTTP_X_FORWARDED_PROTO")
107
- forwarded_scheme = None
108
- if isinstance(forwarded_proto_raw, str) and forwarded_proto_raw.strip():
109
- forwarded_scheme = forwarded_proto_raw.split(",")[0].strip().lower()
110
- forwarded_scheme = forwarded_scheme or origin_parsed.scheme
111
- forwarded_port = (
112
- forwarded_port
113
- if forwarded_port is not None
114
- else _default_port_for_scheme(forwarded_scheme) or origin_port
115
- )
116
-
117
- candidates: list[tuple[str, int]] = []
118
- if req_host:
119
- candidates.append((req_host, int(req_port)))
120
- if forwarded_host:
121
- candidates.append((forwarded_host, int(forwarded_port)))
122
-
123
- if not candidates:
124
- return False, "missing_host"
125
-
126
- for host, port in candidates:
127
- if origin_host == host and origin_port == port:
128
- return True, None
129
-
130
- # Preserve the original mismatch semantics for debugging.
131
- if origin_host not in {host for host, _ in candidates}:
132
- return False, "origin_host_mismatch"
133
- return False, "origin_port_mismatch"
134
-
135
-
136
-class SingletonInstantiationError(RuntimeError):
137
- """Raised when a WebSocketHandler subclass is instantiated directly.
138
-
139
- Handlers must be retrieved via ``get_instance`` to guarantee singleton
140
- semantics and consistent lifecycle behaviour.
141
- """
142
-
143
-
144
-class ConnectionNotFoundError(RuntimeError):
145
- """Raised when attempting to emit to a non-existent WebSocket connection."""
146
-
147
- def __init__(self, sid: str, *, namespace: str | None = None) -> None:
148
- self.sid = sid
149
- self.namespace = namespace
150
- if namespace:
151
- super().__init__(f"Connection not found: namespace={namespace} sid={sid}")
152
- else:
153
- super().__init__(f"Connection not found: {sid}")
154
-
155
-
156
-class WebSocketResult:
157
- """Helper wrapper for standardized handler results.
158
-
159
- Instances are converted to the canonical ``RequestResultItem`` shape by
160
- :class:`WebSocketManager`. Helper constructors enforce payload validation so
161
- handlers no longer need to hand‑craft dictionaries.
162
- """
163
-
164
- __slots__ = ("_ok", "_data", "_error", "_correlation_id", "_duration_ms")
165
-
166
- def __init__(
167
- self,
168
- ok: bool,
169
- data: dict[str, Any] | None = None,
170
- error: dict[str, Any] | None = None,
171
- correlation_id: str | None = None,
172
- duration_ms: float | None = None,
173
- ) -> None:
174
- if ok and error:
175
- raise ValueError("Cannot be both ok and have an error")
176
- if not ok and not error:
177
- raise ValueError("Must either be ok or have an error")
178
- if data is not None and not isinstance(data, dict):
179
- raise TypeError("Data payload must be a dictionary or None")
180
- if error is not None and not isinstance(error, dict):
181
- raise TypeError("Error payload must be a dictionary or None")
182
- if correlation_id is not None and not isinstance(correlation_id, str):
183
- raise TypeError("Correlation ID must be a string or None")
184
- if duration_ms is not None and not isinstance(duration_ms, (int, float)):
185
- raise TypeError("Duration must be a number or None")
186
-
187
- self._ok = bool(ok)
188
- self._data = dict(data) if data is not None else None
189
- self._error = dict(error) if error is not None else None
190
- self._correlation_id = correlation_id
191
- self._duration_ms = float(duration_ms) if duration_ms is not None else None
192
-
193
- @classmethod
194
- def ok(
195
- cls,
196
- data: dict[str, Any] | None = None,
197
- *,
198
- correlation_id: str | None = None,
199
- duration_ms: float | None = None,
200
- ) -> "WebSocketResult":
201
- if data is not None and not isinstance(data, dict):
202
- raise TypeError("WebSocketResult.ok data must be a dict or None")
203
- payload = dict(data) if data is not None else None
204
- return cls(
205
- ok=True,
206
- data=payload,
207
- correlation_id=correlation_id,
208
- duration_ms=duration_ms,
209
- )
210
-
211
- @classmethod
212
- def error(
213
- cls,
214
- *,
215
- code: str,
216
- message: str,
217
- details: Any | None = None,
218
- correlation_id: str | None = None,
219
- duration_ms: float | None = None,
220
- ) -> "WebSocketResult":
221
- if not isinstance(code, str) or not code.strip():
222
- raise ValueError("Error code must be a non-empty string")
223
- if not isinstance(message, str) or not message.strip():
224
- raise ValueError("Error message must be a non-empty string")
225
-
226
- error_payload: dict[str, Any] = {"code": code, "error": message}
227
- if details is not None:
228
- error_payload["details"] = details
229
- return cls(
230
- ok=False,
231
- error=error_payload,
232
- correlation_id=correlation_id,
233
- duration_ms=duration_ms,
234
- )
235
-
236
- def as_result(
237
- self,
238
- *,
239
- handler_id: str,
240
- fallback_correlation_id: str | None,
241
- duration_ms: float | None = None,
242
- ) -> dict[str, Any]:
243
- result: dict[str, Any] = {
244
- "handlerId": handler_id,
245
- "ok": self._ok,
246
- }
247
-
248
- effective_duration = (
249
- self._duration_ms if self._duration_ms is not None else duration_ms
250
- )
251
- if effective_duration is not None:
252
- result["durationMs"] = round(effective_duration, 4)
253
-
254
- correlation = (
255
- self._correlation_id
256
- if self._correlation_id is not None
257
- else fallback_correlation_id
258
- )
259
- if correlation is not None:
260
- result["correlationId"] = correlation
261
-
262
- if self._ok:
263
- result["data"] = dict(self._data) if self._data is not None else {}
264
- else:
265
- result["error"] = dict(self._error) if self._error is not None else {
266
- "code": "INTERNAL_ERROR",
267
- "error": "Internal server error",
268
- }
269
- return result
270
-
271
-
272
-class WebSocketHandler(ABC):
273
- """Base class for WebSocket event handlers.
274
-
275
- The interface mirrors :class:`helpers.api.ApiHandler` with declarative
276
- security configuration and lifecycle hooks. Handlers are namespace-wide:
277
- every inbound event for the bound namespace is dispatched to
278
- :meth:`process_event`, which decides whether and how to respond.
279
- """
280
-
281
- _instances: dict[type["WebSocketHandler"], "WebSocketHandler"] = {}
282
- _construction_tokens: dict[type["WebSocketHandler"], bool] = {}
283
- _singleton_lock = threading.RLock()
284
-
285
- def __init__(self, socketio: socketio.AsyncServer, lock: threading.RLock) -> None:
286
- """Create a handler bound to the shared Socket.IO instance."""
287
-
288
- cls = self.__class__
289
- if not WebSocketHandler._construction_tokens.get(cls):
290
- raise SingletonInstantiationError(
291
- f"{cls.__name__} must be instantiated via {cls.__name__}.get_instance()"
292
- )
293
-
294
- self.socketio: socketio.AsyncServer = socketio
295
- self.lock: threading.RLock = lock
296
- self._manager: Optional[WebSocketManager] = None
297
- self._namespace: str | None = None
298
-
299
- @classmethod
300
- def get_instance(
301
- cls,
302
- socketio: socketio.AsyncServer | None = None,
303
- lock: threading.RLock | None = None,
304
- *args: Any,
305
- **kwargs: Any,
306
- ) -> "WebSocketHandler":
307
- """Return the singleton instance for ``cls``.
308
-
309
- Args:
310
- socketio: Shared AsyncServer instance (required on first call).
311
- lock: Shared threading lock (required on first call).
312
- *args: Optional subclass-specific constructor args.
313
- **kwargs: Optional subclass-specific constructor kwargs.
314
- """
315
-
316
- if cls is WebSocketHandler:
317
- raise TypeError("WebSocketHandler must be subclassed before use")
318
-
319
- with WebSocketHandler._singleton_lock:
320
- instance = WebSocketHandler._instances.get(cls)
321
- if instance is not None:
322
- return instance
323
-
324
- if socketio is None or lock is None:
325
- raise ValueError(
326
- f"{cls.__name__}.get_instance() requires socketio and lock on first call"
327
- )
328
-
329
- WebSocketHandler._construction_tokens[cls] = True
330
- try:
331
- instance = cls(socketio, lock, *args, **kwargs)
332
- finally:
333
- WebSocketHandler._construction_tokens.pop(cls, None)
334
-
335
- WebSocketHandler._instances[cls] = instance
336
- return instance
337
-
338
- @classmethod
339
- def _reset_instance_for_testing(cls) -> None:
340
- """Reset the cached singleton instance (testing helper)."""
341
-
342
- with WebSocketHandler._singleton_lock:
343
- WebSocketHandler._instances.pop(cls, None)
344
- WebSocketHandler._construction_tokens.pop(cls, None)
345
-
346
- @classmethod
347
- def validate_event_type(cls, event_type: str) -> str:
348
- """Validate a runtime event name before dispatch."""
349
-
350
- if not isinstance(event_type, str):
351
- raise TypeError("Event type must be a string")
352
- if not _EVENT_NAME_PATTERN.fullmatch(event_type):
353
- raise ValueError(
354
- f"Invalid event type '{event_type}' – must match lowercase_snake_case"
355
- )
356
- if event_type in _RESERVED_EVENT_NAMES:
357
- raise ValueError(
358
- f"Event type '{event_type}' is reserved by Socket.IO and cannot be used"
359
- )
360
- return event_type
361
-
362
- @classmethod
363
- def requires_auth(cls) -> bool:
364
- """Return whether an authenticated Flask session is required."""
365
-
366
- return True
367
-
368
- @classmethod
369
- def requires_csrf(cls) -> bool:
370
- """Return whether CSRF validation is required for the handler.
371
-
372
- This mirrors ApiHandler.requires_csrf(): by default, authenticated
373
- WebSocket handlers also require CSRF validation during the Socket.IO
374
- connect step.
375
- """
376
-
377
- return cls.requires_auth()
378
-
379
- async def on_connect(self, sid: str) -> None:
380
- """Lifecycle hook invoked when a client connects."""
381
-
382
- return None
383
-
384
- async def on_disconnect(self, sid: str) -> None:
385
- """Lifecycle hook invoked when a client disconnects."""
386
-
387
- return None
388
-
389
- @abstractmethod
390
- async def process_event(
391
- self,
392
- event_type: str,
393
- data: dict[str, Any],
394
- sid: str,
395
- ) -> dict[str, Any] | WebSocketResult | None:
396
- """Process an incoming event dispatched to the handler.
397
-
398
- Returning ``None`` indicates fire-and-forget semantics. Returning a
399
- dictionary includes the payload in the Socket.IO acknowledgement.
400
- """
401
-
402
- def bind_manager(self, manager: WebSocketManager, *, namespace: str) -> None:
403
- """Associate this handler instance with the shared WebSocket manager."""
404
-
405
- self._manager = manager
406
- self._namespace = namespace
407
-
408
- @property
409
- def namespace(self) -> str:
410
- if not self._namespace:
411
- raise RuntimeError("WebSocketHandler is missing namespace binding")
412
- return self._namespace
413
-
414
- @property
415
- def manager(self) -> WebSocketManager:
416
- """Return the bound WebSocket manager.
417
-
418
- Raises:
419
- RuntimeError: If the handler has not been registered yet.
420
- """
421
-
422
- if not self._manager:
423
- raise RuntimeError("WebSocketHandler is not registered with a manager")
424
- return self._manager
425
-
426
- @property
427
- def identifier(self) -> str:
428
- """Return a stable identifier used in aggregated responses."""
429
-
430
- return f"{self.__class__.__module__}.{self.__class__.__name__}"
431
-
432
- async def emit_to(
433
- self,
434
- sid: str,
435
- event_type: str,
436
- data: dict[str, Any],
437
- *,
438
- correlation_id: str | None = None,
439
- ) -> None:
440
- """Emit an event to a specific connection or buffer it if offline."""
441
- await self.manager.emit_to(
442
- self.namespace,
443
- sid,
444
- event_type,
445
- data,
446
- handler_id=self.identifier,
447
- correlation_id=correlation_id,
448
- )
449
-
450
- async def broadcast(
451
- self,
452
- event_type: str,
453
- data: dict[str, Any],
454
- *,
455
- exclude_sids: str | Iterable[str] | None = None,
456
- correlation_id: str | None = None,
457
- ) -> None:
458
- """Broadcast an event to all connections, optionally excluding one."""
459
- await self.manager.broadcast(
460
- self.namespace,
461
- event_type,
462
- data,
463
- exclude_sids=exclude_sids,
464
- handler_id=self.identifier,
465
- correlation_id=correlation_id,
466
- )
467
-
468
- # ------------------------------------------------------------------
469
- # Convenience wrappers for standardized result helpers
470
- # ------------------------------------------------------------------
471
-
472
- @staticmethod
473
- def result_ok(
474
- data: dict[str, Any] | None = None,
475
- *,
476
- correlation_id: str | None = None,
477
- duration_ms: float | None = None,
478
- ) -> WebSocketResult:
479
- """Return a standardized success result."""
480
-
481
- return WebSocketResult.ok(
482
- data=data,
483
- correlation_id=correlation_id,
484
- duration_ms=duration_ms,
485
- )
486
-
487
- @staticmethod
488
- def result_error(
489
- *,
490
- code: str,
491
- message: str,
492
- details: Any | None = None,
493
- correlation_id: str | None = None,
494
- duration_ms: float | None = None,
495
- ) -> WebSocketResult:
496
- """Return a standardized error result."""
497
-
498
- return WebSocketResult.error(
499
- code=code,
500
- message=message,
501
- details=details,
502
- correlation_id=correlation_id,
503
- duration_ms=duration_ms,
504
- )
505
-
506
- async def request(
507
- self,
508
- sid: str,
509
- event_type: str,
510
- data: dict[str, Any],
511
- *,
512
- timeout_ms: int = 0,
513
- include_handlers: Iterable[str] | None = None,
514
- ) -> dict[str, Any]:
515
- """Send a request-response event to a specific connection and aggregate results.
516
-
517
- Returns a payload shaped as ``{"correlationId": str, "results": RequestResultItem[]}``.
518
- """
519
-
520
- return await self.manager.request_for_sid(
521
- namespace=self.namespace,
522
- sid=sid,
523
- event_type=event_type,
524
- data=data,
525
- timeout_ms=timeout_ms,
526
- handler_id=self.identifier,
527
- include_handlers=set(include_handlers) if include_handlers else None,
528
- )
529
-
530
- async def request_all(
531
- self,
532
- event_type: str,
533
- data: dict[str, Any],
534
- *,
535
- timeout_ms: int = 0,
536
- exclude_handlers: Iterable[str] | None = None,
537
- ) -> list[dict[str, Any]]:
538
- """Fan a request out to every active connection and aggregate responses.
539
-
540
- Each entry in the returned list is ``{"sid": str, "correlationId": str, "results": RequestResultItem[]}``.
541
- """
542
-
543
- return await self.manager.route_event_all(
544
- self.namespace,
545
- event_type=event_type,
546
- data=data,
547
- timeout_ms=timeout_ms,
548
- exclude_handlers=set(exclude_handlers) if exclude_handlers else None,
549
- handler_id=self.identifier,
550
- )
helpers/websocket_manager.py
deleted
-1189
@@ -1,1189 +0,0 @@
1
-from __future__ import annotations
2
-
3
-import asyncio, os
4
-import time
5
-import threading
6
-from collections import defaultdict, deque
7
-from dataclasses import dataclass, field
8
-from datetime import datetime, timedelta, timezone
9
-from typing import Any, Callable, Deque, Dict, Iterable, List, Optional, Set
10
-
11
-import socketio
12
-import uuid
13
-
14
-from helpers.defer import DeferredTask
15
-from helpers.print_style import PrintStyle
16
-from helpers import runtime
17
-from helpers.websocket import ConnectionNotFoundError, WebSocketHandler, WebSocketResult
18
-from helpers.state_monitor import _ws_debug_enabled
19
-
20
-BUFFER_MAX_SIZE = 100
21
-BUFFER_TTL = timedelta(hours=1)
22
-_shared_websocket_manager: WebSocketManager | None = None
23
-
24
-
25
-async def send_data(
26
- event_name: str,
27
- data: dict[str, Any],
28
- endpoint_name: str = "/webui",
29
- connection_id: str | None = None,
30
-) -> None:
31
- manager = get_shared_websocket_manager()
32
- print(f"Sending data to {endpoint_name}/{event_name} with data {data}")
33
- await manager.send_data(endpoint_name, event_name, data, connection_id)
34
-
35
-
36
-def _utcnow() -> datetime:
37
- return datetime.now(timezone.utc)
38
-
39
-
40
-def set_shared_websocket_manager(manager: "WebSocketManager") -> None:
41
- global _shared_websocket_manager
42
- _shared_websocket_manager = manager
43
-
44
-
45
-def get_shared_websocket_manager() -> "WebSocketManager":
46
- manager = _shared_websocket_manager
47
- if manager is None:
48
- raise RuntimeError("Shared WebSocketManager has not been initialized")
49
- return manager
50
-
51
-
52
-@dataclass
53
-class BufferedEvent:
54
- event_type: str
55
- data: dict[str, Any]
56
- handler_id: str | None = None
57
- correlation_id: str | None = None
58
- timestamp: datetime = field(default_factory=_utcnow)
59
-
60
-
61
-@dataclass
62
-class ConnectionInfo:
63
- namespace: str
64
- sid: str
65
- connected_at: datetime = field(default_factory=_utcnow)
66
- last_activity: datetime = field(default_factory=_utcnow)
67
-
68
-
69
-ConnectionIdentity = tuple[str, str] # (namespace, sid)
70
-
71
-
72
-@dataclass
73
-class _HandlerExecution:
74
- handler: WebSocketHandler
75
- value: Any
76
- duration_ms: float | None
77
-
78
-
79
-DIAGNOSTIC_EVENT = "ws_dev_console_event"
80
-LIFECYCLE_CONNECT_EVENT = "ws_lifecycle_connect"
81
-LIFECYCLE_DISCONNECT_EVENT = "ws_lifecycle_disconnect"
82
-
83
-
84
-class WebSocketManager:
85
- def __init__(self, socketio: socketio.AsyncServer, lock) -> None:
86
- self.socketio = socketio
87
- self.lock = lock
88
- self.handlers: defaultdict[str, List[WebSocketHandler]] = defaultdict(list)
89
- self.connections: Dict[ConnectionIdentity, ConnectionInfo] = {}
90
- self.buffers: defaultdict[ConnectionIdentity, Deque[BufferedEvent]] = (
91
- defaultdict(deque)
92
- )
93
- self._known_sids: Set[ConnectionIdentity] = set()
94
- self._identifier: str = f"{self.__class__.__module__}.{self.__class__.__name__}"
95
- # Session tracking (single-user default)
96
- self.user_to_sids: defaultdict[str, Set[ConnectionIdentity]] = defaultdict(set)
97
- self.sid_to_user: Dict[ConnectionIdentity, str | None] = {}
98
- self._ALL_USERS_BUCKET = "allUsers"
99
- self._server_restart_enabled: bool = False
100
- self._diagnostic_watchers: Set[ConnectionIdentity] = set()
101
- self._diagnostics_enabled: bool = runtime.is_development()
102
- self._dispatcher_loop: asyncio.AbstractEventLoop | None = None
103
- self._handler_worker: DeferredTask | None = None
104
-
105
- # Internal: development-only debug logging to avoid noise in production
106
- def _debug(self, message: str) -> None:
107
- value = os.getenv("A0_WS_DEBUG", "").strip().lower()
108
- if value in {"1", "true", "yes", "on"}:
109
- PrintStyle.debug(message)
110
-
111
- def _ensure_dispatcher_loop(self) -> None:
112
- if self._dispatcher_loop is None:
113
- try:
114
- self._dispatcher_loop = asyncio.get_running_loop()
115
- except RuntimeError:
116
- return
117
-
118
- def _get_handler_worker(self) -> DeferredTask:
119
- if self._handler_worker is None:
120
- self._handler_worker = DeferredTask(thread_name="WebSocketHandlers")
121
- return self._handler_worker
122
-
123
- async def _run_on_dispatcher_loop(self, coro: Any) -> Any:
124
- self._ensure_dispatcher_loop()
125
- dispatcher_loop = self._dispatcher_loop
126
- if dispatcher_loop is None:
127
- return await coro
128
- if dispatcher_loop.is_closed():
129
- try:
130
- coro.close()
131
- except Exception: # pragma: no cover - best-effort cleanup
132
- pass
133
- raise RuntimeError("Dispatcher event loop is closed")
134
-
135
- try:
136
- running_loop = asyncio.get_running_loop()
137
- except RuntimeError:
138
- running_loop = None
139
-
140
- if running_loop is dispatcher_loop:
141
- return await coro
142
-
143
- future = asyncio.run_coroutine_threadsafe(coro, dispatcher_loop)
144
- return await asyncio.wrap_future(future)
145
-
146
- def _diagnostics_active(self) -> bool:
147
- if not self._diagnostics_enabled:
148
- return False
149
- with self.lock:
150
- return bool(self._diagnostic_watchers)
151
-
152
- def _copy_diagnostic_watchers(self) -> list[ConnectionIdentity]:
153
- with self.lock:
154
- return list(self._diagnostic_watchers)
155
-
156
- def register_diagnostic_watcher(self, namespace: str, sid: str) -> bool:
157
- if not self._diagnostics_enabled:
158
- return False
159
- identity: ConnectionIdentity = (namespace, sid)
160
- with self.lock:
161
- if identity not in self.connections:
162
- return False
163
- self._diagnostic_watchers.add(identity)
164
- return True
165
-
166
- def unregister_diagnostic_watcher(self, namespace: str, sid: str) -> None:
167
- identity: ConnectionIdentity = (namespace, sid)
168
- with self.lock:
169
- self._diagnostic_watchers.discard(identity)
170
-
171
- def _timestamp(self) -> str:
172
- return _utcnow().isoformat(timespec="milliseconds").replace("+00:00", "Z")
173
-
174
- def _summarize_payload(self, payload: dict[str, Any] | None) -> dict[str, Any]:
175
- if not isinstance(payload, dict):
176
- return {}
177
- summary: dict[str, Any] = {}
178
- for key in list(payload.keys())[:5]:
179
- value = payload[key]
180
- if isinstance(value, (str, int, float, bool)) or value is None:
181
- preview = value
182
- elif isinstance(value, dict):
183
- preview = f"dict({len(value)})"
184
- elif isinstance(value, list):
185
- preview = f"list({len(value)})"
186
- else:
187
- preview = value.__class__.__name__
188
- summary[key] = preview
189
- summary["__sizeBytes__"] = len(str(payload).encode("utf-8"))
190
- return summary
191
-
192
- def _summarize_results(self, results: List[dict[str, Any]]) -> dict[str, Any]:
193
- summary = {"ok": 0, "error": 0, "handlers": []}
194
- for result in results:
195
- handler_id = result.get("handlerId")
196
- ok = bool(result.get("ok"))
197
- if ok:
198
- summary["ok"] += 1
199
- else:
200
- summary["error"] += 1
201
- summary["handlers"].append(
202
- {
203
- "handlerId": handler_id,
204
- "ok": ok,
205
- "errorCode": (result.get("error") or {}).get("code"),
206
- "durationMs": result.get("durationMs"),
207
- }
208
- )
209
- summary["handlerCount"] = len(summary["handlers"])
210
- return summary
211
-
212
- async def _publish_diagnostic_event(
213
- self, payload: dict[str, Any] | Callable[[], dict[str, Any]]
214
- ) -> None:
215
- if not self._diagnostics_enabled:
216
- return
217
- watchers = self._copy_diagnostic_watchers()
218
- if not watchers:
219
- return
220
- effective_payload = payload() if callable(payload) else payload
221
- if (
222
- isinstance(effective_payload, dict)
223
- and "sourceNamespace" not in effective_payload
224
- ):
225
- origin = effective_payload.get("namespace")
226
- if isinstance(origin, str) and origin.strip():
227
- effective_payload = {
228
- **effective_payload,
229
- "sourceNamespace": origin.strip(),
230
- }
231
-
232
- async def _emit_to_watcher(identity: ConnectionIdentity) -> None:
233
- namespace, sid = identity
234
- try:
235
- await self.emit_to(
236
- namespace,
237
- sid,
238
- DIAGNOSTIC_EVENT,
239
- effective_payload,
240
- handler_id=self._identifier,
241
- diagnostic=True,
242
- )
243
- except ConnectionNotFoundError:
244
- self.unregister_diagnostic_watcher(namespace, sid)
245
-
246
- await asyncio.gather(*(_emit_to_watcher(identity) for identity in watchers))
247
-
248
- def _schedule_lifecycle_broadcast(
249
- self, namespace: str, event_type: str, payload: dict[str, Any]
250
- ) -> None:
251
- async def _broadcast() -> None:
252
- try:
253
- await self.broadcast(
254
- namespace,
255
- event_type,
256
- payload,
257
- diagnostic=True,
258
- )
259
- except Exception as exc: # pragma: no cover - diagnostic
260
- self._debug(f"Failed to broadcast lifecycle event {event_type}: {exc}")
261
-
262
- asyncio.create_task(_broadcast())
263
-
264
- def _normalize_handler_filter(self, value: Any, field_name: str) -> Set[str] | None:
265
- if value is None:
266
- return None
267
- if isinstance(value, str):
268
- return {value}
269
- try:
270
- iterator = iter(value)
271
- except TypeError as exc: # pragma: no cover - defensive
272
- raise ValueError(
273
- f"{field_name} must be an array of handler identifiers"
274
- ) from exc
275
-
276
- normalized: Set[str] = set()
277
- for item in iterator:
278
- if not isinstance(item, str):
279
- raise ValueError(
280
- f"{field_name} values must be handler identifier strings"
281
- )
282
- normalized.add(item)
283
- return normalized
284
-
285
- def _normalize_sid_filter(self, value: str | Iterable[str] | None) -> Set[str]:
286
- if value is None:
287
- return set()
288
- if isinstance(value, str):
289
- return {value}
290
- normalized: Set[str] = set()
291
- for item in value:
292
- normalized.add(str(item))
293
- return normalized
294
-
295
- def _select_handlers(
296
- self,
297
- namespace: str,
298
- *,
299
- include: Set[str] | None,
300
- exclude: Set[str] | None,
301
- ) -> tuple[list[WebSocketHandler], Set[str]]:
302
- registered = self.handlers.get(namespace, [])
303
- available_ids = {handler.identifier for handler in registered}
304
-
305
- if include is not None:
306
- unknown = include - available_ids
307
- if unknown:
308
- raise ValueError(
309
- f"Unknown handler(s) in includeHandlers for namespace '{namespace}': "
310
- f"{', '.join(sorted(unknown))}"
311
- )
312
- if exclude is not None:
313
- unknown = exclude - available_ids
314
- if unknown:
315
- raise ValueError(
316
- f"Unknown handler(s) in excludeHandlers for namespace '{namespace}': "
317
- f"{', '.join(sorted(unknown))}"
318
- )
319
-
320
- selected: list[WebSocketHandler] = []
321
- for handler in registered:
322
- ident = handler.identifier
323
- if include is not None and ident not in include:
324
- continue
325
- if exclude is not None and ident in exclude:
326
- continue
327
- selected.append(handler)
328
-
329
- return selected, available_ids
330
-
331
- def _resolve_correlation_id(self, payload: dict[str, Any]) -> str:
332
- value = payload.get("correlationId")
333
- if isinstance(value, str) and value.strip():
334
- correlation_id = value.strip()
335
- else:
336
- correlation_id = uuid.uuid4().hex
337
- payload["correlationId"] = correlation_id
338
- return correlation_id
339
-
340
- def register_handlers(
341
- self, handlers_by_namespace: dict[str, Iterable[WebSocketHandler]]
342
- ) -> None:
343
- for namespace, handlers in handlers_by_namespace.items():
344
- for handler in handlers:
345
- handler.bind_manager(self, namespace=namespace)
346
- if _ws_debug_enabled():
347
- PrintStyle.info(
348
- "Registered WebSocket handler %s namespace=%s"
349
- % (handler.identifier, namespace)
350
- )
351
- existing = self.handlers.get(namespace, [])
352
- if handler in existing:
353
- PrintStyle.warning(
354
- f"Duplicate handler registration for namespace '{namespace}'"
355
- )
356
- self.handlers[namespace].append(handler)
357
- self._debug(
358
- f"Registered handler {handler.identifier} namespace={namespace}"
359
- )
360
-
361
- def iter_event_types(self, namespace: str) -> Iterable[str]:
362
- return []
363
-
364
- def iter_namespaces(self) -> list[str]:
365
- return list(self.handlers.keys())
366
-
367
- async def _invoke_handler(
368
- self,
369
- handler: WebSocketHandler,
370
- event_type: str,
371
- payload: dict[str, Any],
372
- sid: str,
373
- ) -> _HandlerExecution:
374
- instrument = self._diagnostics_active()
375
- start = time.perf_counter() if instrument else None
376
- try:
377
- value = await self._get_handler_worker().execute_inside(
378
- handler.process_event, event_type, payload, sid
379
- )
380
- except Exception as exc: # pragma: no cover - handled by caller
381
- duration_ms = (
382
- (time.perf_counter() - start) * 1000 if start is not None else None
383
- )
384
- return _HandlerExecution(handler, exc, duration_ms)
385
- duration_ms = (
386
- (time.perf_counter() - start) * 1000 if start is not None else None
387
- )
388
- return _HandlerExecution(handler, value, duration_ms)
389
-
390
- async def handle_connect(
391
- self, namespace: str, sid: str, user_id: str | None = None
392
- ) -> None:
393
- self._ensure_dispatcher_loop()
394
- user_bucket = user_id or "single_user"
395
- identity: ConnectionIdentity = (namespace, sid)
396
- with self.lock:
397
- self.connections[identity] = ConnectionInfo(namespace=namespace, sid=sid)
398
- self._known_sids.add(identity)
399
- self.sid_to_user[identity] = user_bucket
400
- self.user_to_sids[self._ALL_USERS_BUCKET].add(identity)
401
- self.user_to_sids[user_bucket].add(identity)
402
- connection_count = sum(
403
- 1 for conn_identity in self.connections if conn_identity[0] == namespace
404
- )
405
- if _ws_debug_enabled():
406
- PrintStyle.info(f"WebSocket connected: namespace={namespace} sid={sid}")
407
- await self._run_lifecycle(namespace, lambda h: h.on_connect(sid))
408
- await self._flush_buffer(identity)
409
- if self._server_restart_enabled:
410
- await self.emit_to(
411
- namespace,
412
- sid,
413
- "server_restart",
414
- {
415
- "emittedAt": _utcnow()
416
- .isoformat(timespec="milliseconds")
417
- .replace("+00:00", "Z"),
418
- "runtimeId": runtime.get_runtime_id(),
419
- },
420
- handler_id=self._identifier,
421
- )
422
- if _ws_debug_enabled():
423
- PrintStyle.info(
424
- f"server_restart broadcast emitted to namespace={namespace} sid={sid}"
425
- )
426
- lifecycle_payload = {
427
- "namespace": namespace,
428
- "sid": sid,
429
- "connectionCount": connection_count,
430
- "timestamp": self._timestamp(),
431
- }
432
- await self._publish_diagnostic_event(
433
- {
434
- "kind": "lifecycle",
435
- "event": "connect",
436
- **lifecycle_payload,
437
- }
438
- )
439
- self._schedule_lifecycle_broadcast(
440
- namespace, LIFECYCLE_CONNECT_EVENT, lifecycle_payload
441
- )
442
-
443
- async def handle_disconnect(self, namespace: str, sid: str) -> None:
444
- self._ensure_dispatcher_loop()
445
- identity: ConnectionIdentity = (namespace, sid)
446
- with self.lock:
447
- self.connections.pop(identity, None)
448
- # session tracking cleanup
449
- user_bucket = self.sid_to_user.pop(identity, None)
450
- if self._ALL_USERS_BUCKET in self.user_to_sids:
451
- self.user_to_sids[self._ALL_USERS_BUCKET].discard(identity)
452
- if not self.user_to_sids[self._ALL_USERS_BUCKET]:
453
- self.user_to_sids.pop(self._ALL_USERS_BUCKET, None)
454
- if user_bucket and user_bucket in self.user_to_sids:
455
- self.user_to_sids[user_bucket].discard(identity)
456
- if not self.user_to_sids[user_bucket]:
457
- self.user_to_sids.pop(user_bucket, None)
458
- connection_count = sum(
459
- 1 for conn_identity in self.connections if conn_identity[0] == namespace
460
- )
461
- self.unregister_diagnostic_watcher(namespace, sid)
462
- PrintStyle.info(f"WebSocket disconnected: namespace={namespace} sid={sid}")
463
- await self._run_lifecycle(namespace, lambda h: h.on_disconnect(sid))
464
- lifecycle_payload = {
465
- "namespace": namespace,
466
- "sid": sid,
467
- "connectionCount": connection_count,
468
- "timestamp": self._timestamp(),
469
- }
470
- await self._publish_diagnostic_event(
471
- {
472
- "kind": "lifecycle",
473
- "event": "disconnect",
474
- **lifecycle_payload,
475
- }
476
- )
477
- self._schedule_lifecycle_broadcast(
478
- namespace, LIFECYCLE_DISCONNECT_EVENT, lifecycle_payload
479
- )
480
-
481
- async def route_event(
482
- self,
483
- namespace: str,
484
- event_type: str,
485
- data: dict[str, Any],
486
- sid: str,
487
- ack: Optional[Callable[[Any], None]] = None,
488
- *,
489
- include_handlers: Set[str] | None = None,
490
- exclude_handlers: Set[str] | None = None,
491
- allow_exclude: bool = False,
492
- handler_id: str | None = None,
493
- ) -> dict[str, Any]:
494
- self._ensure_dispatcher_loop()
495
- incoming = dict(data or {})
496
- correlation_id = self._resolve_correlation_id(incoming)
497
- self._debug(
498
- f"Routing event namespace={namespace} '{event_type}' sid={sid} correlation={correlation_id}"
499
- )
500
-
501
- include_meta_raw = incoming.pop("includeHandlers", None)
502
- exclude_meta_raw = incoming.pop("excludeHandlers", None)
503
-
504
- if "data" in incoming and isinstance(incoming.get("data"), dict):
505
- handler_payload = dict(incoming.get("data") or {})
506
- if "excludeSids" in incoming:
507
- handler_payload["excludeSids"] = incoming.get("excludeSids")
508
- else:
509
- handler_payload = dict(incoming)
510
-
511
- handler_payload["correlationId"] = correlation_id
512
-
513
- try:
514
- include_meta = self._normalize_handler_filter(
515
- include_meta_raw, "includeHandlers"
516
- )
517
- except ValueError as exc:
518
- error = self._build_error_result(
519
- handler_id=handler_id or self._identifier,
520
- code="INVALID_FILTER",
521
- message=str(exc),
522
- correlation_id=correlation_id,
523
- )
524
- if ack:
525
- ack({"correlationId": correlation_id, "results": [error]})
526
- return {"correlationId": correlation_id, "results": [error]}
527
-
528
- try:
529
- exclude_meta = self._normalize_handler_filter(
530
- exclude_meta_raw, "excludeHandlers"
531
- )
532
- except ValueError as exc:
533
- error = self._build_error_result(
534
- handler_id=handler_id or self._identifier,
535
- code="INVALID_FILTER",
536
- message=str(exc),
537
- correlation_id=correlation_id,
538
- )
539
- payload_error = {"correlationId": correlation_id, "results": [error]}
540
- if ack:
541
- ack(payload_error)
542
- return payload_error
543
-
544
- if exclude_meta_raw is not None and not allow_exclude:
545
- error = self._build_error_result(
546
- handler_id=handler_id or self._identifier,
547
- code="INVALID_FILTER",
548
- message="excludeHandlers is not supported for this operation",
549
- correlation_id=correlation_id,
550
- )
551
- if ack:
552
- ack({"correlationId": correlation_id, "results": [error]})
553
- return {"correlationId": correlation_id, "results": [error]}
554
-
555
- if include_handlers is not None and include_meta is not None:
556
- if include_handlers != include_meta:
557
- error = self._build_error_result(
558
- handler_id=handler_id or self._identifier,
559
- code="INVALID_FILTER",
560
- message="Conflicting includeHandlers filters supplied",
561
- correlation_id=correlation_id,
562
- )
563
- if ack:
564
- ack({"correlationId": correlation_id, "results": [error]})
565
- return {"correlationId": correlation_id, "results": [error]}
566
-
567
- if allow_exclude and exclude_handlers is not None and exclude_meta is not None:
568
- if exclude_handlers != exclude_meta:
569
- error = self._build_error_result(
570
- handler_id=handler_id or self._identifier,
571
- code="INVALID_FILTER",
572
- message="Conflicting excludeHandlers filters supplied",
573
- correlation_id=correlation_id,
574
- )
575
- if ack:
576
- ack({"correlationId": correlation_id, "results": [error]})
577
- return {"correlationId": correlation_id, "results": [error]}
578
-
579
- include = include_handlers or include_meta
580
- exclude = exclude_handlers or (exclude_meta if allow_exclude else None)
581
-
582
- try:
583
- WebSocketHandler.validate_event_type(event_type)
584
- except (TypeError, ValueError) as exc:
585
- error = self._build_error_result(
586
- handler_id=handler_id or self._identifier,
587
- code="INVALID_EVENT",
588
- message=str(exc),
589
- correlation_id=correlation_id,
590
- )
591
- if ack:
592
- ack({"correlationId": correlation_id, "results": [error]})
593
- return {"correlationId": correlation_id, "results": [error]}
594
-
595
- registered = self.handlers.get(namespace, [])
596
- if not registered:
597
- PrintStyle.warning(f"No handlers registered for namespace '{namespace}'")
598
- error = self._build_error_result(
599
- handler_id=handler_id or self._identifier,
600
- code="NO_HANDLERS",
601
- message=f"No handler for namespace '{namespace}'",
602
- correlation_id=correlation_id,
603
- )
604
- if ack:
605
- ack({"correlationId": correlation_id, "results": [error]})
606
- return {"correlationId": correlation_id, "results": [error]}
607
-
608
- try:
609
- selected_handlers, _ = self._select_handlers(
610
- namespace, include=include, exclude=exclude
611
- )
612
- except ValueError as exc:
613
- error = self._build_error_result(
614
- handler_id=handler_id or self._identifier,
615
- code="INVALID_FILTER",
616
- message=str(exc),
617
- correlation_id=correlation_id,
618
- )
619
- if ack:
620
- ack({"correlationId": correlation_id, "results": [error]})
621
- return {"correlationId": correlation_id, "results": [error]}
622
-
623
- if not selected_handlers:
624
- error = self._build_error_result(
625
- handler_id=handler_id or self._identifier,
626
- code="NO_HANDLERS",
627
- message=f"No handler for '{event_type}' after applying filters",
628
- correlation_id=correlation_id,
629
- )
630
- if ack:
631
- ack({"correlationId": correlation_id, "results": [error]})
632
- return {"correlationId": correlation_id, "results": [error]}
633
-
634
- with self.lock:
635
- info = self.connections.get((namespace, sid))
636
- if info:
637
- info.last_activity = _utcnow()
638
-
639
- executions = await asyncio.gather(
640
- *[
641
- self._invoke_handler(handler, event_type, dict(handler_payload), sid)
642
- for handler in selected_handlers
643
- ]
644
- )
645
-
646
- results: List[dict[str, Any]] = []
647
- for execution in executions:
648
- handler = execution.handler
649
- value = execution.value
650
- duration_ms = execution.duration_ms
651
-
652
- if isinstance(value, Exception): # pragma: no cover - defensive logging
653
- PrintStyle.error(
654
- f"Error in handler {handler.identifier} for '{event_type}' (correlation {correlation_id}): {value}"
655
- )
656
- results.append(
657
- self._build_error_result(
658
- handler_id=handler.identifier,
659
- code="HANDLER_ERROR",
660
- message="Internal server error",
661
- details=str(value),
662
- correlation_id=correlation_id,
663
- duration_ms=duration_ms,
664
- )
665
- )
666
- continue
667
-
668
- if isinstance(value, WebSocketResult):
669
- results.append(
670
- value.as_result(
671
- handler_id=handler.identifier,
672
- fallback_correlation_id=correlation_id,
673
- duration_ms=duration_ms,
674
- )
675
- )
676
- continue
677
-
678
- if value is None:
679
- helper_result = WebSocketResult(ok=True)
680
- elif isinstance(value, dict):
681
- helper_result = WebSocketResult(ok=True, data=value)
682
- else:
683
- helper_result = WebSocketResult(ok=True, data={"result": value})
684
-
685
- results.append(
686
- helper_result.as_result(
687
- handler_id=handler.identifier,
688
- fallback_correlation_id=correlation_id,
689
- duration_ms=duration_ms,
690
- )
691
- )
692
-
693
- await self._publish_diagnostic_event(
694
- lambda: {
695
- "kind": "inbound",
696
- "sourceNamespace": namespace,
697
- "namespace": namespace,
698
- "eventType": event_type,
699
- "sid": sid,
700
- "correlationId": correlation_id,
701
- "timestamp": self._timestamp(),
702
- "handlerCount": len(selected_handlers),
703
- "durationMs": sum((exec.duration_ms or 0.0) for exec in executions),
704
- "resultSummary": self._summarize_results(results),
705
- "payloadSummary": self._summarize_payload(handler_payload),
706
- }
707
- )
708
-
709
- response_payload = {"correlationId": correlation_id, "results": results}
710
- if ack:
711
- ack(response_payload)
712
- self._debug(
713
- f"Completed event namespace={namespace} '{event_type}' sid={sid} correlation={correlation_id}"
714
- )
715
- return response_payload
716
-
717
- async def request_for_sid(
718
- self,
719
- *,
720
- namespace: str,
721
- sid: str,
722
- event_type: str,
723
- data: dict[str, Any],
724
- timeout_ms: int = 0,
725
- handler_id: str | None = None,
726
- include_handlers: Set[str] | None = None,
727
- ) -> dict[str, Any]:
728
- payload = dict(data or {})
729
- correlation_id = self._resolve_correlation_id(payload)
730
-
731
- with self.lock:
732
- connected = (namespace, sid) in self.connections
733
- if not connected:
734
- return {
735
- "correlationId": correlation_id,
736
- "results": [
737
- self._build_error_result(
738
- handler_id=handler_id or self._identifier,
739
- code="CONNECTION_NOT_FOUND",
740
- message=f"Connection '{sid}' not found in namespace '{namespace}'",
741
- correlation_id=correlation_id,
742
- )
743
- ],
744
- }
745
-
746
- async def _invoke() -> dict[str, Any]:
747
- return await self.route_event(
748
- namespace,
749
- event_type,
750
- payload,
751
- sid,
752
- include_handlers=include_handlers,
753
- handler_id=handler_id,
754
- )
755
-
756
- if timeout_ms and timeout_ms > 0:
757
- try:
758
- return await asyncio.wait_for(_invoke(), timeout=timeout_ms / 1000)
759
- except asyncio.TimeoutError:
760
- PrintStyle.warning(
761
- f"request timeout for sid {sid} event '{event_type}'"
762
- )
763
- return {
764
- "correlationId": correlation_id,
765
- "results": [
766
- self._build_error_result(
767
- handler_id=handler_id or self._identifier,
768
- code="TIMEOUT",
769
- message="Request timeout",
770
- correlation_id=correlation_id,
771
- )
772
- ],
773
- }
774
- return await _invoke()
775
-
776
- async def route_event_all(
777
- self,
778
- namespace: str,
779
- event_type: str,
780
- data: dict[str, Any],
781
- *,
782
- timeout_ms: int = 0,
783
- exclude_handlers: Set[str] | None = None,
784
- handler_id: str | None = None,
785
- ) -> list[dict[str, Any]]:
786
- """Fan-out a request to all active connections and aggregate responses."""
787
-
788
- base_payload = dict(data or {})
789
- exclude_meta_raw = base_payload.pop("excludeHandlers", None)
790
- exclude_combined: Set[str] | None = exclude_handlers
791
- correlation_id = self._resolve_correlation_id(base_payload)
792
-
793
- if exclude_meta_raw is not None:
794
- try:
795
- exclude_meta = self._normalize_handler_filter(
796
- exclude_meta_raw, "excludeHandlers"
797
- )
798
- except ValueError as exc:
799
- error = self._build_error_result(
800
- handler_id=handler_id or self._identifier,
801
- code="INVALID_FILTER",
802
- message=str(exc),
803
- correlation_id=correlation_id,
804
- )
805
- return [
806
- {
807
- "sid": "__invalid__",
808
- "correlationId": correlation_id,
809
- "results": [error],
810
- }
811
- ]
812
-
813
- if exclude_combined is None:
814
- exclude_combined = exclude_meta
815
- elif exclude_meta is not None and exclude_combined != exclude_meta:
816
- error = self._build_error_result(
817
- handler_id=handler_id or self._identifier,
818
- code="INVALID_FILTER",
819
- message="Conflicting excludeHandlers filters supplied",
820
- correlation_id=correlation_id,
821
- )
822
- return [
823
- {
824
- "sid": "__invalid__",
825
- "correlationId": correlation_id,
826
- "results": [error],
827
- }
828
- ]
829
-
830
- self._debug(
831
- f"Starting requestAll namespace={namespace} for '{event_type}' correlation={correlation_id}"
832
- )
833
-
834
- with self.lock:
835
- active_sids = [
836
- conn_identity[1]
837
- for conn_identity in self.connections.keys()
838
- if conn_identity[0] == namespace
839
- ]
840
- if not active_sids:
841
- self._debug(
842
- f"No active connections for requestAll namespace={namespace} '{event_type}' correlation={correlation_id}"
843
- )
844
- return []
845
-
846
- timeout_seconds = timeout_ms / 1000 if timeout_ms and timeout_ms > 0 else None
847
-
848
- async def _invoke_for_sid(target_sid: str) -> dict[str, Any]:
849
- async def _dispatch() -> dict[str, Any]:
850
- return await self.route_event(
851
- namespace,
852
- event_type,
853
- base_payload,
854
- target_sid,
855
- allow_exclude=True,
856
- exclude_handlers=exclude_combined,
857
- handler_id=handler_id,
858
- )
859
-
860
- if timeout_seconds is None:
861
- return await _dispatch()
862
-
863
- try:
864
- task = asyncio.create_task(_dispatch())
865
- return await asyncio.wait_for(
866
- asyncio.shield(task), timeout=timeout_seconds
867
- )
868
- except asyncio.TimeoutError:
869
- PrintStyle.warning(
870
- f"requestAll timeout for sid {target_sid} correlation={correlation_id}"
871
- )
872
- # Ensure any late exceptions are observed so asyncio does not log
873
- # "Task exception was never retrieved".
874
- try:
875
- task.add_done_callback(lambda t: t.exception()) # type: ignore[arg-type]
876
- except Exception: # pragma: no cover - defensive
877
- pass
878
- return {
879
- "correlationId": correlation_id,
880
- "results": [
881
- self._build_error_result(
882
- handler_id=handler_id or self._identifier,
883
- code="TIMEOUT",
884
- message="Request timeout",
885
- correlation_id=correlation_id,
886
- )
887
- ],
888
- }
889
-
890
- tasks = {sid: asyncio.create_task(_invoke_for_sid(sid)) for sid in active_sids}
891
-
892
- aggregated: list[dict[str, Any]] = []
893
- for sid, task in tasks.items():
894
- result = await task
895
- if isinstance(result, dict):
896
- aggregated.append(
897
- {
898
- "sid": sid,
899
- "correlationId": result.get("correlationId", correlation_id),
900
- "results": result.get("results", []),
901
- }
902
- )
903
- else:
904
- aggregated.append(
905
- {
906
- "sid": sid,
907
- "correlationId": correlation_id,
908
- "results": result,
909
- }
910
- )
911
-
912
- self._debug(
913
- f"Completed requestAll namespace={namespace} for '{event_type}' correlation={correlation_id}"
914
- )
915
- return aggregated
916
-
917
- def _wrap_envelope(
918
- self,
919
- handler_id: str | None,
920
- data: dict[str, Any],
921
- *,
922
- correlation_id: str | None = None,
923
- ) -> dict[str, Any]:
924
- hid = handler_id or self._identifier
925
- ts = _utcnow().isoformat(timespec="milliseconds").replace("+00:00", "Z")
926
- event_id = str(uuid.uuid4())
927
- correlation = correlation_id or str(uuid.uuid4())
928
- return {
929
- "handlerId": hid,
930
- "eventId": event_id,
931
- "correlationId": correlation,
932
- "ts": ts,
933
- "data": data or {},
934
- }
935
-
936
- async def emit_to(
937
- self,
938
- namespace: str,
939
- sid: str,
940
- event_type: str,
941
- data: dict[str, Any],
942
- *,
943
- handler_id: str | None = None,
944
- correlation_id: str | None = None,
945
- diagnostic: bool = False,
946
- ) -> None:
947
- envelope = self._wrap_envelope(
948
- handler_id,
949
- data,
950
- correlation_id=correlation_id,
951
- )
952
- delivered = False
953
- buffered = False
954
- identity: ConnectionIdentity = (namespace, sid)
955
-
956
- with self.lock:
957
- connected = identity in self.connections
958
- known = identity in self._known_sids or identity in self.buffers
959
-
960
- if connected:
961
- self._debug(
962
- "Emit to namespace=%s sid=%s event=%s eventId=%s correlationId=%s handlerId=%s"
963
- % (
964
- namespace,
965
- sid,
966
- event_type,
967
- envelope.get("eventId"),
968
- envelope.get("correlationId"),
969
- envelope.get("handlerId"),
970
- )
971
- )
972
- await self._run_on_dispatcher_loop(
973
- self.socketio.emit(event_type, envelope, to=sid, namespace=namespace)
974
- )
975
- delivered = True
976
- else:
977
- if not known:
978
- raise ConnectionNotFoundError(sid, namespace=namespace)
979
- with self.lock:
980
- self._buffer_event(
981
- identity,
982
- event_type,
983
- data,
984
- handler_id,
985
- envelope["correlationId"],
986
- )
987
- buffered = True
988
-
989
- if not diagnostic:
990
- await self._publish_diagnostic_event(
991
- lambda: {
992
- "kind": "outbound",
993
- "direction": "emit_to",
994
- "eventType": event_type,
995
- "namespace": namespace,
996
- "sid": sid,
997
- "correlationId": envelope["correlationId"],
998
- "handlerId": envelope["handlerId"],
999
- "timestamp": self._timestamp(),
1000
- "delivered": delivered,
1001
- "buffered": buffered,
1002
- "payloadSummary": self._summarize_payload(data),
1003
- }
1004
- )
1005
-
1006
- async def send_data(
1007
- self,
1008
- endpoint_name: str,
1009
- event_name: str,
1010
- data: dict[str, Any],
1011
- connection_id: str | None = None,
1012
- ) -> None:
1013
- if connection_id is not None:
1014
- await self.emit_to(endpoint_name, connection_id, event_name, data)
1015
- return
1016
- await self.broadcast(endpoint_name, event_name, data)
1017
-
1018
- async def broadcast(
1019
- self,
1020
- namespace: str,
1021
- event_type: str,
1022
- data: dict[str, Any],
1023
- *,
1024
- exclude_sids: str | Iterable[str] | None = None,
1025
- handler_id: str | None = None,
1026
- correlation_id: str | None = None,
1027
- diagnostic: bool = False,
1028
- ) -> None:
1029
- excluded = self._normalize_sid_filter(exclude_sids)
1030
-
1031
- targets: list[str] = []
1032
- with self.lock:
1033
- current_identities = list(self.connections.keys())
1034
- for conn_identity in current_identities:
1035
- if conn_identity[0] != namespace:
1036
- continue
1037
- sid = conn_identity[1]
1038
- if sid in excluded:
1039
- continue
1040
- targets.append(sid)
1041
- await self.emit_to(
1042
- namespace,
1043
- sid,
1044
- event_type,
1045
- data,
1046
- handler_id=handler_id,
1047
- correlation_id=correlation_id,
1048
- diagnostic=diagnostic,
1049
- )
1050
-
1051
- if not diagnostic:
1052
- await self._publish_diagnostic_event(
1053
- lambda: {
1054
- "kind": "outbound",
1055
- "direction": "broadcast",
1056
- "eventType": event_type,
1057
- "namespace": namespace,
1058
- "targets": targets[:10],
1059
- "targetCount": len(targets),
1060
- "correlationId": correlation_id,
1061
- "handlerId": handler_id or self._identifier,
1062
- "timestamp": self._timestamp(),
1063
- "payloadSummary": self._summarize_payload(data),
1064
- }
1065
- )
1066
-
1067
- async def _run_lifecycle(
1068
- self, namespace: str, fn: Callable[[WebSocketHandler], Any]
1069
- ) -> None:
1070
- seen: Set[WebSocketHandler] = set()
1071
- coros: list[Any] = []
1072
- for handler in self.handlers.get(namespace, []):
1073
- if handler in seen:
1074
- continue
1075
- seen.add(handler)
1076
- coros.append(self._get_handler_worker().execute_inside(fn, handler))
1077
- if coros:
1078
- await asyncio.gather(*coros, return_exceptions=True)
1079
-
1080
- def _buffer_event(
1081
- self,
1082
- identity: ConnectionIdentity,
1083
- event_type: str,
1084
- data: dict[str, Any],
1085
- handler_id: str | None,
1086
- correlation_id: str | None,
1087
- ) -> None:
1088
- namespace, sid = identity
1089
- buffer = self.buffers[identity]
1090
- buffer.append(
1091
- BufferedEvent(
1092
- event_type=event_type,
1093
- data=data,
1094
- handler_id=handler_id,
1095
- correlation_id=correlation_id,
1096
- )
1097
- )
1098
- while len(buffer) > BUFFER_MAX_SIZE:
1099
- dropped = buffer.popleft()
1100
- PrintStyle.warning(
1101
- f"Dropping buffered event '{dropped.event_type}' for namespace={namespace} sid={sid} (overflow)"
1102
- )
1103
- self._debug(
1104
- f"Buffered event namespace={namespace} '{event_type}' sid={sid} (queue length={len(buffer)})"
1105
- )
1106
-
1107
- async def _flush_buffer(self, identity: ConnectionIdentity) -> None:
1108
- self._ensure_dispatcher_loop()
1109
- buffer = self.buffers.get(identity)
1110
- if not buffer:
1111
- return
1112
- namespace, sid = identity
1113
- now = _utcnow()
1114
- delivered = 0
1115
- while buffer:
1116
- event = buffer.popleft()
1117
- if now - event.timestamp > BUFFER_TTL:
1118
- self._debug(
1119
- f"Discarding expired buffered event '{event.event_type}' for sid {sid}"
1120
- )
1121
- continue
1122
- envelope = self._wrap_envelope(
1123
- event.handler_id,
1124
- event.data,
1125
- correlation_id=event.correlation_id,
1126
- )
1127
- self._debug(
1128
- "Flush to sid=%s event=%s eventId=%s correlationId=%s handlerId=%s"
1129
- % (
1130
- sid,
1131
- event.event_type,
1132
- envelope.get("eventId"),
1133
- envelope.get("correlationId"),
1134
- envelope.get("handlerId"),
1135
- )
1136
- )
1137
- await self._run_on_dispatcher_loop(
1138
- self.socketio.emit(
1139
- event.event_type, envelope, to=sid, namespace=namespace
1140
- )
1141
- )
1142
- delivered += 1
1143
- if identity in self.buffers:
1144
- self.buffers.pop(identity, None)
1145
- if delivered:
1146
- PrintStyle.info(
1147
- f"Flushed {delivered} buffered event(s) to namespace={namespace} sid={sid}"
1148
- )
1149
-
1150
- def _build_error_result(
1151
- self,
1152
- *,
1153
- handler_id: str | None = None,
1154
- code: str,
1155
- message: str,
1156
- details: str | None = None,
1157
- correlation_id: str | None = None,
1158
- duration_ms: float | None = None,
1159
- ) -> dict[str, Any]:
1160
- error_payload = {"code": code, "error": message}
1161
- if details:
1162
- error_payload["details"] = details
1163
- result: dict[str, Any] = {
1164
- "handlerId": handler_id or self._identifier,
1165
- "ok": False,
1166
- "error": error_payload,
1167
- }
1168
- if correlation_id is not None:
1169
- result["correlationId"] = correlation_id
1170
- if duration_ms is not None:
1171
- result["durationMs"] = round(duration_ms, 4)
1172
- return result
1173
-
1174
- # Session tracking helpers (single-user defaults)
1175
- def get_sids_for_user(self, user: str | None = None) -> list[str]:
1176
- """Return SIDs for a user; single-user default returns all active SIDs."""
1177
- with self.lock:
1178
- bucket = self._ALL_USERS_BUCKET if user is None else user
1179
- return list(self.user_to_sids.get(bucket, set())) # type: ignore
1180
-
1181
- def get_user_for_sid(self, sid: str) -> str | None:
1182
- """Return user identifier for a SID or None."""
1183
- with self.lock:
1184
- return self.sid_to_user.get(sid) # type: ignore
1185
-
1186
- def set_server_restart_broadcast(self, enabled: bool) -> None:
1187
- """Enable or disable automatic server restart broadcasts."""
1188
-
1189
- self._server_restart_enabled = bool(enabled)
helpers/websocket_namespace_discovery.py
deleted
-186
@@ -1,186 +0,0 @@
1
-from __future__ import annotations
2
-
3
-import importlib.util
4
-import inspect
5
-import os
6
-from dataclasses import dataclass
7
-from types import ModuleType
8
-from typing import Iterable
9
-
10
-from helpers.files import get_abs_path
11
-from helpers.print_style import PrintStyle
12
-from helpers.websocket import WebSocketHandler
13
-
14
-
15
-@dataclass(frozen=True)
16
-class NamespaceDiscovery:
17
- namespace: str
18
- handler_classes: tuple[type[WebSocketHandler], ...]
19
- source_files: tuple[str, ...]
20
-
21
-
22
-def _to_namespace(entry_name: str) -> str:
23
- if entry_name == "_default":
24
- return "/"
25
- stripped = entry_name[: -len("_handler")] if entry_name.endswith("_handler") else entry_name
26
- if not stripped:
27
- raise ValueError(f"Invalid handler entry name: {entry_name!r}")
28
- return f"/{stripped}"
29
-
30
-
31
-def _unique_module_name(file_path: str) -> str:
32
- # Use a stable, unique module name derived from the relative path to avoid
33
- # collisions when importing different files with the same basename.
34
- rel_path = os.path.relpath(file_path, get_abs_path("."))
35
- rel_no_ext = os.path.splitext(rel_path)[0]
36
- safe = "".join(ch if ch.isalnum() else "_" for ch in rel_no_ext)
37
- return f"a0_ws_ns_{safe}"
38
-
39
-
40
-def _import_module(file_path: str) -> ModuleType:
41
- abs_path = get_abs_path(file_path)
42
- module_name = _unique_module_name(abs_path)
43
- spec = importlib.util.spec_from_file_location(module_name, abs_path)
44
- if spec is None or spec.loader is None:
45
- raise ImportError(f"Could not load module from {abs_path}")
46
- module = importlib.util.module_from_spec(spec)
47
- spec.loader.exec_module(module)
48
- return module
49
-
50
-
51
-def _get_handler_classes(module: ModuleType) -> list[type[WebSocketHandler]]:
52
- discovered: list[type[WebSocketHandler]] = []
53
- for _name, cls in inspect.getmembers(module, inspect.isclass):
54
- if cls is WebSocketHandler:
55
- continue
56
- if not issubclass(cls, WebSocketHandler):
57
- continue
58
- if cls.__module__ != module.__name__:
59
- continue
60
- discovered.append(cls)
61
- return discovered
62
-
63
-
64
-def discover_websocket_namespaces(
65
- *,
66
- handlers_folder: str = "python/websocket_handlers",
67
- include_root_default: bool = True,
68
-) -> list[NamespaceDiscovery]:
69
- """
70
- Discover websocket namespaces from first-level filesystem entries.
71
-
72
- Supported entries:
73
- - File entry: `*_handler.py` defines an application namespace.
74
- - Folder entry: `<name>/` or `<name>_handler/` defines an application namespace and loads
75
- `*.py` files one level deep (ignores `__init__.py` and ignores deeper nesting).
76
- - Reserved root mapping: `_default.py` maps to `/` when `include_root_default=True`.
77
- """
78
-
79
- abs_folder = get_abs_path(handlers_folder)
80
- entries: list[NamespaceDiscovery] = []
81
-
82
- try:
83
- filenames = sorted(os.listdir(abs_folder))
84
- except FileNotFoundError:
85
- PrintStyle.warning(f"WebSocket handlers folder not found: {abs_folder}")
86
- return []
87
-
88
- for entry in filenames:
89
- entry_path = os.path.join(abs_folder, entry)
90
-
91
- # Folder entries define namespaces and can host multiple handler modules.
92
- if os.path.isdir(entry_path):
93
- if entry.startswith("__"):
94
- continue
95
- namespace = _to_namespace(entry)
96
-
97
- handler_classes: list[type[WebSocketHandler]] = []
98
- source_files: list[str] = []
99
-
100
- try:
101
- child_names = sorted(os.listdir(entry_path))
102
- except FileNotFoundError:
103
- continue
104
-
105
- for child in child_names:
106
- if not child.endswith(".py"):
107
- continue
108
- if child == "__init__.py":
109
- continue
110
- child_path = os.path.join(entry_path, child)
111
- if not os.path.isfile(child_path):
112
- # Ignore deeper nesting.
113
- continue
114
-
115
- module = _import_module(child_path)
116
- discovered = _get_handler_classes(module)
117
- if not discovered:
118
- raise RuntimeError(
119
- f"WebSocket handler module {child_path} defines no WebSocketHandler subclasses"
120
- )
121
- if len(discovered) > 1:
122
- raise RuntimeError(
123
- f"WebSocket handler module {child_path} defines multiple WebSocketHandler subclasses: "
124
- f"{', '.join(sorted(cls.__name__ for cls in discovered))}"
125
- )
126
- handler_classes.append(discovered[0])
127
- source_files.append(child_path)
128
-
129
- if not handler_classes:
130
- PrintStyle.warning(
131
- f"WebSocket handlers folder entry '{entry_path}' is empty; treating namespace '{namespace}' as unregistered"
132
- )
133
- continue
134
-
135
- entries.append(
136
- NamespaceDiscovery(
137
- namespace=namespace,
138
- handler_classes=tuple(handler_classes),
139
- source_files=tuple(source_files),
140
- )
141
- )
142
- continue
143
-
144
- # File entries define namespaces.
145
- if not entry.endswith(".py"):
146
- continue
147
- if entry == "__init__.py":
148
- continue
149
-
150
- if entry == "_default.py":
151
- if not include_root_default:
152
- continue
153
- entry_name = "_default"
154
- else:
155
- if not entry.endswith("_handler.py"):
156
- continue
157
- entry_name = entry[: -len("_handler.py")]
158
-
159
- namespace = _to_namespace(entry_name)
160
- module_path = os.path.join(abs_folder, entry)
161
-
162
- module = _import_module(module_path)
163
- handler_classes = _get_handler_classes(module)
164
- if not handler_classes:
165
- raise RuntimeError(
166
- f"WebSocket handler module {module_path} defines no WebSocketHandler subclasses"
167
- )
168
- if len(handler_classes) > 1:
169
- raise RuntimeError(
170
- f"WebSocket handler module {module_path} defines multiple WebSocketHandler subclasses: "
171
- f"{', '.join(sorted(cls.__name__ for cls in handler_classes))}"
172
- )
173
-
174
- entries.append(
175
- NamespaceDiscovery(
176
- namespace=namespace,
177
- handler_classes=(handler_classes[0],),
178
- source_files=(module_path,),
179
- )
180
- )
181
-
182
- return entries
183
-
184
-
185
-def iter_discovered_namespaces(discoveries: Iterable[NamespaceDiscovery]) -> list[str]:
186
- return [entry.namespace for entry in discoveries]
python/websocket_handlers/_default.py
deleted
-26
@@ -1,26 +0,0 @@
1
-from __future__ import annotations
2
-
3
-from typing import Any
4
-
5
-from helpers.websocket import WebSocketHandler, WebSocketResult
6
-
7
-
8
-class RootDefaultHandler(WebSocketHandler):
9
- """Reserved root (`/`) namespace diagnostics-only handler.
10
-
11
- Root is intentionally *not* used for application traffic. This handler exists to support
12
- optional low-risk diagnostics on `/` without making root behave like a global namespace.
13
- """
14
-
15
- @classmethod
16
- def requires_auth(cls) -> bool:
17
- return False
18
-
19
- @classmethod
20
- def requires_csrf(cls) -> bool:
21
- return False
22
-
23
- async def process_event(
24
- self, event_type: str, data: dict[str, Any], sid: str
25
- ) -> dict[str, Any] | WebSocketResult | None:
26
- return {"ok": True, "namespace": self.namespace, "sid": sid, "echo": data}
python/websocket_handlers/dev_websocket_test_handler.py
deleted
-117
@@ -1,117 +0,0 @@
1
-from __future__ import annotations
2
-
3
-import asyncio
4
-from typing import Any, Dict
5
-
6
-from helpers.print_style import PrintStyle
7
-from helpers import runtime
8
-from helpers.websocket import WebSocketHandler, WebSocketResult
9
-
10
-
11
-class DevWebsocketTestHandler(WebSocketHandler):
12
- """Test harness handler powering the developer WebSocket validation component."""
13
-
14
- async def process_event(
15
- self, event_type: str, data: Dict[str, Any], sid: str
16
- ) -> dict[str, Any] | WebSocketResult | None:
17
- if event_type == "ws_event_console_subscribe":
18
- if not runtime.is_development():
19
- return self.result_error(
20
- code="NOT_AVAILABLE",
21
- message="Event console is available only in development mode",
22
- )
23
- registered = self.manager.register_diagnostic_watcher(self.namespace, sid)
24
- if not registered:
25
- return self.result_error(
26
- code="SUBSCRIBE_FAILED",
27
- message="Unable to subscribe to diagnostics",
28
- )
29
- return self.result_ok(
30
- {"status": "subscribed", "timestamp": data.get("requestedAt")}
31
- )
32
-
33
- if event_type == "ws_event_console_unsubscribe":
34
- self.manager.unregister_diagnostic_watcher(self.namespace, sid)
35
- return self.result_ok({"status": "unsubscribed"})
36
-
37
- if event_type == "ws_tester_emit":
38
- message = data.get("message", "emit")
39
- payload = {
40
- "message": message,
41
- "echo": True,
42
- "timestamp": data.get("timestamp"),
43
- }
44
- await self.broadcast("ws_tester_broadcast", payload)
45
- PrintStyle.info(f"Harness emit broadcasted message='{message}'")
46
- return None
47
-
48
- if event_type == "ws_tester_request":
49
- value = data.get("value")
50
- response = {
51
- "echo": value,
52
- "handler": self.identifier,
53
- "status": "ok",
54
- }
55
- PrintStyle.debug("Harness request responded with echo %s", value)
56
- return self.result_ok(
57
- response,
58
- correlation_id=data.get("correlationId"),
59
- )
60
-
61
- if event_type == "ws_tester_request_delayed":
62
- delay_ms = int(data.get("delay_ms", 0))
63
- await asyncio.sleep(delay_ms / 1000)
64
- PrintStyle.warning(
65
- "Harness delayed request finished after %s ms", delay_ms
66
- )
67
- return self.result_ok(
68
- {
69
- "status": "delayed",
70
- "delay_ms": delay_ms,
71
- "handler": self.identifier,
72
- },
73
- correlation_id=data.get("correlationId"),
74
- )
75
-
76
- if event_type == "ws_tester_trigger_persistence":
77
- phase = data.get("phase", "unknown")
78
- payload = {
79
- "phase": phase,
80
- "handler": self.identifier,
81
- }
82
- await self.emit_to(sid, "ws_tester_persistence", payload)
83
- PrintStyle.info(f"Harness persistence event phase='{phase}' -> {sid}")
84
- return None
85
-
86
- if event_type == "ws_tester_request_all":
87
- marker = data.get("marker")
88
- PrintStyle.debug(
89
- "Harness requestAll invoked by %s marker='%s'", sid, marker
90
- )
91
- exclude_handlers = data.get("excludeHandlers")
92
- aggregated = await self.request_all(
93
- "ws_tester_request",
94
- data,
95
- timeout_ms=2_000,
96
- exclude_handlers=exclude_handlers,
97
- )
98
- return self.result_ok(
99
- {"results": aggregated},
100
- correlation_id=data.get("correlationId"),
101
- )
102
-
103
- if event_type == "ws_tester_broadcast_demo_trigger":
104
- payload = {
105
- "demo": True,
106
- "requested_at": data.get("requested_at"),
107
- }
108
- await self.broadcast("ws_tester_broadcast_demo", payload)
109
- PrintStyle.info("Harness broadcast demo event dispatched")
110
- return None
111
-
112
- PrintStyle.warning(f"Harness received unknown event '{event_type}'")
113
- return self.result_error(
114
- code="HARNESS_UNKNOWN_EVENT",
115
- message="Unhandled event",
116
- details=event_type,
117
- )
python/websocket_handlers/hello_handler.py
deleted
-15
@@ -1,15 +0,0 @@
1
-from __future__ import annotations
2
-
3
-from helpers.print_style import PrintStyle
4
-from helpers.websocket import WebSocketHandler
5
-
6
-
7
-class HelloHandler(WebSocketHandler):
8
- """Sample handler used for foundational testing."""
9
-
10
- async def process_event(self, event_type: str, data: dict, sid: str):
11
- name = data.get("name") or "stranger"
12
- PrintStyle.info(f"hello_request from {sid} ({name})")
13
- return {"message": f"Hello, {name}!", "handler": self.identifier}
14
-
15
-
python/websocket_handlers/webui_handler.py
deleted
-34
@@ -1,34 +0,0 @@
1
-from helpers.websocket import WebSocketHandler, WebSocketResult
2
-from helpers import extension
3
-
4
-
5
-class WebuiHandler(WebSocketHandler):
6
- async def on_connect(self, sid: str) -> None:
7
- await extension.call_extensions_async(
8
- "webui_ws_connect", agent=None, instance=self, sid=sid
9
- )
10
-
11
- async def on_disconnect(self, sid: str) -> None:
12
- await extension.call_extensions_async(
13
- "webui_ws_disconnect", agent=None, instance=self, sid=sid
14
- )
15
-
16
- async def process_event(
17
- self, event_type: str, data: dict, sid: str
18
- ) -> dict | WebSocketResult | None:
19
- response_data: dict = {}
20
-
21
- await extension.call_extensions_async(
22
- "webui_ws_event",
23
- agent=None,
24
- instance=self,
25
- sid=sid,
26
- event_type=event_type,
27
- data=data,
28
- response_data=response_data,
29
- )
30
-
31
- return self.result_ok(
32
- response_data,
33
- correlation_id=data.get("correlationId"),
34
- )
tests/test_websocket_client_api_surface.py
deleted
-41
@@ -1,41 +0,0 @@
1
-import re
2
-import sys
3
-from pathlib import Path
4
-
5
-
6
-PROJECT_ROOT = Path(__file__).resolve().parents[1]
7
-if str(PROJECT_ROOT) not in sys.path:
8
- sys.path.insert(0, str(PROJECT_ROOT))
9
-
10
-
11
-def _get_named_exports(source: str) -> set[str]:
12
- exports: set[str] = set()
13
-
14
- exports.update(re.findall(r"^export\s+function\s+([A-Za-z0-9_]+)\s*\(", source, flags=re.M))
15
- exports.update(re.findall(r"^export\s+const\s+([A-Za-z0-9_]+)\s*=", source, flags=re.M))
16
- exports.update(re.findall(r"^export\s+class\s+([A-Za-z0-9_]+)\s*[\{:]", source, flags=re.M))
17
-
18
- for m in re.findall(r"^export\s*\{([^}]+)\}\s*;?", source, flags=re.M):
19
- for item in m.split(","):
20
- item = item.strip()
21
- if not item:
22
- continue
23
- # Handle: `foo as bar`
24
- parts = item.split()
25
- if len(parts) >= 3 and parts[-2] == "as":
26
- exports.add(parts[-1])
27
- else:
28
- exports.add(parts[0])
29
-
30
- return exports
31
-
32
-
33
-def test_websocket_js_exports_minimal_namespaced_api_surface() -> None:
34
- source = (PROJECT_ROOT / "webui" / "js" / "websocket.js").read_text(encoding="utf-8")
35
- exports = _get_named_exports(source)
36
-
37
- assert "createNamespacedClient" in exports
38
- assert "getNamespacedClient" in exports
39
-
40
- assert "broadcast" not in exports
41
- assert "requestAll" not in exports
tests/test_websocket_csrf.py
deleted
-51
@@ -1,51 +0,0 @@
1
-import sys
2
-from pathlib import Path
3
-
4
-PROJECT_ROOT = Path(__file__).resolve().parents[1]
5
-if str(PROJECT_ROOT) not in sys.path:
6
- sys.path.insert(0, str(PROJECT_ROOT))
7
-
8
-from helpers.websocket import validate_ws_origin
9
-
10
-
11
-def test_validate_ws_origin_allows_same_origin_with_explicit_port():
12
- ok, reason = validate_ws_origin(
13
- {
14
- "HTTP_ORIGIN": "http://localhost:5000",
15
- "HTTP_HOST": "localhost:5000",
16
- }
17
- )
18
- assert ok is True
19
- assert reason is None
20
-
21
-
22
-def test_validate_ws_origin_allows_default_https_port_without_explicit_port():
23
- ok, reason = validate_ws_origin(
24
- {
25
- "HTTP_ORIGIN": "https://example.com",
26
- "HTTP_HOST": "example.com",
27
- }
28
- )
29
- assert ok is True
30
- assert reason is None
31
-
32
-
33
-def test_validate_ws_origin_rejects_missing_origin():
34
- ok, reason = validate_ws_origin(
35
- {
36
- "HTTP_HOST": "localhost:5000",
37
- }
38
- )
39
- assert ok is False
40
- assert reason == "missing_origin"
41
-
42
-
43
-def test_validate_ws_origin_rejects_cross_origin():
44
- ok, reason = validate_ws_origin(
45
- {
46
- "HTTP_ORIGIN": "http://evil.test",
47
- "HTTP_HOST": "localhost:5000",
48
- }
49
- )
50
- assert ok is False
51
- assert reason == "origin_host_mismatch"
tests/test_websocket_handlers.py
deleted
-176
@@ -1,176 +0,0 @@
1
-import sys
2
-import threading
3
-from pathlib import Path
4
-
5
-import pytest
6
-
7
-PROJECT_ROOT = Path(__file__).resolve().parents[1]
8
-if str(PROJECT_ROOT) not in sys.path:
9
- sys.path.insert(0, str(PROJECT_ROOT))
10
-
11
-from helpers.websocket import (
12
- WebSocketHandler,
13
- WebSocketResult,
14
- SingletonInstantiationError,
15
-)
16
-
17
-
18
-class _FakeSocketIO:
19
- async def emit(self, *_args, **_kwargs): # pragma: no cover - helper stub
20
- return None
21
-
22
- async def disconnect(self, *_args, **_kwargs): # pragma: no cover - helper stub
23
- return None
24
-
25
-
26
-class _TestHandler(WebSocketHandler):
27
- @classmethod
28
- def get_event_types(cls) -> list[str]:
29
- return ["test_event"]
30
-
31
- async def process_event(self, event_type: str, data: dict, sid: str) -> None:
32
- return None
33
-
34
-
35
-def _make_handler() -> _TestHandler:
36
- _TestHandler._reset_instance_for_testing()
37
- return _TestHandler.get_instance(_FakeSocketIO(), threading.RLock())
38
-
39
-
40
-def test_websocket_result_ok_clones_payload():
41
- payload = {"value": 1}
42
- result = WebSocketResult.ok(payload)
43
-
44
- assert result.as_result(
45
- handler_id="handler",
46
- fallback_correlation_id="corr",
47
- )["data"] == payload
48
-
49
- payload["value"] = 2
50
- assert result.as_result(
51
- handler_id="handler",
52
- fallback_correlation_id="corr",
53
- )["data"] == {"value": 1}
54
-
55
-
56
-def test_websocket_result_error_contains_metadata():
57
- result = WebSocketResult.error(
58
- code="E_TEST",
59
- message="failure",
60
- details="additional",
61
- correlation_id="corr",
62
- duration_ms=12.5,
63
- )
64
-
65
- as_payload = result.as_result(handler_id="handler", fallback_correlation_id=None)
66
- assert as_payload["ok"] is False
67
- assert as_payload["error"] == {
68
- "code": "E_TEST",
69
- "error": "failure",
70
- "details": "additional",
71
- }
72
- assert as_payload["correlationId"] == "corr"
73
- assert as_payload["durationMs"] == pytest.approx(12.5, rel=1e-3)
74
-
75
-
76
-def test_websocket_result_applies_fallback_correlation_and_duration():
77
- result = WebSocketResult.ok(duration_ms=5.4321)
78
- payload = result.as_result(
79
- handler_id="handler",
80
- fallback_correlation_id="corr-fallback",
81
- )
82
- assert payload["correlationId"] == "corr-fallback"
83
- assert payload["durationMs"] == pytest.approx(5.4321, rel=1e-3)
84
-
85
-
86
-def test_handler_result_helpers_return_websocket_result_instances():
87
- handler = _make_handler()
88
-
89
- ok_result = handler.result_ok({"foo": "bar"}, correlation_id="cid")
90
- assert isinstance(ok_result, WebSocketResult)
91
- ok_payload = ok_result.as_result(
92
- handler_id="handler",
93
- fallback_correlation_id=None,
94
- )
95
- assert ok_payload["ok"] is True
96
- assert ok_payload["data"] == {"foo": "bar"}
97
- assert ok_payload["correlationId"] == "cid"
98
-
99
- err_result = handler.result_error(
100
- code="E_BAD",
101
- message="boom",
102
- details="missing",
103
- correlation_id="err",
104
- )
105
- assert isinstance(err_result, WebSocketResult)
106
- err_payload = err_result.as_result(
107
- handler_id="handler",
108
- fallback_correlation_id=None,
109
- )
110
- assert err_payload["ok"] is False
111
- assert err_payload["error"] == {
112
- "code": "E_BAD",
113
- "error": "boom",
114
- "details": "missing",
115
- }
116
- assert err_payload["correlationId"] == "err"
117
-
118
-
119
-def test_result_error_requires_error_payload():
120
- with pytest.raises(ValueError):
121
- WebSocketResult(ok=False)
122
-
123
- with pytest.raises(ValueError):
124
- WebSocketResult.error(code="", message="boom")
125
-
126
-
127
-def test_handler_direct_instantiation_disallowed():
128
- with pytest.raises(SingletonInstantiationError):
129
- _TestHandler(_FakeSocketIO(), threading.RLock())
130
-
131
-
132
-def test_get_instance_returns_singleton():
133
- _TestHandler._reset_instance_for_testing()
134
- socketio = _FakeSocketIO()
135
- lock = threading.RLock()
136
- first = _TestHandler.get_instance(socketio, lock)
137
- second = _TestHandler.get_instance(None, None)
138
- assert first is second
139
-
140
-
141
-@pytest.mark.asyncio
142
-async def test_state_sync_handler_registers_and_routes_state_request():
143
- from helpers.websocket_manager import WebSocketManager
144
- from python.websocket_handlers.webui_handler import WebuiHandler
145
- from helpers.state_monitor import _reset_state_monitor_for_testing
146
-
147
- _reset_state_monitor_for_testing()
148
- WebuiHandler._reset_instance_for_testing()
149
-
150
- socketio = _FakeSocketIO()
151
- lock = threading.RLock()
152
- manager = WebSocketManager(socketio, lock)
153
- handler = WebuiHandler.get_instance(socketio, lock)
154
- namespace = "/webui"
155
- manager.register_handlers({namespace: [handler]})
156
- await manager.handle_connect(namespace, "sid-1")
157
-
158
- response = await manager.route_event(
159
- namespace,
160
- "state_request",
161
- {
162
- "correlationId": "smoke-1",
163
- "ts": "2025-12-28T00:00:00.000Z",
164
- "data": {
165
- "context": None,
166
- "log_from": 0,
167
- "notifications_from": 0,
168
- "timezone": "UTC",
169
- },
170
- },
171
- "sid-1",
172
- )
173
-
174
- assert response["correlationId"] == "smoke-1"
175
- assert response["results"] and response["results"][0]["ok"] is True
176
- await manager.handle_disconnect(namespace, "sid-1")
tests/test_websocket_harness.py
deleted
-173
@@ -1,173 +0,0 @@
1
-import sys
2
-import threading
3
-from pathlib import Path
4
-from typing import Any
5
-
6
-import pytest
7
-
8
-PROJECT_ROOT = Path(__file__).resolve().parents[1]
9
-if str(PROJECT_ROOT) not in sys.path:
10
- sys.path.insert(0, str(PROJECT_ROOT))
11
-
12
-from helpers.websocket_manager import WebSocketManager
13
-from python.websocket_handlers.dev_websocket_test_handler import (
14
- DevWebsocketTestHandler,
15
-)
16
-
17
-NAMESPACE = "/dev_websocket_test"
18
-
19
-
20
-class FakeSocketIOServer:
21
- def __init__(self) -> None:
22
- from unittest.mock import AsyncMock
23
-
24
- self.emit = AsyncMock()
25
- self.disconnect = AsyncMock()
26
-
27
-
28
-async def _create_manager() -> tuple[WebSocketManager, DevWebsocketTestHandler, FakeSocketIOServer]:
29
- socketio = FakeSocketIOServer()
30
- manager = WebSocketManager(socketio, threading.RLock())
31
- DevWebsocketTestHandler._reset_instance_for_testing()
32
- handler = DevWebsocketTestHandler.get_instance(socketio, threading.RLock())
33
- manager.register_handlers({NAMESPACE: [handler]})
34
- await manager.handle_connect(NAMESPACE, "sid-primary")
35
- return manager, handler, socketio
36
-
37
-
38
-@pytest.mark.asyncio
39
-async def test_harness_emit_broadcasts_to_active_connections():
40
- manager, _handler, socketio = await _create_manager()
41
-
42
- await manager.route_event(
43
- NAMESPACE,
44
- "ws_tester_emit",
45
- {"message": "emit-check", "timestamp": "2025-10-29T12:00:00Z"},
46
- "sid-primary",
47
- )
48
-
49
- socketio.emit.assert_awaited()
50
- emit_calls = [(call.args, call.kwargs) for call in socketio.emit.await_args_list]
51
- match = next((c for c in emit_calls if c[0] and c[0][0] == "ws_tester_broadcast"), None)
52
- assert match is not None
53
- args, kwargs = match
54
- envelope = args[1]
55
- assert envelope["handlerId"].endswith("DevWebsocketTestHandler")
56
- assert envelope["data"]["message"] == "emit-check"
57
- assert kwargs == {"to": "sid-primary", "namespace": NAMESPACE}
58
-
59
-
60
-@pytest.mark.asyncio
61
-async def test_harness_request_returns_per_handler_result():
62
- manager, _handler, _socketio = await _create_manager()
63
-
64
- response = await manager.route_event(
65
- NAMESPACE,
66
- "ws_tester_request",
67
- {"value": 42},
68
- "sid-primary",
69
- )
70
-
71
- assert isinstance(response, dict)
72
- assert response["results"]
73
- first = response["results"][0]
74
- assert first["ok"] is True
75
- assert first["data"]["echo"] == 42
76
- assert response["correlationId"]
77
- assert first["handlerId"].endswith("DevWebsocketTestHandler")
78
- assert first["correlationId"] == response["correlationId"]
79
-
80
-
81
-@pytest.mark.asyncio
82
-async def test_harness_request_delayed_waits_for_sleep(monkeypatch):
83
- manager, _handler, _socketio = await _create_manager()
84
-
85
- calls: list[float] = []
86
-
87
- async def _fake_sleep(delay: float) -> None: # pragma: no cover - helper
88
- calls.append(delay)
89
-
90
- monkeypatch.setattr(
91
- "python.websocket_handlers.dev_websocket_test_handler.asyncio.sleep",
92
- _fake_sleep,
93
- )
94
-
95
- await manager.route_event(
96
- NAMESPACE,
97
- "ws_tester_request_delayed",
98
- {"delay_ms": 1500},
99
- "sid-primary",
100
- )
101
-
102
- assert calls == [1.5]
103
-
104
-
105
-@pytest.mark.asyncio
106
-async def test_harness_persistence_emit_targets_requesting_sid():
107
- manager, _handler, socketio = await _create_manager()
108
-
109
- await manager.route_event(
110
- NAMESPACE,
111
- "ws_tester_trigger_persistence",
112
- {"phase": "after"},
113
- "sid-primary",
114
- )
115
-
116
- socketio.emit.assert_awaited()
117
- emit_calls = [(call.args, call.kwargs) for call in socketio.emit.await_args_list]
118
- match = next((c for c in emit_calls if c[0] and c[0][0] == "ws_tester_persistence"), None)
119
- assert match is not None
120
- args, kwargs = match
121
- payload = args[1]
122
- assert payload["handlerId"] == _handler.identifier
123
- assert payload["data"] == {"phase": "after", "handler": _handler.identifier}
124
- assert kwargs == {"to": "sid-primary", "namespace": NAMESPACE}
125
-
126
-
127
-@pytest.mark.asyncio
128
-async def test_harness_request_all_aggregates_all_connections():
129
- manager, _handler, _socketio = await _create_manager()
130
- await manager.handle_connect(NAMESPACE, "sid-secondary")
131
-
132
- response = await manager.route_event(
133
- NAMESPACE,
134
- "ws_tester_request_all",
135
- {"marker": "aggregate"},
136
- "sid-primary",
137
- )
138
-
139
- assert response["results"] and response["results"][0]["ok"] is True
140
- data = response["results"][0]["data"]
141
- aggregated = data.get("results") or data.get("result")
142
- assert isinstance(aggregated, list)
143
- by_sid: dict[str, Any] = {entry["sid"]: entry["results"] for entry in aggregated}
144
- assert set(by_sid.keys()) == {"sid-primary", "sid-secondary"}
145
- for results in by_sid.values():
146
- assert results and results[0]["ok"] is True
147
- payload = results[0]["data"]
148
- assert payload["handler"].endswith("DevWebsocketTestHandler")
149
- assert results[0]["handlerId"].endswith("DevWebsocketTestHandler")
150
- assert results[0]["correlationId"] == response["results"][0]["correlationId"]
151
- assert response["correlationId"]
152
-
153
-
154
-@pytest.mark.asyncio
155
-async def test_harness_request_all_respects_exclude_handlers():
156
- manager, handler, _socketio = await _create_manager()
157
- await manager.handle_connect(NAMESPACE, "sid-secondary")
158
-
159
- response = await manager.route_event(
160
- NAMESPACE,
161
- "ws_tester_request_all",
162
- {
163
- "marker": "exclude",
164
- "excludeHandlers": [handler.identifier],
165
- },
166
- "sid-primary",
167
- )
168
-
169
- assert response["correlationId"]
170
- first = response["results"][0]
171
- assert first["ok"] is False
172
- assert first["error"]["code"] == "INVALID_FILTER"
173
- assert "excludeHandlers" in first["error"]["error"]
tests/test_websocket_manager.py
deleted
-853
@@ -1,853 +0,0 @@
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.websocket import ConnectionNotFoundError, WebSocketHandler, WebSocketResult
16
-from helpers.websocket_manager import (
17
- WebSocketManager,
18
- BUFFER_TTL,
19
- DIAGNOSTIC_EVENT,
20
- LIFECYCLE_CONNECT_EVENT,
21
- LIFECYCLE_DISCONNECT_EVENT,
22
-)
23
-
24
-NAMESPACE = "/test"
25
-
26
-
27
-class FakeSocketIOServer:
28
- def __init__(self):
29
- self.emit = AsyncMock()
30
- self.disconnect = AsyncMock()
31
-
32
-
33
-class DummyHandler(WebSocketHandler):
34
- def __init__(self, socketio, lock, results):
35
- super().__init__(socketio, lock)
36
- self.results = results
37
-
38
- @classmethod
39
- def get_event_types(cls) -> list[str]:
40
- return ["dummy"]
41
-
42
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
43
- response = {"sid": sid, "data": data}
44
- self.results.append(response)
45
- return response
46
-
47
-
48
-@pytest.mark.asyncio
49
-async def test_connect_disconnect_updates_registry():
50
- socketio = FakeSocketIOServer()
51
- manager = WebSocketManager(socketio, threading.RLock())
52
-
53
- await manager.handle_connect(NAMESPACE, "abc")
54
- assert (NAMESPACE, "abc") in manager.connections
55
-
56
- await manager.handle_disconnect(NAMESPACE, "abc")
57
- assert (NAMESPACE, "abc") not in manager.connections
58
-
59
-
60
-@pytest.mark.asyncio
61
-async def test_server_restart_broadcast_emitted_when_enabled():
62
- socketio = FakeSocketIOServer()
63
- manager = WebSocketManager(socketio, threading.RLock())
64
- manager.set_server_restart_broadcast(True)
65
-
66
- await manager.handle_connect(NAMESPACE, "sid-restart")
67
-
68
- socketio.emit.assert_awaited()
69
- args, kwargs = socketio.emit.await_args_list[0]
70
- assert args[0] == "server_restart"
71
- envelope = args[1]
72
- assert envelope["handlerId"] == manager._identifier # noqa: SLF001
73
- assert envelope["data"]["runtimeId"]
74
- assert kwargs == {"to": "sid-restart", "namespace": NAMESPACE}
75
-
76
-
77
-@pytest.mark.asyncio
78
-async def test_server_restart_broadcast_skipped_when_disabled():
79
- socketio = FakeSocketIOServer()
80
- manager = WebSocketManager(socketio, threading.RLock())
81
- manager.set_server_restart_broadcast(False)
82
-
83
- await manager.handle_connect(NAMESPACE, "sid-no-restart")
84
-
85
- assert socketio.emit.await_count == 0
86
-
87
-
88
-@pytest.mark.asyncio
89
-async def test_broadcast_performance_smoke(monkeypatch):
90
- socketio = FakeSocketIOServer()
91
- manager = WebSocketManager(socketio, threading.RLock())
92
-
93
- for idx in range(50):
94
- await manager.handle_connect(NAMESPACE, f"sid-{idx}")
95
-
96
- import time
97
-
98
- start = time.perf_counter()
99
- await manager.broadcast(NAMESPACE, "perf_event", {"ok": True})
100
- duration_ms = (time.perf_counter() - start) * 1000
101
-
102
- assert socketio.emit.await_count == 50
103
- assert duration_ms < 300
104
-
105
-
106
-@pytest.mark.asyncio
107
-async def test_route_event_invokes_handler_and_ack():
108
- socketio = FakeSocketIOServer()
109
- manager = WebSocketManager(socketio, threading.RLock())
110
-
111
- results = []
112
- DummyHandler._reset_instance_for_testing()
113
- handler = DummyHandler.get_instance(socketio, threading.RLock(), results)
114
- manager.register_handlers({NAMESPACE: [handler]})
115
- await manager.handle_connect(NAMESPACE, "sid-1")
116
-
117
- response = await manager.route_event(NAMESPACE, "dummy", {"foo": "bar"}, "sid-1")
118
-
119
- assert results[0]["sid"] == "sid-1"
120
- assert results[0]["data"]["foo"] == "bar"
121
- assert "correlationId" in results[0]["data"]
122
-
123
- assert isinstance(response, dict)
124
- assert "correlationId" in response
125
- assert isinstance(response["results"], list)
126
- entry = response["results"][0]
127
- assert entry["ok"] is True
128
- assert entry["data"]["sid"] == "sid-1"
129
- assert entry["data"]["data"]["foo"] == "bar"
130
-
131
-
132
-@pytest.mark.asyncio
133
-async def test_route_event_no_handler_returns_standard_error():
134
- socketio = FakeSocketIOServer()
135
- manager = WebSocketManager(socketio, threading.RLock())
136
- await manager.handle_connect(NAMESPACE, "sid-1")
137
-
138
- response = await manager.route_event(NAMESPACE, "missing", {}, "sid-1")
139
-
140
- assert len(response["results"]) == 1
141
- result = response["results"][0]
142
- assert result["handlerId"].endswith("WebSocketManager")
143
- assert result["ok"] is False
144
- assert result["error"]["code"] == "NO_HANDLERS"
145
- assert (
146
- result["error"]["error"]
147
- == f"No handler for namespace '{NAMESPACE}' event 'missing'"
148
- )
149
-
150
-
151
-@pytest.mark.asyncio
152
-async def test_route_event_all_returns_empty_when_no_connections():
153
- socketio = FakeSocketIOServer()
154
- manager = WebSocketManager(socketio, threading.RLock())
155
-
156
- results = await manager.route_event_all(NAMESPACE, "event", {}, timeout_ms=1000)
157
-
158
- assert results == []
159
-
160
-
161
-@pytest.mark.asyncio
162
-async def test_route_event_all_aggregates_results():
163
- socketio = FakeSocketIOServer()
164
- manager = WebSocketManager(socketio, threading.RLock())
165
-
166
- class EchoHandler(WebSocketHandler):
167
- @classmethod
168
- def get_event_types(cls) -> list[str]:
169
- return ["multi"]
170
-
171
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
172
- return {"sid": sid, "echo": data}
173
-
174
- EchoHandler._reset_instance_for_testing()
175
- handler = EchoHandler.get_instance(socketio, threading.RLock())
176
- manager.register_handlers({NAMESPACE: [handler]})
177
-
178
- await manager.handle_connect(NAMESPACE, "sid-1")
179
- await manager.handle_connect(NAMESPACE, "sid-2")
180
-
181
- aggregated = await manager.route_event_all(
182
- NAMESPACE, "multi", {"value": 42}, timeout_ms=1000
183
- )
184
-
185
- assert len(aggregated) == 2
186
- by_sid = {entry["sid"]: entry for entry in aggregated}
187
- assert by_sid["sid-1"]["results"][0]["ok"] is True
188
- payload_sid1 = by_sid["sid-1"]["results"][0]["data"]
189
- assert payload_sid1["sid"] == "sid-1"
190
- assert payload_sid1["echo"]["value"] == 42
191
- assert "correlationId" in payload_sid1["echo"]
192
- assert by_sid["sid-2"]["results"][0]["ok"] is True
193
- payload_sid2 = by_sid["sid-2"]["results"][0]["data"]
194
- assert payload_sid2["sid"] == "sid-2"
195
- assert payload_sid2["echo"]["value"] == 42
196
- assert by_sid["sid-1"]["correlationId"]
197
-
198
-
199
-@pytest.mark.asyncio
200
-async def test_route_event_all_timeout_marks_error():
201
- socketio = FakeSocketIOServer()
202
- manager = WebSocketManager(socketio, threading.RLock())
203
-
204
- class SlowHandler(WebSocketHandler):
205
- @classmethod
206
- def get_event_types(cls) -> list[str]:
207
- return ["slow"]
208
-
209
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
210
- await asyncio.sleep(0.2)
211
- return {"status": "done"}
212
-
213
- SlowHandler._reset_instance_for_testing()
214
- handler = SlowHandler.get_instance(socketio, threading.RLock())
215
- manager.register_handlers({NAMESPACE: [handler]})
216
- await manager.handle_connect(NAMESPACE, "sid-1")
217
-
218
- aggregated = await manager.route_event_all(NAMESPACE, "slow", {}, timeout_ms=50)
219
-
220
- assert len(aggregated) == 1
221
- first_entry = aggregated[0]
222
- result = first_entry["results"][0]
223
- assert result["ok"] is False
224
- assert result["error"] == {"code": "TIMEOUT", "error": "Request timeout"}
225
- assert first_entry["correlationId"]
226
-
227
-
228
-@pytest.mark.asyncio
229
-async def test_route_event_exception_standardizes_error_payload():
230
- socketio = FakeSocketIOServer()
231
- manager = WebSocketManager(socketio, threading.RLock())
232
-
233
- class FailingHandler(WebSocketHandler):
234
- @classmethod
235
- def get_event_types(cls) -> list[str]:
236
- return ["boom"]
237
-
238
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
239
- raise RuntimeError("kaboom")
240
-
241
- FailingHandler._reset_instance_for_testing()
242
- handler = FailingHandler.get_instance(socketio, threading.RLock())
243
- manager.register_handlers({NAMESPACE: [handler]})
244
- await manager.handle_connect(NAMESPACE, "sid-1")
245
-
246
- response = await manager.route_event(NAMESPACE, "boom", {}, "sid-1")
247
-
248
- assert len(response["results"]) == 1
249
- result = response["results"][0]
250
- assert result["handlerId"].endswith("FailingHandler")
251
- assert result["ok"] is False
252
- assert result["error"]["code"] == "HANDLER_ERROR"
253
- assert result["error"]["error"] == "Internal server error"
254
- assert "details" in result["error"]
255
-
256
-
257
-@pytest.mark.asyncio
258
-async def test_route_event_offloads_blocking_handlers():
259
- socketio = FakeSocketIOServer()
260
- manager = WebSocketManager(socketio, threading.RLock())
261
-
262
- class BlockingHandler(WebSocketHandler):
263
- @classmethod
264
- def get_event_types(cls) -> list[str]:
265
- return ["block"]
266
-
267
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
268
- time.sleep(0.2)
269
- return {"status": "done"}
270
-
271
- BlockingHandler._reset_instance_for_testing()
272
- handler = BlockingHandler.get_instance(socketio, threading.RLock())
273
- manager.register_handlers({NAMESPACE: [handler]})
274
- await manager.handle_connect(NAMESPACE, "sid-1")
275
-
276
- route_task = asyncio.create_task(
277
- manager.route_event(NAMESPACE, "block", {}, "sid-1")
278
- )
279
- await asyncio.sleep(0)
280
-
281
- t0 = time.perf_counter()
282
- await asyncio.sleep(0.05)
283
- elapsed = time.perf_counter() - t0
284
- assert elapsed < 0.15
285
-
286
- response = await route_task
287
- assert response["results"]
288
-
289
-
290
-@pytest.mark.asyncio
291
-async def test_route_event_unwraps_ts_data_envelope_and_preserves_correlation_id():
292
- socketio = FakeSocketIOServer()
293
- manager = WebSocketManager(socketio, threading.RLock())
294
-
295
- results: list[dict[str, Any]] = []
296
- DummyHandler._reset_instance_for_testing()
297
- handler = DummyHandler.get_instance(socketio, threading.RLock(), results)
298
- manager.register_handlers({NAMESPACE: [handler]})
299
- await manager.handle_connect(NAMESPACE, "sid-1")
300
-
301
- response = await manager.route_event(
302
- NAMESPACE,
303
- "dummy",
304
- {
305
- "correlationId": "client-1",
306
- "ts": "2025-10-29T12:00:00.000Z",
307
- "data": {"value": 123},
308
- },
309
- "sid-1",
310
- )
311
-
312
- assert response["correlationId"] == "client-1"
313
- assert len(results) == 1
314
- handler_payload = results[0]["data"]
315
- assert handler_payload["value"] == 123
316
- assert handler_payload["correlationId"] == "client-1"
317
- assert "ts" not in handler_payload
318
- assert "data" not in handler_payload
319
-
320
-
321
-@pytest.mark.asyncio
322
-async def test_emit_to_unknown_sid_raises_error():
323
- socketio = FakeSocketIOServer()
324
- manager = WebSocketManager(socketio, threading.RLock())
325
-
326
- with pytest.raises(ConnectionNotFoundError):
327
- await manager.emit_to(NAMESPACE, "unknown", "event", {})
328
-
329
-
330
-@pytest.mark.asyncio
331
-async def test_emit_to_known_disconnected_sid_buffers():
332
- socketio = FakeSocketIOServer()
333
- manager = WebSocketManager(socketio, threading.RLock())
334
- await manager.handle_connect(NAMESPACE, "sid-1")
335
- await manager.handle_disconnect(NAMESPACE, "sid-1")
336
-
337
- await manager.emit_to(
338
- NAMESPACE, "sid-1", "event", {"a": 1}, correlation_id="corr-1"
339
- )
340
-
341
- assert (NAMESPACE, "sid-1") in manager.buffers
342
- buffered = list(manager.buffers[(NAMESPACE, "sid-1")])
343
- assert len(buffered) == 1
344
- assert buffered[0].event_type == "event"
345
- assert buffered[0].data == {"a": 1}
346
- assert buffered[0].correlation_id == "corr-1"
347
-
348
-
349
-@pytest.mark.asyncio
350
-async def test_buffer_overflow_drops_oldest(monkeypatch):
351
- socketio = FakeSocketIOServer()
352
- manager = WebSocketManager(socketio, threading.RLock())
353
-
354
- await manager.handle_connect(NAMESPACE, "offline")
355
- await manager.handle_disconnect(NAMESPACE, "offline")
356
-
357
- monkeypatch.setattr("helpers.websocket_manager.BUFFER_MAX_SIZE", 2)
358
-
359
- await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 0})
360
- await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 1})
361
- await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 2})
362
-
363
- buffer = manager.buffers[(NAMESPACE, "offline")]
364
- assert len(buffer) == 2
365
- assert buffer[0].data["idx"] == 1
366
- assert buffer[1].data["idx"] == 2
367
-
368
-
369
-@pytest.mark.asyncio
370
-async def test_expired_buffer_entries_are_discarded(monkeypatch):
371
- socketio = FakeSocketIOServer()
372
- manager = WebSocketManager(socketio, threading.RLock())
373
-
374
- await manager.handle_connect(NAMESPACE, "sid-expired")
375
- await manager.handle_disconnect(NAMESPACE, "sid-expired")
376
-
377
- from datetime import timedelta, timezone, datetime
378
-
379
- past = datetime.now(timezone.utc) - (BUFFER_TTL + timedelta(seconds=5))
380
- future = past + BUFFER_TTL + timedelta(seconds=10)
381
-
382
- await manager.emit_to(NAMESPACE, "sid-expired", "event", {"a": 1})
383
- manager.buffers[(NAMESPACE, "sid-expired")][0].timestamp = past
384
-
385
- socketio.emit.reset_mock()
386
-
387
- monkeypatch.setattr(
388
- "helpers.websocket_manager._utcnow",
389
- lambda: future,
390
- )
391
- await manager.handle_connect(NAMESPACE, "sid-expired")
392
-
393
- assert socketio.emit.await_count == 0
394
- assert (NAMESPACE, "sid-expired") not in manager.buffers
395
-
396
-
397
-@pytest.mark.asyncio
398
-async def test_flush_buffer_delivers_and_logs(monkeypatch):
399
- socketio = FakeSocketIOServer()
400
- manager = WebSocketManager(socketio, threading.RLock())
401
- await manager.handle_connect(NAMESPACE, "sid-1")
402
- await manager.handle_disconnect(NAMESPACE, "sid-1")
403
-
404
- await manager.emit_to(NAMESPACE, "sid-1", "event", {"a": 1})
405
-
406
- await manager.handle_connect(NAMESPACE, "sid-1")
407
-
408
- assert len(socketio.emit.await_args_list) == 1
409
- awaited_call = socketio.emit.await_args_list[0]
410
- assert awaited_call.args[0] == "event"
411
- envelope = awaited_call.args[1]
412
- assert envelope["data"] == {"a": 1}
413
- assert "eventId" in envelope and "handlerId" in envelope and "ts" in envelope
414
- assert awaited_call.kwargs == {"to": "sid-1", "namespace": NAMESPACE}
415
- assert (NAMESPACE, "sid-1") not in manager.buffers
416
-
417
-
418
-@pytest.mark.asyncio
419
-async def test_broadcast_excludes_multiple_sids():
420
- socketio = FakeSocketIOServer()
421
- manager = WebSocketManager(socketio, threading.RLock())
422
-
423
- for sid in ("sid-1", "sid-2", "sid-3"):
424
- await manager.handle_connect(NAMESPACE, sid)
425
-
426
- await manager.broadcast(
427
- NAMESPACE,
428
- "event",
429
- {"foo": "bar"},
430
- exclude_sids={"sid-1", "sid-3"},
431
- handler_id="custom.broadcast",
432
- correlation_id="corr-b",
433
- )
434
-
435
- assert len(socketio.emit.await_args_list) == 1
436
- awaited_call = socketio.emit.await_args_list[0]
437
- assert awaited_call.args[0] == "event"
438
- envelope = awaited_call.args[1]
439
- assert envelope["data"] == {"foo": "bar"}
440
- assert envelope["handlerId"] == "custom.broadcast"
441
- assert envelope["correlationId"] == "corr-b"
442
- assert "eventId" in envelope and "ts" in envelope
443
- assert awaited_call.kwargs == {"to": "sid-2", "namespace": NAMESPACE}
444
-
445
-
446
-@pytest.mark.asyncio
447
-async def test_emit_to_wraps_envelope_with_metadata():
448
- socketio = FakeSocketIOServer()
449
- manager = WebSocketManager(socketio, threading.RLock())
450
- await manager.handle_connect(NAMESPACE, "sid-meta")
451
-
452
- await manager.emit_to(
453
- NAMESPACE,
454
- "sid-meta",
455
- "meta_event",
456
- {"payload": True},
457
- handler_id="custom.handler",
458
- correlation_id="corr-meta",
459
- )
460
-
461
- socketio.emit.assert_awaited_once()
462
- args, kwargs = socketio.emit.await_args_list[0]
463
- assert args[0] == "meta_event"
464
- envelope = args[1]
465
- assert envelope["handlerId"] == "custom.handler"
466
- assert envelope["correlationId"] == "corr-meta"
467
- assert envelope["data"] == {"payload": True}
468
- assert kwargs == {"to": "sid-meta", "namespace": NAMESPACE}
469
-
470
-
471
-@pytest.mark.asyncio
472
-async def test_timestamps_are_timezone_aware():
473
- socketio = FakeSocketIOServer()
474
- manager = WebSocketManager(socketio, threading.RLock())
475
-
476
- await manager.handle_connect(NAMESPACE, "sid-utc")
477
- info = manager.connections[(NAMESPACE, "sid-utc")]
478
-
479
- assert info.connected_at.tzinfo is not None
480
- assert info.last_activity.tzinfo is not None
481
-
482
- with patch("helpers.websocket_manager._utcnow") as mocked_now:
483
- mocked_now.return_value = info.last_activity
484
- await manager.route_event(NAMESPACE, "unknown", {}, "sid-utc")
485
- assert info.last_activity.tzinfo is not None
486
-
487
-class DuplicateHandler(WebSocketHandler):
488
- @classmethod
489
- def get_event_types(cls) -> list[str]:
490
- return ["dup_event"]
491
-
492
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
493
- return {"handledBy": self.identifier}
494
-
495
-
496
-class AnotherDuplicateHandler(WebSocketHandler):
497
- @classmethod
498
- def get_event_types(cls) -> list[str]:
499
- return ["dup_event"]
500
-
501
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
502
- return {"handledBy": self.identifier}
503
-
504
-
505
-def test_register_handlers_warns_on_duplicates(monkeypatch):
506
- socketio = FakeSocketIOServer()
507
- manager = WebSocketManager(socketio, threading.RLock())
508
-
509
- warnings: list[str] = []
510
-
511
- def capture_warning(message: str) -> None:
512
- warnings.append(message)
513
-
514
- monkeypatch.setattr(
515
- "helpers.print_style.PrintStyle.warning", staticmethod(capture_warning)
516
- )
517
-
518
- DuplicateHandler._reset_instance_for_testing()
519
- AnotherDuplicateHandler._reset_instance_for_testing()
520
- handler_a = DuplicateHandler.get_instance(socketio, threading.RLock())
521
- handler_b = AnotherDuplicateHandler.get_instance(socketio, threading.RLock())
522
-
523
- manager.register_handlers({NAMESPACE: [handler_a, handler_b]})
524
-
525
- assert any("Duplicate handler registration" in msg for msg in warnings)
526
-
527
-
528
-class NonDictHandler(WebSocketHandler):
529
- @classmethod
530
- def get_event_types(cls) -> list[str]:
531
- return ["non_dict"]
532
-
533
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
534
- return "raw-value"
535
-
536
-
537
-@pytest.mark.asyncio
538
-async def test_route_event_standardizes_success_payload():
539
- socketio = FakeSocketIOServer()
540
- manager = WebSocketManager(socketio, threading.RLock())
541
-
542
- NonDictHandler._reset_instance_for_testing()
543
- handler = NonDictHandler.get_instance(socketio, threading.RLock())
544
- manager.register_handlers({NAMESPACE: [handler]})
545
-
546
- response = await manager.route_event(NAMESPACE, "non_dict", {}, "sid-123")
547
-
548
- assert len(response["results"]) == 1
549
- assert response["results"][0]["ok"] is True
550
- assert response["results"][0]["data"] == {"result": "raw-value"}
551
-
552
-
553
-class ErrorHandler(WebSocketHandler):
554
- @classmethod
555
- def get_event_types(cls) -> list[str]:
556
- return ["boom"]
557
-
558
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
559
- raise RuntimeError("BOOM")
560
-
561
-
562
-class ResultHandler(WebSocketHandler):
563
- @classmethod
564
- def get_event_types(cls) -> list[str]: # pragma: no cover - simple declaration
565
- return ["result_event", "result_error"]
566
-
567
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
568
- if event_type == "result_event":
569
- return WebSocketResult.ok({"sid": sid}, correlation_id="explicit", duration_ms=1.234)
570
- return WebSocketResult.error(
571
- code="E_RESULT",
572
- message="boom",
573
- details="test",
574
- )
575
-
576
-
577
-@pytest.mark.asyncio
578
-async def test_route_event_standardizes_error_payload():
579
- socketio = FakeSocketIOServer()
580
- manager = WebSocketManager(socketio, threading.RLock())
581
-
582
- ErrorHandler._reset_instance_for_testing()
583
- handler = ErrorHandler.get_instance(socketio, threading.RLock())
584
- manager.register_handlers({NAMESPACE: [handler]})
585
-
586
- response = await manager.route_event(NAMESPACE, "boom", {}, "sid-123")
587
-
588
- assert len(response["results"]) == 1
589
- payload = response["results"][0]
590
- assert payload["ok"] is False
591
- assert payload["error"]["code"] == "HANDLER_ERROR"
592
- assert payload["error"]["error"] == "Internal server error"
593
-
594
-
595
-@pytest.mark.asyncio
596
-async def test_route_event_accepts_websocket_result_instances():
597
- socketio = FakeSocketIOServer()
598
- manager = WebSocketManager(socketio, threading.RLock())
599
-
600
- ResultHandler._reset_instance_for_testing()
601
- handler = ResultHandler.get_instance(socketio, threading.RLock())
602
- manager.register_handlers({NAMESPACE: [handler]})
603
-
604
- response = await manager.route_event(NAMESPACE, "result_event", {}, "sid-123")
605
-
606
- assert response["results"]
607
- payload = response["results"][0]
608
- assert payload["ok"] is True
609
- assert payload["data"] == {"sid": "sid-123"}
610
- assert payload["correlationId"] == "explicit"
611
- assert payload["durationMs"] == pytest.approx(1.234, rel=1e-3)
612
-
613
-
614
-@pytest.mark.asyncio
615
-async def test_route_event_preserves_websocket_result_errors():
616
- socketio = FakeSocketIOServer()
617
- manager = WebSocketManager(socketio, threading.RLock())
618
-
619
- ResultHandler._reset_instance_for_testing()
620
- handler = ResultHandler.get_instance(socketio, threading.RLock())
621
- manager.register_handlers({NAMESPACE: [handler]})
622
-
623
- response = await manager.route_event(NAMESPACE, "result_error", {}, "sid-123")
624
-
625
- payload = response["results"][0]
626
- assert payload["ok"] is False
627
- assert payload["error"] == {"code": "E_RESULT", "error": "boom", "details": "test"}
628
-
629
-
630
-class AlphaFilterHandler(WebSocketHandler):
631
- @classmethod
632
- def get_event_types(cls) -> list[str]:
633
- return ["filter_event"]
634
-
635
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
636
- return {"handledBy": self.identifier, "sid": sid}
637
-
638
-
639
-class BetaFilterHandler(WebSocketHandler):
640
- @classmethod
641
- def get_event_types(cls) -> list[str]:
642
- return ["filter_event"]
643
-
644
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
645
- return {"handledBy": self.identifier, "sid": sid}
646
-
647
-
648
-@pytest.mark.asyncio
649
-async def test_route_event_include_handlers_filters_results():
650
- socketio = FakeSocketIOServer()
651
- manager = WebSocketManager(socketio, threading.RLock())
652
-
653
- AlphaFilterHandler._reset_instance_for_testing()
654
- BetaFilterHandler._reset_instance_for_testing()
655
- alpha = AlphaFilterHandler.get_instance(socketio, threading.RLock())
656
- beta = BetaFilterHandler.get_instance(socketio, threading.RLock())
657
- manager.register_handlers({NAMESPACE: [alpha, beta]})
658
- await manager.handle_connect(NAMESPACE, "sid-filter")
659
-
660
- response = await manager.route_event(
661
- NAMESPACE,
662
- "filter_event",
663
- {
664
- "includeHandlers": [alpha.identifier],
665
- "payload": True,
666
- },
667
- "sid-filter",
668
- )
669
-
670
- assert response["correlationId"]
671
- results = response["results"]
672
- assert len(results) == 1
673
- assert results[0]["handlerId"] == alpha.identifier
674
- assert results[0]["data"]["handledBy"] == alpha.identifier
675
-
676
-
677
-@pytest.mark.asyncio
678
-async def test_route_event_rejects_exclude_handlers_without_permission():
679
- socketio = FakeSocketIOServer()
680
- manager = WebSocketManager(socketio, threading.RLock())
681
-
682
- AlphaFilterHandler._reset_instance_for_testing()
683
- handler = AlphaFilterHandler.get_instance(socketio, threading.RLock())
684
- manager.register_handlers({NAMESPACE: [handler]})
685
- await manager.handle_connect(NAMESPACE, "sid-exclude")
686
-
687
- response = await manager.route_event(
688
- NAMESPACE,
689
- "filter_event",
690
- {"excludeHandlers": [handler.identifier]},
691
- "sid-exclude",
692
- )
693
-
694
- result = response["results"][0]
695
- assert result["error"]["code"] == "INVALID_FILTER"
696
- assert "excludeHandlers" in result["error"]["error"]
697
-
698
-
699
-@pytest.mark.asyncio
700
-async def test_route_event_all_respects_exclude_handlers():
701
- socketio = FakeSocketIOServer()
702
- manager = WebSocketManager(socketio, threading.RLock())
703
-
704
- AlphaFilterHandler._reset_instance_for_testing()
705
- BetaFilterHandler._reset_instance_for_testing()
706
- alpha = AlphaFilterHandler.get_instance(socketio, threading.RLock())
707
- beta = BetaFilterHandler.get_instance(socketio, threading.RLock())
708
- manager.register_handlers({NAMESPACE: [alpha, beta]})
709
-
710
- await manager.handle_connect(NAMESPACE, "sid-a")
711
- await manager.handle_connect(NAMESPACE, "sid-b")
712
-
713
- aggregated = await manager.route_event_all(
714
- NAMESPACE,
715
- "filter_event",
716
- {"excludeHandlers": [beta.identifier]},
717
- handler_id="test.manager",
718
- )
719
-
720
- assert aggregated
721
- for entry in aggregated:
722
- assert entry["correlationId"]
723
- assert entry["results"]
724
- assert all(result["handlerId"] == alpha.identifier for result in entry["results"])
725
-
726
-
727
-@pytest.mark.asyncio
728
-async def test_route_event_preserves_correlation_id():
729
- socketio = FakeSocketIOServer()
730
- manager = WebSocketManager(socketio, threading.RLock())
731
-
732
- results = []
733
- DummyHandler._reset_instance_for_testing()
734
- handler = DummyHandler.get_instance(socketio, threading.RLock(), results)
735
- manager.register_handlers({NAMESPACE: [handler]})
736
- await manager.handle_connect(NAMESPACE, "sid-correlation")
737
-
738
- response = await manager.route_event(
739
- NAMESPACE,
740
- "dummy",
741
- {"foo": "bar", "correlationId": "manual-correlation"},
742
- "sid-correlation",
743
- )
744
-
745
- assert response["correlationId"] == "manual-correlation"
746
- result = response["results"][0]
747
- assert result["correlationId"] == "manual-correlation"
748
-
749
-
750
-@pytest.mark.asyncio
751
-async def test_request_preserves_explicit_correlation_id():
752
- socketio = FakeSocketIOServer()
753
- manager = WebSocketManager(socketio, threading.RLock())
754
-
755
- DummyHandler._reset_instance_for_testing()
756
- handler = DummyHandler.get_instance(socketio, threading.RLock(), [])
757
- manager.register_handlers({NAMESPACE: [handler]})
758
- await manager.handle_connect(NAMESPACE, "sid-request")
759
-
760
- response = await manager.request_for_sid(
761
- namespace=NAMESPACE,
762
- sid="sid-request",
763
- event_type="dummy",
764
- data={"payload": True, "correlationId": "req-correlation"},
765
- handler_id="tester",
766
- )
767
-
768
- assert response["correlationId"] == "req-correlation"
769
- result = response["results"][0]
770
- assert result["correlationId"] == "req-correlation"
771
-
772
-
773
-@pytest.mark.asyncio
774
-async def test_request_all_entries_include_correlation_id():
775
- socketio = FakeSocketIOServer()
776
- manager = WebSocketManager(socketio, threading.RLock())
777
-
778
- DummyHandler._reset_instance_for_testing()
779
- handler = DummyHandler.get_instance(socketio, threading.RLock(), [])
780
- manager.register_handlers({NAMESPACE: [handler]})
781
-
782
- await manager.handle_connect(NAMESPACE, "sid-1")
783
- await manager.handle_connect(NAMESPACE, "sid-2")
784
-
785
- aggregated = await manager.route_event_all(
786
- NAMESPACE,
787
- "dummy",
788
- {"value": 1, "correlationId": "agg-correlation"},
789
- )
790
-
791
- assert aggregated
792
- for entry in aggregated:
793
- assert entry["correlationId"] == "agg-correlation"
794
- assert entry["results"]
795
- assert entry["results"][0]["correlationId"] == "agg-correlation"
796
-
797
-
798
-def test_debug_logging_respects_runtime_flag(monkeypatch):
799
- socketio = FakeSocketIOServer()
800
- manager = WebSocketManager(socketio, threading.RLock())
801
-
802
- logs: list[str] = []
803
-
804
- def capture(message: str) -> None:
805
- logs.append(message)
806
-
807
- monkeypatch.setattr("helpers.print_style.PrintStyle.debug", staticmethod(capture))
808
- monkeypatch.setattr("helpers.websocket_manager.runtime.is_development", lambda: False)
809
-
810
- manager._debug("should-not-log") # noqa: SLF001
811
- assert logs == []
812
-
813
- monkeypatch.setattr("helpers.websocket_manager.runtime.is_development", lambda: True)
814
- manager._debug("should-log") # noqa: SLF001
815
- assert logs == ["should-log"]
816
-
817
-
818
-@pytest.mark.asyncio
819
-async def test_diagnostic_event_emitted_for_inbound():
820
- socketio = FakeSocketIOServer()
821
- manager = WebSocketManager(socketio, threading.RLock())
822
-
823
- results: list[dict[str, Any]] = []
824
- DummyHandler._reset_instance_for_testing()
825
- handler = DummyHandler.get_instance(socketio, threading.RLock(), results)
826
- manager.register_handlers({NAMESPACE: [handler]})
827
-
828
- await manager.handle_connect(NAMESPACE, "observer")
829
- assert manager.register_diagnostic_watcher(NAMESPACE, "observer") is True
830
- await manager.handle_connect(NAMESPACE, "sid-client")
831
-
832
- await manager.route_event(NAMESPACE, "dummy", {"payload": "value"}, "sid-client")
833
-
834
- emitted_events = [call.args[0] for call in socketio.emit.await_args_list]
835
- assert DIAGNOSTIC_EVENT in emitted_events
836
-
837
-
838
-@pytest.mark.asyncio
839
-async def test_lifecycle_events_broadcast(monkeypatch):
840
- socketio = FakeSocketIOServer()
841
- manager = WebSocketManager(socketio, threading.RLock())
842
-
843
- broadcast_mock = AsyncMock()
844
- monkeypatch.setattr(manager, "broadcast", broadcast_mock)
845
-
846
- await manager.handle_connect(NAMESPACE, "sid-life")
847
- await asyncio.sleep(0)
848
- await manager.handle_disconnect(NAMESPACE, "sid-life")
849
- await asyncio.sleep(0)
850
-
851
- events = [call.args[1] for call in broadcast_mock.await_args_list]
852
- assert LIFECYCLE_CONNECT_EVENT in events
853
- assert LIFECYCLE_DISCONNECT_EVENT in events
tests/test_websocket_namespace_discovery.py
deleted
-225
@@ -1,225 +0,0 @@
1
-import asyncio
2
-import contextlib
3
-import socket
4
-import sys
5
-from pathlib import Path
6
-from typing import Any, AsyncIterator
7
-
8
-import pytest
9
-
10
-PROJECT_ROOT = Path(__file__).resolve().parents[1]
11
-if str(PROJECT_ROOT) not in sys.path:
12
- sys.path.insert(0, str(PROJECT_ROOT))
13
-
14
-
15
-@contextlib.asynccontextmanager
16
-async def _run_asgi_app(app: Any) -> AsyncIterator[str]:
17
- import uvicorn
18
-
19
- sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
20
- sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
21
- sock.bind(("127.0.0.1", 0))
22
- sock.listen(128)
23
-
24
- port = sock.getsockname()[1]
25
-
26
- config = uvicorn.Config(
27
- app,
28
- host="127.0.0.1",
29
- port=port,
30
- log_level="warning",
31
- access_log=False,
32
- lifespan="off",
33
- )
34
- server = uvicorn.Server(config)
35
- server.install_signal_handlers = lambda: None # type: ignore[method-assign]
36
-
37
- task = asyncio.create_task(server.serve(sockets=[sock]))
38
- try:
39
- while not server.started:
40
- await asyncio.sleep(0.01)
41
- yield f"http://127.0.0.1:{port}"
42
- finally:
43
- server.should_exit = True
44
- try:
45
- await asyncio.wait_for(task, timeout=5)
46
- finally:
47
- sock.close()
48
-
49
-
50
-def _write_handler_module(path: Path, class_name: str, event_type: str) -> None:
51
- path.write_text(
52
- "\n".join(
53
- [
54
- "from __future__ import annotations",
55
- "",
56
- "from typing import Any",
57
- "",
58
- "from helpers.websocket import WebSocketHandler",
59
- "",
60
- f"class {class_name}(WebSocketHandler):",
61
- " @classmethod",
62
- " def requires_auth(cls) -> bool:",
63
- " return False",
64
- "",
65
- " @classmethod",
66
- " def requires_csrf(cls) -> bool:",
67
- " return False",
68
- "",
69
- " @classmethod",
70
- " def get_event_types(cls) -> list[str]:",
71
- f" return ['{event_type}']",
72
- "",
73
- " async def process_event(self, event_type: str, data: dict[str, Any], sid: str):",
74
- " return {'ok': True}",
75
- "",
76
- ]
77
- ),
78
- encoding="utf-8",
79
- )
80
-
81
-
82
-def test_discovery_supports_folder_entries_and_ignores_deeper_nesting(tmp_path: Path) -> None:
83
- from helpers.websocket_namespace_discovery import discover_websocket_namespaces
84
-
85
- folder = tmp_path / "orders"
86
- folder.mkdir()
87
- _write_handler_module(folder / "orders.py", "OrdersHandler", "orders_request")
88
-
89
- # Deeper nesting must be ignored (and must not be imported).
90
- nested = folder / "nested"
91
- nested.mkdir()
92
- (nested / "boom.py").write_text("raise RuntimeError('should-not-import')\n", encoding="utf-8")
93
-
94
- discoveries = discover_websocket_namespaces(handlers_folder=str(tmp_path), include_root_default=False)
95
- by_ns = {d.namespace: d for d in discoveries}
96
-
97
- assert "/orders" in by_ns
98
- entry = by_ns["/orders"]
99
- assert [cls.__name__ for cls in entry.handler_classes] == ["OrdersHandler"]
100
-
101
-
102
-def test_discovery_folder_suffix_handler_stripped(tmp_path: Path) -> None:
103
- from helpers.websocket_namespace_discovery import discover_websocket_namespaces
104
-
105
- folder = tmp_path / "sales_handler"
106
- folder.mkdir()
107
- _write_handler_module(folder / "main.py", "SalesHandler", "sales_request")
108
-
109
- discoveries = discover_websocket_namespaces(handlers_folder=str(tmp_path), include_root_default=False)
110
- namespaces = {d.namespace for d in discoveries}
111
- assert "/sales" in namespaces
112
-
113
-
114
-def test_discovery_empty_folder_warns_and_treats_namespace_unregistered(tmp_path: Path, monkeypatch) -> None:
115
- from flask import Flask
116
- import socketio
117
-
118
- from helpers.websocket_manager import WebSocketManager
119
- from helpers.websocket_namespace_discovery import discover_websocket_namespaces
120
- from run_ui import configure_websocket_namespaces
121
-
122
- empty = tmp_path / "empty"
123
- empty.mkdir()
124
- (empty / "__init__.py").write_text("# init\n", encoding="utf-8")
125
-
126
- warnings: list[str] = []
127
-
128
- def _warn(message: str) -> None:
129
- warnings.append(message)
130
-
131
- monkeypatch.setattr("helpers.print_style.PrintStyle.warning", staticmethod(_warn))
132
-
133
- discoveries = discover_websocket_namespaces(handlers_folder=str(tmp_path), include_root_default=False)
134
- assert "/empty" not in {d.namespace for d in discoveries}
135
- assert any("empty" in msg.lower() for msg in warnings)
136
-
137
- # Integration check: treat as unregistered -> UNKNOWN_NAMESPACE connect_error.
138
- app = Flask("test_empty_folder_unregistered")
139
- app.secret_key = "test-secret"
140
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
141
- lock = __import__("threading").RLock()
142
- manager = WebSocketManager(sio, lock)
143
-
144
- handlers_by_namespace: dict[str, list[Any]] = {}
145
- for discovery in discoveries:
146
- handlers_by_namespace[discovery.namespace] = [
147
- cls.get_instance(sio, lock) for cls in discovery.handler_classes
148
- ]
149
-
150
- configure_websocket_namespaces(
151
- webapp=app,
152
- socketio_server=sio,
153
- websocket_manager=manager,
154
- handlers_by_namespace=handlers_by_namespace,
155
- )
156
-
157
- asgi_app = socketio.ASGIApp(sio)
158
- async def _run() -> None:
159
- async with _run_asgi_app(asgi_app) as base_url:
160
- client = socketio.AsyncClient()
161
- connect_error_fut: asyncio.Future[Any] = asyncio.get_running_loop().create_future()
162
-
163
- async def _on_connect_error(data: Any) -> None:
164
- if not connect_error_fut.done():
165
- connect_error_fut.set_result(data)
166
-
167
- client.on("connect_error", _on_connect_error, namespace="/empty")
168
- try:
169
- with pytest.raises(socketio.exceptions.ConnectionError):
170
- await client.connect(base_url, namespaces=["/empty"])
171
- err = await asyncio.wait_for(connect_error_fut, timeout=2)
172
- assert err["message"] == "UNKNOWN_NAMESPACE"
173
- assert err["data"]["namespace"] == "/empty"
174
- finally:
175
- try:
176
- await client.disconnect()
177
- except Exception:
178
- pass
179
-
180
- asyncio.run(_run())
181
-
182
-
183
-def test_discovery_invalid_modules_fail_fast_with_descriptive_errors(tmp_path: Path) -> None:
184
- from helpers.websocket_namespace_discovery import discover_websocket_namespaces
185
-
186
- # 0 handlers in a *_handler.py module
187
- (tmp_path / "bad_handler.py").write_text(
188
- "class NotAHandler:\n pass\n", encoding="utf-8"
189
- )
190
- with pytest.raises(RuntimeError) as excinfo:
191
- discover_websocket_namespaces(handlers_folder=str(tmp_path), include_root_default=False)
192
- assert "defines no WebSocketHandler subclasses" in str(excinfo.value)
193
-
194
- # 2+ handlers in a *_handler.py module
195
- tmp_path.joinpath("bad_handler.py").unlink()
196
- (tmp_path / "two_handler.py").write_text(
197
- "\n".join(
198
- [
199
- "from helpers.websocket import WebSocketHandler",
200
- "class A(WebSocketHandler):",
201
- " @classmethod",
202
- " def requires_auth(cls): return False",
203
- " @classmethod",
204
- " def requires_csrf(cls): return False",
205
- " @classmethod",
206
- " def get_event_types(cls): return ['two_a']",
207
- " async def process_event(self, event_type, data, sid): return {'ok': True}",
208
- "class B(WebSocketHandler):",
209
- " @classmethod",
210
- " def requires_auth(cls): return False",
211
- " @classmethod",
212
- " def requires_csrf(cls): return False",
213
- " @classmethod",
214
- " def get_event_types(cls): return ['two_b']",
215
- " async def process_event(self, event_type, data, sid): return {'ok': True}",
216
- "",
217
- ]
218
- ),
219
- encoding="utf-8",
220
- )
221
- with pytest.raises(RuntimeError) as excinfo2:
222
- discover_websocket_namespaces(handlers_folder=str(tmp_path), include_root_default=False)
223
- message = str(excinfo2.value)
224
- assert "defines multiple WebSocketHandler subclasses" in message
225
- assert "A" in message and "B" in message
tests/test_websocket_namespace_security.py
deleted
-464
@@ -1,464 +0,0 @@
1
-import asyncio
2
-import contextlib
3
-import socket
4
-import sys
5
-import threading
6
-from pathlib import Path
7
-from typing import Any, AsyncIterator
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
-
16
-@contextlib.asynccontextmanager
17
-async def _run_asgi_app(app: Any) -> AsyncIterator[str]:
18
- import uvicorn
19
-
20
- sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
21
- sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
22
- sock.bind(("127.0.0.1", 0))
23
- sock.listen(128)
24
-
25
- port = sock.getsockname()[1]
26
-
27
- config = uvicorn.Config(
28
- app,
29
- host="127.0.0.1",
30
- port=port,
31
- log_level="warning",
32
- access_log=False,
33
- lifespan="off",
34
- )
35
- server = uvicorn.Server(config)
36
- server.install_signal_handlers = lambda: None # type: ignore[method-assign]
37
-
38
- task = asyncio.create_task(server.serve(sockets=[sock]))
39
- try:
40
- while not server.started:
41
- await asyncio.sleep(0.01)
42
- yield f"http://127.0.0.1:{port}"
43
- finally:
44
- server.should_exit = True
45
- try:
46
- await asyncio.wait_for(task, timeout=5)
47
- finally:
48
- sock.close()
49
-
50
-
51
-def _make_session_cookie(app: Any, data: dict[str, Any]) -> str:
52
- from flask.sessions import SecureCookieSessionInterface
53
-
54
- serializer = SecureCookieSessionInterface().get_signing_serializer(app)
55
- assert serializer is not None
56
- return serializer.dumps(data)
57
-
58
-
59
-@pytest.mark.asyncio
60
-async def test_connect_security_is_computed_per_namespace_and_enforced(monkeypatch) -> None:
61
- from flask import Flask
62
- import socketio
63
-
64
- from helpers.websocket import WebSocketHandler
65
- from helpers.websocket_manager import WebSocketManager
66
- from helpers import runtime
67
- from run_ui import configure_websocket_namespaces
68
-
69
- class OpenHandler(WebSocketHandler):
70
- @classmethod
71
- def requires_auth(cls) -> bool:
72
- return False
73
-
74
- @classmethod
75
- def requires_csrf(cls) -> bool:
76
- return False
77
-
78
- @classmethod
79
- def get_event_types(cls) -> list[str]:
80
- return ["open_ping"]
81
-
82
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str) -> dict[str, Any]:
83
- return {"ok": True}
84
-
85
- class SecureHandler(WebSocketHandler):
86
- @classmethod
87
- def get_event_types(cls) -> list[str]:
88
- return ["secure_ping"]
89
-
90
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str) -> dict[str, Any]:
91
- return {"ok": True}
92
-
93
- OpenHandler._reset_instance_for_testing()
94
- SecureHandler._reset_instance_for_testing()
95
-
96
- monkeypatch.setattr("helpers.login.get_credentials_hash", lambda: "hash")
97
-
98
- webapp = Flask("test_websocket_namespace_security")
99
- webapp.secret_key = "test-secret"
100
-
101
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
102
- lock = threading.RLock()
103
- manager = WebSocketManager(sio, lock)
104
- handlers_by_namespace = {
105
- "/open": [OpenHandler.get_instance(sio, lock)],
106
- "/secure": [SecureHandler.get_instance(sio, lock)],
107
- }
108
-
109
- configure_websocket_namespaces(
110
- webapp=webapp,
111
- socketio_server=sio,
112
- websocket_manager=manager,
113
- handlers_by_namespace=handlers_by_namespace,
114
- )
115
-
116
- asgi_app = socketio.ASGIApp(sio)
117
-
118
- async with _run_asgi_app(asgi_app) as base_url:
119
- # Open namespace should not require auth/csrf (but Origin validation is always enforced).
120
- open_client = socketio.AsyncClient()
121
- await open_client.connect(
122
- base_url,
123
- namespaces=["/open"],
124
- headers={"Origin": base_url},
125
- wait_timeout=2,
126
- )
127
- try:
128
- res = await open_client.call("open_ping", {}, namespace="/open", timeout=2)
129
- assert isinstance(res, dict)
130
- assert res.get("results")
131
- res_unhandled = await open_client.call("unhandled_event", {"x": 1}, namespace="/open", timeout=2)
132
- assert res_unhandled["results"]
133
- assert res_unhandled["results"][0]["ok"] is False
134
- assert res_unhandled["results"][0]["error"]["code"] == "NO_HANDLERS"
135
- finally:
136
- await open_client.disconnect()
137
-
138
- # Secure namespace rejects without valid session+csrf when credentials are configured.
139
- secure_client = socketio.AsyncClient()
140
- with pytest.raises(socketio.exceptions.ConnectionError):
141
- await secure_client.connect(
142
- base_url,
143
- namespaces=["/secure"],
144
- headers={"Origin": base_url},
145
- wait_timeout=2,
146
- )
147
- await secure_client.disconnect()
148
-
149
- # Secure namespace accepts valid session + auth csrf_token + runtime-scoped csrf cookie.
150
- csrf_token = "csrf-1"
151
- session_cookie = _make_session_cookie(
152
- webapp,
153
- {
154
- "authentication": "hash",
155
- "csrf_token": csrf_token,
156
- "user_id": "u1",
157
- },
158
- )
159
- session_cookie_name = webapp.config.get("SESSION_COOKIE_NAME", "session")
160
- csrf_cookie_name = f"csrf_token_{runtime.get_runtime_id()}"
161
- cookie_header = f"{session_cookie_name}={session_cookie}; {csrf_cookie_name}={csrf_token}"
162
-
163
- secure_client_ok = socketio.AsyncClient()
164
- await secure_client_ok.connect(
165
- base_url,
166
- namespaces=["/secure"],
167
- headers={"Origin": base_url, "Cookie": cookie_header},
168
- auth={"csrf_token": csrf_token},
169
- wait_timeout=2,
170
- )
171
- try:
172
- res2 = await secure_client_ok.call("secure_ping", {}, namespace="/secure", timeout=2)
173
- assert isinstance(res2, dict)
174
- assert res2.get("results")
175
- finally:
176
- await secure_client_ok.disconnect()
177
-
178
-
179
-@pytest.mark.asyncio
180
-async def test_unknown_namespace_rejected_with_deterministic_connect_error_payload() -> None:
181
- from flask import Flask
182
- import socketio
183
-
184
- from helpers.websocket import WebSocketHandler
185
- from helpers.websocket_manager import WebSocketManager
186
- from run_ui import configure_websocket_namespaces
187
-
188
- class OpenHandler(WebSocketHandler):
189
- @classmethod
190
- def requires_auth(cls) -> bool:
191
- return False
192
-
193
- @classmethod
194
- def requires_csrf(cls) -> bool:
195
- return False
196
-
197
- @classmethod
198
- def get_event_types(cls) -> list[str]:
199
- return ["open_ping"]
200
-
201
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str) -> dict[str, Any]:
202
- return {"ok": True}
203
-
204
- OpenHandler._reset_instance_for_testing()
205
-
206
- webapp = Flask("test_unknown_namespace_rejection")
207
- webapp.secret_key = "test-secret"
208
-
209
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
210
- lock = threading.RLock()
211
- manager = WebSocketManager(sio, lock)
212
-
213
- configure_websocket_namespaces(
214
- webapp=webapp,
215
- socketio_server=sio,
216
- websocket_manager=manager,
217
- handlers_by_namespace={"/open": [OpenHandler.get_instance(sio, lock)]},
218
- )
219
-
220
- asgi_app = socketio.ASGIApp(sio)
221
-
222
- async with _run_asgi_app(asgi_app) as base_url:
223
- client = socketio.AsyncClient()
224
- connect_error_fut: asyncio.Future[Any] = asyncio.get_running_loop().create_future()
225
-
226
- async def _on_connect_error(data: Any) -> None:
227
- if not connect_error_fut.done():
228
- connect_error_fut.set_result(data)
229
-
230
- client.on("connect_error", _on_connect_error, namespace="/unknown")
231
-
232
- try:
233
- with pytest.raises(socketio.exceptions.ConnectionError):
234
- await client.connect(base_url, namespaces=["/unknown"])
235
-
236
- err = await asyncio.wait_for(connect_error_fut, timeout=2)
237
- assert err["message"] == "UNKNOWN_NAMESPACE"
238
- assert err["data"] == {"code": "UNKNOWN_NAMESPACE", "namespace": "/unknown"}
239
- finally:
240
- try:
241
- await client.disconnect()
242
- except Exception:
243
- pass
244
-
245
-
246
-@pytest.mark.asyncio
247
-async def test_secure_namespace_rejects_missing_auth_even_with_valid_csrf(monkeypatch) -> None:
248
- from flask import Flask
249
- import socketio
250
-
251
- from helpers.websocket import WebSocketHandler
252
- from helpers.websocket_manager import WebSocketManager
253
- from helpers import runtime
254
- from run_ui import configure_websocket_namespaces
255
-
256
- class SecureHandler(WebSocketHandler):
257
- @classmethod
258
- def get_event_types(cls) -> list[str]:
259
- return ["secure_ping"]
260
-
261
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str) -> dict[str, Any]:
262
- return {"ok": True}
263
-
264
- SecureHandler._reset_instance_for_testing()
265
-
266
- monkeypatch.setattr("helpers.login.get_credentials_hash", lambda: "hash")
267
-
268
- webapp = Flask("test_ws_secure_missing_auth")
269
- webapp.secret_key = "test-secret"
270
-
271
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
272
- lock = threading.RLock()
273
- manager = WebSocketManager(sio, lock)
274
- handlers_by_namespace = {
275
- "/secure": [SecureHandler.get_instance(sio, lock)],
276
- }
277
-
278
- configure_websocket_namespaces(
279
- webapp=webapp,
280
- socketio_server=sio,
281
- websocket_manager=manager,
282
- handlers_by_namespace=handlers_by_namespace,
283
- )
284
-
285
- asgi_app = socketio.ASGIApp(sio)
286
-
287
- async with _run_asgi_app(asgi_app) as base_url:
288
- csrf_token = "csrf-auth-missing"
289
- session_cookie = _make_session_cookie(
290
- webapp,
291
- {
292
- "csrf_token": csrf_token,
293
- "user_id": "u1",
294
- },
295
- )
296
- session_cookie_name = webapp.config.get("SESSION_COOKIE_NAME", "session")
297
- csrf_cookie_name = f"csrf_token_{runtime.get_runtime_id()}"
298
- cookie_header = f"{session_cookie_name}={session_cookie}; {csrf_cookie_name}={csrf_token}"
299
-
300
- client = socketio.AsyncClient()
301
- with pytest.raises(socketio.exceptions.ConnectionError):
302
- await client.connect(
303
- base_url,
304
- namespaces=["/secure"],
305
- headers={"Origin": base_url, "Cookie": cookie_header},
306
- auth={"csrf_token": csrf_token},
307
- wait_timeout=2,
308
- )
309
- await client.disconnect()
310
-
311
-
312
-@pytest.mark.asyncio
313
-async def test_secure_namespace_rejects_invalid_csrf_cookie(monkeypatch) -> None:
314
- from flask import Flask
315
- import socketio
316
-
317
- from helpers.websocket import WebSocketHandler
318
- from helpers.websocket_manager import WebSocketManager
319
- from helpers import runtime
320
- from run_ui import configure_websocket_namespaces
321
-
322
- class SecureHandler(WebSocketHandler):
323
- @classmethod
324
- def get_event_types(cls) -> list[str]:
325
- return ["secure_ping"]
326
-
327
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str) -> dict[str, Any]:
328
- return {"ok": True}
329
-
330
- SecureHandler._reset_instance_for_testing()
331
-
332
- monkeypatch.setattr("helpers.login.get_credentials_hash", lambda: "hash")
333
-
334
- webapp = Flask("test_ws_secure_invalid_csrf")
335
- webapp.secret_key = "test-secret"
336
-
337
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
338
- lock = threading.RLock()
339
- manager = WebSocketManager(sio, lock)
340
- handlers_by_namespace = {
341
- "/secure": [SecureHandler.get_instance(sio, lock)],
342
- }
343
-
344
- configure_websocket_namespaces(
345
- webapp=webapp,
346
- socketio_server=sio,
347
- websocket_manager=manager,
348
- handlers_by_namespace=handlers_by_namespace,
349
- )
350
-
351
- asgi_app = socketio.ASGIApp(sio)
352
-
353
- async with _run_asgi_app(asgi_app) as base_url:
354
- csrf_token = "csrf-good"
355
- session_cookie = _make_session_cookie(
356
- webapp,
357
- {
358
- "authentication": "hash",
359
- "csrf_token": csrf_token,
360
- "user_id": "u1",
361
- },
362
- )
363
- session_cookie_name = webapp.config.get("SESSION_COOKIE_NAME", "session")
364
- csrf_cookie_name = f"csrf_token_{runtime.get_runtime_id()}"
365
- cookie_header = f"{session_cookie_name}={session_cookie}; {csrf_cookie_name}=csrf-bad"
366
-
367
- client = socketio.AsyncClient()
368
- with pytest.raises(socketio.exceptions.ConnectionError):
369
- await client.connect(
370
- base_url,
371
- namespaces=["/secure"],
372
- headers={"Origin": base_url, "Cookie": cookie_header},
373
- auth={"csrf_token": csrf_token},
374
- wait_timeout=2,
375
- )
376
- await client.disconnect()
377
-
378
-
379
-@pytest.mark.asyncio
380
-async def test_csrf_required_without_auth_is_enforced(monkeypatch) -> None:
381
- from flask import Flask
382
- import socketio
383
-
384
- from helpers.websocket import WebSocketHandler
385
- from helpers.websocket_manager import WebSocketManager
386
- from helpers import runtime
387
- from run_ui import configure_websocket_namespaces
388
-
389
- class CsrfOnlyHandler(WebSocketHandler):
390
- @classmethod
391
- def requires_auth(cls) -> bool:
392
- return False
393
-
394
- @classmethod
395
- def requires_csrf(cls) -> bool:
396
- return True
397
-
398
- @classmethod
399
- def get_event_types(cls) -> list[str]:
400
- return ["csrf_only_ping"]
401
-
402
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str) -> dict[str, Any]:
403
- return {"ok": True}
404
-
405
- CsrfOnlyHandler._reset_instance_for_testing()
406
-
407
- monkeypatch.setattr("helpers.login.get_credentials_hash", lambda: None)
408
-
409
- webapp = Flask("test_ws_csrf_only")
410
- webapp.secret_key = "test-secret"
411
-
412
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
413
- lock = threading.RLock()
414
- manager = WebSocketManager(sio, lock)
415
- handlers_by_namespace = {
416
- "/csrf_only": [CsrfOnlyHandler.get_instance(sio, lock)],
417
- }
418
-
419
- configure_websocket_namespaces(
420
- webapp=webapp,
421
- socketio_server=sio,
422
- websocket_manager=manager,
423
- handlers_by_namespace=handlers_by_namespace,
424
- )
425
-
426
- asgi_app = socketio.ASGIApp(sio)
427
-
428
- async with _run_asgi_app(asgi_app) as base_url:
429
- client = socketio.AsyncClient()
430
- with pytest.raises(socketio.exceptions.ConnectionError):
431
- await client.connect(
432
- base_url,
433
- namespaces=["/csrf_only"],
434
- headers={"Origin": base_url},
435
- wait_timeout=2,
436
- )
437
- await client.disconnect()
438
-
439
- csrf_token = "csrf-only"
440
- session_cookie = _make_session_cookie(
441
- webapp,
442
- {
443
- "csrf_token": csrf_token,
444
- "user_id": "u1",
445
- },
446
- )
447
- session_cookie_name = webapp.config.get("SESSION_COOKIE_NAME", "session")
448
- csrf_cookie_name = f"csrf_token_{runtime.get_runtime_id()}"
449
- cookie_header = f"{session_cookie_name}={session_cookie}; {csrf_cookie_name}={csrf_token}"
450
-
451
- client_ok = socketio.AsyncClient()
452
- await client_ok.connect(
453
- base_url,
454
- namespaces=["/csrf_only"],
455
- headers={"Origin": base_url, "Cookie": cookie_header},
456
- auth={"csrf_token": csrf_token},
457
- wait_timeout=2,
458
- )
459
- try:
460
- res = await client_ok.call("csrf_only_ping", {}, namespace="/csrf_only", timeout=2)
461
- assert isinstance(res, dict)
462
- assert res.get("results")
463
- finally:
464
- await client_ok.disconnect()
tests/test_websocket_namespaces.py
deleted
-497
@@ -1,497 +0,0 @@
1
-import asyncio
2
-import contextlib
3
-import socket
4
-import sys
5
-import threading
6
-from pathlib import Path
7
-from typing import Any, AsyncIterator
8
-from unittest.mock import AsyncMock
9
-
10
-import pytest
11
-
12
-PROJECT_ROOT = Path(__file__).resolve().parents[1]
13
-if str(PROJECT_ROOT) not in sys.path:
14
- sys.path.insert(0, str(PROJECT_ROOT))
15
-
16
-from helpers.state_monitor import StateMonitor
17
-from helpers.websocket_manager import WebSocketManager
18
-
19
-
20
-class FakeSocketIOServer:
21
- def __init__(self) -> None:
22
- self.emit = AsyncMock()
23
- self.disconnect = AsyncMock()
24
-
25
-
26
-@contextlib.asynccontextmanager
27
-async def _run_asgi_app(app: Any) -> AsyncIterator[str]:
28
- import uvicorn
29
-
30
- sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
31
- sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
32
- sock.bind(("127.0.0.1", 0))
33
- sock.listen(128)
34
-
35
- port = sock.getsockname()[1]
36
-
37
- config = uvicorn.Config(
38
- app,
39
- host="127.0.0.1",
40
- port=port,
41
- log_level="warning",
42
- access_log=False,
43
- lifespan="off",
44
- )
45
- server = uvicorn.Server(config)
46
- server.install_signal_handlers = lambda: None # type: ignore[method-assign]
47
-
48
- task = asyncio.create_task(server.serve(sockets=[sock]))
49
- try:
50
- while not server.started:
51
- await asyncio.sleep(0.01)
52
- yield f"http://127.0.0.1:{port}"
53
- finally:
54
- server.should_exit = True
55
- try:
56
- await asyncio.wait_for(task, timeout=5)
57
- finally:
58
- sock.close()
59
-
60
-
61
-@pytest.mark.asyncio
62
-async def test_manager_identity_is_namespace_and_sid_allows_same_sid_across_namespaces() -> None:
63
- socketio = FakeSocketIOServer()
64
- manager = WebSocketManager(socketio, threading.RLock())
65
- # Avoid flakiness from lifecycle broadcasts scheduled via asyncio.create_task.
66
- manager._schedule_lifecycle_broadcast = lambda *_args, **_kwargs: None # type: ignore[assignment]
67
-
68
- sid = "shared-sid"
69
- ns_a = "/a"
70
- ns_b = "/b"
71
-
72
- await manager.handle_connect(ns_a, sid)
73
- await manager.handle_connect(ns_b, sid)
74
-
75
- assert (ns_a, sid) in manager.connections
76
- assert (ns_b, sid) in manager.connections
77
-
78
- await manager.handle_disconnect(ns_a, sid)
79
- assert (ns_a, sid) not in manager.connections
80
- assert (ns_b, sid) in manager.connections
81
-
82
- await manager.emit_to(ns_a, sid, "test_event", {"value": 1}, correlation_id="corr-1")
83
-
84
- assert (ns_a, sid) in manager.buffers
85
- assert (ns_b, sid) not in manager.buffers
86
- assert socketio.emit.await_count == 0
87
-
88
-
89
-def test_state_monitor_tracks_two_identities_for_same_sid_across_namespaces() -> None:
90
- monitor = StateMonitor()
91
- sid = "shared-sid"
92
- monitor.register_sid("/a", sid)
93
- monitor.register_sid("/b", sid)
94
-
95
- debug = monitor._debug_state()
96
- assert ("/a", sid) in debug["identities"]
97
- assert ("/b", sid) in debug["identities"]
98
-
99
-
100
-@pytest.mark.asyncio
101
-async def test_namespace_isolation_state_sync_vs_dev_websocket_test() -> None:
102
- """
103
- CONTRACT.INVARIANT.NS.ISOLATION: no cross-namespace delivery for application events.
104
-
105
- Acceptance proof for `/webui` vs `/dev_websocket_test` namespaces.
106
- """
107
-
108
- from flask import Flask
109
- import socketio
110
-
111
- from helpers.websocket import WebSocketHandler
112
- from helpers.websocket_manager import WebSocketManager
113
- from run_ui import configure_websocket_namespaces
114
-
115
- class StateHandler(WebSocketHandler):
116
- @classmethod
117
- def requires_auth(cls) -> bool:
118
- return False
119
-
120
- @classmethod
121
- def requires_csrf(cls) -> bool:
122
- return False
123
-
124
- @classmethod
125
- def get_event_types(cls) -> list[str]:
126
- return ["state_request"]
127
-
128
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
129
- await self.emit_to(sid, "state_push", {"source": "state_sync"})
130
- return {"ok": True}
131
-
132
- class DevHandler(WebSocketHandler):
133
- @classmethod
134
- def requires_auth(cls) -> bool:
135
- return False
136
-
137
- @classmethod
138
- def requires_csrf(cls) -> bool:
139
- return False
140
-
141
- @classmethod
142
- def get_event_types(cls) -> list[str]:
143
- return ["ws_tester_emit"]
144
-
145
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
146
- await self.broadcast("ws_tester_broadcast", {"source": "dev_websocket_test"})
147
- return None
148
-
149
- StateHandler._reset_instance_for_testing()
150
- DevHandler._reset_instance_for_testing()
151
-
152
- webapp = Flask("test_namespace_isolation")
153
- webapp.secret_key = "test-secret"
154
-
155
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
156
- lock = threading.RLock()
157
- manager = WebSocketManager(sio, lock)
158
-
159
- configure_websocket_namespaces(
160
- webapp=webapp,
161
- socketio_server=sio,
162
- websocket_manager=manager,
163
- handlers_by_namespace={
164
- "/webui": [StateHandler.get_instance(sio, lock)],
165
- "/dev_websocket_test": [DevHandler.get_instance(sio, lock)],
166
- },
167
- )
168
-
169
- asgi_app = socketio.ASGIApp(sio)
170
-
171
- async with _run_asgi_app(asgi_app) as base_url:
172
- client = socketio.AsyncClient()
173
-
174
- state_push_state = asyncio.Event()
175
- state_push_dev = asyncio.Event()
176
- tester_broadcast_dev = asyncio.Event()
177
- tester_broadcast_state = asyncio.Event()
178
-
179
- async def _on_state_push_state(_payload: Any) -> None:
180
- state_push_state.set()
181
-
182
- async def _on_state_push_dev(_payload: Any) -> None:
183
- state_push_dev.set()
184
-
185
- async def _on_tester_broadcast_dev(_payload: Any) -> None:
186
- tester_broadcast_dev.set()
187
-
188
- async def _on_tester_broadcast_state(_payload: Any) -> None:
189
- tester_broadcast_state.set()
190
-
191
- client.on("state_push", _on_state_push_state, namespace="/webui")
192
- client.on("state_push", _on_state_push_dev, namespace="/dev_websocket_test")
193
- client.on("ws_tester_broadcast", _on_tester_broadcast_dev, namespace="/dev_websocket_test")
194
- client.on("ws_tester_broadcast", _on_tester_broadcast_state, namespace="/webui")
195
-
196
- await client.connect(
197
- base_url,
198
- namespaces=["/webui", "/dev_websocket_test"],
199
- headers={"Origin": base_url},
200
- wait_timeout=2,
201
- )
202
- try:
203
- await client.call("state_request", {"context": None}, namespace="/webui", timeout=2)
204
- await asyncio.wait_for(state_push_state.wait(), timeout=2)
205
- await asyncio.sleep(0.05)
206
- assert state_push_dev.is_set() is False
207
-
208
- await client.emit("ws_tester_emit", {"message": "hi"}, namespace="/dev_websocket_test")
209
- await asyncio.wait_for(tester_broadcast_dev.wait(), timeout=2)
210
- await asyncio.sleep(0.05)
211
- assert tester_broadcast_state.is_set() is False
212
- finally:
213
- await client.disconnect()
214
-
215
-
216
-@pytest.mark.asyncio
217
-async def test_diagnostics_include_source_namespace_and_deliver_on_dev_namespace_only() -> None:
218
- """
219
- CONTRACT.Diagnostics: dev console diagnostics are delivered on `/dev_websocket_test`,
220
- but must include `sourceNamespace` identifying the origin namespace.
221
- """
222
-
223
- from helpers.websocket import WebSocketHandler
224
- from helpers.websocket_manager import DIAGNOSTIC_EVENT, WebSocketManager
225
-
226
- class DummyHandler(WebSocketHandler):
227
- @classmethod
228
- def get_event_types(cls) -> list[str]:
229
- return ["dummy_event"]
230
-
231
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
232
- return {"ok": True}
233
-
234
- DummyHandler._reset_instance_for_testing()
235
-
236
- socketio = FakeSocketIOServer()
237
- manager = WebSocketManager(socketio, threading.RLock())
238
- manager._schedule_lifecycle_broadcast = lambda *_args, **_kwargs: None # type: ignore[assignment]
239
-
240
- ns_state = "/webui"
241
- ns_dev = "/dev_websocket_test"
242
-
243
- handler = DummyHandler.get_instance(socketio, threading.RLock())
244
- manager.register_handlers({ns_state: [handler]})
245
-
246
- await manager.handle_connect(ns_dev, "sid-watcher")
247
- await manager.handle_connect(ns_state, "sid-client")
248
- assert manager.register_diagnostic_watcher(ns_dev, "sid-watcher") is True
249
-
250
- socketio.emit.reset_mock()
251
-
252
- await manager.route_event(ns_state, "dummy_event", {"payload": True}, "sid-client")
253
-
254
- calls = [(call.args, call.kwargs) for call in socketio.emit.await_args_list]
255
- diagnostic = next((c for c in calls if c[0] and c[0][0] == DIAGNOSTIC_EVENT), None)
256
- assert diagnostic is not None
257
-
258
- args, kwargs = diagnostic
259
- envelope = args[1]
260
- assert kwargs == {"to": "sid-watcher", "namespace": ns_dev}
261
- assert envelope["data"]["sourceNamespace"] == ns_state
262
-
263
-
264
-def test_namespace_discovery_maps_core_handlers_to_expected_namespaces() -> None:
265
- """
266
- US1 regression: ensure discovery assigns core handlers to their dedicated namespaces
267
- (no cross-registration).
268
- """
269
-
270
- from helpers.websocket_namespace_discovery import discover_websocket_namespaces
271
-
272
- discoveries = discover_websocket_namespaces(
273
- handlers_folder="python/websocket_handlers",
274
- include_root_default=True,
275
- )
276
- by_namespace = {entry.namespace: entry for entry in discoveries}
277
-
278
- assert "/webui" in by_namespace
279
- assert "/dev_websocket_test" in by_namespace
280
-
281
- state_cls_names = [cls.__name__ for cls in by_namespace["/webui"].handler_classes]
282
- dev_cls_names = [cls.__name__ for cls in by_namespace["/dev_websocket_test"].handler_classes]
283
-
284
- assert state_cls_names == ["WebuiHandler"]
285
- assert dev_cls_names == ["DevWebsocketTestHandler"]
286
-
287
-
288
-def test_run_ui_builds_namespace_handler_map_without_cross_registration() -> None:
289
- from run_ui import _build_websocket_handlers_by_namespace
290
-
291
- handlers_by_namespace = _build_websocket_handlers_by_namespace(object(), threading.RLock())
292
-
293
- assert "/webui" in handlers_by_namespace
294
- assert "/dev_websocket_test" in handlers_by_namespace
295
-
296
- assert all(
297
- handler.__class__.__name__ != "DevWebsocketTestHandler"
298
- for handler in handlers_by_namespace["/webui"]
299
- )
300
- assert all(
301
- handler.__class__.__name__ != "WebuiHandler"
302
- for handler in handlers_by_namespace["/dev_websocket_test"]
303
- )
304
-
305
-
306
-@pytest.mark.asyncio
307
-async def test_route_event_dispatches_only_within_connected_namespace_and_results_are_scoped() -> None:
308
- """
309
- CONTRACT.NS.ROUTING: inbound routing is restricted to handlers in the connected namespace.
310
- """
311
-
312
- from helpers.websocket import WebSocketHandler
313
-
314
- socketio = FakeSocketIOServer()
315
- manager = WebSocketManager(socketio, threading.RLock())
316
- manager._schedule_lifecycle_broadcast = lambda *_args, **_kwargs: None # type: ignore[assignment]
317
-
318
- ns_state = "/webui"
319
- ns_dev = "/dev_websocket_test"
320
-
321
- calls: list[str] = []
322
-
323
- class StatePingHandler(WebSocketHandler):
324
- @classmethod
325
- def get_event_types(cls) -> list[str]:
326
- return ["route_test"]
327
-
328
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
329
- calls.append(f"state:{sid}")
330
- return {"ns": "state"}
331
-
332
- class DevPingHandler(WebSocketHandler):
333
- @classmethod
334
- def get_event_types(cls) -> list[str]:
335
- return ["route_test"]
336
-
337
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
338
- calls.append(f"dev:{sid}")
339
- return {"ns": "dev"}
340
-
341
- StatePingHandler._reset_instance_for_testing()
342
- DevPingHandler._reset_instance_for_testing()
343
-
344
- state_handler = StatePingHandler.get_instance(socketio, threading.RLock())
345
- dev_handler = DevPingHandler.get_instance(socketio, threading.RLock())
346
-
347
- manager.register_handlers({ns_state: [state_handler], ns_dev: [dev_handler]})
348
- await manager.handle_connect(ns_state, "sid-state")
349
- await manager.handle_connect(ns_dev, "sid-dev")
350
-
351
- res_state = await manager.route_event(ns_state, "route_test", {"x": 1}, "sid-state")
352
- assert {item["handlerId"] for item in res_state["results"]} == {state_handler.identifier}
353
- assert res_state["results"][0]["data"]["ns"] == "state"
354
-
355
- res_dev = await manager.route_event(ns_dev, "route_test", {"x": 2}, "sid-dev")
356
- assert {item["handlerId"] for item in res_dev["results"]} == {dev_handler.identifier}
357
- assert res_dev["results"][0]["data"]["ns"] == "dev"
358
-
359
- assert calls == ["state:sid-state", "dev:sid-dev"]
360
-
361
-
362
-@pytest.mark.asyncio
363
-async def test_lifecycle_broadcasts_deliver_only_within_the_namespace() -> None:
364
- """
365
- CONTRACT.NS.DELIVERY: lifecycle broadcasts are namespace-scoped.
366
- """
367
-
368
- from helpers.websocket_manager import (
369
- LIFECYCLE_CONNECT_EVENT,
370
- LIFECYCLE_DISCONNECT_EVENT,
371
- )
372
-
373
- socketio = FakeSocketIOServer()
374
- manager = WebSocketManager(socketio, threading.RLock())
375
-
376
- ns_state = "/webui"
377
- ns_dev = "/dev_websocket_test"
378
-
379
- # Connect events should broadcast only within their namespace.
380
- await manager.handle_connect(ns_state, "sid-state-1")
381
- await asyncio.sleep(0)
382
- state_connect_calls = [
383
- call
384
- for call in socketio.emit.await_args_list
385
- if call.args and call.args[0] == LIFECYCLE_CONNECT_EVENT
386
- ]
387
- assert state_connect_calls
388
- assert all(call.kwargs.get("namespace") == ns_state for call in state_connect_calls)
389
-
390
- socketio.emit.reset_mock()
391
- await manager.handle_connect(ns_dev, "sid-dev-1")
392
- await asyncio.sleep(0)
393
- dev_connect_calls = [
394
- call
395
- for call in socketio.emit.await_args_list
396
- if call.args and call.args[0] == LIFECYCLE_CONNECT_EVENT
397
- ]
398
- assert dev_connect_calls
399
- assert all(call.kwargs.get("namespace") == ns_dev for call in dev_connect_calls)
400
-
401
- # Disconnect broadcasts go to remaining peers in that namespace only.
402
- socketio.emit.reset_mock()
403
- await manager.handle_connect(ns_state, "sid-state-2")
404
- await manager.handle_connect(ns_dev, "sid-dev-2")
405
- socketio.emit.reset_mock()
406
-
407
- await manager.handle_disconnect(ns_state, "sid-state-2")
408
- await asyncio.sleep(0)
409
- state_disconnect_calls = [
410
- call
411
- for call in socketio.emit.await_args_list
412
- if call.args and call.args[0] == LIFECYCLE_DISCONNECT_EVENT
413
- ]
414
- assert state_disconnect_calls
415
- assert all(call.kwargs.get("namespace") == ns_state for call in state_disconnect_calls)
416
- assert all(call.kwargs.get("to") == "sid-state-1" for call in state_disconnect_calls)
417
-
418
-
419
-@pytest.mark.asyncio
420
-async def test_request_semantics_no_handlers_and_timeouts_are_namespace_scoped_and_order_insensitive() -> None:
421
- """
422
- CONTRACT.REQUEST.RESULTS + CONTRACT.REQUEST.RESULTS.ORDERING + CONTRACT.NS.ROUTING.
423
- """
424
-
425
- from helpers.websocket import WebSocketHandler
426
-
427
- socketio = FakeSocketIOServer()
428
- manager = WebSocketManager(socketio, threading.RLock())
429
- manager._schedule_lifecycle_broadcast = lambda *_args, **_kwargs: None # type: ignore[assignment]
430
-
431
- ns_state = "/webui"
432
- ns_dev = "/dev_websocket_test"
433
-
434
- class Alpha(WebSocketHandler):
435
- @classmethod
436
- def get_event_types(cls) -> list[str]:
437
- return ["multi", "slow"]
438
-
439
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
440
- if event_type == "slow":
441
- await asyncio.sleep(0.2)
442
- return {"alpha": True}
443
- return {"alpha": True}
444
-
445
- class Beta(WebSocketHandler):
446
- @classmethod
447
- def get_event_types(cls) -> list[str]:
448
- return ["multi"]
449
-
450
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
451
- return {"beta": True}
452
-
453
- Alpha._reset_instance_for_testing()
454
- Beta._reset_instance_for_testing()
455
- alpha = Alpha.get_instance(socketio, threading.RLock())
456
- beta = Beta.get_instance(socketio, threading.RLock())
457
-
458
- manager.register_handlers({ns_state: [alpha, beta]})
459
- await manager.handle_connect(ns_state, "sid-a")
460
- await manager.handle_connect(ns_state, "sid-b")
461
- await manager.handle_connect(ns_dev, "sid-dev")
462
-
463
- # Unknown event name -> NO_HANDLERS (no hang), scoped to the namespace.
464
- no_handler = await manager.route_event(ns_dev, "missing_event", {"x": 1}, "sid-dev")
465
- assert no_handler["results"][0]["ok"] is False
466
- assert no_handler["results"][0]["error"]["code"] == "NO_HANDLERS"
467
- assert ns_dev in no_handler["results"][0]["error"]["error"]
468
-
469
- # Unknown event name in a namespace that *does* have other handlers -> NO_HANDLERS.
470
- unhandled_in_state = await manager.route_event(ns_state, "unknown_event", {"x": 1}, "sid-a")
471
- assert unhandled_in_state["results"][0]["ok"] is False
472
- assert unhandled_in_state["results"][0]["error"]["code"] == "NO_HANDLERS"
473
- assert ns_state in unhandled_in_state["results"][0]["error"]["error"]
474
-
475
- # Known event name in the wrong namespace -> NO_HANDLERS (no cross-namespace fallback).
476
- wrong_namespace = await manager.route_event(ns_dev, "multi", {"x": 1}, "sid-dev")
477
- assert wrong_namespace["results"][0]["ok"] is False
478
- assert wrong_namespace["results"][0]["error"]["code"] == "NO_HANDLERS"
479
- assert ns_dev in wrong_namespace["results"][0]["error"]["error"]
480
-
481
- # Order-insensitive results[]: both handlers must be present regardless of ordering.
482
- multi = await manager.route_event(ns_state, "multi", {"x": 1}, "sid-a")
483
- handler_ids = {item["handlerId"] for item in multi["results"]}
484
- assert handler_ids == {alpha.identifier, beta.identifier}
485
-
486
- # Timeout results are represented as TIMEOUT items and scoped to the namespace.
487
- aggregated = await manager.route_event_all(ns_state, "slow", {"x": 1}, timeout_ms=50)
488
- assert len(aggregated) == 2 # only state namespace connections
489
- assert {entry["sid"] for entry in aggregated} == {"sid-a", "sid-b"}
490
- for entry in aggregated:
491
- assert entry["results"]
492
- assert entry["results"][0]["ok"] is False
493
- assert entry["results"][0]["error"]["code"] == "TIMEOUT"
494
-
495
- # Allow the underlying slow route_event coroutines to complete so pytest's event loop
496
- # teardown does not cancel them mid-flight (avoids noisy InvalidStateError callbacks).
497
- await asyncio.sleep(0.3)
tests/test_websocket_namespaces_integration.py
deleted
-113
@@ -1,113 +0,0 @@
1
-import asyncio
2
-import contextlib
3
-import socket
4
-from typing import Any, AsyncIterator
5
-
6
-import pytest
7
-
8
-
9
-@contextlib.asynccontextmanager
10
-async def _run_asgi_app(app: Any) -> AsyncIterator[str]:
11
- import uvicorn
12
-
13
- sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
14
- sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
15
- sock.bind(("127.0.0.1", 0))
16
- sock.listen(128)
17
-
18
- port = sock.getsockname()[1]
19
-
20
- config = uvicorn.Config(
21
- app,
22
- host="127.0.0.1",
23
- port=port,
24
- log_level="warning",
25
- access_log=False,
26
- lifespan="off",
27
- )
28
- server = uvicorn.Server(config)
29
- server.install_signal_handlers = lambda: None # type: ignore[method-assign]
30
-
31
- task = asyncio.create_task(server.serve(sockets=[sock]))
32
- try:
33
- while not server.started:
34
- await asyncio.sleep(0.01)
35
- yield f"http://127.0.0.1:{port}"
36
- finally:
37
- server.should_exit = True
38
- try:
39
- await asyncio.wait_for(task, timeout=5)
40
- finally:
41
- sock.close()
42
-
43
-
44
-@pytest.mark.asyncio
45
-async def test_unregistered_namespace_connection_fails_with_unknown_namespace_connect_error() -> None:
46
- """
47
- US5 integration: unregistered namespace connections fail deterministically with a structured
48
- connect_error payload (UNKNOWN_NAMESPACE), independent of python-socketio defaults.
49
- """
50
-
51
- from flask import Flask
52
- import socketio
53
-
54
- from helpers.websocket import WebSocketHandler
55
- from helpers.websocket_manager import WebSocketManager
56
- from run_ui import configure_websocket_namespaces
57
-
58
- class OpenHandler(WebSocketHandler):
59
- @classmethod
60
- def requires_auth(cls) -> bool:
61
- return False
62
-
63
- @classmethod
64
- def requires_csrf(cls) -> bool:
65
- return False
66
-
67
- @classmethod
68
- def get_event_types(cls) -> list[str]:
69
- return ["open_ping"]
70
-
71
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
72
- return {"ok": True}
73
-
74
- OpenHandler._reset_instance_for_testing()
75
-
76
- webapp = Flask("test_ws_namespaces_integration")
77
- webapp.secret_key = "test-secret"
78
-
79
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
80
- lock = __import__("threading").RLock()
81
- manager = WebSocketManager(sio, lock)
82
-
83
- configure_websocket_namespaces(
84
- webapp=webapp,
85
- socketio_server=sio,
86
- websocket_manager=manager,
87
- handlers_by_namespace={"/open": [OpenHandler.get_instance(sio, lock)]},
88
- )
89
-
90
- asgi_app = socketio.ASGIApp(sio)
91
-
92
- async with _run_asgi_app(asgi_app) as base_url:
93
- client = socketio.AsyncClient()
94
- connect_error_fut: asyncio.Future[Any] = asyncio.get_running_loop().create_future()
95
-
96
- async def _on_connect_error(data: Any) -> None:
97
- if not connect_error_fut.done():
98
- connect_error_fut.set_result(data)
99
-
100
- client.on("connect_error", _on_connect_error, namespace="/unknown")
101
-
102
- try:
103
- with pytest.raises(socketio.exceptions.ConnectionError):
104
- await client.connect(base_url, namespaces=["/unknown"])
105
-
106
- err = await asyncio.wait_for(connect_error_fut, timeout=2)
107
- assert err["message"] == "UNKNOWN_NAMESPACE"
108
- assert err["data"] == {"code": "UNKNOWN_NAMESPACE", "namespace": "/unknown"}
109
- finally:
110
- try:
111
- await client.disconnect()
112
- except Exception:
113
- pass
tests/test_websocket_root_namespace.py
deleted
-183
@@ -1,183 +0,0 @@
1
-import asyncio
2
-import contextlib
3
-import socket
4
-from typing import Any, AsyncIterator
5
-
6
-import pytest
7
-
8
-
9
-@contextlib.asynccontextmanager
10
-async def _run_asgi_app(app: Any) -> AsyncIterator[str]:
11
- import uvicorn
12
-
13
- sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
14
- sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
15
- sock.bind(("127.0.0.1", 0))
16
- sock.listen(128)
17
-
18
- port = sock.getsockname()[1]
19
-
20
- config = uvicorn.Config(
21
- app,
22
- host="127.0.0.1",
23
- port=port,
24
- log_level="warning",
25
- access_log=False,
26
- lifespan="off",
27
- )
28
- server = uvicorn.Server(config)
29
- server.install_signal_handlers = lambda: None # type: ignore[method-assign]
30
-
31
- task = asyncio.create_task(server.serve(sockets=[sock]))
32
- try:
33
- while not server.started:
34
- await asyncio.sleep(0.01)
35
- yield f"http://127.0.0.1:{port}"
36
- finally:
37
- server.should_exit = True
38
- try:
39
- await asyncio.wait_for(task, timeout=5)
40
- finally:
41
- sock.close()
42
-
43
-
44
-@pytest.mark.asyncio
45
-async def test_root_namespace_request_style_calls_resolve_with_no_handlers() -> None:
46
- """
47
- CONTRACT.INVARIANT.NS.ROOT.UNHANDLED: root (`/`) is reserved and unhandled for application
48
- events by default, but request-style calls must not hang (NO_HANDLERS).
49
- """
50
-
51
- from flask import Flask
52
- import socketio
53
-
54
- from helpers.websocket import WebSocketHandler
55
- from helpers.websocket_manager import WebSocketManager
56
- from run_ui import configure_websocket_namespaces
57
-
58
- app = Flask("test_ws_root_namespace")
59
- app.secret_key = "test-secret"
60
-
61
- calls: list[str] = []
62
-
63
- class HelloHandler(WebSocketHandler):
64
- @classmethod
65
- def requires_auth(cls) -> bool:
66
- return False
67
-
68
- @classmethod
69
- def requires_csrf(cls) -> bool:
70
- return False
71
-
72
- @classmethod
73
- def get_event_types(cls) -> list[str]:
74
- return ["hello_request"]
75
-
76
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
77
- calls.append(sid)
78
- return {"hello": True}
79
-
80
- HelloHandler._reset_instance_for_testing()
81
-
82
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
83
- lock = __import__("threading").RLock()
84
- manager = WebSocketManager(sio, lock)
85
-
86
- configure_websocket_namespaces(
87
- webapp=app,
88
- socketio_server=sio,
89
- websocket_manager=manager,
90
- handlers_by_namespace={
91
- "/webui": [HelloHandler.get_instance(sio, lock)],
92
- },
93
- )
94
-
95
- asgi_app = socketio.ASGIApp(sio)
96
-
97
- async with _run_asgi_app(asgi_app) as base_url:
98
- client = socketio.AsyncClient()
99
- await client.connect(
100
- base_url,
101
- namespaces=["/"],
102
- headers={"Origin": base_url},
103
- wait_timeout=2,
104
- )
105
- try:
106
- res_unknown = await client.call("unknown_event", {"x": 1}, namespace="/", timeout=2)
107
- assert res_unknown["results"][0]["ok"] is False
108
- assert res_unknown["results"][0]["error"]["code"] == "NO_HANDLERS"
109
-
110
- res_known_elsewhere = await client.call("hello_request", {"name": "x"}, namespace="/", timeout=2)
111
- assert res_known_elsewhere["results"][0]["ok"] is False
112
- assert res_known_elsewhere["results"][0]["error"]["code"] == "NO_HANDLERS"
113
- assert calls == []
114
- finally:
115
- await client.disconnect()
116
-
117
-
118
-@pytest.mark.asyncio
119
-async def test_root_namespace_fire_and_forget_does_not_invoke_application_handlers() -> None:
120
- """
121
- Fire-and-forget emits on `/` must not invoke any application handler.
122
- """
123
-
124
- from flask import Flask
125
- import socketio
126
-
127
- from helpers.websocket import WebSocketHandler
128
- from helpers.websocket_manager import WebSocketManager
129
- from run_ui import configure_websocket_namespaces
130
-
131
- app = Flask("test_ws_root_fire_and_forget")
132
- app.secret_key = "test-secret"
133
-
134
- calls: list[str] = []
135
-
136
- class SideEffectHandler(WebSocketHandler):
137
- @classmethod
138
- def requires_auth(cls) -> bool:
139
- return False
140
-
141
- @classmethod
142
- def requires_csrf(cls) -> bool:
143
- return False
144
-
145
- @classmethod
146
- def get_event_types(cls) -> list[str]:
147
- return ["hello_request"]
148
-
149
- async def process_event(self, event_type: str, data: dict[str, Any], sid: str):
150
- calls.append(sid)
151
- return {"ok": True}
152
-
153
- SideEffectHandler._reset_instance_for_testing()
154
-
155
- sio = socketio.AsyncServer(async_mode="asgi", cors_allowed_origins="*", namespaces="*")
156
- lock = __import__("threading").RLock()
157
- manager = WebSocketManager(sio, lock)
158
-
159
- configure_websocket_namespaces(
160
- webapp=app,
161
- socketio_server=sio,
162
- websocket_manager=manager,
163
- handlers_by_namespace={
164
- "/webui": [SideEffectHandler.get_instance(sio, lock)],
165
- },
166
- )
167
-
168
- asgi_app = socketio.ASGIApp(sio)
169
-
170
- async with _run_asgi_app(asgi_app) as base_url:
171
- client = socketio.AsyncClient()
172
- await client.connect(
173
- base_url,
174
- namespaces=["/"],
175
- headers={"Origin": base_url},
176
- wait_timeout=2,
177
- )
178
- try:
179
- await client.emit("hello_request", {"name": "x"}, namespace="/")
180
- await asyncio.sleep(0.1)
181
- assert calls == []
182
- finally:
183
- await client.disconnect()
tests/websocket_namespace_test_utils.py
deleted
-56
@@ -1,56 +0,0 @@
1
-from __future__ import annotations
2
-
3
-from dataclasses import dataclass
4
-from typing import Any
5
-from unittest.mock import AsyncMock
6
-
7
-
8
-ConnectionIdentity = tuple[str, str] # (namespace, sid)
9
-
10
-
11
-def nsid(namespace: str, sid: str) -> ConnectionIdentity:
12
- return (namespace, sid)
13
-
14
-
15
-@dataclass(frozen=True)
16
-class SocketIOCall:
17
- args: tuple[Any, ...]
18
- kwargs: dict[str, Any]
19
-
20
- @property
21
- def namespace(self) -> str | None:
22
- value = self.kwargs.get("namespace")
23
- if value is None:
24
- return None
25
- if not isinstance(value, str):
26
- raise TypeError(f"Expected namespace to be str, got {type(value).__name__}")
27
- return value
28
-
29
-
30
-class FakeSocketIOServer:
31
- """
32
- Test double for python-socketio AsyncServer.
33
-
34
- Captures calls and surfaces the optional Socket.IO namespace dimension via recorded kwargs.
35
- """
36
-
37
- def __init__(self) -> None:
38
- self._emit_calls: list[SocketIOCall] = []
39
- self._disconnect_calls: list[SocketIOCall] = []
40
-
41
- self.emit = AsyncMock(side_effect=self._emit)
42
- self.disconnect = AsyncMock(side_effect=self._disconnect)
43
-
44
- async def _emit(self, *args: Any, **kwargs: Any) -> None:
45
- self._emit_calls.append(SocketIOCall(args=args, kwargs=dict(kwargs)))
46
-
47
- async def _disconnect(self, *args: Any, **kwargs: Any) -> None:
48
- self._disconnect_calls.append(SocketIOCall(args=args, kwargs=dict(kwargs)))
49
-
50
- @property
51
- def emit_calls(self) -> list[SocketIOCall]:
52
- return self._emit_calls
53
-
54
- @property
55
- def disconnect_calls(self) -> list[SocketIOCall]:
56
- return self._disconnect_calls