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