Rebuild test suite & update documentation

keyboardstaff committed Mar 26, 2026 at 01:12 UTC dae37ee0f6f0a2a91a57a36ac50273b84d3236d6
12 files changed +1372 -130
AGENTS.md
+4 -6
@@ -71,12 +71,10 @@ When running in Docker, Agent Zero uses two distinct Python runtimes to isolate
71 ├── initialize.py # Framework initialization logic
72 ├── models.py # LLM provider configurations
73 ├── run_ui.py # WebUI server entry point
74 -├── python/
75 -│ ├── api/ # API Handlers (ApiHandler subclasses)
76 -│ ├── extensions/ # Backend lifecycle extensions
77 -│ ├── helpers/ # Shared Python utilities (plugins, files, etc.)
78 -│ ├── tools/ # Agent tools (Tool subclasses)
79 -│ └── websocket_handlers/# WebSocket event handlers
74 +├── api/ # API Handlers (ApiHandler subclasses) + WsHandler subclasses (ws_*.py)
75 +├── extensions/ # Backend lifecycle extensions
76 +├── helpers/ # Shared Python utilities (plugins, files, etc.)
77 +├── tools/ # Agent tools (Tool subclasses)
78 ├── webui/
79 │ ├── components/ # Alpine.js components
80 │ ├── js/ # Core frontend logic (modals, stores, etc.)
docs/developer/websockets.md
+30 -38
@@ -23,10 +23,9 @@ This guide consolidates everything you need to design, implement, and troublesho
23 ## Architecture at a Glance
24
25 - **Runtime (`run_ui.py`)** – boots `python-socketio.AsyncServer` inside an ASGI stack served by Uvicorn. Flask routes are mounted via `uvicorn.middleware.wsgi.WSGIMiddleware`, and Flask + Socket.IO share the same process so session cookies and CSRF semantics stay aligned.
26 -- **Singleton handlers** – every `WebSocketHandler` subclass exposes `get_instance()` and is registered exactly once. Direct instantiation raises `SingletonInstantiationError`, keeping shared state and lifecycle hooks deterministic.
27 -- **Dispatcher offload** – handler entrypoints (`process_event`, `on_connect`, `on_disconnect`) run in a background worker loop (via `DeferredTask`) so blocking handlers cannot stall the Socket.IO transport. Socket.IO emits/disconnects are marshalled back to the dispatcher loop. Diagnostic timing and payload summaries are only built when Event Console watchers are subscribed (development mode).
28 -- **`python/helpers/websocket_manager.py`** – orchestrates routing, buffering, aggregation, metadata envelopes, and session tracking. Think of it as the “switchboard” for every WebSocket event.
29 -- **`python/helpers/websocket.py`** – base class for application handlers. Provides lifecycle hooks, helper methods (`emit_to`, `broadcast`, `request`, `request_all`) and identifier metadata.
26 +- **Handler base class** – every handler derives from `WsHandler` (defined in `helpers/ws.py`) and implements `process(event, data, sid)`. Handlers are instantiated directly and registered with the manager.
27 +- **Dispatcher offload** – handler entrypoints (`process`, `on_connect`, `on_disconnect`) run in a background worker loop (via `DeferredTask`) so blocking handlers cannot stall the Socket.IO transport. Socket.IO emits/disconnects are marshalled back to the dispatcher loop. Diagnostic timing and payload summaries are only built when Event Console watchers are subscribed (development mode).
28 +- **`helpers/ws_manager.py`** – orchestrates routing, buffering, aggregation, metadata envelopes, and session tracking. Think of it as the "switchboard" for every WebSocket event.
29 - **`webui/js/websocket.js`** – frontend singleton exposing a minimal client API (`emit`, `request`, `on`, `off`) with lazy connection management and development-only logging (no client-side `broadcast()` or `requestAll()` helpers).
30 - **Developer Harness (`webui/components/settings/developer/websocket-test-store.js`)** – manual and automatic validation suite for emit/request flows, timeout behaviour (including the default unlimited wait), correlation ID propagation, envelope metadata, subscription persistence across reconnect, and development-mode diagnostics.
31 - **Specs & Contracts** – canonical definitions live under `specs/003-websocket-event-handlers/`. This guide references those documents but focuses on applied usage.
@@ -38,7 +37,7 @@ This guide consolidates everything you need to design, implement, and troublesho
37 | Term | Where it Appears | Meaning |
38 |------|------------------|---------|
39 | `sid` | Socket.IO | Connection identifier for a Socket.IO namespace connection. With only the root namespace (`/`), each tab has one `sid`. When connecting to multiple namespaces, a tab has one `sid` per namespace. Treat connection identity as `(namespace, sid)`. |
41 -| `handlerId` | Manager Envelope | Fully-qualified Python class name (e.g., `python.websocket_handlers.notifications.NotificationHandler`). Used for result aggregation and logging. |
40 +| `handlerId` | Manager Envelope | Fully-qualified Python class name (e.g., `api.ws_webui.WsWebui`). Used for result aggregation and logging. |
41 | `eventId` | Manager Envelope | UUIDv4 generated for every server→client delivery. Unique per emission. Useful when correlating broadcast fan-out or diagnosing duplicates. |
42 | `correlationId` | Bidirectional flows | Thread that ties together request, response, and any follow-up events. Client may supply one; otherwise the manager generates and echoes it everywhere. |
43 | `data` | Envelope payload | Application payload you define. Always a JSON-serialisable object. |
@@ -53,7 +52,7 @@ Useful mental model: **client ↔ manager ↔ handler**. The manager normalises
52
53 1. **Lazy Connect** – `/js/websocket.js` connects only when a consumer uses the client API (e.g., `emit`, `request`, `on`). Consumers may still explicitly `await websocket.connect()` to block UI until the socket is ready.
54 2. **Handshake** – Socket.IO connects using the existing Flask session cookie and a CSRF token provided via the Socket.IO `auth` payload (`csrf_token`). The token is obtained from `GET /csrf_token` (see `/js/api.js#getCsrfToken()`), which also sets the runtime-scoped cookie `csrf_token_{runtime_id}`. The server validates an **Origin allowlist** (RFC 6455 / OWASP CSWSH baseline) and then checks handler requirements (`requires_auth`, `requires_csrf`) before accepting.
56 -3. **Lifecycle Hooks** – After acceptance, `WebSocketHandler.on_connect(sid)` fires for every registered handler. Use it for initial emits, state bookkeeping, or session tracking.
55 +3. **Lifecycle Hooks** – After acceptance, `WsHandler.on_connect(sid)` fires for every registered handler. Use it for initial emits, state bookkeeping, or session tracking.
56 4. **Normal Operation** – Client emits events. Manager routes them to the appropriate handlers, gathers results, and wraps outbound deliveries in the mandatory envelope.
57 5. **Disconnection & Buffering** – If a tab goes away without a graceful disconnect, fire-and-forget events accumulate (max 100). On reconnect, the manager flushes the buffer via `emit_to`. Request flows respond with explicit `CONNECTION_NOT_FOUND` errors.
58 6. **Reconnection Attempts** – Socket.IO handles reconnect attempts; the manager continues to buffer fire-and-forget events (up to 1 hour) for temporarily disconnected SIDs and flushes them on reconnect.
@@ -70,13 +69,13 @@ Agent Zero can also push poll-shaped state snapshots over the WebSocket bus, rep
69 ### Thinking in Roles
70
71 - **Client** (frontend) is the page that imports `/js/websocket.js`. It acts as both a **producer** (calling `emit`, `request`) and a **consumer** (subscribing with `on`).
73 -- **Manager** (`WebSocketManager`) sits server-side and routes everything. It resolves correlation IDs, wraps envelopes, and fans out results.
74 -- **Handler** (`WebSocketHandler`) executes the application logic. Each handler may emit additional events back to the client or initiate its own requests to connected SIDs.
72 +- **Manager** (`WsManager`) sits server-side and routes everything. It resolves correlation IDs, wraps envelopes, and fans out results.
73 +- **Handler** (`WsHandler`) executes the application logic. Each handler may emit additional events back to the client or initiate its own requests to connected SIDs.
74
75 ### Flow Overview (by Operation)
76
77 ```
79 -Client emit() ───▶ Manager route_event() ───▶ Handler.process_event()
78 +Client emit() ───▶ Manager route_event() ───▶ Handler.process()
79 │ │ └──(fire-and-forget, no ack)
80 └── throws if └── validates payload + routes by namespace/event type
81 not connected updates last_activity
@@ -137,21 +136,20 @@ These expanded flows complement the operation matrix later in the guide, ensurin
136
137 ### 1. Handler Discovery & Setup
138
140 -Handlers are discovered deterministically from `python/websocket_handlers/`:
139 +Handlers are `WsHandler` subclasses discovered from `api/ws_*.py`:
140
142 -- **File entry**: `python/websocket_handlers/webui_handler.py` → namespace `/webui`
143 -- **Folder entry**: `python/websocket_handlers/orders/` or `python/websocket_handlers/orders_handler/` → namespace `/orders` (loads `*.py` one level deep; ignores `__init__.py` and deeper nesting)
144 -- **Reserved root**: `python/websocket_handlers/_default.py` → namespace `/` (diagnostics-only by default)
141 +- **Example**: `api/ws_webui.py` → handles WebUI events
142 +- **Dev test**: `api/ws_dev_test.py` → developer harness handler
143 +- **Hello**: `api/ws_hello.py` → minimal example handler
144
146 -Create handler modules under the appropriate namespace entry and inherit from `WebSocketHandler`.
145 +Create new handler files as `api/ws_<name>.py` and inherit from `WsHandler`.
146
147 ```python
149 -from helpers.websocket import WebSocketHandler
148 +from helpers.ws import WsHandler
149
151 -class DashboardHandler(WebSocketHandler):
152 - @classmethod
153 - def get_event_types(cls) -> list[str]:
154 - return ["dashboard_refresh", "dashboard_push"]
150 +class WsMyFeature(WsHandler):
151 + HANDLER_ID = "my_feature"
152 + HANDLED_EVENTS = ["my_event_a", "my_event_b"]
153
154 async def process_event(self, event_type: str, data: dict[str, Any], sid: str) -> dict | None:
155 if event_type == "dashboard_refresh":
@@ -226,7 +224,7 @@ if not results:
224
225 ### 5. Session Tracking Helpers
226
229 -`WebSocketManager` maintains lightweight mappings that you can use from handlers:
227 +`WsManager` maintains lightweight mappings that you can use from handlers:
228
229 ```python
230 all_sids = self.manager.get_sids_for_user() # today: every active sid
@@ -470,21 +468,15 @@ results.forEach(({ handlerId, ok, data, error }) => {
468 Server (two handlers listening to the same event):
469
470 ```python
473 -class TaskMetrics(WebSocketHandler):
474 - @classmethod
475 - def get_event_types(cls) -> list[str]:
476 - return ["refresh_metrics"]
471 +class TaskMetrics(WsHandler):
472
478 - async def process_event(self, event_type: str, data: dict, sid: str) -> dict | None:
473 + async def process(self, event: str, data: dict, sid: str) -> dict | None:
474 stats = await self._load_task_metrics(data["duration"])
475 return {"metrics": stats}
476
482 -class HostMetrics(WebSocketHandler):
483 - @classmethod
484 - def get_event_types(cls) -> list[str]:
485 - return ["refresh_metrics"]
477 +class HostMetrics(WsHandler):
478
487 - async def process_event(self, event_type: str, data: dict, sid: str) -> dict | None:
479 + async def process(self, event: str, data: dict, sid: str) -> dict | None:
480 return {"metrics": await self._load_host_metrics(data["duration"])}
481 ```
482
@@ -550,7 +542,7 @@ The manager validates the payload, resolves/creates `correlationId`, and passes
542
543 ```json
544 {
553 - "handlerId": "python.websocket_handlers.notifications.NotificationHandler",
545 + "handlerId": "api.ws_webui.WsWebui",
546 "eventId": "b7e2a9cd-2857-4f7a-8bf4-12a736cb6720",
547 "correlationId": "caller-supplied-or-generated",
548 "ts": "2025-10-31T13:13:37.123Z",
@@ -579,8 +571,8 @@ The manager validates the payload, resolves/creates `correlationId`, and passes
571 ### WebSocket Event Console
572
573 - Location: `Settings → Developer → WebSocket Event Console`.
582 -- Enabling capture calls `websocket.request("ws_event_console_subscribe", { requestedAt })`. The handler (`DevWebsocketTestHandler`) refuses the subscription outside development mode and registers the SID as a **diagnostic watcher** by calling `WebSocketManager.register_diagnostic_watcher`. Only connected SIDs can subscribe.
583 -- Disabling capture calls `websocket.request("ws_event_console_unsubscribe", {})`. Disconnecting also triggers `WebSocketManager.unregister_diagnostic_watcher`, so stranded watchers never accumulate.
574 +- Enabling capture calls `websocket.request("ws_event_console_subscribe", { requestedAt })`. The handler (`DevWebsocketTestHandler`) refuses the subscription outside development mode and registers the SID as a **diagnostic watcher** by calling `WsManager.register_diagnostic_watcher`. Only connected SIDs can subscribe.
575 +- Disabling capture calls `websocket.request("ws_event_console_unsubscribe", {})`. Disconnecting also triggers `WsManager.unregister_diagnostic_watcher`, so stranded watchers never accumulate.
576 - While at least one watcher exists, the manager streams `ws_dev_console_event` envelopes (documented in `contracts/event-schemas.md`). Each payload contains:
577 - `kind`: `"inbound" | "outbound" | "lifecycle"`
578 - `eventType`, `sid`, `targets[]`, delivery/buffer flags
@@ -596,7 +588,7 @@ The manager validates the payload, resolves/creates `correlationId`, and passes
588
589 ### Instrumentation & Logging
590
599 -- `WebSocketManager` offloads handler execution via `DeferredTask` and may record `durationMs` when development diagnostics are active (Event Console watchers subscribed). These metrics flow into the Event Console stream (and may also appear in `request()` / `request_all()` results), keeping steady-state overhead near zero when diagnostics are closed.
591 +- `WsManager` offloads handler execution via `DeferredTask` and may record `durationMs` when development diagnostics are active (Event Console watchers subscribed). These metrics flow into the Event Console stream (and may also appear in `request()` / `request_all()` results), keeping steady-state overhead near zero when diagnostics are closed.
592 - Lifecycle events capture `connectionCount`, ISO8601 timestamps, and SID so dashboards can correlate UI behaviour with connection churn.
593 - Backend logging: use `PrintStyle.debug/info/warning` and always include `handlerId`, `eventType`, `sid`, and `correlationId`. The manager already logs connection events, missing handlers, and buffer overflows.
594 - Frontend logging: `websocket.debugLog()` mirrors backend debug messages but only when `window.runtimeInfo.isDevelopment` is true.
@@ -661,7 +653,7 @@ The manager validates the payload, resolves/creates `correlationId`, and passes
653 - [`frontend-api.md`](../specs/003-websocket-event-handlers/contracts/frontend-api.md)
654 - [`event-schemas.md`](../specs/003-websocket-event-handlers/contracts/event-schemas.md)
655 - [`security-contract.md`](../specs/003-websocket-event-handlers/contracts/security-contract.md)
664 -- **Implementation Reference** – Inspect `python/helpers/websocket_manager.py`, `python/helpers/websocket.py`, `webui/js/websocket.js`, and the developer harness in `webui/components/settings/developer/websocket-test-store.js` for concrete examples.
656 +- **Implementation Reference** – Inspect `helpers/ws_manager.py`, `helpers/ws.py`, `webui/js/websocket.js`, and the developer harness in `webui/components/settings/developer/websocket-test-store.js` for concrete examples.
657
658 > **Tip:** When extending the infrastructure (new metadata) start by updating the contracts, sync the manager/frontend helpers, and then document the change here so producers and consumers stay in lockstep.
659
@@ -671,10 +663,10 @@ The WebSocket stack standardizes backend error codes returned in `RequestResultI
663
664 | Code | Scope | Meaning | Typical Remediation | Example Payload |
665 |------|-------|---------|---------------------|-----------------|
674 -| `NO_HANDLERS` | Manager routing | No handler is registered for the requested `eventType`. | Register a handler for the event or correct the event name. | `{ "handlerId": "WebSocketManager", "ok": false, "error": { "code": "NO_HANDLERS", "error": "No handler for 'missing'" } }` |
666 +| `NO_HANDLERS` | Manager routing | No handler is registered for the requested `eventType`. | Register a handler for the event or correct the event name. | `{ "handlerId": "WsManager", "ok": false, "error": { "code": "NO_HANDLERS", "error": "No handler for 'missing'" } }` |
667 | `TIMEOUT` | Aggregated or single request | The request exceeded `timeoutMs`. | Increase `timeoutMs`, reduce handler processing time, or split work. | `{ "handlerId": "ExampleHandler", "ok": false, "error": { "code": "TIMEOUT", "error": "Request timeout" } }` |
676 -| `CONNECTION_NOT_FOUND` | Single‑sid request | Target `sid` is not connected/known. | Use an active `sid` or retry after reconnect. | `{ "handlerId": "WebSocketManager", "ok": false, "error": { "code": "CONNECTION_NOT_FOUND", "error": "Connection 'sid-123' not found" } }` |
677 -| `HARNESS_UNKNOWN_EVENT` | Developer harness | Harness test handler received an unsupported event name. | Update harness sources or disable the step before running automation. | `{ "handlerId": "python.websocket_handlers.dev_websocket_test_handler.DevWebsocketTestHandler", "ok": false, "error": { "code": "HARNESS_UNKNOWN_EVENT", "error": "Unhandled event", "details": "ws_tester_foo" } }` |
668 +| `CONNECTION_NOT_FOUND` | Single‑sid request | Target `sid` is not connected/known. | Use an active `sid` or retry after reconnect. | `{ "handlerId": "WsManager", "ok": false, "error": { "code": "CONNECTION_NOT_FOUND", "error": "Connection 'sid-123' not found" } }` |
669 +| `HARNESS_UNKNOWN_EVENT` | Developer harness | Harness test handler received an unsupported event name. | Update harness sources or disable the step before running automation. | `{ "handlerId": "api.ws_dev_test.WsDevTest", "ok": false, "error": { "code": "HARNESS_UNKNOWN_EVENT", "error": "Unhandled event", "details": "ws_tester_foo" } }` |
670
671 Notes
672 - Error payload shape follows the contract documented in `contracts/event-schemas.md` (`RequestResultItem.error`).
knowledge/main/about/architecture.md
+1 -1
@@ -64,4 +64,4 @@ The plugin system (`python/helpers/plugins.py`) discovers plugins from `plugins/
64
65 The web UI is built with Alpine.js and ES module components. The main shell is `webui/index.html`. Components are in `webui/components/`. Frontend state is managed via Alpine stores defined with `createStore` from `/js/AlpineStore.js`.
66
67 -Real-time communication uses Socket.io WebSockets. The backend WebSocket handlers are in `python/websocket_handlers/`. API handlers are in `python/api/`, each deriving from `ApiHandler` in `python/helpers/api.py`.
67 +Real-time communication uses Socket.io WebSockets via a unified `/ws` namespace. WebSocket handlers (WsHandler subclasses) are in `api/ws_*.py`. The connection manager is in `helpers/ws_manager.py`. API handlers are in `api/`, each deriving from `ApiHandler` in `helpers/api.py`.
tests/test_multi_tab_isolation.py
+2 -2
@@ -18,7 +18,7 @@ async def test_state_monitor_per_sid_isolation_independent_snapshots_seq_and_cur
18 snapshot_calls: list[dict[str, object]] = []
19 emitted: list[dict[str, object]] = []
20
21 - namespace = "/webui"
21 + namespace = "/ws"
22
23 async def fake_build_snapshot_from_request(*, request):
24 context = request.context
@@ -133,7 +133,7 @@ async def test_state_monitor_mark_dirty_for_context_scopes_to_active_context():
133 from helpers.state_snapshot import StateRequestV1
134
135 monitor = StateMonitor(debounce_seconds=60.0)
136 - namespace = "/webui"
136 + namespace = "/ws"
137 monitor.register_sid(namespace, "sid-a")
138 monitor.register_sid(namespace, "sid-b")
139
tests/test_state_monitor.py
+1 -1
@@ -13,7 +13,7 @@ async def test_state_monitor_debounce_coalesces_without_postponing_and_cleanup_c
13 from helpers.state_monitor import StateMonitor
14 from helpers.state_snapshot import StateRequestV1
15
16 - namespace = "/webui"
16 + namespace = "/ws"
17 monitor = StateMonitor(debounce_seconds=10.0)
18 monitor.register_sid(namespace, "sid-1")
19 monitor.bind_manager(type("FakeManager", (), {"_dispatcher_loop": None})())
tests/test_state_sync_handler.py
+44 -66
@@ -10,9 +10,9 @@ PROJECT_ROOT = Path(__file__).resolve().parents[1]
10 if str(PROJECT_ROOT) not in sys.path:
11 sys.path.insert(0, str(PROJECT_ROOT))
12
13 -from helpers.websocket_manager import WebSocketManager
13 +from helpers.ws_manager import WsManager
14
15 -NAMESPACE = "/webui"
15 +NAMESPACE = "/ws"
16
17
18 class FakeSocketIOServer:
@@ -23,94 +23,76 @@ class FakeSocketIOServer:
23 self.disconnect = AsyncMock()
24
25
26 -async def _create_manager() -> WebSocketManager:
27 - socketio = FakeSocketIOServer()
28 - manager = WebSocketManager(socketio, threading.RLock())
29 -
30 - from python.websocket_handlers.webui_handler import WebuiHandler
26 +async def _create_manager() -> tuple[WsManager, "WsWebui"]:
27 + from api.ws_webui import WsWebui
28 from helpers.state_monitor import _reset_state_monitor_for_testing
29
30 + socketio = FakeSocketIOServer()
31 + lock = threading.RLock()
32 + manager = WsManager(socketio, lock)
33 +
34 _reset_state_monitor_for_testing()
34 - WebuiHandler._reset_instance_for_testing()
35 - handler = WebuiHandler.get_instance(socketio, threading.RLock())
36 - manager.register_handlers({NAMESPACE: [handler]})
35 + handler = WsWebui(socketio, lock, manager=manager, namespace=NAMESPACE)
36 await manager.handle_connect(NAMESPACE, "sid-1")
38 - return manager
37 + await handler.on_connect("sid-1")
38 + return manager, handler
39
40
41 -async def _create_manager_with_socketio() -> tuple[WebSocketManager, FakeSocketIOServer]:
42 - socketio = FakeSocketIOServer()
43 - manager = WebSocketManager(socketio, threading.RLock())
44 -
45 - from python.websocket_handlers.webui_handler import WebuiHandler
41 +async def _create_manager_with_socketio() -> tuple[WsManager, "WsWebui", FakeSocketIOServer]:
42 + from api.ws_webui import WsWebui
43 from helpers.state_monitor import _reset_state_monitor_for_testing
44
45 + socketio = FakeSocketIOServer()
46 + lock = threading.RLock()
47 + manager = WsManager(socketio, lock)
48 +
49 _reset_state_monitor_for_testing()
49 - WebuiHandler._reset_instance_for_testing()
50 - handler = WebuiHandler.get_instance(socketio, threading.RLock())
51 - manager.register_handlers({NAMESPACE: [handler]})
50 + handler = WsWebui(socketio, lock, manager=manager, namespace=NAMESPACE)
51 await manager.handle_connect(NAMESPACE, "sid-1")
53 - return manager, socketio
52 + await handler.on_connect("sid-1")
53 + return manager, handler, socketio
54
55
56 @pytest.mark.asyncio
57 async def test_state_request_success_returns_wire_level_shape_and_contract_payload():
58 - manager = await _create_manager()
58 + _manager, handler = await _create_manager()
59
60 - response = await manager.route_event(
61 - NAMESPACE,
60 + result = await handler.process(
61 "state_request",
62 {
63 "correlationId": "client-1",
65 - "ts": "2025-12-28T00:00:00.000Z",
66 - "data": {
67 - "context": None,
68 - "log_from": 0,
69 - "notifications_from": 0,
70 - "timezone": "UTC",
71 - },
64 + "context": None,
65 + "log_from": 0,
66 + "notifications_from": 0,
67 + "timezone": "UTC",
68 },
69 "sid-1",
70 )
71
76 - assert response["correlationId"] == "client-1"
77 - assert isinstance(response.get("results"), list)
78 - assert response["results"]
79 -
80 - first = response["results"][0]
81 - assert first["ok"] is True
82 - assert first["correlationId"] == "client-1"
83 - assert isinstance(first.get("data"), dict)
84 - assert set(first["data"].keys()) >= {"runtime_epoch", "seq_base"}
85 - assert isinstance(first["data"]["runtime_epoch"], str) and first["data"]["runtime_epoch"]
86 - assert isinstance(first["data"]["seq_base"], int)
72 + assert isinstance(result, dict)
73 + assert set(result.keys()) >= {"runtime_epoch", "seq_base"}
74 + assert isinstance(result["runtime_epoch"], str) and result["runtime_epoch"]
75 + assert isinstance(result["seq_base"], int)
76
77
78 @pytest.mark.asyncio
79 async def test_state_request_invalid_payload_returns_invalid_request_error():
91 - manager = await _create_manager()
80 + _manager, handler = await _create_manager()
81
93 - response = await manager.route_event(
94 - NAMESPACE,
82 + result = await handler.process(
83 "state_request",
84 {
85 "correlationId": "client-2",
98 - "ts": "2025-12-28T00:00:00.000Z",
99 - "data": {
100 - "context": None,
101 - "log_from": -1,
102 - "notifications_from": 0,
103 - "timezone": "UTC",
104 - },
86 + "context": None,
87 + "log_from": -1,
88 + "notifications_from": 0,
89 + "timezone": "UTC",
90 },
91 "sid-1",
92 )
93
109 - assert response["correlationId"] == "client-2"
110 - assert response["results"]
111 - first = response["results"][0]
112 - assert first["ok"] is False
113 - assert first["error"]["code"] == "INVALID_REQUEST"
94 + assert isinstance(result, dict)
95 + assert result.get("code") == "INVALID_REQUEST"
96
97
98 @pytest.mark.asyncio
@@ -118,7 +100,7 @@ async def test_state_push_gating_and_initial_snapshot_delivery():
100 from helpers.state_monitor import get_state_monitor
101 from helpers.state_snapshot import validate_snapshot_schema_v1
102
121 - manager, socketio = await _create_manager_with_socketio()
103 + manager, handler, socketio = await _create_manager_with_socketio()
104
105 push_ready = asyncio.Event()
106 captured: dict[str, object] = {}
@@ -136,18 +118,14 @@ async def test_state_push_gating_and_initial_snapshot_delivery():
118 assert not push_ready.is_set()
119
120 start = time.monotonic()
139 - await manager.route_event(
140 - NAMESPACE,
121 + await handler.process(
122 "state_request",
123 {
124 "correlationId": "client-gating",
144 - "ts": "2025-12-28T00:00:00.000Z",
145 - "data": {
146 - "context": None,
147 - "log_from": 0,
148 - "notifications_from": 0,
149 - "timezone": "UTC",
150 - },
125 + "context": None,
126 + "log_from": 0,
127 + "notifications_from": 0,
128 + "timezone": "UTC",
129 },
130 "sid-1",
131 )
tests/test_state_sync_welcome_screen.py
+14 -16
@@ -10,9 +10,9 @@ PROJECT_ROOT = Path(__file__).resolve().parents[1]
10 if str(PROJECT_ROOT) not in sys.path:
11 sys.path.insert(0, str(PROJECT_ROOT))
12
13 -from helpers.websocket_manager import WebSocketManager
13 +from helpers.ws_manager import WsManager
14
15 -NAMESPACE = "/webui"
15 +NAMESPACE = "/ws"
16
17
18 class FakeSocketIOServer:
@@ -32,16 +32,18 @@ async def test_state_sync_handshake_and_initial_snapshot_work_with_no_selected_c
32
33 from helpers.state_snapshot import validate_snapshot_schema_v1
34 from helpers.state_monitor import _reset_state_monitor_for_testing
35 - from python.websocket_handlers.webui_handler import WebuiHandler
35 + from api.ws_webui import WsWebui
36
37 socketio = FakeSocketIOServer()
38 - manager = WebSocketManager(socketio, threading.RLock())
38 + manager = WsManager(socketio, threading.RLock())
39
40 _reset_state_monitor_for_testing()
41 - WebuiHandler._reset_instance_for_testing()
42 - handler = WebuiHandler.get_instance(socketio, threading.RLock())
43 - manager.register_handlers({NAMESPACE: [handler]})
41 +
42 + lock = threading.RLock()
43 + handler = WsWebui(socketio, lock, manager=manager, namespace=NAMESPACE)
44 +
45 await manager.handle_connect(NAMESPACE, "sid-1")
46 + await handler.on_connect("sid-1")
47
48 push_ready = asyncio.Event()
49 captured: dict[str, object] = {}
@@ -54,18 +56,14 @@ async def test_state_sync_handshake_and_initial_snapshot_work_with_no_selected_c
56 socketio.emit.side_effect = _emit
57
58 start = time.monotonic()
57 - await manager.route_event(
58 - NAMESPACE,
59 + result = await handler.process(
60 "state_request",
61 {
62 "correlationId": "client-welcome",
62 - "ts": "2026-01-05T00:00:00.000Z",
63 - "data": {
64 - "context": None, # welcome screen (no selected chat)
65 - "log_from": 0,
66 - "notifications_from": 0,
67 - "timezone": "UTC",
68 - },
63 + "context": None,
64 + "log_from": 0,
65 + "notifications_from": 0,
66 + "timezone": "UTC",
67 },
68 "sid-1",
69 )
tests/test_ws_client_api_surface.py new
+41
@@ -0,0 +1,41 @@
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_ws_csrf.py new
+51
@@ -0,0 +1,51 @@
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.ws 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_ws_handlers.py new
+113
@@ -0,0 +1,113 @@
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.ws_manager import WsResult
12 +
13 +
14 +class _FakeSocketIO:
15 + async def emit(self, *_args, **_kwargs): # pragma: no cover - helper stub
16 + return None
17 +
18 + async def disconnect(self, *_args, **_kwargs): # pragma: no cover - helper stub
19 + return None
20 +
21 +
22 +def test_ws_result_ok_clones_payload():
23 + payload = {"value": 1}
24 + result = WsResult.ok(payload)
25 +
26 + assert result.as_result(
27 + handler_id="handler",
28 + fallback_correlation_id="corr",
29 + )["data"] == payload
30 +
31 + payload["value"] = 2
32 + assert result.as_result(
33 + handler_id="handler",
34 + fallback_correlation_id="corr",
35 + )["data"] == {"value": 1}
36 +
37 +
38 +def test_ws_result_error_contains_metadata():
39 + result = WsResult.error(
40 + code="E_TEST",
41 + message="failure",
42 + details="additional",
43 + correlation_id="corr",
44 + duration_ms=12.5,
45 + )
46 +
47 + as_payload = result.as_result(handler_id="handler", fallback_correlation_id=None)
48 + assert as_payload["ok"] is False
49 + assert as_payload["error"] == {
50 + "code": "E_TEST",
51 + "error": "failure",
52 + "details": "additional",
53 + }
54 + assert as_payload["correlationId"] == "corr"
55 + assert as_payload["durationMs"] == pytest.approx(12.5, rel=1e-3)
56 +
57 +
58 +def test_ws_result_applies_fallback_correlation_and_duration():
59 + result = WsResult.ok(duration_ms=5.4321)
60 + payload = result.as_result(
61 + handler_id="handler",
62 + fallback_correlation_id="corr-fallback",
63 + )
64 + assert payload["correlationId"] == "corr-fallback"
65 + assert payload["durationMs"] == pytest.approx(5.4321, rel=1e-3)
66 +
67 +
68 +def test_result_error_requires_error_payload():
69 + with pytest.raises(ValueError):
70 + WsResult(ok=False)
71 +
72 + with pytest.raises(ValueError):
73 + WsResult.error(code="", message="boom")
74 +
75 +
76 +@pytest.mark.asyncio
77 +async def test_state_sync_handler_registers_and_routes_state_request():
78 + from helpers.ws_manager import WsManager
79 + from api.ws_webui import WsWebui
80 + from helpers.state_monitor import _reset_state_monitor_for_testing
81 +
82 + _reset_state_monitor_for_testing()
83 +
84 + socketio = _FakeSocketIO()
85 + lock = threading.RLock()
86 + manager = WsManager(socketio, lock)
87 +
88 + namespace = "/ws"
89 + handler = WsWebui(socketio, lock, manager=manager, namespace=namespace)
90 +
91 + # Register connection with manager (sets up dispatcher loop)
92 + await manager.handle_connect(namespace, "sid-1")
93 + # Trigger StateMonitor binding via extension
94 + await handler.on_connect("sid-1")
95 +
96 + result = await handler.process(
97 + "state_request",
98 + {
99 + "correlationId": "smoke-1",
100 + "context": None,
101 + "log_from": 0,
102 + "notifications_from": 0,
103 + "timezone": "UTC",
104 + },
105 + "sid-1",
106 + )
107 +
108 + assert result is not None
109 + assert "runtime_epoch" in result
110 + assert result.get("seq_base") == 1
111 +
112 + await handler.on_disconnect("sid-1")
113 + await manager.handle_disconnect(namespace, "sid-1")
tests/test_ws_manager.py new
+795
@@ -0,0 +1,795 @@
1 +import asyncio
2 +import sys
3 +import threading
4 +import time
5 +from pathlib import Path
6 +from typing import Any
7 +from unittest.mock import AsyncMock, patch
8 +
9 +import pytest
10 +
11 +PROJECT_ROOT = Path(__file__).resolve().parents[1]
12 +if str(PROJECT_ROOT) not in sys.path:
13 + sys.path.insert(0, str(PROJECT_ROOT))
14 +
15 +from helpers.ws import ConnectionNotFoundError, WsHandler
16 +from helpers.ws_manager import (
17 + WsManager,
18 + WsResult,
19 + BUFFER_TTL,
20 + DIAGNOSTIC_EVENT,
21 + LIFECYCLE_CONNECT_EVENT,
22 + LIFECYCLE_DISCONNECT_EVENT,
23 +)
24 +
25 +NAMESPACE = "/test"
26 +
27 +
28 +class FakeSocketIOServer:
29 + def __init__(self):
30 + self.emit = AsyncMock()
31 + self.disconnect = AsyncMock()
32 +
33 +
34 +class DummyHandler(WsHandler):
35 + def __init__(self, socketio, lock, results=None):
36 + super().__init__(socketio, lock)
37 + self.results = results if results is not None else []
38 +
39 + async def process(self, event: str, data: dict[str, Any], sid: str):
40 + response = {"sid": sid, "data": data}
41 + self.results.append(response)
42 + return response
43 +
44 +
45 +@pytest.mark.asyncio
46 +async def test_connect_disconnect_updates_registry():
47 + socketio = FakeSocketIOServer()
48 + manager = WsManager(socketio, threading.RLock())
49 +
50 + await manager.handle_connect(NAMESPACE, "abc")
51 + assert (NAMESPACE, "abc") in manager.connections
52 +
53 + await manager.handle_disconnect(NAMESPACE, "abc")
54 + assert (NAMESPACE, "abc") not in manager.connections
55 +
56 +
57 +@pytest.mark.asyncio
58 +async def test_server_restart_broadcast_emitted_when_enabled():
59 + socketio = FakeSocketIOServer()
60 + manager = WsManager(socketio, threading.RLock())
61 + manager.set_server_restart_broadcast(True)
62 +
63 + await manager.handle_connect(NAMESPACE, "sid-restart")
64 +
65 + socketio.emit.assert_awaited()
66 + args, kwargs = socketio.emit.await_args_list[0]
67 + assert args[0] == "server_restart"
68 + envelope = args[1]
69 + assert envelope["handlerId"] == manager._identifier # noqa: SLF001
70 + assert envelope["data"]["runtimeId"]
71 + assert kwargs == {"to": "sid-restart", "namespace": NAMESPACE}
72 +
73 +
74 +@pytest.mark.asyncio
75 +async def test_server_restart_broadcast_skipped_when_disabled():
76 + socketio = FakeSocketIOServer()
77 + manager = WsManager(socketio, threading.RLock())
78 + manager.set_server_restart_broadcast(False)
79 +
80 + await manager.handle_connect(NAMESPACE, "sid-no-restart")
81 +
82 + assert socketio.emit.await_count == 0
83 +
84 +
85 +@pytest.mark.asyncio
86 +async def test_broadcast_performance_smoke(monkeypatch):
87 + socketio = FakeSocketIOServer()
88 + manager = WsManager(socketio, threading.RLock())
89 +
90 + for idx in range(50):
91 + await manager.handle_connect(NAMESPACE, f"sid-{idx}")
92 +
93 + # Drain lifecycle broadcast tasks queued by handle_connect and reset the
94 + # mock so we only measure the explicit broadcast below.
95 + for _ in range(55):
96 + await asyncio.sleep(0)
97 + socketio.emit.reset_mock()
98 +
99 + import time
100 +
101 + start = time.perf_counter()
102 + await manager.broadcast(NAMESPACE, "perf_event", {"ok": True})
103 + duration_ms = (time.perf_counter() - start) * 1000
104 +
105 + assert socketio.emit.await_count == 50
106 + assert duration_ms < 300
107 +
108 +
109 +@pytest.mark.asyncio
110 +async def test_route_event_invokes_handler_and_ack():
111 + socketio = FakeSocketIOServer()
112 + manager = WsManager(socketio, threading.RLock())
113 +
114 + results = []
115 + handler = DummyHandler(socketio, threading.RLock(), results)
116 + manager.register_handlers({NAMESPACE: [handler]})
117 + await manager.handle_connect(NAMESPACE, "sid-1")
118 +
119 + response = await manager.route_event(NAMESPACE, "dummy", {"foo": "bar"}, "sid-1")
120 +
121 + assert results[0]["sid"] == "sid-1"
122 + assert results[0]["data"]["foo"] == "bar"
123 + assert "correlationId" in results[0]["data"]
124 +
125 + assert isinstance(response, dict)
126 + assert "correlationId" in response
127 + assert isinstance(response["results"], list)
128 + entry = response["results"][0]
129 + assert entry["ok"] is True
130 + assert entry["data"]["sid"] == "sid-1"
131 + assert entry["data"]["data"]["foo"] == "bar"
132 +
133 +
134 +@pytest.mark.asyncio
135 +async def test_route_event_no_handler_returns_standard_error():
136 + socketio = FakeSocketIOServer()
137 + manager = WsManager(socketio, threading.RLock())
138 + await manager.handle_connect(NAMESPACE, "sid-1")
139 +
140 + response = await manager.route_event(NAMESPACE, "missing", {}, "sid-1")
141 +
142 + assert len(response["results"]) == 1
143 + result = response["results"][0]
144 + assert result["handlerId"].endswith("WsManager")
145 + assert result["ok"] is False
146 + assert result["error"]["code"] == "NO_HANDLERS"
147 + assert (
148 + result["error"]["error"]
149 + == f"No handler for namespace '{NAMESPACE}'"
150 + )
151 +
152 +
153 +@pytest.mark.asyncio
154 +async def test_route_event_all_returns_empty_when_no_connections():
155 + socketio = FakeSocketIOServer()
156 + manager = WsManager(socketio, threading.RLock())
157 +
158 + results = await manager.route_event_all(NAMESPACE, "event", {}, timeout_ms=1000)
159 +
160 + assert results == []
161 +
162 +
163 +@pytest.mark.asyncio
164 +async def test_route_event_all_aggregates_results():
165 + socketio = FakeSocketIOServer()
166 + manager = WsManager(socketio, threading.RLock())
167 +
168 + class EchoHandler(WsHandler):
169 + async def process(self, event: str, data: dict[str, Any], sid: str):
170 + return {"sid": sid, "echo": data}
171 +
172 + handler = EchoHandler(socketio, threading.RLock())
173 + manager.register_handlers({NAMESPACE: [handler]})
174 +
175 + await manager.handle_connect(NAMESPACE, "sid-1")
176 + await manager.handle_connect(NAMESPACE, "sid-2")
177 +
178 + aggregated = await manager.route_event_all(
179 + NAMESPACE, "multi", {"value": 42}, timeout_ms=1000
180 + )
181 +
182 + assert len(aggregated) == 2
183 + by_sid = {entry["sid"]: entry for entry in aggregated}
184 + assert by_sid["sid-1"]["results"][0]["ok"] is True
185 + payload_sid1 = by_sid["sid-1"]["results"][0]["data"]
186 + assert payload_sid1["sid"] == "sid-1"
187 + assert payload_sid1["echo"]["value"] == 42
188 + assert "correlationId" in payload_sid1["echo"]
189 + assert by_sid["sid-2"]["results"][0]["ok"] is True
190 + payload_sid2 = by_sid["sid-2"]["results"][0]["data"]
191 + assert payload_sid2["sid"] == "sid-2"
192 + assert payload_sid2["echo"]["value"] == 42
193 + assert by_sid["sid-1"]["correlationId"]
194 +
195 +
196 +@pytest.mark.asyncio
197 +async def test_route_event_all_timeout_marks_error():
198 + socketio = FakeSocketIOServer()
199 + manager = WsManager(socketio, threading.RLock())
200 +
201 + class SlowHandler(WsHandler):
202 + async def process(self, event: str, data: dict[str, Any], sid: str):
203 + await asyncio.sleep(0.2)
204 + return {"status": "done"}
205 +
206 + handler = SlowHandler(socketio, threading.RLock())
207 + manager.register_handlers({NAMESPACE: [handler]})
208 + await manager.handle_connect(NAMESPACE, "sid-1")
209 +
210 + aggregated = await manager.route_event_all(NAMESPACE, "slow", {}, timeout_ms=50)
211 +
212 + assert len(aggregated) == 1
213 + first_entry = aggregated[0]
214 + result = first_entry["results"][0]
215 + assert result["ok"] is False
216 + assert result["error"] == {"code": "TIMEOUT", "error": "Request timeout"}
217 + assert first_entry["correlationId"]
218 +
219 +
220 +@pytest.mark.asyncio
221 +async def test_route_event_exception_standardizes_error_payload():
222 + socketio = FakeSocketIOServer()
223 + manager = WsManager(socketio, threading.RLock())
224 +
225 + class FailingHandler(WsHandler):
226 + async def process(self, event: str, data: dict[str, Any], sid: str):
227 + raise RuntimeError("kaboom")
228 +
229 + handler = FailingHandler(socketio, threading.RLock())
230 + manager.register_handlers({NAMESPACE: [handler]})
231 + await manager.handle_connect(NAMESPACE, "sid-1")
232 +
233 + response = await manager.route_event(NAMESPACE, "boom", {}, "sid-1")
234 +
235 + assert len(response["results"]) == 1
236 + result = response["results"][0]
237 + assert result["handlerId"].endswith("FailingHandler")
238 + assert result["ok"] is False
239 + assert result["error"]["code"] == "HANDLER_ERROR"
240 + assert result["error"]["error"] == "Internal server error"
241 + assert "details" in result["error"]
242 +
243 +
244 +@pytest.mark.asyncio
245 +async def test_route_event_offloads_blocking_handlers():
246 + socketio = FakeSocketIOServer()
247 + manager = WsManager(socketio, threading.RLock())
248 +
249 + class BlockingHandler(WsHandler):
250 + async def process(self, event: str, data: dict[str, Any], sid: str):
251 + time.sleep(0.2)
252 + return {"status": "done"}
253 +
254 + handler = BlockingHandler(socketio, threading.RLock())
255 + manager.register_handlers({NAMESPACE: [handler]})
256 + await manager.handle_connect(NAMESPACE, "sid-1")
257 +
258 + route_task = asyncio.create_task(
259 + manager.route_event(NAMESPACE, "block", {}, "sid-1")
260 + )
261 + await asyncio.sleep(0)
262 +
263 + t0 = time.perf_counter()
264 + await asyncio.sleep(0.05)
265 + elapsed = time.perf_counter() - t0
266 + assert elapsed < 0.15
267 +
268 + response = await route_task
269 + assert response["results"]
270 +
271 +
272 +@pytest.mark.asyncio
273 +async def test_route_event_unwraps_ts_data_envelope_and_preserves_correlation_id():
274 + socketio = FakeSocketIOServer()
275 + manager = WsManager(socketio, threading.RLock())
276 +
277 + results: list[dict[str, Any]] = []
278 + handler = DummyHandler(socketio, threading.RLock(), results)
279 + manager.register_handlers({NAMESPACE: [handler]})
280 + await manager.handle_connect(NAMESPACE, "sid-1")
281 +
282 + response = await manager.route_event(
283 + NAMESPACE,
284 + "dummy",
285 + {
286 + "correlationId": "client-1",
287 + "ts": "2025-10-29T12:00:00.000Z",
288 + "data": {"value": 123},
289 + },
290 + "sid-1",
291 + )
292 +
293 + assert response["correlationId"] == "client-1"
294 + assert len(results) == 1
295 + handler_payload = results[0]["data"]
296 + assert handler_payload["value"] == 123
297 + assert handler_payload["correlationId"] == "client-1"
298 + assert "ts" not in handler_payload
299 + assert "data" not in handler_payload
300 +
301 +
302 +@pytest.mark.asyncio
303 +async def test_emit_to_unknown_sid_raises_error():
304 + socketio = FakeSocketIOServer()
305 + manager = WsManager(socketio, threading.RLock())
306 +
307 + with pytest.raises(ConnectionNotFoundError):
308 + await manager.emit_to(NAMESPACE, "unknown", "event", {})
309 +
310 +
311 +@pytest.mark.asyncio
312 +async def test_emit_to_known_disconnected_sid_buffers():
313 + socketio = FakeSocketIOServer()
314 + manager = WsManager(socketio, threading.RLock())
315 + await manager.handle_connect(NAMESPACE, "sid-1")
316 + await manager.handle_disconnect(NAMESPACE, "sid-1")
317 +
318 + await manager.emit_to(
319 + NAMESPACE, "sid-1", "event", {"a": 1}, correlation_id="corr-1"
320 + )
321 +
322 + assert (NAMESPACE, "sid-1") in manager.buffers
323 + buffered = list(manager.buffers[(NAMESPACE, "sid-1")])
324 + assert len(buffered) == 1
325 + assert buffered[0].event_type == "event"
326 + assert buffered[0].data == {"a": 1}
327 + assert buffered[0].correlation_id == "corr-1"
328 +
329 +
330 +@pytest.mark.asyncio
331 +async def test_buffer_overflow_drops_oldest(monkeypatch):
332 + socketio = FakeSocketIOServer()
333 + manager = WsManager(socketio, threading.RLock())
334 +
335 + await manager.handle_connect(NAMESPACE, "offline")
336 + await manager.handle_disconnect(NAMESPACE, "offline")
337 +
338 + monkeypatch.setattr("helpers.ws_manager.BUFFER_MAX_SIZE", 2)
339 +
340 + await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 0})
341 + await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 1})
342 + await manager.emit_to(NAMESPACE, "offline", "event", {"idx": 2})
343 +
344 + buffer = manager.buffers[(NAMESPACE, "offline")]
345 + assert len(buffer) == 2
346 + assert buffer[0].data["idx"] == 1
347 + assert buffer[1].data["idx"] == 2
348 +
349 +
350 +@pytest.mark.asyncio
351 +async def test_expired_buffer_entries_are_discarded(monkeypatch):
352 + socketio = FakeSocketIOServer()
353 + manager = WsManager(socketio, threading.RLock())
354 +
355 + await manager.handle_connect(NAMESPACE, "sid-expired")
356 + await manager.handle_disconnect(NAMESPACE, "sid-expired")
357 +
358 + from datetime import timedelta, timezone, datetime
359 +
360 + past = datetime.now(timezone.utc) - (BUFFER_TTL + timedelta(seconds=5))
361 + future = past + BUFFER_TTL + timedelta(seconds=10)
362 +
363 + await manager.emit_to(NAMESPACE, "sid-expired", "event", {"a": 1})
364 + manager.buffers[(NAMESPACE, "sid-expired")][0].timestamp = past
365 +
366 + socketio.emit.reset_mock()
367 +
368 + monkeypatch.setattr(
369 + "helpers.ws_manager._utcnow",
370 + lambda: future,
371 + )
372 + await manager.handle_connect(NAMESPACE, "sid-expired")
373 +
374 + assert socketio.emit.await_count == 0
375 + assert (NAMESPACE, "sid-expired") not in manager.buffers
376 +
377 +
378 +@pytest.mark.asyncio
379 +async def test_flush_buffer_delivers_and_logs(monkeypatch):
380 + socketio = FakeSocketIOServer()
381 + manager = WsManager(socketio, threading.RLock())
382 + await manager.handle_connect(NAMESPACE, "sid-1")
383 + await manager.handle_disconnect(NAMESPACE, "sid-1")
384 +
385 + await manager.emit_to(NAMESPACE, "sid-1", "event", {"a": 1})
386 +
387 + await manager.handle_connect(NAMESPACE, "sid-1")
388 +
389 + assert len(socketio.emit.await_args_list) == 1
390 + awaited_call = socketio.emit.await_args_list[0]
391 + assert awaited_call.args[0] == "event"
392 + envelope = awaited_call.args[1]
393 + assert envelope["data"] == {"a": 1}
394 + assert "eventId" in envelope and "handlerId" in envelope and "ts" in envelope
395 + assert awaited_call.kwargs == {"to": "sid-1", "namespace": NAMESPACE}
396 + assert (NAMESPACE, "sid-1") not in manager.buffers
397 +
398 +
399 +@pytest.mark.asyncio
400 +async def test_broadcast_excludes_multiple_sids():
401 + socketio = FakeSocketIOServer()
402 + manager = WsManager(socketio, threading.RLock())
403 +
404 + for sid in ("sid-1", "sid-2", "sid-3"):
405 + await manager.handle_connect(NAMESPACE, sid)
406 +
407 + # Drain lifecycle broadcast tasks from handle_connect
408 + for _ in range(10):
409 + await asyncio.sleep(0)
410 + socketio.emit.reset_mock()
411 +
412 + await manager.broadcast(
413 + NAMESPACE,
414 + "event",
415 + {"foo": "bar"},
416 + exclude_sids={"sid-1", "sid-3"},
417 + handler_id="custom.broadcast",
418 + correlation_id="corr-b",
419 + )
420 +
421 + assert len(socketio.emit.await_args_list) == 1
422 + awaited_call = socketio.emit.await_args_list[0]
423 + assert awaited_call.args[0] == "event"
424 + envelope = awaited_call.args[1]
425 + assert envelope["data"] == {"foo": "bar"}
426 + assert envelope["handlerId"] == "custom.broadcast"
427 + assert envelope["correlationId"] == "corr-b"
428 + assert "eventId" in envelope and "ts" in envelope
429 + assert awaited_call.kwargs == {"to": "sid-2", "namespace": NAMESPACE}
430 +
431 +
432 +@pytest.mark.asyncio
433 +async def test_emit_to_wraps_envelope_with_metadata():
434 + socketio = FakeSocketIOServer()
435 + manager = WsManager(socketio, threading.RLock())
436 + await manager.handle_connect(NAMESPACE, "sid-meta")
437 +
438 + await manager.emit_to(
439 + NAMESPACE,
440 + "sid-meta",
441 + "meta_event",
442 + {"payload": True},
443 + handler_id="custom.handler",
444 + correlation_id="corr-meta",
445 + )
446 +
447 + socketio.emit.assert_awaited_once()
448 + args, kwargs = socketio.emit.await_args_list[0]
449 + assert args[0] == "meta_event"
450 + envelope = args[1]
451 + assert envelope["handlerId"] == "custom.handler"
452 + assert envelope["correlationId"] == "corr-meta"
453 + assert envelope["data"] == {"payload": True}
454 + assert kwargs == {"to": "sid-meta", "namespace": NAMESPACE}
455 +
456 +
457 +@pytest.mark.asyncio
458 +async def test_timestamps_are_timezone_aware():
459 + socketio = FakeSocketIOServer()
460 + manager = WsManager(socketio, threading.RLock())
461 +
462 + await manager.handle_connect(NAMESPACE, "sid-utc")
463 + info = manager.connections[(NAMESPACE, "sid-utc")]
464 +
465 + assert info.connected_at.tzinfo is not None
466 + assert info.last_activity.tzinfo is not None
467 +
468 + with patch("helpers.ws_manager._utcnow") as mocked_now:
469 + mocked_now.return_value = info.last_activity
470 + await manager.route_event(NAMESPACE, "unknown", {}, "sid-utc")
471 + assert info.last_activity.tzinfo is not None
472 +
473 +class DuplicateHandler(WsHandler):
474 + async def process(self, event: str, data: dict[str, Any], sid: str):
475 + return {"handledBy": self.identifier}
476 +
477 +
478 +class AnotherDuplicateHandler(WsHandler):
479 + async def process(self, event: str, data: dict[str, Any], sid: str):
480 + return {"handledBy": self.identifier}
481 +
482 +
483 +def test_register_handlers_warns_on_duplicates(monkeypatch):
484 + socketio = FakeSocketIOServer()
485 + manager = WsManager(socketio, threading.RLock())
486 +
487 + warnings: list[str] = []
488 +
489 + def capture_warning(message: str) -> None:
490 + warnings.append(message)
491 +
492 + monkeypatch.setattr(
493 + "helpers.print_style.PrintStyle.warning", staticmethod(capture_warning)
494 + )
495 +
496 + handler_a = DuplicateHandler(socketio, threading.RLock())
497 +
498 + manager.register_handlers({NAMESPACE: [handler_a, handler_a]})
499 +
500 + assert any("Duplicate handler registration" in msg for msg in warnings)
501 +
502 +
503 +class NonDictHandler(WsHandler):
504 + async def process(self, event: str, data: dict[str, Any], sid: str):
505 + return "raw-value"
506 +
507 +
508 +@pytest.mark.asyncio
509 +async def test_route_event_standardizes_success_payload():
510 + socketio = FakeSocketIOServer()
511 + manager = WsManager(socketio, threading.RLock())
512 +
513 + handler = NonDictHandler(socketio, threading.RLock())
514 + manager.register_handlers({NAMESPACE: [handler]})
515 +
516 + response = await manager.route_event(NAMESPACE, "non_dict", {}, "sid-123")
517 +
518 + assert len(response["results"]) == 1
519 + assert response["results"][0]["ok"] is True
520 + assert response["results"][0]["data"] == {"result": "raw-value"}
521 +
522 +
523 +class ErrorHandler(WsHandler):
524 + async def process(self, event: str, data: dict[str, Any], sid: str):
525 + raise RuntimeError("BOOM")
526 +
527 +
528 +class ResultHandler(WsHandler):
529 + async def process(self, event: str, data: dict[str, Any], sid: str):
530 + if event == "result_event":
531 + return WsResult.ok({"sid": sid}, correlation_id="explicit", duration_ms=1.234)
532 + return WsResult.error(
533 + code="E_RESULT",
534 + message="boom",
535 + details="test",
536 + )
537 +
538 +
539 +@pytest.mark.asyncio
540 +async def test_route_event_standardizes_error_payload():
541 + socketio = FakeSocketIOServer()
542 + manager = WsManager(socketio, threading.RLock())
543 +
544 + handler = ErrorHandler(socketio, threading.RLock())
545 + manager.register_handlers({NAMESPACE: [handler]})
546 +
547 + response = await manager.route_event(NAMESPACE, "boom", {}, "sid-123")
548 +
549 + assert len(response["results"]) == 1
550 + payload = response["results"][0]
551 + assert payload["ok"] is False
552 + assert payload["error"]["code"] == "HANDLER_ERROR"
553 + assert payload["error"]["error"] == "Internal server error"
554 +
555 +
556 +@pytest.mark.asyncio
557 +async def test_route_event_accepts_websocket_result_instances():
558 + socketio = FakeSocketIOServer()
559 + manager = WsManager(socketio, threading.RLock())
560 +
561 + handler = ResultHandler(socketio, threading.RLock())
562 + manager.register_handlers({NAMESPACE: [handler]})
563 +
564 + response = await manager.route_event(NAMESPACE, "result_event", {}, "sid-123")
565 +
566 + assert response["results"]
567 + payload = response["results"][0]
568 + assert payload["ok"] is True
569 + assert payload["data"] == {"sid": "sid-123"}
570 + assert payload["correlationId"] == "explicit"
571 + assert payload["durationMs"] == pytest.approx(1.234, rel=1e-3)
572 +
573 +
574 +@pytest.mark.asyncio
575 +async def test_route_event_preserves_websocket_result_errors():
576 + socketio = FakeSocketIOServer()
577 + manager = WsManager(socketio, threading.RLock())
578 +
579 + handler = ResultHandler(socketio, threading.RLock())
580 + manager.register_handlers({NAMESPACE: [handler]})
581 +
582 + response = await manager.route_event(NAMESPACE, "result_error", {}, "sid-123")
583 +
584 + payload = response["results"][0]
585 + assert payload["ok"] is False
586 + assert payload["error"] == {"code": "E_RESULT", "error": "boom", "details": "test"}
587 +
588 +
589 +class AlphaFilterHandler(WsHandler):
590 + async def process(self, event: str, data: dict[str, Any], sid: str):
591 + return {"handledBy": self.identifier, "sid": sid}
592 +
593 +
594 +class BetaFilterHandler(WsHandler):
595 + async def process(self, event: str, data: dict[str, Any], sid: str):
596 + return {"handledBy": self.identifier, "sid": sid}
597 +
598 +
599 +@pytest.mark.asyncio
600 +async def test_route_event_include_handlers_filters_results():
601 + socketio = FakeSocketIOServer()
602 + manager = WsManager(socketio, threading.RLock())
603 +
604 + alpha = AlphaFilterHandler(socketio, threading.RLock())
605 + beta = BetaFilterHandler(socketio, threading.RLock())
606 + manager.register_handlers({NAMESPACE: [alpha, beta]})
607 + await manager.handle_connect(NAMESPACE, "sid-filter")
608 +
609 + response = await manager.route_event(
610 + NAMESPACE,
611 + "filter_event",
612 + {
613 + "includeHandlers": [alpha.identifier],
614 + "payload": True,
615 + },
616 + "sid-filter",
617 + )
618 +
619 + assert response["correlationId"]
620 + results = response["results"]
621 + assert len(results) == 1
622 + assert results[0]["handlerId"] == alpha.identifier
623 + assert results[0]["data"]["handledBy"] == alpha.identifier
624 +
625 +
626 +@pytest.mark.asyncio
627 +async def test_route_event_rejects_exclude_handlers_without_permission():
628 + socketio = FakeSocketIOServer()
629 + manager = WsManager(socketio, threading.RLock())
630 +
631 + handler = AlphaFilterHandler(socketio, threading.RLock())
632 + manager.register_handlers({NAMESPACE: [handler]})
633 + await manager.handle_connect(NAMESPACE, "sid-exclude")
634 +
635 + response = await manager.route_event(
636 + NAMESPACE,
637 + "filter_event",
638 + {"excludeHandlers": [handler.identifier]},
639 + "sid-exclude",
640 + )
641 +
642 + result = response["results"][0]
643 + assert result["error"]["code"] == "INVALID_FILTER"
644 + assert "excludeHandlers" in result["error"]["error"]
645 +
646 +
647 +@pytest.mark.asyncio
648 +async def test_route_event_all_respects_exclude_handlers():
649 + socketio = FakeSocketIOServer()
650 + manager = WsManager(socketio, threading.RLock())
651 +
652 + alpha = AlphaFilterHandler(socketio, threading.RLock())
653 + beta = BetaFilterHandler(socketio, threading.RLock())
654 + manager.register_handlers({NAMESPACE: [alpha, beta]})
655 +
656 + await manager.handle_connect(NAMESPACE, "sid-a")
657 + await manager.handle_connect(NAMESPACE, "sid-b")
658 +
659 + aggregated = await manager.route_event_all(
660 + NAMESPACE,
661 + "filter_event",
662 + {"excludeHandlers": [beta.identifier]},
663 + handler_id="test.manager",
664 + )
665 +
666 + assert aggregated
667 + for entry in aggregated:
668 + assert entry["correlationId"]
669 + assert entry["results"]
670 + assert all(result["handlerId"] == alpha.identifier for result in entry["results"])
671 +
672 +
673 +@pytest.mark.asyncio
674 +async def test_route_event_preserves_correlation_id():
675 + socketio = FakeSocketIOServer()
676 + manager = WsManager(socketio, threading.RLock())
677 +
678 + results = []
679 + handler = DummyHandler(socketio, threading.RLock(), results)
680 + manager.register_handlers({NAMESPACE: [handler]})
681 + await manager.handle_connect(NAMESPACE, "sid-correlation")
682 +
683 + response = await manager.route_event(
684 + NAMESPACE,
685 + "dummy",
686 + {"foo": "bar", "correlationId": "manual-correlation"},
687 + "sid-correlation",
688 + )
689 +
690 + assert response["correlationId"] == "manual-correlation"
691 + result = response["results"][0]
692 + assert result["correlationId"] == "manual-correlation"
693 +
694 +
695 +@pytest.mark.asyncio
696 +async def test_request_preserves_explicit_correlation_id():
697 + socketio = FakeSocketIOServer()
698 + manager = WsManager(socketio, threading.RLock())
699 +
700 + handler = DummyHandler(socketio, threading.RLock())
701 + manager.register_handlers({NAMESPACE: [handler]})
702 + await manager.handle_connect(NAMESPACE, "sid-request")
703 +
704 + response = await manager.request_for_sid(
705 + namespace=NAMESPACE,
706 + sid="sid-request",
707 + event_type="dummy",
708 + data={"payload": True, "correlationId": "req-correlation"},
709 + handler_id="tester",
710 + )
711 +
712 + assert response["correlationId"] == "req-correlation"
713 + result = response["results"][0]
714 + assert result["correlationId"] == "req-correlation"
715 +
716 +
717 +@pytest.mark.asyncio
718 +async def test_request_all_entries_include_correlation_id():
719 + socketio = FakeSocketIOServer()
720 + manager = WsManager(socketio, threading.RLock())
721 +
722 + handler = DummyHandler(socketio, threading.RLock())
723 + manager.register_handlers({NAMESPACE: [handler]})
724 +
725 + await manager.handle_connect(NAMESPACE, "sid-1")
726 + await manager.handle_connect(NAMESPACE, "sid-2")
727 +
728 + aggregated = await manager.route_event_all(
729 + NAMESPACE,
730 + "dummy",
731 + {"value": 1, "correlationId": "agg-correlation"},
732 + )
733 +
734 + assert aggregated
735 + for entry in aggregated:
736 + assert entry["correlationId"] == "agg-correlation"
737 + assert entry["results"]
738 + assert entry["results"][0]["correlationId"] == "agg-correlation"
739 +
740 +
741 +def test_debug_logging_respects_runtime_flag(monkeypatch):
742 + socketio = FakeSocketIOServer()
743 + manager = WsManager(socketio, threading.RLock())
744 +
745 + logs: list[str] = []
746 +
747 + def capture(message: str) -> None:
748 + logs.append(message)
749 +
750 + monkeypatch.setattr("helpers.print_style.PrintStyle.debug", staticmethod(capture))
751 + monkeypatch.setenv("A0_WS_DEBUG", "")
752 +
753 + manager._debug("should-not-log") # noqa: SLF001
754 + assert logs == []
755 +
756 + monkeypatch.setenv("A0_WS_DEBUG", "1")
757 + manager._debug("should-log") # noqa: SLF001
758 + assert logs == ["should-log"]
759 +
760 +
761 +@pytest.mark.asyncio
762 +async def test_diagnostic_event_emitted_for_inbound():
763 + socketio = FakeSocketIOServer()
764 + manager = WsManager(socketio, threading.RLock())
765 +
766 + results: list[dict[str, Any]] = []
767 + handler = DummyHandler(socketio, threading.RLock(), results)
768 + manager.register_handlers({NAMESPACE: [handler]})
769 +
770 + await manager.handle_connect(NAMESPACE, "observer")
771 + assert manager.register_diagnostic_watcher(NAMESPACE, "observer") is True
772 + await manager.handle_connect(NAMESPACE, "sid-client")
773 +
774 + await manager.route_event(NAMESPACE, "dummy", {"payload": "value"}, "sid-client")
775 +
776 + emitted_events = [call.args[0] for call in socketio.emit.await_args_list]
777 + assert DIAGNOSTIC_EVENT in emitted_events
778 +
779 +
780 +@pytest.mark.asyncio
781 +async def test_lifecycle_events_broadcast(monkeypatch):
782 + socketio = FakeSocketIOServer()
783 + manager = WsManager(socketio, threading.RLock())
784 +
785 + broadcast_mock = AsyncMock()
786 + monkeypatch.setattr(manager, "broadcast", broadcast_mock)
787 +
788 + await manager.handle_connect(NAMESPACE, "sid-life")
789 + await asyncio.sleep(0)
790 + await manager.handle_disconnect(NAMESPACE, "sid-life")
791 + await asyncio.sleep(0)
792 +
793 + events = [call.args[1] for call in broadcast_mock.await_args_list]
794 + assert LIFECYCLE_CONNECT_EVENT in events
795 + assert LIFECYCLE_DISCONNECT_EVENT in events
tests/test_ws_security.py new
+276
@@ -0,0 +1,276 @@
1 +import sys
2 +from pathlib import Path
3 +from unittest.mock import patch
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.ws import WsHandler, _SecurityContext, _check_security
12 +
13 +
14 +# Handler variants with different security flag combinations
15 +
16 +class _OpenHandler(WsHandler):
17 + @classmethod
18 + def requires_auth(cls) -> bool:
19 + return False
20 +
21 + @classmethod
22 + def requires_csrf(cls) -> bool:
23 + return False
24 +
25 + async def process(self, event, data, sid):
26 + return None
27 +
28 +
29 +class _AuthOnlyHandler(WsHandler):
30 + @classmethod
31 + def requires_auth(cls) -> bool:
32 + return True
33 +
34 + @classmethod
35 + def requires_csrf(cls) -> bool:
36 + return False
37 +
38 + async def process(self, event, data, sid):
39 + return None
40 +
41 +
42 +class _CsrfHandler(WsHandler):
43 + """Default: requires_auth=True, requires_csrf=True (via default)."""
44 +
45 + async def process(self, event, data, sid):
46 + return None
47 +
48 +
49 +class _LoopbackHandler(WsHandler):
50 + @classmethod
51 + def requires_loopback(cls) -> bool:
52 + return True
53 +
54 + @classmethod
55 + def requires_auth(cls) -> bool:
56 + return False
57 +
58 + @classmethod
59 + def requires_csrf(cls) -> bool:
60 + return False
61 +
62 + async def process(self, event, data, sid):
63 + return None
64 +
65 +
66 +class _ApiKeyHandler(WsHandler):
67 + @classmethod
68 + def requires_api_key(cls) -> bool:
69 + return True
70 +
71 + @classmethod
72 + def requires_auth(cls) -> bool:
73 + return False
74 +
75 + @classmethod
76 + def requires_csrf(cls) -> bool:
77 + return False
78 +
79 + async def process(self, event, data, sid):
80 + return None
81 +
82 +
83 +# Helper to build _SecurityContext quickly
84 +
85 +def _ctx(
86 + *,
87 + auth_hash=None,
88 + csrf_token="tok",
89 + client_csrf_token="tok",
90 + csrf_cookie="tok",
91 + remote_addr="127.0.0.1",
92 + api_key=None,
93 +) -> _SecurityContext:
94 + return _SecurityContext(
95 + auth_hash=auth_hash,
96 + csrf_token=csrf_token,
97 + client_csrf_token=client_csrf_token,
98 + csrf_cookie=csrf_cookie,
99 + remote_addr=remote_addr,
100 + api_key=api_key,
101 + )
102 +
103 +
104 +# Open handler (no security)
105 +
106 +def test_open_handler_always_passes():
107 + assert _check_security(_OpenHandler, _ctx()) is None
108 +
109 +
110 +def test_open_handler_passes_even_without_tokens():
111 + assert _check_security(_OpenHandler, _ctx(csrf_token=None, client_csrf_token=None, csrf_cookie=None)) is None
112 +
113 +
114 +# Loopback
115 +
116 +def test_loopback_allows_127_0_0_1():
117 + assert _check_security(_LoopbackHandler, _ctx(remote_addr="127.0.0.1")) is None
118 +
119 +
120 +def test_loopback_allows_ipv6():
121 + assert _check_security(_LoopbackHandler, _ctx(remote_addr="::1")) is None
122 +
123 +
124 +def test_loopback_rejects_remote():
125 + result = _check_security(_LoopbackHandler, _ctx(remote_addr="192.168.1.50"))
126 + assert result is not None
127 + assert result["code"] == "FORBIDDEN"
128 +
129 +
130 +def test_loopback_rejects_none():
131 + result = _check_security(_LoopbackHandler, _ctx(remote_addr=None))
132 + assert result is not None
133 + assert result["code"] == "FORBIDDEN"
134 +
135 +
136 +# Auth
137 +
138 +@patch("helpers.login.get_credentials_hash", return_value="hashed123")
139 +def test_auth_passes_with_matching_hash(_mock):
140 + result = _check_security(_AuthOnlyHandler, _ctx(auth_hash="hashed123"))
141 + assert result is None
142 +
143 +
144 +@patch("helpers.login.get_credentials_hash", return_value="hashed123")
145 +def test_auth_rejects_wrong_hash(_mock):
146 + result = _check_security(_AuthOnlyHandler, _ctx(auth_hash="wrong"))
147 + assert result is not None
148 + assert result["code"] == "AUTH_REQUIRED"
149 +
150 +
151 +@patch("helpers.login.get_credentials_hash", return_value="hashed123")
152 +def test_auth_rejects_missing_hash(_mock):
153 + result = _check_security(_AuthOnlyHandler, _ctx(auth_hash=None))
154 + assert result is not None
155 + assert result["code"] == "AUTH_REQUIRED"
156 +
157 +
158 +@patch("helpers.login.get_credentials_hash", return_value=None)
159 +def test_auth_passes_when_no_credentials_configured(_mock):
160 + """When no password is set (get_credentials_hash returns None/empty),
161 + auth check should pass regardless of the client hash."""
162 + result = _check_security(_AuthOnlyHandler, _ctx(auth_hash=None))
163 + assert result is None
164 +
165 +
166 +# CSRF
167 +
168 +@patch("helpers.login.get_credentials_hash", return_value=None)
169 +def test_csrf_passes_with_all_tokens_matching(_mock):
170 + result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token="abc", csrf_cookie="abc"))
171 + assert result is None
172 +
173 +
174 +@patch("helpers.login.get_credentials_hash", return_value=None)
175 +def test_csrf_rejects_missing_server_token(_mock):
176 + result = _check_security(_CsrfHandler, _ctx(csrf_token=None, client_csrf_token="abc", csrf_cookie="abc"))
177 + assert result is not None
178 + assert result["code"] == "CSRF_MISSING"
179 +
180 +
181 +@patch("helpers.login.get_credentials_hash", return_value=None)
182 +def test_csrf_rejects_missing_client_token(_mock):
183 + result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token=None, csrf_cookie="abc"))
184 + assert result is not None
185 + assert result["code"] == "CSRF_INVALID"
186 +
187 +
188 +@patch("helpers.login.get_credentials_hash", return_value=None)
189 +def test_csrf_rejects_mismatched_client_token(_mock):
190 + result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token="xyz", csrf_cookie="abc"))
191 + assert result is not None
192 + assert result["code"] == "CSRF_INVALID"
193 +
194 +
195 +@patch("helpers.login.get_credentials_hash", return_value=None)
196 +def test_csrf_rejects_mismatched_cookie(_mock):
197 + result = _check_security(_CsrfHandler, _ctx(csrf_token="abc", client_csrf_token="abc", csrf_cookie="wrong"))
198 + assert result is not None
199 + assert result["code"] == "CSRF_COOKIE"
200 +
201 +
202 +# API Key
203 +
204 +@patch("helpers.settings.get_settings", return_value={"mcp_server_token": "secret-key-123"})
205 +def test_api_key_passes_with_correct_key(_mock):
206 + result = _check_security(_ApiKeyHandler, _ctx(api_key="secret-key-123"))
207 + assert result is None
208 +
209 +
210 +@patch("helpers.settings.get_settings", return_value={"mcp_server_token": "secret-key-123"})
211 +def test_api_key_rejects_wrong_key(_mock):
212 + result = _check_security(_ApiKeyHandler, _ctx(api_key="wrong-key"))
213 + assert result is not None
214 + assert result["code"] == "API_KEY_REQUIRED"
215 +
216 +
217 +@patch("helpers.settings.get_settings", return_value={"mcp_server_token": "secret-key-123"})
218 +def test_api_key_rejects_missing_key(_mock):
219 + result = _check_security(_ApiKeyHandler, _ctx(api_key=None))
220 + assert result is not None
221 + assert result["code"] == "API_KEY_REQUIRED"
222 +
223 +
224 +# Combined flags
225 +
226 +class _FullSecurityHandler(WsHandler):
227 + @classmethod
228 + def requires_loopback(cls) -> bool:
229 + return True
230 +
231 + @classmethod
232 + def requires_auth(cls) -> bool:
233 + return True
234 +
235 + @classmethod
236 + def requires_api_key(cls) -> bool:
237 + return True
238 +
239 + async def process(self, event, data, sid):
240 + return None
241 +
242 +
243 +@patch("helpers.login.get_credentials_hash", return_value="hash")
244 +@patch("helpers.settings.get_settings", return_value={"mcp_server_token": "key"})
245 +def test_full_security_passes_when_all_match(_mock_settings, _mock_login):
246 + result = _check_security(
247 + _FullSecurityHandler,
248 + _ctx(
249 + remote_addr="127.0.0.1",
250 + auth_hash="hash",
251 + csrf_token="tok",
252 + client_csrf_token="tok",
253 + csrf_cookie="tok",
254 + api_key="key",
255 + ),
256 + )
257 + assert result is None
258 +
259 +
260 +@patch("helpers.login.get_credentials_hash", return_value="hash")
261 +@patch("helpers.settings.get_settings", return_value={"mcp_server_token": "key"})
262 +def test_full_security_fails_at_first_check_loopback(_mock_settings, _mock_login):
263 + """Loopback check runs first; if it fails, later checks don't matter."""
264 + result = _check_security(
265 + _FullSecurityHandler,
266 + _ctx(
267 + remote_addr="10.0.0.1",
268 + auth_hash="hash",
269 + csrf_token="tok",
270 + client_csrf_token="tok",
271 + csrf_cookie="tok",
272 + api_key="key",
273 + ),
274 + )
275 + assert result is not None
276 + assert result["code"] == "FORBIDDEN"