diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 652845f6268..db36500935c 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -22,6 +22,7 @@ from litellm._uuid import uuid from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import delete_cached_project_object from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup from litellm.proxy.management.teams.access import is_team_admin from litellm.proxy.management_endpoints.common_utils import _set_object_metadata_field from litellm.proxy.management_endpoints.team_admin_field_permissions import team_admin_may_manage_projects @@ -768,11 +769,14 @@ async def update_project( ) if data.team_id is not None and data.team_id != existing_project.team_id: - mismatched_key_count: Final = await _verification_token_table(prisma_client).count( - where={ - "project_id": data.project_id, - "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], - } + mismatched_key_count: Final = await bounded_db_lookup( + prisma_client.writer_db.litellm_verificationtoken.count( + where={ + "project_id": data.project_id, + "OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}], + } + ), + name="project_key_ownership", ) if mismatched_key_count > 0: raise HTTPException( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 88cda0c526e..4b678a48eeb 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -42,6 +42,7 @@ from litellm.constants import ( from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.models.credentials import CredentialItem +from litellm.models.project import LiteLLM_ProjectTable from litellm.proxy._experimental.mcp_server.db import ( rotate_mcp_server_credentials_master_key, rotate_mcp_user_credentials_master_key, @@ -83,6 +84,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage from litellm.proxy.management.teams.access import TEAM_ADMIN_ONLY, TEAM_OR_ORG_ADMIN, is_team_admin @@ -133,13 +135,12 @@ from litellm.proxy.utils import ( handle_exception_on_proxy, is_valid_api_key, ) -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import BaseRepository, record_to_dict from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.config_repository import ConfigParam, ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.prisma_protocols import TableActions -from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import ( DeletedVerificationTokenRepository, DeprecatedVerificationTokenRepository, @@ -1817,7 +1818,15 @@ async def _check_key_project_team( key_team_id: str | None, prisma_client: PrismaClient, ) -> None: - project_obj: Final = await ProjectRepository(prisma_client).find_by_id(project_id) + project_record: Final = await bounded_db_lookup( + prisma_client.writer_db.litellm_projecttable.find_unique(where={"project_id": project_id}), + name="project", + ) + project_obj: Final = ( + LiteLLM_ProjectTable.model_validate(record_to_dict(project_record)) + if project_record is not None + else None + ) if project_obj is None: raise HTTPException( diff --git a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index 83ad594c458..b5046244c31 100644 --- a/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/unit/enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -1236,6 +1236,7 @@ def _project_update_mocks(monkeypatch, stored_metadata: dict) -> mock.MagicMock: mock_prisma.jsonify_object = lambda data: data mock_prisma.db.litellm_projecttable.find_unique = mock.AsyncMock(return_value=existing_row) mock_prisma.db.litellm_projecttable.update = mock.AsyncMock(return_value=mock.MagicMock()) + mock_prisma.writer_db = mock_prisma.db monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True) monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma) @@ -1270,12 +1271,14 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock( return_value=LiteLLM_TeamTable(team_id=destination_team_id) ) + mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0) + mock_prisma.writer_db = mock.MagicMock() async def count_teamless_keys(*, where: Mapping[str, object]) -> int: conditions: Final = where.get("OR") return int(isinstance(conditions, list) and {"team_id": None} in conditions) - mock_prisma.db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=count_teamless_keys) + mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=count_teamless_keys) with pytest.raises(ProxyException) as error: await _run_project_update(project_id, team_id=destination_team_id) @@ -1288,7 +1291,8 @@ async def test_update_project_rejects_move_when_attached_teamless_key_exists( } assert error.value.code == "400" assert expected_detail["error"] in error.value.message - mock_prisma.db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once() + mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited() mock_prisma.db.litellm_projecttable.update.assert_not_awaited() diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 42b5652144a..8a3e31f1644 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -1,3 +1,4 @@ +import asyncio from collections.abc import Mapping from contextlib import ExitStack from typing import Final @@ -44,6 +45,7 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, project_cache_key +from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_key_project_team, @@ -20543,18 +20545,24 @@ def _configure_key_endpoints( user_api_key_cache: UserApiKeyCache, ) -> AsyncMock: mock_prisma_client: Final = _make_generate_mock_prisma() - mock_prisma_client.writer_db = mock_prisma_client.db project_obj: Final = user_api_key_cache.get_cache( key=project_cache_key(_OWNED_PROJECT), model_type=LiteLLM_ProjectTableCachedObj, ) + project_row: Final = ( + LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=project_obj.team_id) + if project_obj is not None + else None + ) mock_prisma_client.db.litellm_projecttable = MagicMock() mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock( - return_value=( - LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=project_obj.team_id) - if project_obj is not None - else None - ) + return_value=project_row + ) + mock_prisma_client.writer_db = MagicMock() + mock_prisma_client.writer_db.litellm_teamtable = mock_prisma_client.db.litellm_teamtable + mock_prisma_client.writer_db.litellm_projecttable = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=project_row ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) @@ -20604,6 +20612,10 @@ async def test_key_generation_uses_database_project_team_when_cache_is_stale( prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) prisma_client.db.litellm_projecttable = MagicMock() prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_KEY_TEAM) + ) + prisma_client.writer_db.litellm_projecttable = MagicMock() + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) ) monkeypatch.setattr(litellm, "default_key_generate_params", None) @@ -20638,6 +20650,36 @@ async def test_key_generation_uses_database_project_team_when_cache_is_stale( assert error.value.detail == expected_detail +@pytest.mark.asyncio +async def test_key_generation_fails_when_writer_project_lookup_exceeds_deadline( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + stalled_lookup: Final = asyncio.Event() + monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05) + + async def never_returns_project(*, where: Mapping[str, object]) -> None: + assert where == {"project_id": _OWNED_PROJECT} + await stalled_lookup.wait() + + prisma_client.writer_db.litellm_projecttable.find_unique = never_returns_project + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(DBLookupDeadlineExceeded) as error: + await asyncio.wait_for( + _common_key_generation_helper( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ), + timeout=1.0, + ) + + assert error.value.lookup == "project" + + @pytest.mark.parametrize("project_team_id", ["team-a", None]) @pytest.mark.asyncio async def test_key_generation_accepts_same_team_and_unowned_projects( @@ -20903,7 +20945,7 @@ async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_change request_fields: dict[str, str], ) -> None: prisma_client: Final = MagicMock() - prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable( project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM, @@ -20925,7 +20967,7 @@ async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_change @pytest.mark.asyncio async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> None: prisma_client: Final = MagicMock() - prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable( project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM, @@ -20956,8 +20998,9 @@ async def test_regenerate_checks_project_team_ownership( ) -> None: existing_key: Final = LiteLLM_VerificationToken(token="abc123", team_id=key_team_id) mock_prisma_client: Final = _make_regenerate_mock_prisma() - mock_prisma_client.db.litellm_projecttable = MagicMock() - mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock( + mock_prisma_client.writer_db = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b") ) user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") @@ -21004,7 +21047,7 @@ async def test_regenerate_checks_project_team_ownership( @pytest.mark.asyncio async def test_key_project_team_validation_uses_project_missing_404() -> None: prisma_client: Final = MagicMock() - prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=None) + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=None) with pytest.raises(HTTPException) as error: await _check_key_project_team(