diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 6e4a2dde685..6877aea63cf 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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( diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py index e443d9ce136..ffd438bd015 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py @@ -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,