diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 1e10834752a..e35256b56c5 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1379,6 +1379,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: @@ -1484,10 +1488,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)) @@ -1555,7 +1565,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 @@ -1831,7 +1853,7 @@ async def _check_key_project_team( if project_obj.team_id is None or project_obj.team_id == key_team_id: return - raise HTTPException( + raise KeyProjectTeamMismatchError( status_code=400, detail={ "error": ( @@ -1843,6 +1865,10 @@ async def _check_key_project_team( ) +class KeyProjectTeamMismatchError(HTTPException): + pass + + async def _check_key_project_team_on_mutation( data: UpdateKeyRequest | RegenerateKeyRequest, existing_key_row: LiteLLM_VerificationToken, diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index b4da1a5a3f5..b1765082a62 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -89,6 +89,13 @@ def _object_permission_table_rows() -> list[dict[str, JsonValue]]: ) +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', @@ -337,13 +344,25 @@ def test_service_account_generate_rejects_foreign_team_project_without_writing_k 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]}, + { + "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', @@ -1082,17 +1101,33 @@ def test_key_generate_rejects_foreign_team_project(ownership_gateway: Gateway) - 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, "models": [model]}, + { + "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 read_rows( - 'SELECT token FROM "LiteLLM_VerificationToken" WHERE project_id = %s AND team_id = %s', - (project_b, team_a), - ) == [] + 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: 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 e62be2499be..77ac29755c4 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -56,6 +56,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( _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, @@ -1823,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