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