From 25bb5b48e781c63a8eb18e93b1c41ce0f9aa0313 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 21:26:19 +0000 Subject: [PATCH] fix(proxy): validate org ids and delete organizations atomically Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../organization_endpoints.py | 94 ++++++++------ tests/e2e/coverage_registry/mgmt.yaml | 1 + tests/e2e/e2e_http.py | 16 ++- tests/e2e/management/management_client.py | 10 ++ tests/e2e/management/test_management_e2e.py | 26 ++++ tests/e2e/transport.py | 34 ++++- .../test_organization_endpoints.py | 120 +++++++++++++++++- 7 files changed, 254 insertions(+), 47 deletions(-) diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 24bbd2b4b1f..c39752b3300 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -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", diff --git a/tests/e2e/coverage_registry/mgmt.yaml b/tests/e2e/coverage_registry/mgmt.yaml index e1a840b1239..a95b880324e 100644 --- a/tests/e2e/coverage_registry/mgmt.yaml +++ b/tests/e2e/coverage_registry/mgmt.yaml @@ -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"} diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index e5d50d05c87..17102fbd2c0 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -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, ) ) diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py index 7366695c0d1..7f7172e71e1 100644 --- a/tests/e2e/management/management_client.py +++ b/tests/e2e/management/management_client.py @@ -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( diff --git a/tests/e2e/management/test_management_e2e.py b/tests/e2e/management/test_management_e2e.py index da0fc37aff8..1a0eb70d2e8 100644 --- a/tests/e2e/management/test_management_e2e.py +++ b/tests/e2e/management/test_management_e2e.py @@ -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") diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py index 87aad0d08de..950d85b219f 100644 --- a/tests/e2e/transport.py +++ b/tests/e2e/transport.py @@ -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, 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 3c6afa86c45..fd81379b30d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -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()