Enforce realtime resolved model scopes

This commit is contained in:
Emerson Gomes 2026-06-06 16:34:52 -05:00
parent b94a3061dd
commit 2961910849
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
2 changed files with 269 additions and 10 deletions

View file

@ -3126,19 +3126,90 @@ async def can_key_call_resolved_model(
valid_token: UserAPIKeyAuth,
llm_router: Optional[litellm.Router],
) -> None:
if valid_token.config:
return
if (
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
skip_key_model_check = valid_token.config or (
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,
)
if not skip_key_model_check:
await can_key_call_model(
model=model,
llm_model_list=llm_model_list,
valid_token=valid_token,
llm_router=llm_router,
)
team_object: Optional[LiteLLM_TeamTableCachedObj] = None
team_object_from_lookup = False
if valid_token.team_id is not None:
try:
team_object = await get_team_object(
team_id=valid_token.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=valid_token.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
team_object_from_lookup = True
except Exception:
team_object = LiteLLM_TeamTableCachedObj(
team_id=valid_token.team_id,
models=valid_token.team_models,
blocked=valid_token.team_blocked,
team_alias=valid_token.team_alias,
metadata=valid_token.team_metadata,
object_permission_id=valid_token.team_object_permission_id,
object_permission=valid_token.team_object_permission,
)
if team_object is not None:
try:
await can_team_access_model(
model=model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases,
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
raise
if not await _key_access_group_grants_model(
model=model,
valid_token=valid_token,
team_object=team_object,
llm_router=llm_router,
):
raise
if valid_token.user_id is not None and team_object_from_lookup:
await _check_team_member_model_access(
model=model,
team_object=team_object,
valid_token=valid_token,
llm_router=llm_router,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if valid_token.project_id is not None:
project_object = await get_project_object(
project_id=valid_token.project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if project_object is not None and len(project_object.models) > 0:
can_project_access_model(
model=model,
project_object=project_object,
llm_router=llm_router,
)
def can_org_access_model(

View file

@ -467,6 +467,152 @@ async def test_transcription_sessions_rejects_disallowed_resolved_model(
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
@pytest.mark.asyncio
async def test_transcription_sessions_rejects_disallowed_team_model_scope(
proxy_app,
):
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
team = LiteLLM_TeamTableCachedObj(
team_id="team-a",
models=["gpt-4o-realtime-preview"],
)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user",
team_id="team-a",
models=["*"],
)
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,
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new=AsyncMock(return_value=team),
),
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new=AsyncMock(return_value=None),
),
):
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 "team" in response.text.lower()
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_transcription_sessions_rejects_disallowed_project_model_scope(
proxy_app,
):
from litellm.proxy._types import LiteLLM_ProjectTableCachedObj
project = LiteLLM_ProjectTableCachedObj(
project_id="project-a",
models=["gpt-4o-realtime-preview"],
created_by="test-user",
updated_by="test-user",
)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user",
project_id="project-a",
models=["*"],
)
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,
patch(
"litellm.proxy.auth.auth_checks.get_project_object",
new=AsyncMock(return_value=project),
),
):
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 "project" in response.text.lower()
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_transcription_sessions_rejects_disallowed_team_member_model_scope(
proxy_app,
):
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTableCachedObj,
)
team = LiteLLM_TeamTableCachedObj(team_id="team-a", models=["*"])
membership = LiteLLM_TeamMembership(
user_id="test-user",
team_id="team-a",
litellm_budget_table=LiteLLM_BudgetTable(
allowed_models=["gpt-4o-realtime-preview"],
),
)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user",
team_id="team-a",
models=["*"],
)
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,
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new=AsyncMock(return_value=team),
),
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new=AsyncMock(return_value=membership),
),
):
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 "Team member not allowed to access model" 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
@ -490,6 +636,48 @@ async def test_realtime_transcription_websocket_default_model_checks_key_scope()
assert "not allowed to access model" in close_kwargs["reason"]
@pytest.mark.asyncio
async def test_realtime_transcription_websocket_default_model_checks_team_scope():
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
team = LiteLLM_TeamTableCachedObj(
team_id="team-a",
models=["gpt-4o-realtime-preview"],
)
websocket = MagicMock()
websocket.headers = {}
websocket.close = AsyncMock()
websocket.accept = AsyncMock()
with (
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new=AsyncMock(return_value=team),
),
patch(
"litellm.proxy.auth.auth_checks.get_team_membership",
new=AsyncMock(return_value=None),
),
):
await proxy_server.realtime_websocket_endpoint(
websocket=websocket,
model=None,
intent="transcription",
user_api_key_dict=UserAPIKeyAuth(
user_id="test-user",
team_id="team-a",
models=["*"],
),
)
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,