mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
25bb5b48e7
commit
21c77e9db2
2 changed files with 82 additions and 35 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue