fix(proxy): validate org ids and delete organizations atomically

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-24 21:26:19 +00:00
parent fc29fb513c
commit 25bb5b48e7
7 changed files with 254 additions and 47 deletions

View file

@ -40,7 +40,6 @@ from litellm.proxy.auth.auth_checks import (
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.budget_management_endpoints import (
new_budget,
update_budget,
@ -60,7 +59,7 @@ from litellm.proxy.management_helpers.utils import (
get_new_internal_user_defaults,
management_endpoint_wrapper,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.proxy.utils import PrismaClient
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.organization_repository import OrganizationRepository
@ -206,6 +205,15 @@ class _TransactionTables(Protocol):
@property
def litellm_organizationtable(self) -> "_OrganizationTableClient": ...
@property
def litellm_teamtable(self) -> "_TeamTableClient": ...
@property
def litellm_organizationmembership(self) -> "_OrganizationMembershipTableClient": ...
@property
def litellm_verificationtoken(self) -> "_VerificationTokenTableClient": ...
class _TransactionManager(Protocol):
async def __aenter__(self) -> "_TransactionTables": ...
@ -995,49 +1003,31 @@ async def delete_organization(
detail={"error": "Only proxy admins can delete organizations"},
)
deleted_orgs: Final = []
for organization_id in data.organization_ids:
# delete all teams in the organization
await _table(TeamRepository(prisma_client)).delete_many(where={"organization_id": organization_id})
# delete all members in the organization
await _table(OrganizationMembershipRepository(prisma_client)).delete_many(
where={"organization_id": organization_id}
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
)
await _delete_organization_keys(
organization_id=organization_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# delete the organization
deleted_org = await _table(OrganizationRepository(prisma_client)).delete(
where={"organization_id": organization_id},
include={"members": True, "teams": True, "litellm_budget_table": True},
)
if deleted_org is None:
raise HTTPException(
status_code=404,
detail={"error": f"Organization={organization_id} not found"},
)
deleted_orgs.append(deleted_org)
return deleted_orgs
async def _delete_organization_keys(
organization_id: str,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None,
) -> None:
key_filter: Final[_OrganizationIdFilter] = {"organization_id": organization_id}
keys_to_delete: Final = await _table(VerificationTokenRepository(prisma_client)).find_many(where=key_filter)
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,
)
await _table(VerificationTokenRepository(prisma_client)).delete_many(where=key_filter)
tx_manager: Final[_TransactionManager] = prisma_client.db.tx()
async with tx_manager as tx:
deleted_orgs: Final = await _delete_organizations_in_tx(tx=tx, organization_ids=requested_ids)
await delete_cache_key_objects(
hashed_tokens=hashed_tokens_to_delete,
user_api_key_cache=user_api_key_cache,
@ -1045,6 +1035,34 @@ async def _delete_organization_keys(
)
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
return deleted_orgs
async def _delete_organizations_in_tx(
tx: _TransactionTables,
organization_ids: tuple[str, ...],
) -> tuple["PrismaOrganizationTable | None", ...]:
return tuple(
[
await _delete_organization_in_tx(tx=tx, organization_id=organization_id)
for organization_id in organization_ids
]
)
async def _delete_organization_in_tx(
tx: _TransactionTables,
organization_id: str,
) -> "PrismaOrganizationTable | None":
org_filter: Final[_OrganizationIdFilter] = {"organization_id": organization_id}
await tx.litellm_teamtable.delete_many(where=org_filter)
await tx.litellm_organizationmembership.delete_many(where=org_filter)
await tx.litellm_verificationtoken.delete_many(where=org_filter)
return await tx.litellm_organizationtable.delete(
where=org_filter,
include={"members": True, "teams": True, "litellm_budget_table": True}, # mutable-ok: Prisma include
)
@router.get(
"/organization/list",

View file

@ -49,6 +49,7 @@
- {id: mgmt.organization.new.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "organization_endpoints.py:403", rationale: "Org for multi-tenant isolation"}
- {id: mgmt.organization.update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "organization_endpoints.py:545", rationale: "Org metadata updates persist"}
- {id: mgmt.organization.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "organization_endpoints.py:710", rationale: "Cascades to teams/keys"}
- {id: mgmt.organization.delete.unknown_id_rejects_whole_request, module: mgmt, tier: P1, surface: api, assertions: [unknown_id_rejects_whole_request], source: "organization_endpoints.py", fail_before_fix: proven, rationale: "One unknown id in a batch delete 404s and deletes nothing"}
- {id: mgmt.organization.member_add.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "organization_endpoints.py:835", rationale: "Org member onboarding"}
- {id: mgmt.customer.new.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "customer_endpoints.py:372", rationale: "End-user for spend tracking"}
- {id: mgmt.customer.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "customer_endpoints.py:480", rationale: "Removes from spend tracking"}

View file

@ -642,13 +642,23 @@ def put[R: BaseModel](
return classify(resp, response_type)
def probe(url: URL, *, headers: BaseModel, params: BaseModel, timeout: float = 30.0) -> ProbeResult:
def probe(
url: URL,
*,
headers: BaseModel,
params: BaseModel | None = None,
method: str = "GET",
json: BaseModel | None = None,
timeout: float = 30.0,
) -> ProbeResult:
try:
resp = request_with_retry(
lambda: requests.get(
lambda: requests.request(
method,
str(url),
headers=_headers(headers),
params=params.model_dump(by_alias=True, exclude_none=True),
params=_params(params),
json=wire_body(json) if json is not None else None,
timeout=timeout,
)
)

View file

@ -566,6 +566,16 @@ class ManagementClient:
headers=self.proxy.management_headers(),
)
def delete_orgs_status(self, organization_ids: list[str]) -> ProbeResult:
"""DELETE /organization/delete judged by HTTP outcome only, for requests
the test expects the route to reject."""
return self.proxy.transport.probe(
"/organization/delete",
method="DELETE",
json=OrgDeleteBody(organization_ids=organization_ids),
headers=self.proxy.management_headers(),
)
def create_tag(self, body: TagNewBody) -> None:
_ = unwrap(
self.proxy.transport.post(

View file

@ -632,6 +632,32 @@ class TestOrganizationRoutes:
_ = _poll(client, gone, f"org {org_id} still resolved on /organization/info after /organization/delete")
@pytest.mark.covers("mgmt.organization.delete.unknown_id_rejects_whole_request")
def test_delete_with_unknown_id_rejects_whole_request(
self, client: ManagementClient, resources: ResourceManager
) -> None:
org_id = client.create_org(OrgNewBody(organization_alias=f"e2e-mgmt-org-{unique_marker()}"))
resources.defer(lambda: client.delete_org(org_id))
team_id = client.create_team(
TeamNewBody(team_alias=f"e2e-mgmt-team-{unique_marker()}", organization_id=org_id)
)
resources.defer(lambda: client.delete_team(team_id))
rejected = client.delete_orgs_status([org_id, f"e2e-missing-{unique_marker()}"])
assert rejected.status_code == 404, (
f"/organization/delete with one unknown id must reject the whole request with 404, got "
f"{rejected.status_code}: {rejected.body[:300]}"
)
assert client.org_info_status(org_id).status_code == 200, (
f"org {org_id} must still resolve on /organization/info after the rejected batch delete"
)
team_probe = client.team_info_status(team_id)
assert team_probe.status_code == 200, (
f"team {team_id} inside org {org_id} must still resolve on /team/info after the rejected batch "
f"delete, got {team_probe.status_code}: {team_probe.body[:300]}"
)
class TestTagRoutes:
@pytest.mark.covers("mgmt.tag.new.happy_path")

View file

@ -91,7 +91,15 @@ class Transport(Protocol):
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]: ...
def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult: ...
def probe(
self,
path: str,
*,
params: BaseModel | None = None,
headers: BaseModel | None = None,
method: str = "GET",
json: BaseModel | None = None,
) -> ProbeResult: ...
def upload[R: BaseModel](
self,
@ -253,11 +261,21 @@ class HttpTransport:
) -> AbandonedRequest | StreamingResponse:
return e2e_http.abandon(self._url(path), headers=headers, json=json, after=after)
def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult:
def probe(
self,
path: str,
*,
params: BaseModel | None = None,
headers: BaseModel | None = None,
method: str = "GET",
json: BaseModel | None = None,
) -> ProbeResult:
return e2e_http.probe(
self._url(path),
headers=self.master if headers is None else headers,
params=params,
method=method,
json=json,
timeout=self.request_timeout,
)
@ -439,8 +457,16 @@ class SplitTransport:
) -> AbandonedRequest | StreamingResponse:
return self._route(path).abandon(path, headers=headers, json=json, after=after)
def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult:
return self._route(path).probe(path, params=params, headers=headers)
def probe(
self,
path: str,
*,
params: BaseModel | None = None,
headers: BaseModel | None = None,
method: str = "GET",
json: BaseModel | None = None,
) -> ProbeResult:
return self._route(path).probe(path, params=params, headers=headers, method=method, json=json)
def upload[R: BaseModel](
self,

View file

@ -1520,6 +1520,9 @@ 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")]
)
@ -1528,9 +1531,13 @@ async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monke
jwt_table.cascade(("hashed-org-key",))
return 1
prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=cascading_delete_many)
tx: Final = MagicMock()
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)
tx.litellm_organizationtable.delete = AsyncMock(return_value=MagicMock())
prisma_client.db.tx = MagicMock(return_value=_FakeTxContext(tx))
prisma_client.db.litellm_jwtkeymapping = jwt_table
prisma_client.db.litellm_organizationtable.delete = AsyncMock(return_value=MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True, raising=False)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
@ -1545,3 +1552,112 @@ async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monke
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
assert all(cache.get_cache(key=cache_key) == {"retained": True} for cache_key in kept_cache_keys)
assert jwt_table.rows == (kept_row,)
@pytest.mark.asyncio
async def test_delete_organization_unknown_id_rejects_before_any_delete(monkeypatch):
"""One unknown id in organization_ids must 404 the whole request before any write: the
old per-org loop deleted the earlier orgs first, so [existing, missing] removed the
existing org and then failed (LIT-8570)."""
from litellm.proxy._types import DeleteOrganizationRequest, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.organization_endpoints import delete_organization
prisma_client: Final = AsyncMock()
prisma_client.db.litellm_organizationtable.find_many = AsyncMock(
return_value=[SimpleNamespace(organization_id="org-present")]
)
prisma_client.db.tx = MagicMock(return_value=_FakeTxContext(MagicMock()))
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True, raising=False)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
with pytest.raises(HTTPException) as exc_info:
await delete_organization(
data=DeleteOrganizationRequest(organization_ids=["org-present", "org-missing"]),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
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()
@pytest.mark.asyncio
async def test_delete_organization_writes_inside_tx_and_evicts_after_commit(monkeypatch):
"""Every org-cascade write must run inside prisma_client.db.tx() so a batch cannot
half-commit, and the key/jwt cache eviction must run only after the transaction
has committed (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
events: Final[list[str]] = []
tx: Final = MagicMock()
def record(table):
return AsyncMock(side_effect=lambda **kwargs: events.append(table) or 0)
tx.litellm_teamtable.delete_many = record("teams")
tx.litellm_organizationmembership.delete_many = record("members")
tx.litellm_verificationtoken.delete_many = record("keys")
tx.litellm_organizationtable.delete = AsyncMock(
side_effect=lambda **kwargs: events.append("org") or MagicMock()
)
class _RecordingTxContext:
async def __aenter__(self):
events.append("tx_enter")
return tx
async def __aexit__(self, exc_type, exc, tb):
events.append("tx_exit")
return False
prisma_client: Final = AsyncMock()
prisma_client.db.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=[])
prisma_client.db.litellm_jwtkeymapping = CascadingJWTMappingTable([])
prisma_client.db.tx = MagicMock(return_value=_RecordingTxContext())
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",
AsyncMock(side_effect=lambda **kwargs: events.append("evict_keys")),
)
monkeypatch.setattr(
organization_endpoints,
"evict_and_broadcast",
AsyncMock(side_effect=lambda **kwargs: events.append("broadcast")),
)
await delete_organization(
data=DeleteOrganizationRequest(organization_ids=["org-1", "org-2"]),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert events == [
"tx_enter",
"teams",
"members",
"keys",
"org",
"teams",
"members",
"keys",
"org",
"tx_exit",
"evict_keys",
"broadcast",
]
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()