fix(proxy): read project ownership from the primary database under the lookup deadline

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-02 18:19:01 +00:00
parent 7e55e90847
commit 4964d55e6b
4 changed files with 81 additions and 21 deletions

View file

@ -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(

View file

@ -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(

View file

@ -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()

View file

@ -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(