fix(proxy): validate and fetch org delete rows inside the writer transaction

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-24 23:35:18 +00:00
parent 25bb5b48e7
commit 21c77e9db2
2 changed files with 82 additions and 35 deletions

View file

@ -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(

View file

@ -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()