mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Enforce realtime transcription model access
This commit is contained in:
parent
3c5c83f4e4
commit
b94a3061dd
4 changed files with 96 additions and 0 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue