diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index d134c39c91b..ccdf5057b1c 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -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: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 323b9e434a8..5f49416d186 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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, ) diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 29a14b37ab9..b1765082a62 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -1,11 +1,39 @@ +import json +import os +import re +import signal +import threading +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor from hashlib import sha256 +from pathlib import Path from typing import Final +from uuid import uuid4 +import httpx +import psutil +import psycopg import pytest -from integration._support.client import Gateway, object_value, string_value -from integration._support.database import read_rows +from integration._support.client import ( + JSON_OBJECT, + Gateway, + delete_key_if_present, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from integration._support.database import read_rows, write_rows +from integration._support.process import group_members, owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server from pydantic import JsonValue +from litellm.models.user import LiteLLM_UserTable +from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken + +_OWNED_PROXY_SALT_KEY: Final = "sk-integration-salt" +_OWNERSHIP_MARKER: Final = re.compile(rb"ownership-[0-9a-f]{32}") + def _project_rows(project_id: str) -> list[dict[str, JsonValue]]: return read_rows( @@ -17,6 +45,958 @@ def _project_rows(project_id: str) -> list[dict[str, JsonValue]]: ) +def _key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT token, key_alias, team_id, project_id FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _key_permission_and_budget_ids(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT object_permission_id, budget_id FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _key_state_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(k) AS key_row, to_jsonb(b) AS budget_row FROM "LiteLLM_VerificationToken" AS k ' + 'LEFT JOIN "LiteLLM_BudgetTable" AS b ON b.budget_id = k.budget_id WHERE k.token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _project_state_rows(project_id: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(p) AS project_row, to_jsonb(b) AS budget_row FROM "LiteLLM_ProjectTable" AS p ' + 'LEFT JOIN "LiteLLM_BudgetTable" AS b ON b.budget_id = p.budget_id WHERE p.project_id = %s', + (project_id,), + ) + + +def _object_permission_rows(permission_id: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT * FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s', + (permission_id,), + ) + + +def _object_permission_table_rows() -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(p) AS row FROM "LiteLLM_ObjectPermissionTable" AS p ORDER BY object_permission_id', + (), + ) + + +def _budget_table_rows() -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b ORDER BY budget_id', + (), + ) + + +def _budget_rows(budget_id: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(b) AS row FROM "LiteLLM_BudgetTable" AS b WHERE budget_id = %s', + (budget_id,), + ) + + +def _clear_key_object_permission(key: str, permission_id: str) -> None: + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET object_permission_id = NULL WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + write_rows( + 'DELETE FROM "LiteLLM_ObjectPermissionTable" WHERE object_permission_id = %s', + (permission_id,), + ) + + +def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(d) AS row FROM "LiteLLM_DeletedVerificationToken" AS d WHERE d.token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _deprecated_key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT to_jsonb(d) AS row FROM "LiteLLM_DeprecatedVerificationToken" AS d WHERE d.token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _cli_session_token( + user_id: str, + team_id: str | None, + *, + monkeypatch: pytest.MonkeyPatch, + max_budget: float | None = None, +) -> str: + monkeypatch.setenv("LITELLM_SALT_KEY", _OWNED_PROXY_SALT_KEY) + user: Final = LiteLLM_UserTable( + user_id=user_id, + user_role="internal_user", + teams=[team_id] if team_id is not None else [], + models=[], + max_budget=max_budget, + ) + return ExperimentalUIJWTToken.get_cli_jwt_auth_token( + user_info=user, + team_id=team_id, + team_alias="ownership-team" if team_id is not None else None, + max_budget=max_budget, + ) + + +def _discard_unexpected_key(candidate: Gateway, response: httpx.Response) -> None: + if response.status_code != 200: + return + body: Final = JSON_OBJECT.validate_json(response.content) + candidate.post("/key/delete", {"keys": [string_value(body["key"])]}) + + +def _ownership_chat_reply(identity: str, stream: bool) -> Reply: + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4.1-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ownership ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + return Reply( + content_type="text/event-stream", + chunks=( + f"data: {json.dumps({'id': identity, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4.1-mini', 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': 'ownership'}}]})}\n\n".encode(), + f"data: {json.dumps({'id': identity, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4.1-mini', 'choices': [{'index': 0, 'delta': {'content': ' ok'}, 'finish_reason': 'stop'}], 'usage': {'prompt_tokens': 7, 'completion_tokens': 2, 'total_tokens': 9}})}\n\n".encode(), + b"data: [DONE]\n\n", + ), + ) + + +def _ownership_responses_reply(identity: str, stream: bool) -> Reply: + response: Final = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": f"msg_{identity}", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "ownership ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}, + } + if not stream: + return Reply(body=json.dumps(response).encode()) + events: Final = ( + { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + }, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": f"msg_{identity}", + "output_index": 0, + "content_index": 0, + "delta": "ownership ok", + }, + {"type": "response.completed", "sequence_number": 2, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _ownership_upstream(request: Request) -> Reply: + found: Final = _OWNERSHIP_MARKER.search(request.body) + if found is None: + return Reply(status=400, body=b'{"error":"missing ownership marker"}') + marker: Final = found.group(0).decode() + body: Final = JSON_OBJECT.validate_json(request.body) + stream: Final = body.get("stream") is True + if request.target.endswith("/responses"): + return _ownership_responses_reply(f"resp_{marker}", stream) + return _ownership_chat_reply(f"chatcmpl-{marker}", stream) + + +def _ownership_sse_events(text: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + JSON_OBJECT.validate_json(line[6:]) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + + +def _ownership_request_marker(request: Request) -> str: + found: Final = _OWNERSHIP_MARKER.search(request.body) + assert found is not None, request.body + return found.group(0).decode() + + +def _ownership_request_payload( + path: str, model: str, marker: str, stream: bool +) -> tuple[dict[str, JsonValue], dict[str, str]]: + if path.endswith("/messages"): + return ( + { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": marker}], + "stream": stream, + }, + {"anthropic-version": "2023-06-01"}, + ) + if path.endswith("/responses"): + return {"model": model, "input": marker, "stream": stream}, {} + return {"model": model, "messages": [{"role": "user", "content": marker}], "stream": stream}, {} + + +def _ownership_response_id(response: httpx.Response, path: str) -> str: + if not response.headers.get("content-type", "").startswith("text/event-stream"): + return string_value(JSON_OBJECT.validate_json(response.content)["id"]) + events: Final = _ownership_sse_events(response.text) + identities: Final = ( + tuple( + string_value(object_value(event["message"])["id"]) + for event in events + if event.get("type") == "message_start" + ) + if path.endswith("/messages") + else tuple( + string_value(object_value(event["response"])["id"]) + for event in events + if event.get("type") == "response.completed" + ) + if path.endswith("/responses") + else tuple(string_value(event["id"]) for event in events if "id" in event) + ) + unique_identities: Final = frozenset(identities) + assert len(unique_identities) == 1, response.text + return next(iter(unique_identities)) + + +def _ownership_serving_call( + candidate: Gateway, + key: str, + model: str, + index: int, + marker: str, +) -> tuple[int, str, str]: + stream: Final = index % 2 == 0 + route: Final = index % 3 + paths: Final = ("/v1/chat/completions", "/v1/responses", "/v1/messages") + path: Final = paths[route] + body, headers = _ownership_request_payload(path, model, marker, stream) + response: Final = candidate.request("POST", path, body, key=key, headers=headers) + response.read() + if response.status_code != 200: + return response.status_code, "", response.text + return response.status_code, _ownership_response_id(response, path), response.text + + +def _create_mcp_server(candidate: Gateway, server_id: str, server_name: str, alias: str) -> None: + response: Final = candidate.request( + "POST", + "/v1/mcp/server", + { + "server_id": server_id, + "server_name": server_name, + "alias": alias, + "transport": "sse", + "url": "http://127.0.0.1:9/mcp", + }, + ) + assert response.status_code == 201, response.text + + +def _delete_mcp_server(candidate: Gateway, server_id: str) -> None: + response: Final = candidate.request("DELETE", f"/v1/mcp/server/{server_id}") + assert response.status_code == 202, response.text + + +def test_service_account_generate_rejects_foreign_team_project_without_writing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + alias: Final = f"service-account-{uuid4().hex}" + vector_store: Final = f"vector-store-{uuid4().hex}" + budget_rows_before: Final = _budget_table_rows() + object_permission_rows_before: Final = _object_permission_table_rows() + response: Final = ownership_gateway.request( + "POST", + "/key/service-account/generate", + { + "team_id": team_a, + "project_id": project_b, + "key_alias": alias, + "models": [model], + "soft_budget": 3.5, + "object_permission": {"vector_stores": [vector_store]}, + }, + ) + _discard_unexpected_key(ownership_gateway, response) + assert response.status_code == 400, response.text + assert _budget_table_rows() == budget_rows_before + assert _object_permission_table_rows() == object_permission_rows_before + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND key_alias = %s', + (project_b, alias), + ) + == [] + ) + + +def test_key_generate_unowned_project_accepts_any_team(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + write_rows('UPDATE "LiteLLM_ProjectTable" SET team_id = NULL WHERE project_id = %s', (project,)) + team_key: Final = scenario.key(team_id=team_b, project_id=project, models=[model]) + teamless_key: Final = scenario.key(project_id=project, models=[model]) + chats: Final = tuple( + ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "legacy unowned"}]}, + key=key, + ) + for key in (team_key, teamless_key) + ) + assert all(chat.status_code == 200 for chat in chats), tuple(chat.text for chat in chats) + + +@pytest.mark.parametrize( + ("path", "stream"), + ( + ("/v1/chat/completions", False), + ("/v1/chat/completions", True), + ("/v1/messages", False), + ("/v1/messages", True), + ("/v1/responses", False), + ("/v1/responses", True), + ), +) +def test_legacy_mismatched_key_keeps_serving_and_stays_editable( + ownership_gateway: Gateway, path: str, stream: bool +) -> None: + def respond(request: Request) -> Reply: + return _ownership_upstream(request) + + with wire_server(respond) as upstream: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project, models=[model]) + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET team_id = %s WHERE token = %s', + (team_b, sha256(key.encode()).hexdigest()), + ) + marker: Final = f"ownership-{uuid4().hex}" + body, headers = _ownership_request_payload(path, model, marker, stream) + response: Final = ownership_gateway.request("POST", path, body, key=key, headers=headers) + response.read() + assert response.status_code == 200, response.text + requests: Final = upstream.drain() + assert len(requests) == 1 + expected_target: Final = "/chat/completions" if path.endswith("/chat/completions") else "/responses" + assert requests[0].target.endswith(expected_target), requests[0].target + assert marker.encode() in requests[0].body + assert _ownership_response_id(response, path) != "" + + +def test_stored_team_mismatch_allows_key_edits_and_detach(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project, models=[model]) + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET team_id = %s WHERE token = %s', + (team_b, sha256(key.encode()).hexdigest()), + ) + alias: Final = f"stored-mismatch-{uuid4().hex}" + alias_update: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "key_alias": alias}, + ) + assert alias_update.status_code == 200, alias_update.text + aliased_key: Final = _key_rows(key) + assert aliased_key[0]["key_alias"] == alias + unchanged: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "team_id": team_b}, + ) + assert unchanged.status_code == 200, unchanged.text + assert _key_rows(key) == aliased_key + detached: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "project_id": None}, + ) + assert detached.status_code == 200, detached.text + detached_key: Final = _key_rows(key) + assert detached_key[0]["team_id"] == team_b + assert detached_key[0]["project_id"] is None + + +def test_key_regenerate_cross_team_with_object_permission_writes_nothing(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + key: Final = scenario.key( + team_id=team_a, + project_id=project_a, + models=[model], + object_permission={"vector_stores": ["before"]}, + ) + ids: Final = _key_permission_and_budget_ids(key) + assert len(ids) == 1 + permission_id: Final = string_value(ids[0]["object_permission_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + before: Final = ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + response: Final = ownership_gateway.request( + "POST", + f"/key/{key}/regenerate", + {"project_id": project_b, "object_permission": {"vector_stores": ["after"]}}, + ) + _discard_unexpected_key(ownership_gateway, response) + assert response.status_code == 400, response.text + assert ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == before + + +def test_key_bulk_update_cross_team_with_object_permission_preserves_permission(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key( + team_id=team_a, + project_id=project, + models=[model], + object_permission={"vector_stores": ["before"]}, + ) + ids: Final = _key_permission_and_budget_ids(key) + assert len(ids) == 1 + permission_id: Final = string_value(ids[0]["object_permission_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + before: Final = ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + response: Final = ownership_gateway.request( + "POST", + "/key/bulk_update", + { + "keys": [ + { + "key": key, + "team_id": team_b, + "object_permission": {"vector_stores": ["after"]}, + } + ], + }, + ) + assert response.status_code == 200, response.text + failed_updates: Final = JSON_OBJECT.validate_json(response.content)["failed_updates"] + assert isinstance(failed_updates, list) + assert len(failed_updates) == 1 + failed_update: Final = object_value(failed_updates[0]) + assert f"Project {project} belongs to team {team_a}" in string_value(failed_update["failed_reason"]) + assert ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == before + + +def test_key_update_ambiguous_mcp_permission_error_precedes_project_ownership(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + identifier: Final = f"ambiguous{uuid4().hex}" + first_id: Final = f"mcp{uuid4().hex}" + second_id: Final = f"mcp{uuid4().hex}" + _create_mcp_server(ownership_gateway, first_id, identifier, f"alias{uuid4().hex}") + scenario.cleanups.callback(_delete_mcp_server, ownership_gateway, first_id) + _create_mcp_server(ownership_gateway, second_id, f"name{uuid4().hex}", f"alias{uuid4().hex}") + scenario.cleanups.callback(_delete_mcp_server, ownership_gateway, second_id) + write_rows( + 'UPDATE "LiteLLM_MCPServerTable" SET alias = %s WHERE server_id = %s', + (identifier, second_id), + ) + team_permissions: Final = {"mcp_servers": [first_id, second_id]} + team_a: Final = scenario.team(models=[model], object_permission=team_permissions) + team_b: Final = scenario.team(models=[model], object_permission=team_permissions) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key( + team_id=team_a, + project_id=project, + models=[model], + object_permission={"vector_stores": ["before"]}, + ) + ids: Final = _key_permission_and_budget_ids(key) + assert len(ids) == 1 + permission_id: Final = string_value(ids[0]["object_permission_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + before: Final = ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + response: Final = ownership_gateway.request( + "POST", + "/key/update", + { + "key": key, + "team_id": team_b, + "object_permission": {"mcp_tool_permissions": {identifier: ["tool"]}}, + }, + ) + assert response.status_code == 400, response.text + assert "ambiguous" in response.text.lower(), response.text + assert "project" not in response.text.lower(), response.text + assert ( + _object_permission_rows(permission_id), + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == before + + +def test_team_key_bulk_update_rejects_foreign_team_project(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + key: Final = scenario.key(team_id=team_a, models=[model]) + before: Final = _key_rows(key) + response: Final = ownership_gateway.request( + "POST", + "/team/key/bulk_update", + { + "team_id": team_a, + "key_ids": [sha256(key.encode()).hexdigest()], + "update_fields": {"project_id": project_b, "team_id": team_b}, + }, + ) + assert response.status_code == 422, response.text + assert "project_id" in response.text, response.text + assert "team_id" in response.text, response.text + assert _key_rows(key) == before + + +def test_bulk_update_and_regenerate_new_object_permission_is_served(ownership_gateway: Gateway) -> None: + with wire_server(_ownership_upstream) as upstream: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team: Final = scenario.team(models=[model]) + project: Final = scenario.project(team, models=[model]) + bulk_key: Final = scenario.key(team_id=team, project_id=project, models=[model]) + bulk_response: Final = ownership_gateway.request( + "POST", + "/key/bulk_update", + { + "keys": [ + { + "key": bulk_key, + "object_permission": {"vector_stores": ["bulk"]}, + } + ], + }, + ) + assert bulk_response.status_code == 200, bulk_response.text + bulk_info: Final = ownership_gateway.request("GET", "/key/info", params={"key": bulk_key}) + assert bulk_info.status_code == 200, bulk_info.text + bulk_row: Final = object_value(JSON_OBJECT.validate_json(bulk_info.content)["info"]) + bulk_permission_id: Final = string_value(bulk_row["object_permission_id"]) + assert bulk_permission_id != "" + assert len(_object_permission_rows(bulk_permission_id)) == 1 + scenario.cleanups.callback(_clear_key_object_permission, bulk_key, bulk_permission_id) + bulk_chat: Final = ownership_gateway.chat(model, key=bulk_key, text=f"ownership-{uuid4().hex}") + assert string_value(bulk_chat["id"]) != "" + regenerated: Final = ownership_gateway.request( + "POST", + "/key/generate", + {"team_id": team, "project_id": project, "models": [model]}, + ) + assert regenerated.status_code == 200, regenerated.text + regenerated_key: Final = string_value(JSON_OBJECT.validate_json(regenerated.content)["key"]) + scenario.cleanups.callback(delete_key_if_present, ownership_gateway, regenerated_key) + regeneration: Final = ownership_gateway.request( + "POST", + f"/key/{regenerated_key}/regenerate", + {"object_permission": {"vector_stores": ["regenerated"]}}, + ) + assert regeneration.status_code == 200, regeneration.text + new_key: Final = string_value(JSON_OBJECT.validate_json(regeneration.content)["key"]) + scenario.cleanups.callback(delete_key_if_present, ownership_gateway, new_key) + info: Final = ownership_gateway.request("GET", "/key/info", params={"key": new_key}) + assert info.status_code == 200, info.text + info_row: Final = object_value(JSON_OBJECT.validate_json(info.content)["info"]) + regenerated_permission_id: Final = string_value(info_row["object_permission_id"]) + assert regenerated_permission_id != "" + assert regenerated_permission_id != bulk_permission_id + assert len(_object_permission_rows(regenerated_permission_id)) == 1 + scenario.cleanups.callback(_clear_key_object_permission, new_key, regenerated_permission_id) + chat: Final = ownership_gateway.chat(model, key=new_key, text=f"ownership-{uuid4().hex}") + assert string_value(chat["id"]) != "" + + +def test_ownership_rejections_during_concurrent_traffic_burst(ownership_gateway: Gateway) -> None: + with wire_server(_ownership_upstream) as upstream: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + serving_project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + serving_key: Final = scenario.key(team_id=team_a, project_id=serving_project, models=[model]) + markers: Final = tuple(f"ownership-{uuid4().hex}" for _ in range(30)) + key_before: Final = ( + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) + serving_key_before: Final = ( + _key_rows(serving_key), + _key_permission_and_budget_ids(serving_key), + _deleted_key_rows(serving_key), + _deprecated_key_rows(serving_key), + ) + project_before: Final = _project_state_rows(project_a) + + def serve(index: int) -> tuple[int, str, str]: + return _ownership_serving_call(ownership_gateway, serving_key, model, index, markers[index]) + + def update_key() -> httpx.Response: + return ownership_gateway.request("POST", "/key/update", {"key": key, "team_id": team_b}) + + def bulk_update() -> httpx.Response: + return ownership_gateway.request( + "POST", + "/key/bulk_update", + {"keys": [{"key": key, "team_id": team_b}]}, + ) + + def move_project() -> httpx.Response: + return ownership_gateway.request( + "POST", "/project/update", {"project_id": project_a, "team_id": team_b} + ) + + with ThreadPoolExecutor(max_workers=33) as pool: + serving_futures: Final = tuple(pool.submit(serve, index) for index in range(30)) + ownership_futures: Final = ( + pool.submit(update_key), + pool.submit(bulk_update), + pool.submit(move_project), + ) + serving_results: Final = tuple(future.result(timeout=90) for future in serving_futures) + ownership_results: Final = tuple(future.result(timeout=90) for future in ownership_futures) + assert all(status == 200 for status, _, _ in serving_results), serving_results + assert all(response_id for _, response_id, _ in serving_results), serving_results + assert len({response_id for _, response_id, _ in serving_results}) == 30 + key_update_response: Final = ownership_results[0] + bulk_update_response: Final = ownership_results[1] + project_update_response: Final = ownership_results[2] + assert key_update_response.status_code == 400, key_update_response.text + assert bulk_update_response.status_code == 200, bulk_update_response.text + bulk_result: Final = JSON_OBJECT.validate_json(bulk_update_response.content) + successful_updates: Final = bulk_result["successful_updates"] + failed_updates: Final = bulk_result["failed_updates"] + assert isinstance(successful_updates, list), bulk_update_response.text + assert successful_updates == [], bulk_update_response.text + assert isinstance(failed_updates, list), bulk_update_response.text + assert len(failed_updates) == 1, bulk_update_response.text + failed_update: Final = object_value(failed_updates[0]) + assert f"Project {project_a} belongs to team {team_a}" in string_value(failed_update["failed_reason"]), ( + bulk_update_response.text + ) + assert project_update_response.status_code == 400, project_update_response.text + upstream_requests: Final = upstream.drain() + assert len(upstream_requests) == 30 + received_markers: Final = tuple(_ownership_request_marker(request) for request in upstream_requests) + assert tuple(sorted(received_markers)) == tuple(sorted(markers)) + assert ( + _object_permission_table_rows(), + _key_rows(key), + _key_state_rows(key), + _key_permission_and_budget_ids(key), + _deleted_key_rows(key), + _deprecated_key_rows(key), + ) == key_before + assert ( + _key_rows(serving_key), + _key_permission_and_budget_ids(serving_key), + _deleted_key_rows(serving_key), + _deprecated_key_rows(serving_key), + ) == serving_key_before + assert _project_state_rows(project_a) == project_before + + +def test_ownership_checks_wait_for_row_locks_without_failing_readiness(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + project: Final = scenario.project(team, models=[model], max_budget=7) + key: Final = scenario.key(team_id=team, project_id=project, models=[model]) + alias: Final = f"locked-{uuid4().hex}" + locked: Final = threading.Event() + release: Final = threading.Event() + + def hold_locks() -> None: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute("BEGIN") + connection.execute( + 'SELECT project_id FROM "LiteLLM_ProjectTable" WHERE project_id = %s FOR UPDATE', + (project,), + ) + connection.execute( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE token = %s FOR UPDATE', + (sha256(key.encode()).hexdigest(),), + ) + locked.set() + assert release.wait(timeout=30) + connection.commit() + + with ThreadPoolExecutor(max_workers=3) as pool: + lock_future: Final = pool.submit(hold_locks) + try: + assert locked.wait(timeout=10) + key_future: Final = pool.submit( + ownership_gateway.request, + "POST", + "/key/update", + {"key": key, "key_alias": alias}, + ) + project_future: Final = pool.submit( + ownership_gateway.request, + "POST", + "/project/update", + {"project_id": project, "team_id": team, "max_budget": 19}, + ) + lock_waiters: Final = eventually( + lambda: read_rows( + "SELECT pid FROM pg_stat_activity WHERE datname = current_database() " + "AND wait_event_type = 'Lock' AND cardinality(pg_blocking_pids(pid)) > 0", + (), + ), + lambda rows: len(rows) >= 2, + seconds=10, + ) + assert len(lock_waiters) >= 2 + readiness: Final = eventually( + lambda: ownership_gateway.request("GET", "/health/readiness").status_code, + lambda status: status == 200, + seconds=10, + ) + assert readiness == 200 + finally: + release.set() + key_response: Final = key_future.result(timeout=90) + project_response: Final = project_future.result(timeout=90) + lock_future.result(timeout=30) + assert key_response.status_code == 200, key_response.text + assert project_response.status_code == 200, project_response.text + assert _key_rows(key)[0]["key_alias"] == alias + assert _project_rows(project)[0]["max_budget"] == 19.0 + + +def test_ownership_enforced_after_worker_kill(ownership_gateway: Gateway, tmp_path: Path) -> None: + started: Final = threading.Event() + release_stream: Final = threading.Event() + + def respond(request: Request) -> Reply: + started.set() + reply: Final = _ownership_upstream(request) + return Reply( + status=reply.status, + body=reply.body, + content_type=reply.content_type, + chunks=reply.chunks, + gate_after_first=release_stream, + ) + + with wire_server(respond) as upstream: + with owned_proxy_process( + ownership_gateway, + tmp_path / "worker-kill", + {"LITELLM_SALT_KEY": _OWNED_PROXY_SALT_KEY}, + workers=2, + ) as owned: + with owned.gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{upstream.url}/v1") + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + markers: Final = tuple(f"ownership-{uuid4().hex}" for _ in range(24)) + + def serve(index: int) -> tuple[int, int, str]: + try: + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": markers[index]}], + "stream": True, + }, + key=key, + ) + except httpx.TransportError as error: + return index, 0, str(error) + return index, response.status_code, response.text + + candidate_port: Final = owned.gateway.client.base_url.port + assert candidate_port is not None + workers: Final = tuple( + process + for process in group_members(owned.process.pid) + if process.pid != owned.process.pid + and any( + connection.laddr.port == candidate_port and connection.status == psutil.CONN_LISTEN + for connection in process.net_connections(kind="inet") + ) + ) + assert len(workers) == 2, tuple(process.pid for process in workers) + victim: Final = workers[0] + survivor: Final = workers[1] + with ThreadPoolExecutor(max_workers=24) as pool: + futures: Final = tuple(pool.submit(serve, index) for index in range(24)) + try: + assert started.wait(timeout=10) + os.kill(victim.pid, signal.SIGKILL) + release_stream.set() + psutil.wait_procs((victim,), timeout=10) + assert not psutil.pid_exists(victim.pid), victim.pid + finally: + release_stream.set() + burst: Final = tuple(future.result(timeout=90) for future in futures) + assert psutil.pid_exists(survivor.pid), survivor.pid + assert len(burst) == 24 + burst_requests: Final = upstream.drain() + empty_body_count: Final = sum(not request.body for request in burst_requests) + burst_markers: Final = tuple( + _ownership_request_marker(request) for request in burst_requests if request.body + ) + successful_markers: Final = tuple(markers[index] for index, status, _ in burst if status == 200) + assert burst_markers, f"empty_body_captures={empty_body_count}; burst={burst}" + assert all(burst_markers.count(marker) == 1 for marker in successful_markers), ( + f"empty_body_captures={empty_body_count}; " + f"successful_markers={successful_markers}; upstream_markers={burst_markers}; burst={burst}" + ) + cross_team: Final = owned.gateway.request( + "POST", + "/key/generate", + {"team_id": team_a, "project_id": project_b, "models": [model]}, + ) + _discard_unexpected_key(owned.gateway, cross_team) + assert cross_team.status_code == 400, cross_team.text + marker: Final = f"ownership-{uuid4().hex}" + chat: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + post_kill_requests: Final = upstream.drain() + post_kill_empty_body_count: Final = sum(not request.body for request in post_kill_requests) + post_kill_markers: Final = tuple( + _ownership_request_marker(request) for request in post_kill_requests if request.body + ) + assert post_kill_markers.count(marker) == 1, ( + f"empty_body_captures={post_kill_empty_body_count}; " + f"post_kill_markers={post_kill_markers}; response={chat.text}" + ) + + +@pytest.fixture(scope="module") +def ownership_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + with gateway_from_environment() as gateway: + with owned_proxy( + gateway, + tmp_path_factory.mktemp("project-team-ownership"), + {"LITELLM_SALT_KEY": _OWNED_PROXY_SALT_KEY}, + workers=2, + ) as candidate: + yield candidate + + @pytest.mark.covers("mgmt.project.new.real_route_persists") def test_project_new_persists_real_state(gateway: Gateway) -> None: with gateway.scenario() as scenario: @@ -113,3 +1093,506 @@ def test_project_delete_with_attached_key_refuses_and_preserves_state(gateway: G 'FROM "LiteLLM_VerificationToken" WHERE token = %s', (digest,), ) == key_before + + +def test_key_generate_rejects_foreign_team_project(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + alias: Final = f"cross-team-{uuid4().hex}" + vector_store: Final = f"vector-store-{uuid4().hex}" + budget_rows_before: Final = _budget_table_rows() + object_permission_rows_before: Final = _object_permission_table_rows() + cross_team: Final = ownership_gateway.request( + "POST", + "/key/generate", + { + "team_id": team_a, + "project_id": project_b, + "key_alias": alias, + "models": [model], + "soft_budget": 3.5, + "object_permission": {"vector_stores": [vector_store]}, + }, + ) + _discard_unexpected_key(ownership_gateway, cross_team) + assert cross_team.status_code == 400, cross_team.text + assert _budget_table_rows() == budget_rows_before + assert _object_permission_table_rows() == object_permission_rows_before + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id = %s', + (project_b, team_a), + ) + == [] + ) + + +def test_key_generate_rejects_missing_team_for_owned_project(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + unbound: Final = ownership_gateway.request( + "POST", "/key/generate", {"project_id": project_b, "models": [model]} + ) + _discard_unexpected_key(ownership_gateway, unbound) + assert unbound.status_code == 400, unbound.text + assert read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id IS NULL', + (project_b,), + ) == [] + + +def test_key_generate_same_team_project_key_can_chat(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + project: Final = scenario.project(team, models=[model]) + generated: Final = ownership_gateway.request( + "POST", "/key/generate", {"team_id": team, "project_id": project, "models": [model]} + ) + assert generated.status_code == 200, generated.text + key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) + scenario.cleanups.callback(scenario.delete_key, key) + generated_rows: Final = _key_rows(key) + assert len(generated_rows) == 1 + assert generated_rows[0]["team_id"] == team + assert generated_rows[0]["project_id"] == project + chat: Final = ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "control"}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + assert object_value(JSON_OBJECT.validate_json(chat.content)["usage"])["total_tokens"] == 40 + + +def test_key_update_rejects_team_change_and_allows_unchanged_values(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model], soft_budget=3.0) + permission_seed: Final = ownership_gateway.request( + "POST", + "/key/update", + {"key": key, "object_permission": {"vector_stores": ["existing-store"]}}, + ) + assert permission_seed.status_code == 200, permission_seed.text + key_permission_state: Final = _key_permission_and_budget_ids(key) + assert len(key_permission_state) == 1 + permission_id: Final = string_value(key_permission_state[0]["object_permission_id"]) + budget_id: Final = string_value(key_permission_state[0]["budget_id"]) + scenario.cleanups.callback(_clear_key_object_permission, key, permission_id) + permission_before: Final = _object_permission_rows(permission_id) + budget_before: Final = _budget_rows(budget_id) + assert len(permission_before) == 1 + assert len(budget_before) == 1 + aliased: Final = ownership_gateway.request("POST", "/key/update", {"key": key, "key_alias": "updated"}) + assert aliased.status_code == 200, aliased.text + after_alias: Final = _key_rows(key) + assert len(after_alias) == 1 + assert after_alias[0]["key_alias"] == "updated" + assert after_alias[0]["team_id"] == team_a + assert after_alias[0]["project_id"] == project_a + unchanged_team: Final = ownership_gateway.request("POST", "/key/update", {"key": key, "team_id": team_a}) + assert unchanged_team.status_code == 200, unchanged_team.text + after_unchanged_team: Final = _key_rows(key) + assert after_unchanged_team == after_alias + before_reassignment: Final = _key_rows(key) + assert len(before_reassignment) == 1 + reassigned: Final = ownership_gateway.request( + "POST", + "/key/update", + { + "key": key, + "team_id": team_b, + "soft_budget": 5.0, + "object_permission": {"vector_stores": ["replacement-store"]}, + }, + ) + assert reassigned.status_code == 400, reassigned.text + assert _key_rows(key) == before_reassignment + assert _object_permission_rows(permission_id) == permission_before + assert _budget_rows(budget_id) == budget_before + + +def test_key_update_can_detach_project_and_change_team(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + detached: Final = ownership_gateway.request( + "POST", "/key/update", {"key": key, "project_id": None, "team_id": team_b} + ) + assert detached.status_code == 200, detached.text + after_detach: Final = _key_rows(key) + assert len(after_detach) == 1 + assert after_detach[0]["project_id"] is None + assert after_detach[0]["team_id"] == team_b + assert after_detach[0]["key_alias"] is None + + +def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_b: Final = scenario.project(team_b, models=[model]) + generated: Final = ownership_gateway.request( + "POST", "/key/generate", {"team_id": team_a, "models": [model]} + ) + assert generated.status_code == 200, generated.text + key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) + scenario.cleanups.callback(scenario.delete_key, key) + before: Final = _key_rows(key) + before_deleted: Final = _deleted_key_rows(key) + before_deprecated: Final = _deprecated_key_rows(key) + assert len(before) == 1 + assert before_deleted == [] + assert before_deprecated == [] + response: Final = ownership_gateway.request( + "POST", f"/key/{key}/regenerate", {"project_id": project_b, "grace_period": "1h"} + ) + _discard_unexpected_key(ownership_gateway, response) + assert response.status_code == 400, response.text + assert _key_rows(key) == before + assert _deleted_key_rows(key) == before_deleted + assert _deprecated_key_rows(key) == before_deprecated + + +def test_key_generation_rejects_missing_project_without_writing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + missing_project_id: Final = f"missing-{uuid4()}" + key_generation: Final = ownership_gateway.request( + "POST", + "/key/generate", + {"team_id": team, "project_id": missing_project_id, "models": [model]}, + ) + + assert key_generation.status_code == 404, key_generation.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (missing_project_id,), + ) + == [] + ) + + service_account_generation: Final = ownership_gateway.request( + "POST", + "/key/service-account/generate", + {"team_id": team, "project_id": missing_project_id, "models": [model]}, + ) + + assert service_account_generation.status_code == 500, service_account_generation.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (missing_project_id,), + ) + == [] + ) + + +def test_key_regenerate_routes_reject_missing_project_without_changing_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + key: Final = scenario.key(team_id=team, models=[model]) + missing_project_id: Final = f"missing-{uuid4()}" + before: Final = _key_rows(key) + responses: Final = ( + ownership_gateway.request( + "POST", + "/key/regenerate", + {"key": key, "project_id": missing_project_id}, + ), + ownership_gateway.request( + "POST", + f"/key/{key}/regenerate", + {"project_id": missing_project_id}, + ), + ) + + for response in responses: + assert response.status_code == 500, response.text + assert _key_rows(key) == before + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (missing_project_id,), + ) + == [] + ) + + chat: Final = ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "rejected regeneration preserves key"}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + + +def test_key_generate_nonmember_organization_error_precedes_project_ownership( + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + organization_id: Final = scenario.organization() + owner_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(owner_team, models=[model]) + caller_id: Final = scenario.user(user_role="internal_user") + caller_token: Final = _cli_session_token(caller_id, None, monkeypatch=monkeypatch) + response: Final = ownership_gateway.request( + "POST", + "/key/generate", + { + "project_id": project, + "organization_id": organization_id, + "models": [model], + }, + key=caller_token, + ) + + assert response.status_code == 403, response.text + assert f"Caller is not a member of organization_id={organization_id}" in response.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (project,), + ) + == [] + ) + + +def test_key_generate_duplicate_alias_error_precedes_project_ownership( + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + caller_team: Final = scenario.team(models=[model]) + owner_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(owner_team, models=[model]) + key_alias: Final = f"duplicate-{uuid4()}" + existing_key: Final = scenario.key(team_id=caller_team, key_alias=key_alias, models=[model]) + caller_id: Final = scenario.member(caller_team, role="admin") + caller_token: Final = _cli_session_token(caller_id, caller_team, monkeypatch=monkeypatch) + existing_alias_rows: Final = read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', + (key_alias,), + ) + response: Final = ownership_gateway.request( + "POST", + "/key/generate", + { + "team_id": caller_team, + "project_id": project, + "key_alias": key_alias, + "models": [model], + }, + key=caller_token, + ) + + assert response.status_code == 400, response.text + assert f"Key with alias '{key_alias}' already exists" in response.text + assert len(existing_alias_rows) == 1 + assert read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', + (key_alias,), + ) == existing_alias_rows + assert len(_key_rows(existing_key)) == 1 + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (project,), + ) + == [] + ) + + +def test_key_generate_budget_ceiling_precedes_project_ownership( + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + caller_team: Final = scenario.team(models=[model]) + owner_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(owner_team, models=[model]) + caller_id: Final = scenario.member(caller_team, role="admin") + caller_token: Final = _cli_session_token( + caller_id, caller_team, monkeypatch=monkeypatch, max_budget=1 + ) + response: Final = ownership_gateway.request( + "POST", + "/key/generate", + { + "team_id": caller_team, + "project_id": project, + "max_budget": 5, + "models": [model], + }, + key=caller_token, + ) + + assert response.status_code == 400, response.text + assert "max_budget (5.0) cannot exceed the caller's own max_budget (1.0)" in response.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (project,), + ) + == [] + ) + + +def test_key_regenerate_output_estimate_admin_error_precedes_project_ownership( + ownership_gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + project_team: Final = scenario.team(models=[model]) + destination_team: Final = scenario.team(models=[model]) + project: Final = scenario.project(project_team, models=[model]) + caller_id: Final = scenario.member(project_team, role="admin") + key: Final = scenario.key(team_id=project_team, user_id=caller_id, models=[model]) + caller_token: Final = _cli_session_token(caller_id, project_team, monkeypatch=monkeypatch) + before: Final = _key_rows(key) + response: Final = ownership_gateway.request( + "POST", + f"/key/{key}/regenerate", + { + "team_id": destination_team, + "project_id": project, + "metadata": {"default_estimated_output_tokens": 1}, + }, + key=caller_token, + ) + + assert response.status_code == 403, response.text + assert "Only proxy admins can set" in response.text + assert _key_rows(key) == before + chat: Final = ownership_gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "rejected regeneration keeps key valid"}]}, + key=key, + ) + assert chat.status_code == 200, chat.text + + +def test_key_bulk_update_rejects_foreign_team_project_and_preserves_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + before: Final = _key_rows(key) + assert len(before) == 1 + assert before[0]["team_id"] == team_a + assert before[0]["project_id"] == project_a + + response: Final = ownership_gateway.request( + "POST", + "/key/bulk_update", + { + "keys": [ + {"key": key, "team_id": team_b}, + {"key": key, "max_budget": 10, "tags": ["bulk-update"]}, + ] + }, + ) + assert response.status_code == 200, response.text + body: Final = JSON_OBJECT.validate_json(response.content) + failed_updates: Final = body["failed_updates"] + successful_updates: Final = body["successful_updates"] + assert isinstance(failed_updates, list), response.text + assert isinstance(successful_updates, list), response.text + assert len(failed_updates) == 1 + assert len(successful_updates) == 1 + failed_update: Final = object_value(failed_updates[0]) + successful_update: Final = object_value(successful_updates[0]) + assert string_value(failed_update["key"]) == key + assert f"Project {project_a} belongs to team {team_a}" in string_value(failed_update["failed_reason"]) + assert string_value(successful_update["key"]) == key + + after: Final = _key_rows(key) + assert len(after) == 1 + assert after[0]["team_id"] == before[0]["team_id"] + assert after[0]["project_id"] == before[0]["project_id"] + + +def test_project_update_rejects_moving_project_with_attached_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + budget: Final = scenario.budget(max_budget=3) + attached_project: Final = scenario.project(team_a, budget_id=budget, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=attached_project, models=[model]) + project_before: Final = _project_rows(attached_project) + key_before: Final = _key_rows(key) + budget_before: Final = _budget_rows(budget) + moved_with_key: Final = ownership_gateway.request( + "POST", + "/project/update", + { + "project_id": attached_project, + "team_id": team_b, + "max_budget": 11, + }, + ) + assert moved_with_key.status_code == 400, moved_with_key.text + assert _project_rows(attached_project) == project_before + assert _key_rows(key) == key_before + assert _budget_rows(budget) == budget_before + + +def test_project_update_allows_moving_project_without_keys(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + unattached_project: Final = scenario.project(team_a, models=[model]) + moved_without_key: Final = ownership_gateway.request( + "POST", "/project/update", {"project_id": unattached_project, "team_id": team_b} + ) + assert moved_without_key.status_code == 200, moved_without_key.text + moved_project: Final = _project_rows(unattached_project) + assert len(moved_project) == 1 + assert moved_project[0]["team_id"] == team_b + assert read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s', + (unattached_project,), + ) == [] + + +def test_project_update_rejects_moving_project_with_teamless_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project, models=[model]) + write_rows( + 'UPDATE "LiteLLM_VerificationToken" SET team_id = NULL WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + project_before: Final = _project_rows(project) + key_before: Final = _key_rows(key) + assert key_before[0]["team_id"] is None + moved: Final = ownership_gateway.request("POST", "/project/update", {"project_id": project, "team_id": team_b}) + assert moved.status_code == 400, moved.text + assert _project_rows(project) == project_before + assert _key_rows(key) == key_before 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 226755e7b8e..12c61fe3853 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 @@ -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): """ 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 15bf4f31445..a34372fe091 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -1,9 +1,11 @@ +import asyncio from collections.abc import Mapping from contextlib import ExitStack from typing import Final from types import SimpleNamespace import json from datetime import datetime, timedelta, timezone +from uuid import UUID import litellm import pytest @@ -16,6 +18,7 @@ from fastapi import HTTPException import inspect +from litellm.models.project import LiteLLM_ProjectTable from litellm.proxy._types import ( GenerateKeyRequest, KeyManagementRoutes, @@ -43,12 +46,17 @@ 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.constants import PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS +from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy.management_endpoints.key_management_endpoints import ( + _check_key_project_team, + _check_key_project_team_on_mutation, _check_org_key_limits, _check_project_key_limits, _check_team_key_limits, _common_key_generation_helper, + KeyProjectTeamMismatchError, _effective_key_after_update, _effective_key_for_generate, _enforce_custom_key_policy, @@ -71,11 +79,13 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( delete_verification_tokens, generate_key_fn, generate_key_helper_fn, + generate_service_account_key_fn, key_aliases, key_generation_check, list_keys, prepare_key_update_data, reset_key_spend_fn, + update_key_fn, validate_key_list_check, validate_key_team_change, ) @@ -1068,208 +1078,146 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch): @pytest.mark.asyncio -async def test_key_update_object_permissions_existing_permission(): - """ - Test updating object permissions when a key already has an existing object_permission_id. - - This test verifies that when updating vector stores for a key that already has an - object_permission_id, the existing LiteLLM_ObjectPermissionTable record is updated - with the new permissions and the object_permission_id remains the same. - """ +async def test_key_update_prepares_existing_object_permission_without_writing(): from unittest.mock import AsyncMock, MagicMock - import pytest - - from litellm.proxy._types import ( - LiteLLM_ObjectPermissionBase, - LiteLLM_VerificationToken, - ) + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_VerificationToken from litellm.proxy.management_endpoints.key_management_endpoints import ( - _handle_update_object_permission, + _prepare_key_update_object_permission, + _write_prepared_key_update_object_permission, ) mock_prisma_client = AsyncMock() - - # Mock existing key with object_permission_id existing_key_row = LiteLLM_VerificationToken( token="test_token_hash", object_permission_id="existing_perm_id_123", user_id="user123", team_id=None, ) - - # Mock existing object permission record - existing_object_permission = MagicMock() - existing_object_permission.model_dump.return_value = { + existing_permission = MagicMock() + existing_permission.model_dump.return_value = { "object_permission_id": "existing_perm_id_123", - "vector_stores": ["old_store_1", "old_store_2"], + "vector_stores": ["old_store"], } + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_permission) + permission_upsert = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=existing_object_permission + object_permission_data = LiteLLM_ObjectPermissionBase(vector_stores=["new_store"]).model_dump( + exclude_unset=True, exclude_none=True ) - - # Mock upsert operation - updated_permission = MagicMock() - updated_permission.object_permission_id = "existing_perm_id_123" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=updated_permission - ) - - # Test data with new object permission - data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["new_store_1", "new_store_2", "new_store_3"] - ).model_dump(exclude_unset=True, exclude_none=True), - "user_id": "user123", - } - - # Call the function - result = await _handle_update_object_permission( - data_json=data_json, + upsert = await _prepare_key_update_object_permission( + object_permission_data=object_permission_data, existing_key_row=existing_key_row, prisma_client=mock_prisma_client, ) - # Verify the object_permission was removed from data_json and object_permission_id was set - assert "object_permission" not in result - assert result["object_permission_id"] == "existing_perm_id_123" - - # Verify database operations were called correctly - mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + assert upsert is not None + assert upsert.object_permission_id == "existing_perm_id_123" + assert upsert.record["vector_stores"] == ["new_store"] + permission_upsert.assert_not_awaited() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_awaited_once_with( where={"object_permission_id": "existing_perm_id_123"} ) - mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + result = await _write_prepared_key_update_object_permission( + data_json={"user_id": "user123"}, + upsert=upsert, + prisma_client=mock_prisma_client, + ) + + assert result == {"user_id": "user123", "object_permission_id": "existing_perm_id_123"} + permission_upsert.assert_awaited_once_with( + where={"object_permission_id": "existing_perm_id_123"}, + data={"create": upsert.record, "update": upsert.record}, + ) @pytest.mark.asyncio -async def test_key_update_object_permissions_no_existing_permission(): - """ - Test creating object permissions when a key has no existing object_permission_id. +async def test_key_update_prepares_json_object_permission_and_upserts_new_row(): + import json + from unittest.mock import AsyncMock - This test verifies that when updating object permissions for a key that has - object_permission_id set to None, a new entry is created in the - LiteLLM_ObjectPermissionTable and the key is updated with the new object_permission_id. - """ - from unittest.mock import AsyncMock, MagicMock - - import pytest - - from litellm.proxy._types import ( - LiteLLM_ObjectPermissionBase, - LiteLLM_VerificationToken, - ) + from litellm.proxy._types import LiteLLM_VerificationToken from litellm.proxy.management_endpoints.key_management_endpoints import ( - _handle_update_object_permission, + _prepare_key_update_object_permission, + _write_prepared_key_update_object_permission, ) mock_prisma_client = AsyncMock() - - existing_key_row_no_perm = LiteLLM_VerificationToken( + existing_key_row = LiteLLM_VerificationToken( token="test_token_hash_2", object_permission_id=None, user_id="user456", team_id=None, ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) + permission_upsert = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert - # Mock find_unique to return None (no existing permission) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) - - # Mock upsert to create new record - new_permission = MagicMock() - new_permission.object_permission_id = "new_perm_id_456" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) - - data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["brand_new_store"] - ).model_dump(exclude_unset=True, exclude_none=True), - "user_id": "user456", - } - - result = await _handle_update_object_permission( - data_json=data_json, - existing_key_row=existing_key_row_no_perm, + upsert = await _prepare_key_update_object_permission( + object_permission_data=json.dumps({"vector_stores": ["brand_new_store"]}), + existing_key_row=existing_key_row, prisma_client=mock_prisma_client, ) - # Verify new object_permission_id was set - assert "object_permission" not in result - assert result["object_permission_id"] == "new_perm_id_456" - # Verify upsert was called to create new record - mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + assert upsert is not None + assert upsert.record["vector_stores"] == ["brand_new_store"] + permission_upsert.assert_not_awaited() + result = await _write_prepared_key_update_object_permission( + data_json={}, + upsert=upsert, + prisma_client=mock_prisma_client, + ) + + assert result["object_permission_id"] == upsert.object_permission_id + permission_upsert.assert_awaited_once_with( + where={"object_permission_id": upsert.object_permission_id}, + data={"create": upsert.record, "update": upsert.record}, + ) @pytest.mark.asyncio -async def test_key_update_object_permissions_missing_permission_record(): - """ - Test creating object permissions when existing object_permission_id record is not found. +async def test_key_update_recreates_missing_object_permission_with_existing_id(): + from unittest.mock import AsyncMock - This test verifies that when updating object permissions for a key that has an - object_permission_id but the corresponding record cannot be found in the database, - a new entry is created in the LiteLLM_ObjectPermissionTable with the new permissions. - """ - from unittest.mock import AsyncMock, MagicMock - - import pytest - - from litellm.proxy._types import ( - LiteLLM_ObjectPermissionBase, - LiteLLM_VerificationToken, - ) + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_VerificationToken from litellm.proxy.management_endpoints.key_management_endpoints import ( - _handle_update_object_permission, + _prepare_key_update_object_permission, + _write_prepared_key_update_object_permission, ) mock_prisma_client = AsyncMock() - - existing_key_row_missing_perm = LiteLLM_VerificationToken( + existing_key_row = LiteLLM_VerificationToken( token="test_token_hash_3", object_permission_id="missing_perm_id_789", user_id="user789", team_id=None, ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) + permission_upsert = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert - # Mock find_unique to return None (permission record not found) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) - - # Mock upsert to create new record - new_permission = MagicMock() - new_permission.object_permission_id = "recreated_perm_id_789" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) - - data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["recreated_store"] - ).model_dump(exclude_unset=True, exclude_none=True), - "user_id": "user789", - } - - result = await _handle_update_object_permission( - data_json=data_json, - existing_key_row=existing_key_row_missing_perm, + upsert = await _prepare_key_update_object_permission( + object_permission_data=LiteLLM_ObjectPermissionBase(vector_stores=["recreated_store"]).model_dump( + exclude_unset=True, exclude_none=True + ), + existing_key_row=existing_key_row, prisma_client=mock_prisma_client, ) - # Verify new object_permission_id was set - assert "object_permission" not in result - assert result["object_permission_id"] == "recreated_perm_id_789" - - # Verify find_unique was called with the missing permission ID - mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( - where={"object_permission_id": "missing_perm_id_789"} + assert upsert is not None + assert upsert.object_permission_id == "missing_perm_id_789" + permission_upsert.assert_not_awaited() + await _write_prepared_key_update_object_permission( + data_json={}, + upsert=upsert, + prisma_client=mock_prisma_client, + ) + permission_upsert.assert_awaited_once_with( + where={"object_permission_id": "missing_perm_id_789"}, + data={"create": upsert.record, "update": upsert.record}, ) - - # Verify upsert was called to create new record - mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() @pytest.mark.asyncio @@ -1876,6 +1824,117 @@ async def test_generate_service_account_works_with_team_id(): ) +def _identity_json_object(value: dict[str, object]) -> dict[str, object]: + return value + + +def _key_generation_prisma_client( + project_record: dict[str, str] | None, +) -> tuple[MagicMock, MagicMock, MagicMock]: + prisma_client: Final = MagicMock() + prisma_client.jsonify_object.side_effect = _identity_json_object + + budget_table: Final = MagicMock() + budget_table.create = AsyncMock(return_value=MagicMock(budget_id="created-budget-id")) + budget_table.delete = AsyncMock() + object_permission_table: Final = MagicMock() + object_permission_table.create = AsyncMock( + return_value=MagicMock(object_permission_id="created-object-permission-id") + ) + object_permission_table.delete = AsyncMock() + + prisma_client.db.litellm_budgettable = budget_table + prisma_client.db.litellm_objectpermissiontable = object_permission_table + prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=project_record) + prisma_client.insert_data = AsyncMock() + return prisma_client, budget_table, object_permission_table + + +@pytest.mark.asyncio +async def test_rejected_key_generation_deletes_created_budget_and_default_permission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prisma_client, budget_table, object_permission_table = _key_generation_prisma_client( + {"project_id": "project-b", "team_id": "team-b"} + ) + monkeypatch.setattr(litellm, "key_generation_settings", None) + monkeypatch.setattr( + litellm, + "default_key_generate_params", + {"object_permission": {"vector_stores": ["default-vector-store"]}}, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + + with pytest.raises(KeyProjectTeamMismatchError) as exc: + await _common_key_generation_helper( + data=GenerateKeyRequest( + budget_id="caller-budget-id", + project_id="project-b", + soft_budget=3.5, + team_id="team-a", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ), + litellm_changed_by=None, + team_table=None, + ) + + assert exc.value.status_code == 400 + budget_table.create.assert_awaited_once() + budget_table.delete.assert_awaited_once_with(where={"budget_id": "created-budget-id"}) + object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]}) + object_permission_table.delete.assert_awaited_once_with( + where={"object_permission_id": "created-object-permission-id"} + ) + prisma_client.insert_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_generation_does_not_delete_rows_for_other_http_exceptions( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prisma_client, budget_table, object_permission_table = _key_generation_prisma_client(None) + monkeypatch.setattr(litellm, "key_generation_settings", None) + monkeypatch.setattr( + litellm, + "default_key_generate_params", + {"object_permission": {"vector_stores": ["default-vector-store"]}}, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + + with pytest.raises(HTTPException) as exc: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id="project-b", + router_settings={"weights": {"gpt-4": {"unknown-deployment": 1.0}}}, + soft_budget=3.5, + team_id="team-a", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ), + litellm_changed_by=None, + team_table=None, + ) + + assert type(exc.value) is HTTPException + assert exc.value.status_code == 400 + budget_table.create.assert_awaited_once() + budget_table.delete.assert_not_awaited() + object_permission_table.create.assert_awaited_once_with(data={"vector_stores": ["default-vector-store"]}) + object_permission_table.delete.assert_not_awaited() + + @pytest.mark.asyncio async def test_generate_key_throttle_rejected_for_non_admin(): """Security regression: a non-admin creating a key must not be able to set @@ -2957,7 +3016,10 @@ async def test_update_key_nonexistent_key_returns_404(monkeypatch): def _setup_update_key_mocks(monkeypatch, mock_prisma_client): monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) - monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + proxy_logging_obj: Final = MagicMock() + proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) monkeypatch.setattr("litellm.store_audit_logs", False) @@ -7197,7 +7259,10 @@ _BULK_UPDATE_TEAM: Final = LiteLLM_TeamTableCachedObj(team_id="team-1") async def _run_bulk_update_on_one_key( - monkeypatch, item_payload: Mapping[str, object], team: LiteLLM_TeamTableCachedObj = _BULK_UPDATE_TEAM + monkeypatch: pytest.MonkeyPatch, + item_payload: Mapping[str, object], + team: LiteLLM_TeamTableCachedObj = _BULK_UPDATE_TEAM, + user_api_key_cache: UserApiKeyCache | None = None, ) -> tuple[BulkUpdateKeyResponse, AsyncMock]: from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys @@ -7213,6 +7278,8 @@ async def _run_bulk_update_on_one_key( ) mock_prisma_client.update_data = AsyncMock(return_value={"data": {"token": _BULK_UPDATE_TOKEN}}) _setup_update_key_mocks(monkeypatch, mock_prisma_client) + if user_api_key_cache is not None: + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) monkeypatch.setattr( "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", AsyncMock(return_value=team) ) @@ -7280,10 +7347,51 @@ async def test_bulk_update_keys_object_permission_is_granted_not_dropped(monkeyp upserted = prisma.db.litellm_objectpermissiontable.upsert.call_args.kwargs["data"]["create"] assert upserted["vector_stores"] == ["vs-1"] written = _written_key_row(prisma) - assert written["object_permission_id"] == "objperm-bulk" + assert written["object_permission_id"] == upserted["object_permission_id"] assert not {"max_budget", "team_id", "budget_id"} & written.keys() +@pytest.mark.asyncio +async def test_bulk_update_keys_invalidates_cache_for_new_object_permission(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy.common_utils.user_api_key_cache import object_permission_cache_key + + permission_id = str(UUID("00000000-0000-0000-0000-000000000001")) + permission_cache_key = object_permission_cache_key(permission_id) + user_api_key_cache = UserApiKeyCache() + user_api_key_cache.set_cache( + key=permission_cache_key, + value=LiteLLM_ObjectPermissionTable(object_permission_id=permission_id, vector_stores=["stale"]), + model_type=LiteLLM_ObjectPermissionTable, + ) + assert ( + user_api_key_cache.get_cache( + key=permission_cache_key, + model_type=LiteLLM_ObjectPermissionTable, + ) + is not None + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.object_permission_utils.uuid.uuid4", + lambda: UUID(permission_id), + ) + + response, prisma = await _run_bulk_update_on_one_key( + monkeypatch, + {"object_permission": {"vector_stores": ["vs-1"]}}, + user_api_key_cache=user_api_key_cache, + ) + + assert response.failed_updates == [] + assert _written_key_row(prisma)["object_permission_id"] == permission_id + assert ( + user_api_key_cache.get_cache( + key=permission_cache_key, + model_type=LiteLLM_ObjectPermissionTable, + ) + is None + ) + + @pytest.mark.asyncio async def test_bulk_update_keys_object_permission_outside_the_team_allowlist_is_refused(monkeypatch): """A bulk item's object_permission is checked against the key's team exactly as /key/update @@ -13159,6 +13267,7 @@ async def test_execute_virtual_key_regeneration_hides_the_untouched_modal_expiry _POLICY_DENIAL_MESSAGE = "key duration must be 7d or less" _POLICY_HASHED_TOKEN = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" _POLICY_GENERATED_KEY = {"key": "sk-test-key", "expires": None, "user_id": "test-user", "team_id": None} +_OBJECT_PERMISSION_ID_AFTER_POLICY = "perm-after-policy" def _seven_day_policy(received: list[CustomKeyPolicyRequest]): @@ -13300,7 +13409,11 @@ async def test_regenerate_without_changes_still_runs_custom_key_policy(data): def _policy_existing_team_key() -> LiteLLM_VerificationToken: return LiteLLM_VerificationToken( - token=_POLICY_HASHED_TOKEN, user_id="test-user", team_id="team-a", max_budget=200.0 + token=_POLICY_HASHED_TOKEN, + user_id="test-user", + team_id="team-a", + max_budget=200.0, + object_permission_id=_OBJECT_PERMISSION_ID_AFTER_POLICY, ) @@ -13446,9 +13559,6 @@ async def test_process_single_key_update_rejects_when_custom_key_policy_denies() assert [policy_request.operation for policy_request in received] == ["update"] -_OBJECT_PERMISSION_ID_AFTER_POLICY = "perm-after-policy" - - def _record_object_permission_writes(mock_prisma_client: AsyncMock, events: list[str]) -> None: mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) @@ -13568,9 +13678,12 @@ async def test_regenerate_writes_the_object_permission_row_only_after_the_policy events: list[str] = [] _record_object_permission_writes(mock_prisma_client, events) data = RegenerateKeyRequest(max_budget=50.0, object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["vs-1"])) + existing_key: Final = _make_regenerate_existing_key().model_copy( + update={"object_permission_id": _OBJECT_PERMISSION_ID_AFTER_POLICY} + ) with _regenerate_policy_mocks(_recording_policy(events, allowed=True), AsyncMock(), AsyncMock()): - await _regenerate_under_policy(mock_prisma_client, _make_regenerate_existing_key(), data) + await _regenerate_under_policy(mock_prisma_client, existing_key, data) _assert_permission_row_written_after_policy( events, mock_prisma_client.db.litellm_verificationtoken.update.await_args.kwargs["data"] @@ -18848,7 +18961,9 @@ def _estimate_key_row(token: str, metadata: dict): return existing_key -def _wire_update_key_fn(monkeypatch, existing_key): +def _wire_update_key_fn( + monkeypatch: pytest.MonkeyPatch, existing_key: LiteLLM_VerificationToken | MagicMock +) -> AsyncMock: mock_prisma_client = AsyncMock() updated_key = MagicMock() updated_key.token = existing_key.token @@ -18877,6 +18992,7 @@ def _wire_update_key_fn(monkeypatch, existing_key): "litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias", _noop, ) + return mock_prisma_client @pytest.mark.asyncio @@ -20111,11 +20227,13 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read( ) == ["attached-model"] -async def _cache_with_project(project_id: str, project_models: list[str]) -> UserApiKeyCache: +async def _cache_with_project( + project_id: str, project_models: list[str], team_id: str | None = "team-lit-5823" +) -> UserApiKeyCache: user_api_key_cache = UserApiKeyCache() await user_api_key_cache.async_set_cache( key=project_cache_key(project_id), - value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id="team-lit-5823", models=project_models), + value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id=team_id, models=project_models), model_type=LiteLLM_ProjectTableCachedObj, ) return user_api_key_cache @@ -20526,9 +20644,11 @@ async def test_project_detachment_preserves_omission_and_other_key_fields(): @pytest.mark.parametrize("project_id", [None, "project-orbit", "project-other", ""]) @pytest.mark.asyncio -async def test_project_detachment_uses_effective_project_for_validation(project_id: str | None): +async def test_project_detachment_uses_effective_project_for_validation_on_unowned_project( + project_id: str | None, +): existing: Final = LiteLLM_VerificationToken(token="project-detach-token", project_id="project-orbit") - cache: Final = await _cache_with_project("project-orbit", ["model-orbit"]) + cache: Final = await _cache_with_project("project-orbit", ["model-orbit"], team_id=None) data: Final = UpdateKeyRequest(key=existing.token, project_id=project_id, models=["model-other"]) if project_id is None: await _validate_update_key_data( @@ -20563,6 +20683,924 @@ async def test_key_creator_cannot_detach_project_without_admin_access(): assert "Only proxy admins, team admins, or org admins" in str(exc.value.detail) +_OWNED_PROJECT: Final = "proj-owned-1" +_OWNERSHIP_PROJECT_TEAM: Final = "ownership-project-team" +_OWNERSHIP_KEY_TEAM: Final = "ownership-key-team" +_OWNERSHIP_DESTINATION_TEAM: Final = "ownership-destination-team" + + +def _make_generate_mock_prisma() -> AsyncMock: + mock_prisma_client: Final = AsyncMock() + mock_prisma_client.insert_data = AsyncMock( + return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None, object_permission=None) + ) + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_verificationtoken = MagicMock() + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) + mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock( + return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None, object_permission=None) + ) + mock_prisma_client.db.litellm_teamtable = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + return mock_prisma_client + + +def _configure_key_endpoints( + monkeypatch: pytest.MonkeyPatch, + user_api_key_cache: UserApiKeyCache, +) -> AsyncMock: + mock_prisma_client: Final = _make_generate_mock_prisma() + 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=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) + return mock_prisma_client + + +@pytest.mark.parametrize("key_team_id", ["team-a", None]) +@pytest.mark.asyncio +async def test_key_generation_rejects_foreign_project_team( + monkeypatch: pytest.MonkeyPatch, + key_team_id: str | None, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=key_team_id), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + + expected_detail: Final = { + "error": ( + f"Project {_OWNED_PROJECT} belongs to team team-b, 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." + ) + } + assert error.value.status_code == 400 + assert error.value.detail == expected_detail + + +@pytest.mark.asyncio +async def test_key_generation_budget_ceiling_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + max_budget=10, + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="litellm_login_session", + user_id="session-user", + team_id=_OWNERSHIP_KEY_TEAM, + is_session_token=True, + max_budget=5, + ), + litellm_changed_by=None, + team_table=LiteLLM_TeamTableCachedObj(team_id=_OWNERSHIP_KEY_TEAM), + ) + + assert error.value.status_code == 400 + assert error.value.detail == {"error": "max_budget (10.0) cannot exceed the caller's own max_budget (5.0)."} + + +@pytest.mark.asyncio +async def test_key_generation_organization_membership_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_usertable = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="key-owner", organization_memberships=[]) + ) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + organization_id="org-not-member", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user", + user_id="key-owner", + is_session_token=True, + ), + litellm_changed_by=None, + team_table=None, + ) + + assert error.value.status_code == 403 + assert error.value.detail == "Caller is not a member of organization_id=org-not-member" + + +@pytest.mark.asyncio +async def test_key_generation_premium_permission_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + permissions={"get_spend_routes": True}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + ), + litellm_changed_by=None, + team_table=None, + ) + + assert error.value.status_code == 500 + assert error.value.detail == {"error": "Internal Server Error."} + mock_prisma_client.insert_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_generation_duplicate_alias_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + key_alias: Final = "duplicate-alias" + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock( + return_value=LiteLLM_VerificationToken(token="existing-token", key_alias=key_alias) + ) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + with pytest.raises(ProxyException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + key_alias=key_alias, + ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + + assert error.value.code == "400" + assert error.value.message == ( + f"Key with alias '{key_alias}' already exists. Unique key aliases across all keys are required." + ) + mock_prisma_client.insert_data.assert_not_awaited() + + +@pytest.mark.parametrize( + ("key_team_id", "expected_status"), + [(_OWNERSHIP_PROJECT_TEAM, 200), (_OWNERSHIP_KEY_TEAM, 400)], +) +@pytest.mark.asyncio +async def test_key_generation_uses_database_project_team_when_cache_is_stale( + monkeypatch: pytest.MonkeyPatch, + key_team_id: str, + expected_status: int, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_KEY_TEAM) + 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) + data: Final = GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id=key_team_id) + user_api_key_dict: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin") + + if expected_status == 200: + response: Final = await _common_key_generation_helper( + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + team_table=None, + ) + assert response.team_id == _OWNERSHIP_PROJECT_TEAM + return + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + team_table=None, + ) + + expected_detail: Final = { + "error": ( + f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, but the key belongs to " + f"{_OWNERSHIP_KEY_TEAM}. A key can only be attached to a project owned by its own team." + ) + } + assert error.value.status_code == 400 + assert error.value.detail == expected_detail + + +@pytest.mark.asyncio +async def test_key_generation_slow_writer_project_lookup_does_not_stall_db_tracker( + 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) + monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.0) + + async def slow_project_lookup(*, where: Mapping[str, object]) -> LiteLLM_ProjectTable: + assert where == {"project_id": _OWNED_PROJECT} + await asyncio.sleep(0) + await asyncio.sleep(0) + return LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + + prisma_client.writer_db.litellm_projecttable.find_unique = slow_project_lookup + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + db_lookup_stall_tracker.clear() + try: + response: Final = 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 response.team_id == _OWNERSHIP_PROJECT_TEAM + assert db_lookup_stall_tracker.stalled_within(PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS) is False + finally: + db_lookup_stall_tracker.clear() + + +@pytest.mark.parametrize("project_team_id", ["team-a", None]) +@pytest.mark.asyncio +async def test_key_generation_accepts_same_team_and_unowned_projects( + monkeypatch: pytest.MonkeyPatch, + project_team_id: str | None, +) -> None: + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=project_team_id + ) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + + response: Final = await _common_key_generation_helper( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT, team_id="team-a"), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + + assert response.team_id == "team-a" + + +@pytest.mark.parametrize(("project_team_id", "expected_status"), [("team-b", None), ("team-c", 400)]) +@pytest.mark.asyncio +async def test_key_generation_uses_default_team_for_project_ownership( + monkeypatch: pytest.MonkeyPatch, + project_team_id: str, + expected_status: int | None, +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=project_team_id) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", {"team_id": "team-b"}) + data: Final = GenerateKeyRequest(project_id=_OWNED_PROJECT) + + if expected_status is not None: + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=data, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + assert error.value.status_code == expected_status + assert "belongs to team team-c, but the key belongs to team-b" in str(error.value.detail) + return + + response: Final = await _common_key_generation_helper( + data=data, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + team_table=None, + ) + assert response.team_id == "team-b" + + +@pytest.mark.asyncio +async def test_service_account_generation_rejects_foreign_project_team(monkeypatch: pytest.MonkeyPatch) -> None: + team_id: Final = "service-account-team" + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=team_id) + ) + + with pytest.raises(HTTPException) as error: + await generate_service_account_key_fn( + data=GenerateKeyRequest(team_id=team_id, project_id=_OWNED_PROJECT), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + ) + + assert error.value.status_code == 400 + assert "belongs to team team-b, but the key belongs to service-account-team" in str( + error.value.detail + ) + + +@pytest.mark.asyncio +async def test_key_generation_default_budget_does_not_reject_project_budget(monkeypatch: pytest.MonkeyPatch) -> None: + user_api_key_cache: Final = UserApiKeyCache() + project: Final = LiteLLM_ProjectTableCachedObj( + project_id=_OWNED_PROJECT, + team_id="team-a", + models=[], + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + await user_api_key_cache.async_set_cache( + key=project_cache_key(_OWNED_PROJECT), + value=project, + model_type=LiteLLM_ProjectTableCachedObj, + ) + _configure_key_endpoints(monkeypatch, user_api_key_cache) + monkeypatch.setattr(litellm, "default_key_generate_params", {"team_id": "team-a", "max_budget": 10.0}) + + response: Final = await generate_key_fn( + data=GenerateKeyRequest(project_id=_OWNED_PROJECT), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + ) + + assert response.team_id == "team-a" + + +@pytest.mark.parametrize( + "request_fields", + [{"key_alias": "renamed"}, {"team_id": _OWNERSHIP_KEY_TEAM}], +) +@pytest.mark.asyncio +async def test_key_update_allows_legacy_project_mismatch_when_team_is_unchanged( + monkeypatch: pytest.MonkeyPatch, + request_fields: dict[str, str], +) -> None: + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_KEY_TEAM) + ) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, + ) + + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-key", **request_fields), + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=user_api_key_cache, + ) + + +@pytest.mark.asyncio +async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: pytest.MonkeyPatch) -> None: + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, + object_permission_id="permission-update", + budget_id="budget-update", + ) + mock_prisma_client: Final = _wire_update_key_fn(monkeypatch, existing_key_row) + existing_permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="permission-update", + vector_stores=["existing-store"], + ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_permission) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert + permission_before: Final = existing_permission.model_dump() + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM).model_copy( + update={"team_members": []} + ) + ), + ) + mock_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.proxy.proxy_server.llm_router", MagicMock()) + mock_request: Final = MagicMock() + mock_request.query_params = {} + + with pytest.raises((HTTPException, ProxyException)) as error: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest( + key="sk-key", + team_id=_OWNERSHIP_DESTINATION_TEAM, + soft_budget=5.0, + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["replacement-store"]), + ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + ) + + assert str(getattr(error.value, "status_code", None) or getattr(error.value, "code", None)) == "400" + expected_detail: Final = ( + f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, " + f"but the key belongs to {_OWNERSHIP_DESTINATION_TEAM}" + ) + assert expected_detail in str(getattr(error.value, "detail", None) or getattr(error.value, "message", None)) + assert existing_permission.model_dump() == permission_before + permission_upsert.assert_not_awaited() + mock_prisma_client.tx.assert_not_called() + mock_prisma_client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_update_ambiguous_permission_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, + object_permission_id="permission-ambiguous", + ) + mock_prisma_client: Final = _wire_update_key_fn(monkeypatch, existing_key_row) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=LiteLLM_ObjectPermissionTable( + object_permission_id="permission-ambiguous", + mcp_tool_permissions={}, + ) + ) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert + mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock( + return_value=[ + MagicMock(server_id="wiki-a-id", alias="wiki", server_name="wiki-a"), + MagicMock(server_id="wiki-b-id", alias="wiki", server_name="wiki-b"), + ] + ) + mock_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.proxy.management_endpoints.key_management_endpoints.get_team_object", + AsyncMock( + return_value=LiteLLM_TeamTable( + team_id=_OWNERSHIP_DESTINATION_TEAM, + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="permission-team", + mcp_servers=["wiki-a-id", "wiki-b-id"], + ), + ) + ), + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + + with pytest.raises((HTTPException, ProxyException)) as error: + await update_key_fn( + request=MagicMock(query_params={}), + data=UpdateKeyRequest( + key="sk-key", + team_id=_OWNERSHIP_DESTINATION_TEAM, + object_permission=LiteLLM_ObjectPermissionBase( + mcp_tool_permissions={"wiki": ["read_wiki_structure"]} + ), + ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin"), + litellm_changed_by=None, + ) + + assert str(getattr(error.value, "status_code", None) or getattr(error.value, "code", None)) == "400" + expected_detail: Final = { + "error": ( + "Ambiguous mcp_tool_permissions key: 'wiki' matches MCP servers ['wiki-a-id', 'wiki-b-id']. " + "Key tool permissions by server_id when servers share a name or alias." + ) + } + assert getattr(error.value, "detail", None) == expected_detail or getattr(error.value, "message", None) == str( + expected_detail + ) + permission_upsert.assert_not_awaited() + mock_prisma_client.writer_db.litellm_projecttable.find_unique.assert_not_awaited() + mock_prisma_client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_update_organization_membership_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + user_id="key-owner", + created_by="key-owner", + team_id=_OWNERSHIP_PROJECT_TEAM, + project_id=_OWNED_PROJECT, + ) + mock_prisma_client: Final = _wire_update_key_fn(monkeypatch, existing_key_row) + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache) + mock_prisma_client.db.litellm_usertable = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="key-owner", organization_memberships=[]) + ) + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + ) + team_object: Final = LiteLLM_TeamTableCachedObj( + team_id=_OWNERSHIP_KEY_TEAM, + members_with_roles=[Member(user_id="key-owner", role="admin")], + team_member_permissions=["/key/update"], + ) + team_lookup: Final = AsyncMock(return_value=team_object) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + team_lookup, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + team_lookup, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + mock_request: Final = MagicMock() + mock_request.query_params = {} + + with pytest.raises(ProxyException) as error: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest( + key="sk-key", + team_id=_OWNERSHIP_KEY_TEAM, + organization_id="org-not-member", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user", + user_id="key-owner", + team_id=_OWNERSHIP_PROJECT_TEAM, + ), + litellm_changed_by=None, + ) + + assert error.value.code == "403" + assert error.value.message == "Caller is not a member of organization_id=org-not-member" + + +@pytest.mark.asyncio +async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fields( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateKeyRequest, + BulkUpdateKeyRequestItem, + ) + + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-bulk-key", + user_id=None, + models=[], + team_id=_OWNERSHIP_PROJECT_TEAM, + project_id=_OWNED_PROJECT, + object_permission_id="permission-bulk", + budget_id="budget-bulk", + ) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key_row) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM) + ) + existing_permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="permission-bulk", + vector_stores=["existing-store"], + ) + permission_before: Final = existing_permission.model_dump() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_permission) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert + budget_update: Final = AsyncMock() + mock_prisma_client.db.litellm_budgettable.update = budget_update + mock_prisma_client.get_data = AsyncMock(return_value=existing_key_row) + updated_key: Final = MagicMock() + updated_key.model_dump.return_value = { + "max_budget": 10.0, + "tags": ["bulk-update"], + "team_id": _OWNERSHIP_PROJECT_TEAM, + "project_id": _OWNED_PROJECT, + } + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_key}) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + request: Final = BulkUpdateKeyRequest( + keys=[ + BulkUpdateKeyRequestItem.model_validate( + { + "key": "sk-bulk-key", + "team_id": _OWNERSHIP_DESTINATION_TEAM, + "soft_budget": 5.0, + "object_permission": LiteLLM_ObjectPermissionBase(vector_stores=["replacement-store"]), + } + ), + BulkUpdateKeyRequestItem(key="sk-bulk-key", max_budget=10.0, tags=["bulk-update"]), + ] + ) + user_api_key_dict: Final = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ) + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + new_callable=AsyncMock, + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM), + ), + patch("litellm.proxy.management_endpoints.common_utils._premium_user_check"), + patch("litellm.proxy.utils._premium_user_check"), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks", + SimpleNamespace(async_key_updated_hook=AsyncMock()), + ), + ): + response: Final = await bulk_update_keys( + data=request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert response.total_requested == 2 + assert [(item.key, item.failed_reason) for item in response.failed_updates] == [ + ( + "sk-bulk-key", + ( + f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, but the key belongs to " + f"{_OWNERSHIP_DESTINATION_TEAM}. A key can only be attached to a project owned by its own team." + ), + ) + ] + assert len(response.successful_updates) == 1 + assert response.successful_updates[0].key == "sk-bulk-key" + assert response.successful_updates[0].key_info["max_budget"] == 10.0 + assert response.successful_updates[0].key_info["tags"] == ["bulk-update"] + assert existing_permission.model_dump() == permission_before + permission_upsert.assert_not_awaited() + budget_update.assert_not_awaited() + mock_prisma_client.tx.assert_not_called() + mock_prisma_client.update_data.assert_awaited_once() + + +@pytest.mark.parametrize( + "request_fields", + [{"key_alias": "renamed"}, {"team_id": _OWNERSHIP_KEY_TEAM}], +) +@pytest.mark.asyncio +async def test_key_team_ownership_mutation_allows_legacy_mismatch_without_changes( + request_fields: dict[str, str], +) -> None: + prisma_client: Final = MagicMock() + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_PROJECT_TEAM, + ) + ) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, + ) + + await _check_key_project_team_on_mutation( + data=UpdateKeyRequest(key="sk-key", **request_fields), + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) + + +@pytest.mark.asyncio +async def test_key_team_ownership_mutation_allows_detach_with_team_change() -> None: + prisma_client: Final = MagicMock() + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_PROJECT_TEAM, + ) + ) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-key", + team_id=_OWNERSHIP_KEY_TEAM, + project_id=_OWNED_PROJECT, + ) + + await _check_key_project_team_on_mutation( + data=UpdateKeyRequest(key="sk-key", project_id=None, team_id="team-b"), + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) + + +@pytest.mark.parametrize( + ("key_team_id", "expected_status", "expected_updates"), + [("team-a", 400, 0), ("team-b", None, 1)], +) +@pytest.mark.asyncio +async def test_regenerate_checks_project_team_ownership( + key_team_id: str, + expected_status: int | None, + expected_updates: int, +) -> None: + existing_key: Final = LiteLLM_VerificationToken( + token="abc123", + team_id=key_team_id, + object_permission_id="permission-regenerate", + ) + mock_prisma_client: Final = _make_regenerate_mock_prisma() + 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") + ) + deleted_history_table: Final = MagicMock() + deleted_history_table.create_many = AsyncMock() + mock_prisma_client.db.litellm_deletedverificationtoken = deleted_history_table + deprecated_key_table: Final = MagicMock() + deprecated_key_table.upsert = AsyncMock() + mock_prisma_client.db.litellm_deprecatedverificationtoken = deprecated_key_table + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=LiteLLM_ObjectPermissionTable( + object_permission_id="permission-regenerate", + vector_stores=["existing-store"], + ) + ) + permission_upsert: Final = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.upsert = permission_upsert + user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") + + async def regenerate() -> None: + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=RegenerateKeyRequest( + project_id=_OWNED_PROJECT, + grace_period="1h" if expected_status is not None else None, + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["replacement-store"]), + ), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=MagicMock(), + ) + + if expected_status is not None: + with pytest.raises(HTTPException) as error: + await regenerate() + assert error.value.status_code == expected_status + assert "belongs to team team-b, but the key belongs to team-a" in str(error.value.detail) + deleted_history_table.create_many.assert_not_awaited() + deprecated_key_table.upsert.assert_not_awaited() + permission_upsert.assert_not_awaited() + else: + await regenerate() + permission_upsert.assert_awaited_once() + + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == expected_updates + + +@pytest.mark.asyncio +async def test_regenerate_output_estimate_admin_error_precedes_project_ownership() -> None: + existing_key: Final = LiteLLM_VerificationToken( + token="abc123", + user_id="key-owner", + team_id=_OWNERSHIP_PROJECT_TEAM, + project_id=_OWNED_PROJECT, + metadata={}, + ) + mock_prisma_client: Final = _make_regenerate_mock_prisma() + mock_prisma_client.writer_db = MagicMock() + mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( + return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id=_OWNERSHIP_PROJECT_TEAM) + ) + + with pytest.raises(HTTPException) as error: + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=RegenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + metadata={"default_estimated_output_tokens": 1}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-user", + user_id="key-owner", + ), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert error.value.status_code == 403 + assert "Only proxy admins can set" in str(error.value.detail) + mock_prisma_client.db.litellm_verificationtoken.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_key_project_team_validation_allows_missing_project() -> None: + prisma_client: Final = MagicMock() + prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock(return_value=None) + + await _check_key_project_team( + project_id=_OWNED_PROJECT, + key_team_id="team-a", + prisma_client=prisma_client, + ) + + prisma_client.writer_db.litellm_projecttable.find_unique.assert_awaited_once_with( + where={"project_id": _OWNED_PROJECT} + ) + + @pytest.mark.asyncio async def test_bulk_update_team_keys_runs_custom_key_policy_per_key(monkeypatch): from litellm.types.proxy.management_endpoints.key_management_endpoints import ( @@ -20812,16 +21850,16 @@ async def test_key_update_invalidates_cached_object_permission(monkeypatch): @pytest.mark.asyncio -async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch): +async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch: pytest.MonkeyPatch): """Regression: regenerating a key with new permissions must not keep serving the old grants.""" from litellm.proxy._types import LiteLLM_ObjectPermissionBase, RegenerateKeyRequest from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key from litellm.proxy.management_endpoints.key_management_endpoints import ( _execute_virtual_key_regeneration, ) - permission_id = "objperm-regenerate" + permission_id = str(UUID("00000000-0000-0000-0000-000000000002")) grants = {"served": ["tool_a"]} def _row(**kwargs): @@ -20838,9 +21876,12 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch return_value=MagicMock(object_permission_id=permission_id) ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.management_helpers.object_permission_utils.uuid.uuid4", + lambda: UUID(permission_id), + ) existing_key = _make_regenerate_existing_key() - existing_key.object_permission_id = permission_id user_api_key_cache = UserApiKeyCache() assert ( await get_object_permission( @@ -20874,9 +21915,7 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch hashed_api_key="abc123", key="abc123", data=RegenerateKeyRequest( - object_permission=LiteLLM_ObjectPermissionBase( - mcp_tool_permissions={"server-1": ["tool_a", "tool_b"]} - ) + object_permission=LiteLLM_ObjectPermissionBase(mcp_tool_permissions={"server-1": ["tool_a", "tool_b"]}) ), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, @@ -20884,6 +21923,13 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch proxy_logging_obj=AsyncMock(), ) + assert ( + user_api_key_cache.get_cache( + key=object_permission_cache_key(permission_id), + model_type=LiteLLM_ObjectPermissionTable, + ) + is None + ) grants["served"] = ["tool_a", "tool_b"] reread = await get_object_permission( object_permission_id=permission_id,