This commit is contained in:
devin-ai-integration[bot] 2026-10-04 21:36:46 +00:00 • committed by GitHub
commit e7e092365f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 3182 additions and 223 deletions

View file

@ -12,7 +12,7 @@ Endpoints for /project operations
import json
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, cast
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import TypeAdapter
@ -29,6 +29,7 @@ from litellm.proxy.management_helpers.utils import (
management_endpoint_wrapper,
)
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
from litellm.repositories.base_repository import record_to_dict
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
@ -55,6 +56,15 @@ def _project_table(prisma_client: PrismaClient) -> TableActions["prisma_models.L
return ProjectRepository(prisma_client).table
def _writer_project_table(
prisma_client: PrismaClient,
) -> TableActions["prisma_models.LiteLLM_ProjectTable"]:
return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper
"TableActions[prisma_models.LiteLLM_ProjectTable]",
prisma_client.writer_db.litellm_projecttable,
)
def _verification_token_table(
prisma_client: PrismaClient,
) -> TableActions["prisma_models.LiteLLM_VerificationToken"]:
@ -791,6 +801,38 @@ async def update_project(
**({"max_budget": None} if "max_budget" in data.model_fields_set and data.max_budget is None else {}),
}
object_permission_data: Final = update_data.pop("object_permission", None)
object_permission_payload: Final = (
_OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data) if object_permission_data else None
)
if data.team_id is not None:
current_project_record: Final = await _writer_project_table(prisma_client).find_unique(
where={"project_id": data.project_id}
)
current_project: Final = (
LiteLLM_ProjectTable.model_validate(record_to_dict(current_project_record))
if current_project_record is not None
else None
)
if current_project is not None and data.team_id != current_project.team_id:
mismatched_key_count: Final = await prisma_client.writer_db.litellm_verificationtoken.count(
where={
"project_id": data.project_id,
"OR": [{"team_id": {"not": data.team_id}}, {"team_id": None}],
}
)
if mismatched_key_count > 0:
raise HTTPException(
status_code=400,
detail={
"error": (
f"Project {data.project_id} has {mismatched_key_count} key(s) that do not belong to "
f"team {data.team_id}. Detach or delete them before moving the project."
)
},
)
if budget_updates and existing_project.budget_id:
# Update existing budget
await _budget_table(prisma_client).update(
@ -804,23 +846,17 @@ async def update_project(
for field in budget_updates.keys():
update_data.pop(field, None)
# Handle object permissions
if "object_permission" in update_data:
object_permission_data = update_data.pop("object_permission")
if object_permission_data:
object_permission_payload: Final = _OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data)
if existing_project.object_permission_id:
# Update existing permission
await _object_permission_table(prisma_client).update(
where={"object_permission_id": existing_project.object_permission_id},
data=object_permission_payload,
)
else:
# Create new permission
created_permission: Final = await _object_permission_table(prisma_client).create(
data=object_permission_payload,
)
update_data["object_permission_id"] = created_permission.object_permission_id
if object_permission_payload is not None:
if existing_project.object_permission_id:
await _object_permission_table(prisma_client).update(
where={"object_permission_id": existing_project.object_permission_id},
data=object_permission_payload,
)
else:
created_permission: Final = await _object_permission_table(prisma_client).create(
data=object_permission_payload,
)
update_data["object_permission_id"] = created_permission.object_permission_id
# Handle metadata fields
for field in LiteLLM_ManagementEndpoint_MetadataFields:

View file

@ -43,6 +43,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,
@ -113,10 +114,11 @@ from litellm.proxy.management_helpers.access_group_key_sync import (
)
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
from litellm.proxy.management_helpers.object_permission_utils import (
ObjectPermissionUpsert,
_set_object_permission,
attach_object_permission_to_dict,
handle_update_object_permission_common,
invalidate_cached_object_permissions,
prepare_object_permission_upsert,
validate_key_mcp_servers_against_team,
validate_key_search_tools_against_team,
validate_key_vector_stores_against_team,
@ -136,11 +138,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.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import (
DeletedVerificationTokenRepository,
@ -264,6 +267,15 @@ def _prisma_table(
)
def _writer_project_table(
prisma_client: PrismaClient,
) -> "TableActions[prisma_models.LiteLLM_ProjectTable]":
return cast( # cast-ok: writer_db exposes generated Prisma tables through a dynamic wrapper
"TableActions[prisma_models.LiteLLM_ProjectTable]",
prisma_client.writer_db.litellm_projecttable,
)
def _deleted_verification_token_table(
prisma_client: PrismaClient,
) -> "TableActions[prisma_models.LiteLLM_DeletedVerificationToken]":
@ -1334,7 +1346,8 @@ async def _common_key_generation_helper(
apply_enterprise_key_management_params,
)
data = apply_enterprise_key_management_params(data, team_table)
enterprise_data: Final[object] = apply_enterprise_key_management_params(data, team_table)
data = GenerateKeyRequest.model_validate(enterprise_data)
except Exception as e:
verbose_proxy_logger.debug(
"litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - %s", e
@ -1369,6 +1382,10 @@ async def _common_key_generation_helper(
)
_budget_id = getattr(_budget, "budget_id", None)
created_budget_id: Final[str | None] = (
_budget_id if prisma_client is not None and data.soft_budget is not None else None
)
# ADD METADATA FIELDS
# Set Management Endpoint Metadata Fields
for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
@ -1474,10 +1491,16 @@ async def _common_key_generation_helper(
for _op_field, _op_default_value in _default_object_permission.items():
_caller_object_permission.setdefault(_op_field, _op_default_value)
should_create_object_permission: Final = prisma_client is not None and isinstance(
data_json.get("object_permission"), dict
)
data_json = await _set_object_permission(
data_json=data_json,
prisma_client=prisma_client,
)
created_object_permission_id: Final[str | None] = (
cast(str | None, data_json.get("object_permission_id")) if should_create_object_permission else None
)
_validate_key_alias_format(key_alias=data_json.get("key_alias", None))
@ -1545,7 +1568,19 @@ async def _common_key_generation_helper(
prisma_client=prisma_client,
)
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router)
try:
response = await generate_key_helper_fn(
request_type="key", **data_json, table_name="key", llm_router=llm_router
)
except KeyProjectTeamMismatchError:
if prisma_client is not None:
if created_object_permission_id is not None:
await ObjectPermissionRepository(prisma_client).table.delete(
where={"object_permission_id": created_object_permission_id}
)
if created_budget_id is not None:
await BudgetRepository(prisma_client).table.delete(where={"budget_id": created_budget_id})
raise
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
@ -1807,6 +1842,59 @@ async def _check_project_key_limits(
)
async def _check_key_project_team(
project_id: str,
key_team_id: str | None,
prisma_client: PrismaClient,
) -> None:
project_record: Final = await _writer_project_table(prisma_client).find_unique(where={"project_id": project_id})
if project_record is None:
return
project_obj: Final = LiteLLM_ProjectTable.model_validate(record_to_dict(project_record))
if project_obj.team_id is None or project_obj.team_id == key_team_id:
return
raise KeyProjectTeamMismatchError(
status_code=400,
detail={
"error": (
f"Project {project_id} belongs to team {project_obj.team_id}, but the key belongs to "
f"{key_team_id if key_team_id is not None else 'no team'}. "
"A key can only be attached to a project owned by its own team."
)
},
)
class KeyProjectTeamMismatchError(HTTPException):
pass
async def _check_key_project_team_on_mutation(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
) -> None:
fields_set: Final = data.model_fields_set
team_changed: Final = "team_id" in fields_set and data.team_id != existing_key_row.team_id
project_changed: Final = "project_id" in fields_set and data.project_id != existing_key_row.project_id
if not team_changed and not project_changed:
return
project_id: Final = data.project_id if "project_id" in fields_set else existing_key_row.project_id
if project_id is None:
return
team_id: Final = data.team_id if "team_id" in fields_set else existing_key_row.team_id
await _check_key_project_team(
project_id=project_id,
key_team_id=team_id,
prisma_client=prisma_client,
)
def check_org_key_model_specific_limits(
keys: Sequence[LiteLLM_VerificationToken],
org_table: LiteLLM_OrganizationTable,
@ -2460,6 +2548,18 @@ async def _update_key_row_with_soft_budget(
return result
async def _update_key_row(
prisma_client: PrismaClient,
key: str,
update_values: Mapping[str, object],
) -> _KeyUpdateResult | None:
key_update_data: Final = MappingProxyType({**update_values, "token": key})
response: Final = await prisma_client.update_data(token=key, data=key_update_data)
if response is None:
return None
return cast("_KeyUpdateResult", response) # cast-ok: key update_data returns token and data
async def prepare_key_update_data(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
@ -2566,27 +2666,42 @@ async def prepare_key_update_data(
return non_default_values
async def _handle_update_object_permission(
data_json: dict,
async def _prepare_key_update_object_permission(
object_permission_data: object,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
) -> dict:
"""Persist the requested object permission row and swap it for its id, only after the key policy allowed the write."""
if "object_permission" not in data_json:
return data_json
) -> ObjectPermissionUpsert | None:
if object_permission_data is None:
return None
object_permission_id: Final = await handle_update_object_permission_common(
data_json=data_json,
parsed_object_permission: Final[object] = (
json.loads(object_permission_data) if isinstance(object_permission_data, str) else object_permission_data
)
permission_data: Final[dict[str, object]] = (
TypeAdapter(dict[str, object]).validate_python(parsed_object_permission)
if isinstance(parsed_object_permission, dict)
else {}
)
return await prepare_object_permission_upsert(
new_object_permission=permission_data,
existing_object_permission_id=existing_key_row.object_permission_id,
prisma_client=prisma_client,
)
# Add the object_permission_id to data_json if one was created/updated
if object_permission_id is not None:
data_json["object_permission_id"] = object_permission_id
verbose_proxy_logger.debug("updated object_permission_id: %s", object_permission_id)
return data_json
async def _write_prepared_key_update_object_permission(
data_json: Mapping[str, object],
upsert: ObjectPermissionUpsert | None,
prisma_client: PrismaClient,
) -> Mapping[str, object]:
if upsert is None:
return data_json
await ObjectPermissionRepository(prisma_client).table.upsert(
where={"object_permission_id": upsert.object_permission_id},
data={"create": upsert.record, "update": upsert.record},
)
return MappingProxyType({**data_json, "object_permission_id": upsert.object_permission_id})
def is_different_team(data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken) -> bool:
@ -2698,6 +2813,23 @@ def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_V
return existing_key_row.token
async def _check_single_key_update_team_permissions(
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient | None,
existing_key_row: LiteLLM_VerificationToken,
user_api_key_cache: UserApiKeyCache,
) -> None:
if prisma_client is None:
return
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
user_api_key_dict=user_api_key_dict,
route=KeyManagementRoutes.KEY_UPDATE,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
async def _process_single_key_update(
update_key_request: UpdateKeyRequest,
user_api_key_dict: UserAPIKeyAuth,
@ -2761,15 +2893,12 @@ async def _process_single_key_update(
entity="key",
)
# Check team member permissions
if prisma_client is not None:
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
user_api_key_dict=user_api_key_dict,
route=KeyManagementRoutes.KEY_UPDATE,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
await _check_single_key_update_team_permissions(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
# Custom key update hook
if user_custom_key_update is not None:
@ -2854,11 +2983,24 @@ async def _process_single_key_update(
detail={"error": "Database not connected"},
)
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
object_permission_upsert: Final = await _prepare_key_update_object_permission(
object_permission_data=non_default_values.get("object_permission"),
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
key_update_values: Final = MappingProxyType(
{field: value for field, value in non_default_values.items() if field != "object_permission"}
)
await _check_key_project_team_on_mutation(
data=key_request,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
update_values: Final = await _write_prepared_key_update_object_permission(
data_json=key_update_values,
upsert=object_permission_upsert,
prisma_client=prisma_client,
)
_data: Final = {**update_values, "token": key_request.key}
response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict
"Mapping[str, object] | None",
@ -2869,7 +3011,7 @@ async def _process_single_key_update(
await invalidate_cached_object_permissions(
object_permission_ids=(
existing_key_row.object_permission_id,
non_default_values.get("object_permission_id"),
update_values.get("object_permission_id"),
),
user_api_key_cache=user_api_key_cache,
)
@ -3527,11 +3669,24 @@ async def update_key_fn(
if prisma_client is None:
raise Exception("Not connected to DB!")
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
object_permission_upsert: Final = await _prepare_key_update_object_permission(
object_permission_data=non_default_values.get("object_permission"),
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
key_update_values: Final = MappingProxyType(
{field: value for field, value in non_default_values.items() if field != "object_permission"}
)
await _check_key_project_team_on_mutation(
data=data,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
update_values: Final = await _write_prepared_key_update_object_permission(
data_json=key_update_values,
upsert=object_permission_upsert,
prisma_client=prisma_client,
)
changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name
response: Final = (
await _update_key_row_with_soft_budget(
@ -3543,7 +3698,11 @@ async def update_key_fn(
changed_by=changed_by,
)
if "soft_budget" in data.model_fields_set
else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key}))
else await _update_key_row(
prisma_client=prisma_client,
key=key,
update_values=update_values,
)
)
# Delete - key from cache, since it's been updated!
@ -3551,7 +3710,7 @@ async def update_key_fn(
await invalidate_cached_object_permissions(
object_permission_ids=(
existing_key_row.object_permission_id,
non_default_values.get("object_permission_id"),
update_values.get("object_permission_id"),
),
user_api_key_cache=user_api_key_cache,
)
@ -4810,6 +4969,13 @@ async def generate_key_helper_fn(
# the LiteLLM_VerificationToken table will increase in size if we don't do this check
return user_data
if project_id is not None:
await _check_key_project_team(
project_id=project_id,
key_team_id=team_id,
prisma_client=prisma_client,
)
## CREATE KEY
verbose_proxy_logger.debug(
"prisma_client: Creating Key= %s",
@ -5613,7 +5779,6 @@ async def _execute_virtual_key_regeneration(
new_token: Final = await get_new_token(data=data)
new_token_hash: Final = hash_token(new_token)
new_token_key_name: Final = abbreviate_api_key(api_key=new_token)
update_data = {"token": new_token_hash, "key_name": new_token_key_name}
non_default_values = {}
if data is not None:
@ -5639,12 +5804,27 @@ async def _execute_virtual_key_regeneration(
request=data if data is not None else RegenerateKeyRequest(),
),
)
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
object_permission_upsert: Final = await _prepare_key_update_object_permission(
object_permission_data=non_default_values.get("object_permission"),
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
update_data.update(update_values)
key_update_values: Final = MappingProxyType(
{field: value for field, value in non_default_values.items() if field != "object_permission"}
)
if data is not None:
await _check_key_project_team_on_mutation(
data=data,
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
update_values: Final = await _write_prepared_key_update_object_permission(
data_json=key_update_values,
upsert=object_permission_upsert,
prisma_client=prisma_client,
)
update_data: Final = MappingProxyType({"token": new_token_hash, "key_name": new_token_key_name, **update_values})
jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data)
# Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash,
@ -5682,7 +5862,7 @@ async def _execute_virtual_key_regeneration(
await invalidate_cached_object_permissions(
object_permission_ids=(
key_in_db.object_permission_id,
non_default_values.get("object_permission_id"),
update_values.get("object_permission_id"),
),
user_api_key_cache=user_api_key_cache,
)

File diff suppressed because it is too large Load diff

View file

@ -1,11 +1,17 @@
import asyncio
import os
import traceback
from collections.abc import Mapping
from typing import Final
from litellm._uuid import uuid
from unittest import mock
from dotenv import load_dotenv
from fastapi import HTTPException, Request
from litellm.constants import PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS
from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker
load_dotenv()
import time
@ -32,6 +38,8 @@ from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.caching.caching import DualCache
from litellm.proxy._types import (
LiteLLM_ProjectTable,
LiteLLM_TeamTable,
NewProjectRequest,
UpdateProjectRequest,
DeleteProjectRequest,
@ -1230,6 +1238,11 @@ 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.MagicMock()
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value={"project_id": "project-update-test", "team_id": None}
)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=0)
monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True)
monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma)
@ -1253,6 +1266,207 @@ def _written_project_data(mock_prisma: mock.MagicMock) -> dict:
return mock_prisma.db.litellm_projecttable.update.await_args.kwargs["data"]
@pytest.mark.asyncio
async def test_update_project_object_permission_validation_precedes_budget_write(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-object-permission-validation"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.budget_id = "budget-project"
budget_table: Final = mock.MagicMock()
budget_table.update = mock.AsyncMock()
mock_prisma.db.litellm_budgettable = budget_table
def jsonify_object_permission_as_string(payload: dict[str, object]) -> dict[str, object]:
return {**payload, "object_permission": '{"vector_stores": ["replacement-store"]}'}
mock_prisma.jsonify_object = jsonify_object_permission_as_string
with pytest.raises(ProxyException) as error:
await _run_project_update(
project_id,
max_budget=50,
object_permission={"vector_stores": ["replacement-store"]},
)
assert error.value.code == "500"
assert "Input should be a valid dictionary" in error.value.message
assert "input_type=str" in error.value.message
budget_table.update.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_project_rejects_move_when_attached_teamless_key_exists(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-teamless-key"
destination_team_id: Final = "team-b"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = "team-a"
mock_prisma.db.litellm_projecttable.find_unique.return_value.budget_id = "budget-project"
budget_table: Final = mock.MagicMock()
budget_table.update = mock.AsyncMock()
mock_prisma.db.litellm_budgettable = budget_table
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()
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-a")
)
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.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,
max_budget=50,
)
expected_detail: Final = {
"error": (
f"Project {project_id} has 1 key(s) that do not belong to team {destination_team_id}. "
"Detach or delete them before moving the project."
)
}
assert error.value.code == "400"
assert expected_detail["error"] in error.value.message
mock_prisma.writer_db.litellm_verificationtoken.count.assert_awaited_once()
mock_prisma.db.litellm_verificationtoken.count.assert_not_awaited()
budget_table.update.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_project_rejects_move_when_writer_team_differs_from_stale_reader(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-replica-lag"
destination_team_id: Final = "team-a"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = destination_team_id
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.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-b")
)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=1)
with pytest.raises(ProxyException) as error:
await _run_project_update(project_id, team_id=destination_team_id)
assert error.value.code == "400"
assert (
f"Project {project_id} has 1 key(s) that do not belong to team {destination_team_id}. "
"Detach or delete them before moving the project."
) in error.value.message
mock_prisma.writer_db.litellm_projecttable.find_unique.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()
@pytest.mark.asyncio
async def test_update_project_slow_writer_ownership_reads_do_not_stall_db_tracker(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-slow-writer"
source_team_id: Final = "team-source"
destination_team_id: Final = "team-destination"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = source_team_id
mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock(
return_value=LiteLLM_TeamTable(team_id=destination_team_id, models=[])
)
async def slow_project_lookup(*, where: Mapping[str, object]) -> LiteLLM_ProjectTable:
assert where == {"project_id": project_id}
await asyncio.sleep(0)
await asyncio.sleep(0)
return LiteLLM_ProjectTable(project_id=project_id, team_id=source_team_id)
async def slow_key_count(*, where: Mapping[str, object]) -> int:
assert where == {
"project_id": project_id,
"OR": [{"team_id": {"not": destination_team_id}}, {"team_id": None}],
}
await asyncio.sleep(0)
await asyncio.sleep(0)
return 0
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(side_effect=slow_project_lookup)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(side_effect=slow_key_count)
monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.0)
db_lookup_stall_tracker.clear()
try:
await _run_project_update(project_id, team_id=destination_team_id)
assert mock_prisma.writer_db.litellm_projecttable.find_unique.await_count == 1
assert mock_prisma.writer_db.litellm_verificationtoken.count.await_count == 1
mock_prisma.db.litellm_projecttable.update.assert_awaited_once()
assert db_lookup_stall_tracker.stalled_within(PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS) is False
finally:
db_lookup_stall_tracker.clear()
@pytest.mark.asyncio
async def test_update_project_team_limit_error_precedes_mismatched_key_guard(
monkeypatch: pytest.MonkeyPatch,
) -> None:
project_id: Final = "project-team-limit-precedence"
source_team_id: Final = "team-source"
destination_team_id: Final = "team-destination"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = source_team_id
mock_prisma.db.litellm_teamtable.find_unique = mock.AsyncMock(
return_value=LiteLLM_TeamTable(
team_id=destination_team_id,
models=["allowed-model"],
)
)
mock_prisma.writer_db.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id=source_team_id)
)
mock_prisma.writer_db.litellm_verificationtoken.count = mock.AsyncMock(return_value=1)
with pytest.raises(ProxyException, match="not in team's allowed models") as error:
await _run_project_update(project_id, team_id=destination_team_id, models=["disallowed-model"])
assert error.value.code == "400"
mock_prisma.writer_db.litellm_verificationtoken.count.assert_not_awaited()
mock_prisma.db.litellm_projecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_project_allows_move_when_no_attached_keys_exist(monkeypatch: pytest.MonkeyPatch) -> None:
project_id: Final = "project-without-keys"
destination_team_id: Final = "team-b"
mock_prisma: Final = _project_update_mocks(monkeypatch, {})
mock_prisma.db.litellm_projecttable.find_unique.return_value.team_id = "team-a"
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.litellm_projecttable.find_unique = mock.AsyncMock(
return_value=LiteLLM_ProjectTable(project_id=project_id, team_id="team-a")
)
await _run_project_update(project_id, team_id=destination_team_id)
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_awaited_once()
@pytest.mark.asyncio
async def test_update_project_clears_model_itpm_limit_sent_as_an_empty_map(monkeypatch):
"""