diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index c39752b3300..1a35ba8aae6 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -14,6 +14,7 @@ Endpoints for /organization operations #### ORGANIZATION MANAGEMENT #### from collections.abc import Mapping, Sequence +from datetime import timedelta from typing import ( TYPE_CHECKING, Annotated, @@ -1004,28 +1005,29 @@ async def delete_organization( ) requested_ids: Final = tuple(dict.fromkeys(data.organization_ids)) - existing_rows: Final = await _table(OrganizationRepository(prisma_client)).find_many( - where={"organization_id": {"in": list(requested_ids)}} # mutable-ok: Prisma filter - ) - existing_ids: Final = frozenset(row.organization_id for row in existing_rows) - missing: Final = tuple(organization_id for organization_id in requested_ids if organization_id not in existing_ids) - if missing: - raise HTTPException( - status_code=404, - detail={"error": f"Organization(s) not found: {', '.join(missing)}"}, # mutable-ok: error envelope - ) - - keys_to_delete: Final = await _table(VerificationTokenRepository(prisma_client)).find_many( - where={"organization_id": {"in": list(requested_ids)}} # mutable-ok: Prisma filter - ) - hashed_tokens_to_delete: Final = tuple(key.token for key in keys_to_delete) - jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens( - hashed_tokens=hashed_tokens_to_delete, - prisma_client=prisma_client, - ) - - tx_manager: Final[_TransactionManager] = prisma_client.db.tx() + tx_manager: Final[_TransactionManager] = prisma_client.db.tx(timeout=timedelta(minutes=2)) async with tx_manager as tx: + existing_rows: Final = await tx.litellm_organizationtable.find_many( + where={"organization_id": {"in": list(requested_ids)}} # mutable-ok: Prisma filter + ) + existing_ids: Final = frozenset(row.organization_id for row in existing_rows) + missing: Final = tuple( + organization_id for organization_id in requested_ids if organization_id not in existing_ids + ) + if missing: + raise HTTPException( + status_code=404, + detail={"error": f"Organization(s) not found: {', '.join(missing)}"}, # mutable-ok: error envelope + ) + + keys_to_delete: Final = await tx.litellm_verificationtoken.find_many( + where={"organization_id": {"in": list(requested_ids)}} # mutable-ok: Prisma filter + ) + hashed_tokens_to_delete: Final = tuple(key.token for key in keys_to_delete) + jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens( + hashed_tokens=hashed_tokens_to_delete, + prisma_client=prisma_client, + ) deleted_orgs: Final = await _delete_organizations_in_tx(tx=tx, organization_ids=requested_ids) await delete_cache_key_objects( diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index fd81379b30d..f5d747ae485 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -1520,18 +1520,18 @@ async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monke cache.set_cache(key=cache_key, value={"retained": True}) prisma_client: Final = AsyncMock() - prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[SimpleNamespace(organization_id="org-doomed")] - ) - prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[SimpleNamespace(token="hashed-org-key")] - ) async def cascading_delete_many(where): jwt_table.cascade(("hashed-org-key",)) return 1 tx: Final = MagicMock() + tx.litellm_organizationtable.find_many = AsyncMock( + return_value=[SimpleNamespace(organization_id="org-doomed")] + ) + tx.litellm_verificationtoken.find_many = AsyncMock( + return_value=[SimpleNamespace(token="hashed-org-key")] + ) tx.litellm_teamtable.delete_many = AsyncMock(return_value=0) tx.litellm_organizationmembership.delete_many = AsyncMock(return_value=0) tx.litellm_verificationtoken.delete_many = AsyncMock(side_effect=cascading_delete_many) @@ -1563,10 +1563,11 @@ async def test_delete_organization_unknown_id_rejects_before_any_delete(monkeypa from litellm.proxy.management_endpoints.organization_endpoints import delete_organization prisma_client: Final = AsyncMock() - prisma_client.db.litellm_organizationtable.find_many = AsyncMock( + tx: Final = MagicMock() + tx.litellm_organizationtable.find_many = AsyncMock( return_value=[SimpleNamespace(organization_id="org-present")] ) - prisma_client.db.tx = MagicMock(return_value=_FakeTxContext(MagicMock())) + prisma_client.db.tx = MagicMock(return_value=_FakeTxContext(tx)) monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True, raising=False) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) @@ -1579,9 +1580,10 @@ async def test_delete_organization_unknown_id_rejects_before_any_delete(monkeypa assert exc_info.value.status_code == 404 assert "org-missing" in str(exc_info.value.detail) - prisma_client.db.tx.assert_not_called() - prisma_client.db.litellm_verificationtoken.find_many.assert_not_called() - prisma_client.db.litellm_organizationtable.delete.assert_not_called() + tx.litellm_teamtable.delete_many.assert_not_called() + tx.litellm_organizationmembership.delete_many.assert_not_called() + tx.litellm_verificationtoken.find_many.assert_not_called() + tx.litellm_organizationtable.delete.assert_not_called() @pytest.mark.asyncio @@ -1616,11 +1618,12 @@ async def test_delete_organization_writes_inside_tx_and_evicts_after_commit(monk events.append("tx_exit") return False - prisma_client: Final = AsyncMock() - prisma_client.db.litellm_organizationtable.find_many = AsyncMock( + tx.litellm_organizationtable.find_many = AsyncMock( return_value=[SimpleNamespace(organization_id="org-1"), SimpleNamespace(organization_id="org-2")] ) - prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + tx.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + + prisma_client: Final = AsyncMock() prisma_client.db.litellm_jwtkeymapping = CascadingJWTMappingTable([]) prisma_client.db.tx = MagicMock(return_value=_RecordingTxContext()) @@ -1661,3 +1664,45 @@ async def test_delete_organization_writes_inside_tx_and_evicts_after_commit(monk prisma_client.db.litellm_teamtable.delete_many.assert_not_called() prisma_client.db.litellm_organizationmembership.delete_many.assert_not_called() prisma_client.db.litellm_verificationtoken.delete_many.assert_not_called() + + +@pytest.mark.asyncio +async def test_delete_organization_tx_failure_evicts_nothing(monkeypatch): + """If a delete inside the transaction raises, the exception must propagate and the + key/jwt cache eviction must not run: evicting entries for keys the rollback kept + would leave them unable to authenticate (LIT-8570).""" + from litellm.proxy._types import DeleteOrganizationRequest, LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.management_endpoints import organization_endpoints + from litellm.proxy.management_endpoints.organization_endpoints import delete_organization + + tx: Final = MagicMock() + tx.litellm_organizationtable.find_many = AsyncMock( + return_value=[SimpleNamespace(organization_id="org-1")] + ) + tx.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + tx.litellm_teamtable.delete_many = AsyncMock(return_value=0) + tx.litellm_organizationmembership.delete_many = AsyncMock(return_value=0) + tx.litellm_verificationtoken.delete_many = AsyncMock(side_effect=RuntimeError("db write failed")) + + prisma_client: Final = AsyncMock() + prisma_client.db.litellm_jwtkeymapping = CascadingJWTMappingTable([]) + prisma_client.db.tx = MagicMock(return_value=_FakeTxContext(tx)) + + evict_keys: Final = AsyncMock() + broadcast: Final = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True, raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None) + monkeypatch.setattr(organization_endpoints, "delete_cache_key_objects", evict_keys) + monkeypatch.setattr(organization_endpoints, "evict_and_broadcast", broadcast) + + with pytest.raises(RuntimeError, match="db write failed"): + await delete_organization( + data=DeleteOrganizationRequest(organization_ids=["org-1"]), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + evict_keys.assert_not_called() + broadcast.assert_not_called()