Add WebSocket authentication via api_key

- When --api_key is set, WebSocket connections require auth too
- Supports Authorization: Bearer <key> header or ?token=<key> 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
This commit is contained in:
Aaron Boxer
2026-04-17 10:21:14 -04:00
committed by Aaron Boxer
parent b648bcb2a4
commit 5334ea0f7a
3 changed files with 63 additions and 3 deletions
+2 -2
View File
@@ -100,8 +100,8 @@ if __name__ == "__main__":
'--api_key', '--api_key',
type=str, type=str,
default=None, default=None,
help='Optional API key for authenticating REST API requests. ' help='Optional API key for authenticating REST API and WebSocket connections. '
'Clients must send "Authorization: Bearer <key>" header.' 'Clients must send "Authorization: Bearer <key>" header or "?token=<key>" query parameter.'
) )
parser.add_argument( parser.add_argument(
'--rate_limit_rpm', '--rate_limit_rpm',
+44
View File
@@ -321,6 +321,7 @@ class TestTranscriptionServerCleanup(unittest.TestCase):
ws = MagicMock() ws = MagicMock()
client = MagicMock() client = MagicMock()
self.server.client_manager.add_client(ws, client) self.server.client_manager.add_client(ws, client)
self.cleanup_server = self.server
self.server.cleanup(ws) self.server.cleanup(ws)
self.assertNotIn(ws, self.server.client_manager.clients) self.assertNotIn(ws, self.server.client_manager.clients)
client.cleanup.assert_called_once() client.cleanup.assert_called_once()
@@ -652,5 +653,48 @@ class TestRateLimiting(unittest.TestCase):
self.assertIn("Rate limit", resp.json()["error"]) 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__": if __name__ == "__main__":
unittest.main() unittest.main()
+17 -1
View File
@@ -769,6 +769,21 @@ class TranscriptionServer:
logging.info(f"✅ OpenAI-Compatible API started on http://0.0.0.0:{rest_port}") logging.info(f"✅ OpenAI-Compatible API started on http://0.0.0.0:{rest_port}")
# Original WebSocket server (always supported) # 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( with serve(
functools.partial( functools.partial(
self.recv_audio, self.recv_audio,
@@ -779,7 +794,8 @@ class TranscriptionServer:
trt_py_session=trt_py_session, trt_py_session=trt_py_session,
), ),
host, host,
port port,
**extra_ws_kwargs,
) as server: ) as server:
server.serve_forever() server.serve_forever()