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:
+2
-2
@@ -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',
|
||||||
|
|||||||
@@ -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
@@ -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()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user