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"