mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 5e19a24738 into 635085ac14
This commit is contained in:
commit
e7e092365f
5 changed files with 3182 additions and 223 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue