Enforce realtime transcription model access

This commit is contained in:
Emerson Gomes 2026-06-06 14:24:14 -05:00
parent 3c5c83f4e4
commit b94a3061dd
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
4 changed files with 96 additions and 0 deletions

View file

@ -3120,6 +3120,27 @@ async def can_key_call_model(
raise
async def can_key_call_resolved_model(
model: str,
llm_model_list: Optional[list],
valid_token: UserAPIKeyAuth,
llm_router: Optional[litellm.Router],
) -> None:
if valid_token.config:
return
if (
isinstance(valid_token.models, list)
and SpecialModelNames.all_team_models.value in valid_token.models
):
return
await can_key_call_model(
model=model,
llm_model_list=llm_model_list,
valid_token=valid_token,
llm_router=llm_router,
)
def can_org_access_model(
model: str,
org_object: Optional[LiteLLM_OrganizationTable],

View file

@ -254,6 +254,7 @@ from litellm.proxy.analytics_endpoints.analytics_endpoints import (
)
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
can_key_call_resolved_model,
get_team_object,
log_db_metrics,
)
@ -9476,6 +9477,16 @@ async def realtime_websocket_endpoint(
)
return
assert route_model is not None
try:
await can_key_call_resolved_model(
model=route_model,
llm_model_list=llm_model_list,
valid_token=user_api_key_dict,
llm_router=llm_router,
)
except ProxyException as e:
await websocket.close(code=1008, reason=e.message[:120])
return
await websocket.accept(**accept_kwargs)
# Only use explicit parameters, not all query params

View file

@ -10,6 +10,7 @@ from fastapi import status as http_status
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
@ -410,6 +411,7 @@ async def create_realtime_transcription_session(
add_litellm_data_to_request,
general_settings,
llm_router,
llm_model_list,
proxy_config,
proxy_logging_obj,
route_request,
@ -423,6 +425,12 @@ async def create_realtime_transcription_session(
req = RealtimeTranscriptionSessionRequest(**body)
model: str = req.resolved_model() or "gpt-realtime-whisper"
await can_key_call_resolved_model(
model=model,
valid_token=user_api_key_dict,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
transcription_session = {k: v for k, v in body.items() if k != "model"}
data = {"model": model, "transcription_session": transcription_session}
@ -464,6 +472,8 @@ async def create_realtime_transcription_session(
"litellm.proxy.realtime_endpoints.create_realtime_transcription_session(): Exception - %s",
str(e),
)
if isinstance(e, ProxyException):
raise e
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", getattr(e, "message", str(e))),

View file

@ -436,6 +436,60 @@ def test_transcription_sessions_requires_auth(proxy_app):
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
@pytest.mark.asyncio
async def test_transcription_sessions_rejects_disallowed_resolved_model(
proxy_app,
):
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user",
models=["gpt-4o-realtime-preview"],
)
try:
client = TestClient(proxy_app, raise_server_exceptions=False)
with (
patch("litellm.proxy.proxy_server.route_request") as mock_route_request,
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
):
mock_logging.post_call_failure_hook = AsyncMock()
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
json={
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
},
)
assert response.status_code == 403
assert "Tried to access gpt-realtime-whisper" in response.text
mock_route_request.assert_not_called()
finally:
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
@pytest.mark.asyncio
async def test_realtime_transcription_websocket_default_model_checks_key_scope():
from litellm.proxy import proxy_server
websocket = MagicMock()
websocket.headers = {}
websocket.close = AsyncMock()
websocket.accept = AsyncMock()
await proxy_server.realtime_websocket_endpoint(
websocket=websocket,
model=None,
intent="transcription",
user_api_key_dict=UserAPIKeyAuth(models=["gpt-4o-realtime-preview"]),
)
websocket.accept.assert_not_awaited()
websocket.close.assert_awaited_once()
_, close_kwargs = websocket.close.call_args
assert close_kwargs["code"] == 1008
assert "not allowed to access model" in close_kwargs["reason"]
@pytest.mark.asyncio
async def test_transcription_sessions_encrypts_client_secret(
proxy_app,