From e53d70d39020dbb8d6a12f992eb5603c3d8210e7 Mon Sep 17 00:00:00 2001 From: fusion44 Date: Sat, 11 Jul 2026 11:23:40 +0200 Subject: [PATCH] fix(api): harden WebSocket auth handshake against bad/disconnecting clients Co-Authored-By: Claude Fable 5 --- app/api/ws_manager.py | 8 ++++- tests/test_ws_manager.py | 75 +++++++++++++++++++++++++++++++++++++++- 2 files changed, 81 insertions(+), 2 deletions(-) diff --git a/app/api/ws_manager.py b/app/api/ws_manager.py index 03fe57b..f0dc6f3 100644 --- a/app/api/ws_manager.py +++ b/app/api/ws_manager.py @@ -3,6 +3,7 @@ import json from fastapi.encoders import jsonable_encoder from loguru import logger +from starlette.websockets import WebSocketDisconnect from app.auth.auth_bearer import JWTBearer @@ -27,10 +28,15 @@ class WebSocketManager: except (asyncio.TimeoutError, TimeoutError): await websocket.close(code=4408) return None, False + except WebSocketDisconnect: + return None, False try: msg = json.loads(raw) - token = msg["token"] if msg.get("type") == "auth" else None + if not isinstance(msg, dict): + token = None + else: + token = msg["token"] if msg.get("type") == "auth" else None except (ValueError, TypeError, KeyError): token = None diff --git a/tests/test_ws_manager.py b/tests/test_ws_manager.py index 4be271a..70ee7b5 100644 --- a/tests/test_ws_manager.py +++ b/tests/test_ws_manager.py @@ -2,6 +2,7 @@ import asyncio import json import pytest +from starlette.websockets import WebSocketDisconnect from app.api.ws_manager import WebSocketManager @@ -9,21 +10,27 @@ from app.api.ws_manager import WebSocketManager class FakeWebSocket: """Minimal stand-in for starlette WebSocket.""" - def __init__(self, incoming=None): + def __init__(self, incoming=None, raise_on_send=False, disconnect_on_receive=False): self.accepted = False self.closed_code = None self.sent = [] self._incoming = list(incoming or []) + self.raise_on_send = raise_on_send + self.disconnect_on_receive = disconnect_on_receive async def accept(self): self.accepted = True async def receive_text(self): + if self.disconnect_on_receive: + raise WebSocketDisconnect(code=1000) if not self._incoming: await asyncio.sleep(3600) # never sends -> block until cancelled return self._incoming.pop(0) async def send_text(self, text): + if self.raise_on_send: + raise RuntimeError("connection closed") self.sent.append(text) async def close(self, code=1000): @@ -71,3 +78,69 @@ async def test_auth_timeout_closes_4408(monkeypatch): assert authed is False assert ws.closed_code == 4408 + + +async def test_non_dict_json_first_frame_closes_4401(monkeypatch): + monkeypatch.setattr( + "app.api.ws_manager.JWTBearer", + lambda: type("B", (), {"verify_jwt": lambda self, jwtoken: True})(), + ) + mgr = WebSocketManager() + ws = FakeWebSocket([json.dumps("42")]) # valid JSON, not a dict + + id_, authed = await mgr.connect(ws) + + assert authed is False + assert id_ is None + assert ws.closed_code == 4401 + + +async def test_client_disconnect_during_auth_returns_unauthenticated(monkeypatch): + mgr = WebSocketManager() + ws = FakeWebSocket(disconnect_on_receive=True) + + id_, authed = await mgr.connect(ws) + + assert authed is False + assert id_ is None + # socket is already gone; connect() must not attempt to close() it + assert ws.closed_code is None + + +async def test_broadcast_to_all_delivers_to_multiple_connections(monkeypatch): + mgr = WebSocketManager() + ws1 = FakeWebSocket() + ws2 = FakeWebSocket() + mgr._connections[1] = ws1 + mgr._connections[2] = ws2 + + await mgr.broadcast_to_all("btc_info", {"blocks": 1}) + + assert json.loads(ws1.sent[0]) == {"event": "btc_info", "data": {"blocks": 1}} + assert json.loads(ws2.sent[0]) == {"event": "btc_info", "data": {"blocks": 1}} + + +async def test_send_to_single_drops_connection_on_send_failure(monkeypatch): + mgr = WebSocketManager() + ws = FakeWebSocket(raise_on_send=True) + mgr._connections[1] = ws + + await mgr.send_to_single(1, "btc_info", {"blocks": 1}) + + assert 1 not in mgr._connections + + +async def test_broadcast_to_all_drops_failing_connection_but_reaches_others( + monkeypatch, +): + mgr = WebSocketManager() + bad_ws = FakeWebSocket(raise_on_send=True) + good_ws = FakeWebSocket() + mgr._connections[1] = bad_ws + mgr._connections[2] = good_ws + + await mgr.broadcast_to_all("btc_info", {"blocks": 1}) + + assert 1 not in mgr._connections + assert 2 in mgr._connections + assert json.loads(good_ws.sent[0]) == {"event": "btc_info", "data": {"blocks": 1}}