From 5334ea0f7a40d1310444f70518237753c7e806e6 Mon Sep 17 00:00:00 2001 From: Aaron Boxer Date: Fri, 17 Apr 2026 10:21:14 -0400 Subject: [PATCH] Add WebSocket authentication via api_key - When --api_key is set, WebSocket connections require auth too - Supports Authorization: Bearer header or ?token= query param - Unauthenticated connections receive HTTP 401 before upgrade - Uses websockets process_request callback (no resource allocation before auth) - Added 5 unit tests for WebSocket auth handler --- run_server.py | 4 ++-- tests/test_server_extended.py | 44 +++++++++++++++++++++++++++++++++++ whisper_live/server.py | 18 +++++++++++++- 3 files changed, 63 insertions(+), 3 deletions(-) diff --git a/run_server.py b/run_server.py index 858b74e..5c271dc 100644 --- a/run_server.py +++ b/run_server.py @@ -100,8 +100,8 @@ if __name__ == "__main__": '--api_key', type=str, default=None, - help='Optional API key for authenticating REST API requests. ' - 'Clients must send "Authorization: Bearer " header.' + help='Optional API key for authenticating REST API and WebSocket connections. ' + 'Clients must send "Authorization: Bearer " header or "?token=" query parameter.' ) parser.add_argument( '--rate_limit_rpm', diff --git a/tests/test_server_extended.py b/tests/test_server_extended.py index 744b87c..611db28 100644 --- a/tests/test_server_extended.py +++ b/tests/test_server_extended.py @@ -321,6 +321,7 @@ class TestTranscriptionServerCleanup(unittest.TestCase): ws = MagicMock() client = MagicMock() self.server.client_manager.add_client(ws, client) + self.cleanup_server = self.server self.server.cleanup(ws) self.assertNotIn(ws, self.server.client_manager.clients) client.cleanup.assert_called_once() @@ -652,5 +653,48 @@ class TestRateLimiting(unittest.TestCase): self.assertIn("Rate limit", resp.json()["error"]) +class TestWebSocketAuth(unittest.TestCase): + """Tests for the WebSocket process_request auth callback.""" + + def _make_auth_handler(self, api_key): + """Build the same auth function the server creates.""" + def _ws_auth(path, request_headers): + auth = request_headers.get("Authorization", "") + token_param = None + if "?" in path: + from urllib.parse import urlparse, parse_qs + parsed = urlparse(path) + token_param = parse_qs(parsed.query).get("token", [None])[0] + if auth == f"Bearer {api_key}" or token_param == api_key: + return None + return (401, [("Content-Type", "text/plain")], b"Unauthorized\n") + return _ws_auth + + def test_valid_bearer_token(self): + handler = self._make_auth_handler("my-secret") + result = handler("/", {"Authorization": "Bearer my-secret"}) + self.assertIsNone(result) + + def test_invalid_bearer_token(self): + handler = self._make_auth_handler("my-secret") + result = handler("/", {"Authorization": "Bearer wrong"}) + self.assertEqual(result[0], 401) + + def test_missing_auth_header(self): + handler = self._make_auth_handler("my-secret") + result = handler("/", {}) + self.assertEqual(result[0], 401) + + def test_valid_query_token(self): + handler = self._make_auth_handler("my-secret") + result = handler("/?token=my-secret", {}) + self.assertIsNone(result) + + def test_invalid_query_token(self): + handler = self._make_auth_handler("my-secret") + result = handler("/?token=wrong", {}) + self.assertEqual(result[0], 401) + + if __name__ == "__main__": unittest.main() diff --git a/whisper_live/server.py b/whisper_live/server.py index c3e8a70..8132119 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -769,6 +769,21 @@ class TranscriptionServer: logging.info(f"✅ OpenAI-Compatible API started on http://0.0.0.0:{rest_port}") # Original WebSocket server (always supported) + extra_ws_kwargs = {} + if api_key: + def _ws_auth(path, request_headers): + auth = request_headers.get("Authorization", "") + token_param = None + # Check query string for token parameter + if "?" in path: + from urllib.parse import urlparse, parse_qs + parsed = urlparse(path) + token_param = parse_qs(parsed.query).get("token", [None])[0] + if auth == f"Bearer {api_key}" or token_param == api_key: + return None # Allow connection + return (401, [("Content-Type", "text/plain")], b"Unauthorized\n") + extra_ws_kwargs["process_request"] = _ws_auth + with serve( functools.partial( self.recv_audio, @@ -779,7 +794,8 @@ class TranscriptionServer: trt_py_session=trt_py_session, ), host, - port + port, + **extra_ws_kwargs, ) as server: server.serve_forever()