mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
7e55e90847
commit
4964d55e6b
4 changed files with 81 additions and 21 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue