fix(proxy): evict jwt_key_mapping cache when user, team, or org deletion removes mapped keys

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
ryan 2026-09-17 23:33:59 +00:00
parent 10616d7407
commit 884a467a7f
8 changed files with 462 additions and 11 deletions

View file

@ -3637,6 +3637,22 @@ async def get_jwt_key_mapping_cache_keys_for_token(
return tuple(jwt_key_mapping_cache_key(m.jwt_claim_name, m.jwt_claim_value, m.jwt_issuer) for m in mappings)
class _TokenInFilter(TypedDict):
token: ReadOnly[Mapping[str, Sequence[str]]]
async def get_jwt_key_mapping_cache_keys_for_tokens(
hashed_tokens: Sequence[str],
prisma_client: PrismaClient,
) -> tuple[str, ...]:
"""Cache keys of every JWT claim mapped to any of the given virtual keys."""
if not hashed_tokens:
return ()
token_filter: Final[_TokenInFilter] = {"token": {"in": tuple(hashed_tokens)}}
mappings: Final = await _jwt_key_mapping_table(JWTKeyMappingRepository(prisma_client)).find_many(where=token_filter)
return tuple(jwt_key_mapping_cache_key(m.jwt_claim_name, m.jwt_claim_value, m.jwt_issuer) for m in mappings)
@log_db_metrics
async def get_jwt_key_mapping_object(
jwt_claim_name: str,

View file

@ -23,12 +23,18 @@ from typing import Any, Final, Literal, Protocol, cast, overload
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import get_team_object, get_user_object
from litellm.proxy.auth.auth_checks import (
delete_cache_key_objects,
get_jwt_key_mapping_cache_keys_for_tokens,
get_team_object,
get_user_object,
)
from litellm.proxy.auth.password_policy import validate_password_policy
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
@ -126,6 +132,10 @@ def _verification_token_table(
return token_table
class _UserIdInFilter(TypedDict):
user_id: ReadOnly[Mapping[str, Sequence[str]]]
def _organization_membership_table(
prisma_client: "PrismaClient | None",
) -> "TableActions[prisma_models.LiteLLM_OrganizationMembership]":
@ -2345,6 +2355,8 @@ async def delete_user(
create_audit_log_for_update,
litellm_proxy_admin_name,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
@ -2471,7 +2483,20 @@ async def delete_user(
# End of Audit logging
## DELETE ASSOCIATED KEYS
await _verification_token_table(prisma_client).delete_many(where={"user_id": {"in": data.user_ids}})
key_filter: Final[_UserIdInFilter] = {"user_id": {"in": data.user_ids}}
keys_to_delete: Final = await _verification_token_table(prisma_client).find_many(where=key_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 _verification_token_table(prisma_client).delete_many(where=key_filter)
await delete_cache_key_objects(
hashed_tokens=hashed_tokens_to_delete,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
## DELETE ASSOCIATED INVITATION LINKS
await _invitation_link_table(prisma_client).delete_many(

View file

@ -26,13 +26,21 @@ from typing import (
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import can_user_call_model, get_user_object
from litellm.proxy.auth.auth_checks import (
can_user_call_model,
delete_cache_key_objects,
get_jwt_key_mapping_cache_keys_for_tokens,
get_user_object,
)
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,
@ -52,7 +60,7 @@ from litellm.proxy.management_helpers.utils import (
get_new_internal_user_defaults,
management_endpoint_wrapper,
)
from litellm.proxy.utils import PrismaClient
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.organization_repository import OrganizationRepository
@ -79,6 +87,7 @@ if TYPE_CHECKING:
)
from prisma.models import LiteLLM_OrganizationTable as PrismaOrganizationTable
from prisma.models import LiteLLM_UserTable as PrismaUserTable
from prisma.models import LiteLLM_VerificationToken as PrismaVerificationToken
async def _enterprise_license_required(
@ -168,9 +177,15 @@ class _TeamTableClient(Protocol):
class _VerificationTokenTableClient(Protocol):
async def find_many(self, where: Mapping[str, object] | None = None) -> "Sequence[PrismaVerificationToken]": ...
async def delete_many(self, where: Mapping[str, object]) -> int: ...
class _OrganizationIdFilter(TypedDict):
organization_id: ReadOnly[str]
class _ObjectPermissionTxClient(Protocol):
async def upsert(
self, where: Mapping[str, object], data: Mapping[str, object]
@ -961,7 +976,7 @@ async def delete_organization(
- organization_ids: List[str] - The organization ids to delete.
"""
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
if prisma_client is None:
raise HTTPException(
@ -983,8 +998,12 @@ async def delete_organization(
await _table(OrganizationMembershipRepository(prisma_client)).delete_many(
where={"organization_id": organization_id}
)
# delete all keys in the organization
await _table(VerificationTokenRepository(prisma_client)).delete_many(where={"organization_id": organization_id})
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},
@ -1000,6 +1019,30 @@ async def delete_organization(
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:
"""Delete the organization's keys and drop every cache entry that still resolves to them,
including the jwt_key_mapping entries the FK cascade removes from the table but not from the cache."""
key_filter: Final[_OrganizationIdFilter] = {"organization_id": organization_id}
keys_to_delete: Final = await _table(VerificationTokenRepository(prisma_client)).find_many(where=key_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)
await delete_cache_key_objects(
hashed_tokens=hashed_tokens_to_delete,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
@router.get(
"/organization/list",
tags=["organization management"],

View file

@ -99,6 +99,7 @@ from litellm.proxy.auth.auth_checks import (
can_org_access_model,
delete_cache_key_objects,
delete_cache_team_object,
get_jwt_key_mapping_cache_keys_for_tokens,
get_org_object,
get_team_membership,
get_team_object,
@ -110,6 +111,7 @@ from litellm.proxy.auth.auth_utils import (
enforce_output_token_estimates_are_admin_only,
)
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.callback_utils import encrypt_callback_vars
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -3524,7 +3526,6 @@ async def team_member_delete(
}'
```
"""
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
if prisma_client is None:
@ -3626,6 +3627,10 @@ async def team_member_delete(
"team_id": data.team_id,
}
)
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
hashed_tokens=tuple(key.token for key in keys_to_delete),
prisma_client=prisma_client,
)
if removed_team_members:
await _team_tx_db(tx).update(
@ -3674,6 +3679,7 @@ async def team_member_delete(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
await evict_and_broadcast(cache_keys=tuple(sorted(user_ids_to_delete)), user_api_key_cache=user_api_key_cache)
for user_id in sorted(user_ids_to_delete):
await invalidate_team_member_spend_state(
@ -4264,6 +4270,10 @@ async def delete_team(
)
keys_to_delete: Final = await _tokens_db(prisma_client).find_many(where={"team_id": {"in": data.team_ids}})
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
hashed_tokens=tuple(key.token for key in keys_to_delete),
prisma_client=prisma_client,
)
if keys_to_delete:
await _persist_deleted_verification_tokens(
@ -4280,6 +4290,7 @@ async def delete_team(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
## DELETE ASSOCIATED BYOK MODELS
# Runs before the team rows are deleted so a mid-flight failure never leaves

View file

@ -2627,6 +2627,9 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
)
# Mock all delete_many calls
mock_prisma_client.db.litellm_verificationtoken.find_many = mocker.AsyncMock(
return_value=[]
)
mock_prisma_client.db.litellm_verificationtoken.delete_many = mocker.AsyncMock(
return_value=0
)
@ -2676,6 +2679,105 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
assert condition[field] == {"in": ["admin-creator"]}
class _JWTMappingRow:
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
self.token = token
self.jwt_claim_name = jwt_claim_name
self.jwt_claim_value = jwt_claim_value
self.jwt_issuer = jwt_issuer
class _CascadingJWTMappingTable:
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
def __init__(self, rows):
self.rows = rows
async def find_many(self, where, **kwargs):
return [row for row in self.rows if row.token in where["token"]["in"]]
def cascade(self, deleted_tokens):
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
@pytest.mark.asyncio
async def test_delete_user_evicts_jwt_key_mapping_cache_of_its_keys(mocker):
"""/user/delete bulk-deletes the user's keys without going through /key/delete, so the
jwt_key_mapping cache entries pointing at those keys must be evicted here too. A surviving
entry keeps resolving the deleted token hash until the mapping cache TTL expires: the deleted
identity is either still served through the stale key cache or 401s on every JWT call, and it is
never re-registered (LIT-5387).
The FK cascade drops the mapping rows with the key rows, so the cache keys have to be read
before the delete: reading them afterwards finds nothing to evict.
"""
from litellm.proxy._types import DeleteUserRequest, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user
global_cache_key: Final = jwt_key_mapping_cache_key("sub", "jwt-user", None)
issuer_cache_key: Final = jwt_key_mapping_cache_key("sub", "jwt-user", "https://issuer.example")
unrelated_cache_key: Final = jwt_key_mapping_cache_key("sub", "other-user", None)
jwt_table: Final = _CascadingJWTMappingTable(
[
_JWTMappingRow("hashed-jwt-key", "sub", "jwt-user"),
_JWTMappingRow("hashed-issuer-key", "sub", "jwt-user", "https://issuer.example"),
_JWTMappingRow("hashed-unrelated-key", "sub", "other-user"),
]
)
cache: Final = UserApiKeyCache()
for cache_key, hashed_token in (
(global_cache_key, "hashed-jwt-key"),
(issuer_cache_key, "hashed-issuer-key"),
(unrelated_cache_key, "hashed-unrelated-key"),
):
cache.set_cache(key=cache_key, value=hashed_token)
cache.set_cache(key=hashed_token, value=UserAPIKeyAuth(token=hashed_token))
user_row: Final = mocker.MagicMock()
user_row.user_id = "jwt-user"
user_row.user_email = "jwt-user@example.com"
user_row.teams = []
user_row.model_dump_json.return_value = "{}"
user_row.model_dump.return_value = {"user_id": "jwt-user", "user_email": "jwt-user@example.com", "teams": []}
mock_prisma_client: Final = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=user_row)
mock_prisma_client.db.litellm_teamtable.find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.db.litellm_jwtkeymapping = jwt_table
mock_prisma_client.db.litellm_verificationtoken.find_many = mocker.AsyncMock(
return_value=[SimpleNamespace(token="hashed-jwt-key"), SimpleNamespace(token="hashed-issuer-key")]
)
async def cascading_delete_many(where):
jwt_table.cascade(("hashed-jwt-key", "hashed-issuer-key"))
return 2
mock_prisma_client.db.litellm_verificationtoken.delete_many = mocker.AsyncMock(side_effect=cascading_delete_many)
mock_prisma_client.db.litellm_invitationlink.delete_many = mocker.AsyncMock(return_value=0)
mock_prisma_client.db.litellm_organizationmembership.delete_many = mocker.AsyncMock(return_value=0)
mock_prisma_client.db.litellm_teammembership.delete_many = mocker.AsyncMock(return_value=0)
mock_prisma_client.db.litellm_usertable.delete_many = mocker.AsyncMock(return_value=1)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # test-quality-ok: substitute the database dependency
mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) # test-quality-ok: exercise a real isolated cache
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", None) # test-quality-ok: delete_user reads it off proxy_server at call time
await delete_user(
data=DeleteUserRequest(user_ids=["jwt-user"]),
user_api_key_dict=UserAPIKeyAuth(user_id="proxy-admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert cache.get_cache(key=global_cache_key) is None
assert cache.get_cache(key=issuer_cache_key) is None
assert cache.get_cache(key="hashed-jwt-key") is None
assert cache.get_cache(key="hashed-issuer-key") is None
assert cache.get_cache(key=unrelated_cache_key) == "hashed-unrelated-key"
assert cache.get_cache(key="hashed-unrelated-key") is not None
assert [row.token for row in jwt_table.rows] == ["hashed-unrelated-key"]
@pytest.mark.asyncio
async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker):
"""Regression: an org admin of org-A must not be able to delete a user

View file

@ -5244,7 +5244,10 @@ async def test_delete_verification_tokens_evicts_jwt_key_mapping_cache(monkeypat
virtual_key_mapping_cache_ttl expires, instead of auto-registering again.
"""
jwt_table = _CascadingJWTMappingTable(
[_JWTMappingRow("hashed-token-1", "email", "user@example.com")]
[
_JWTMappingRow("hashed-token-1", "email", "user@example.com"),
_JWTMappingRow("hashed-token-1", "email", "user@example.com", "https://issuer.example"),
]
)
key1 = LiteLLM_VerificationToken(
@ -5302,7 +5305,10 @@ async def test_delete_verification_tokens_evicts_jwt_key_mapping_cache(monkeypat
),
)
assert recording_evict.cache_keys == (jwt_key_mapping_cache_key("email", "user@example.com", None),)
assert recording_evict.cache_keys == (
jwt_key_mapping_cache_key("email", "user@example.com", None),
jwt_key_mapping_cache_key("email", "user@example.com", "https://issuer.example"),
)
@pytest.mark.asyncio

View file

@ -1,7 +1,7 @@
import asyncio
import json
from litellm._uuid import uuid
from types import MappingProxyType
from types import MappingProxyType, SimpleNamespace
from typing import Final, Mapping, Optional, cast
from unittest.mock import AsyncMock, MagicMock, patch
@ -1438,3 +1438,82 @@ def test_organization_routes_reach_their_handler_with_enterprise_license(monkeyp
assert any(
message in response.text for message in (CommonProxyErrors.db_not_connected_error.value, "No db connected")
)
class _JWTMappingRow:
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
self.token = token
self.jwt_claim_name = jwt_claim_name
self.jwt_claim_value = jwt_claim_value
self.jwt_issuer = jwt_issuer
class _CascadingJWTMappingTable:
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
def __init__(self, rows):
self.rows = rows
async def find_many(self, where, **kwargs):
return [row for row in self.rows if row.token in where["token"]["in"]]
def cascade(self, deleted_tokens):
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
@pytest.mark.asyncio
async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monkeypatch):
"""/organization/delete bulk-deletes the org's keys without going through /key/delete, so the
key objects and the jwt_key_mapping entries (issuer-scoped ones included) pointing at them
must be evicted here, or a deleted key keeps authenticating and a JWT identity keeps resolving
a token hash that no longer exists until the TTLs expire. The FK cascade drops the mapping
rows with the key rows, so the cache keys have to be read before the delete (LIT-5387)."""
from litellm.proxy._types import DeleteOrganizationRequest, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.organization_endpoints import delete_organization
doomed_cache_keys: Final = (
"hashed-org-key",
jwt_key_mapping_cache_key("sub", "svc-account", None),
jwt_key_mapping_cache_key("sub", "svc-account", "https://issuer.example"),
)
kept_cache_keys: Final = ("hashed-other-key", jwt_key_mapping_cache_key("sub", "other-account", None))
kept_row: Final = _JWTMappingRow("hashed-other-key", "sub", "other-account")
jwt_table: Final = _CascadingJWTMappingTable(
[
_JWTMappingRow("hashed-org-key", "sub", "svc-account"),
_JWTMappingRow("hashed-org-key", "sub", "svc-account", "https://issuer.example"),
kept_row,
]
)
cache: Final = UserApiKeyCache()
for cache_key in (*doomed_cache_keys, *kept_cache_keys):
cache.set_cache(key=cache_key, value={"retained": True})
prisma_client: Final = AsyncMock()
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
prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=cascading_delete_many)
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)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None)
await delete_organization(
data=DeleteOrganizationRequest(organization_ids=["org-doomed"]),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
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]

View file

@ -9180,6 +9180,175 @@ async def test_team_member_delete_persists_deleted_keys(monkeypatch):
assert cache.get_cache(key="unrelated-key") == {"retained": True}
class _JWTMappingRow:
def __init__(self, token, jwt_claim_name, jwt_claim_value, jwt_issuer=None):
self.token = token
self.jwt_claim_name = jwt_claim_name
self.jwt_claim_value = jwt_claim_value
self.jwt_issuer = jwt_issuer
class _CascadingJWTMappingTable:
"""Mapping rows that LiteLLM_JWTKeyMapping_token_fkey drops when their key row is deleted."""
def __init__(self, rows):
self.rows = rows
async def find_many(self, where, **kwargs):
return [row for row in self.rows if row.token in where["token"]["in"]]
def cascade(self, deleted_tokens):
self.rows = [row for row in self.rows if row.token not in deleted_tokens]
def _seed_jwt_mapping_cache(cache, mapping_rows):
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
cache_keys = tuple(
jwt_key_mapping_cache_key(row.jwt_claim_name, row.jwt_claim_value, row.jwt_issuer) for row in mapping_rows
)
for cache_key, row in zip(cache_keys, mapping_rows):
cache.set_cache(key=cache_key, value=row.token)
return cache_keys
@pytest.mark.asyncio
async def test_team_member_delete_evicts_jwt_key_mapping_cache_of_the_keys_it_deletes(monkeypatch):
"""The member's team keys are deleted in bulk here, not through /key/delete, so the
jwt_key_mapping cache entries pointing at them must be evicted here too, or every JWT call
from that identity resolves the deleted token hash and 401s until the mapping TTL expires.
The FK cascade drops the mapping rows with the key rows, so the cache keys have to be read
before the delete (LIT-5387)."""
from litellm.proxy._types import TeamMemberDeleteRequest
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.key_management_endpoints import LiteLLM_VerificationToken
doomed_rows: Final = (
_JWTMappingRow("hashed-token-1", "sub", "user-123"),
_JWTMappingRow("hashed-token-1", "sub", "user-123", "https://issuer.example"),
)
kept_row: Final = _JWTMappingRow("hashed-other-key", "sub", "user-999")
jwt_table: Final = _CascadingJWTMappingTable([*doomed_rows, kept_row])
team = LiteLLM_TeamTable(
team_id="team-1",
team_alias="test-team",
members_with_roles=[Member(user_id="user-123", role="admin")],
metadata={},
model_max_budget={},
model_spend={},
)
key1 = LiteLLM_VerificationToken(token="hashed-token-1", user_id="user-123", team_id="team-1")
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(
return_value=[MagicMock(user_id="user-123", teams=["team-1"])]
)
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[key1])
async def cascading_delete_many(where):
jwt_table.cascade(("hashed-token-1",))
mock_prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=cascading_delete_many)
mock_prisma_client.db.litellm_jwtkeymapping = jwt_table
_wire_member_delete_tx(mock_prisma_client)
cache: Final = UserApiKeyCache()
doomed_cache_keys: Final = _seed_jwt_mapping_cache(cache, doomed_rows)
(kept_cache_key,) = _seed_jwt_mapping_cache(cache, (kept_row,))
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
monkeypatch.setattr("litellm.proxy.management_endpoints.team_endpoints._is_user_team_admin", lambda **kwargs: True)
await team_member_delete(
data=TeamMemberDeleteRequest(team_id="team-1", user_id="user-123"),
user_api_key_dict=UserAPIKeyAuth(
user_id="admin-user", api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN.value
),
)
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
assert cache.get_cache(key=kept_cache_key) == "hashed-other-key"
assert jwt_table.rows == [kept_row]
@pytest.mark.asyncio
async def test_delete_team_evicts_jwt_key_mapping_cache_of_the_keys_it_deletes(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""Same contract as /team/member_delete for the bulk key delete in /team/delete: the
jwt_key_mapping cache entries of the team's keys, issuer-scoped ones included, are gone
after the delete while entries pointing at other keys survive (LIT-5387)."""
from litellm.proxy._types import DeleteTeamRequest, LiteLLM_VerificationToken
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
doomed_rows: Final = (
_JWTMappingRow("hashed-doomed-key", "sub", "svc-account"),
_JWTMappingRow("hashed-doomed-key", "sub", "svc-account", "https://issuer.example"),
)
kept_row: Final = _JWTMappingRow("hashed-unrelated-key", "sub", "svc-account", "https://other-issuer.example")
jwt_table: Final = _CascadingJWTMappingTable([*doomed_rows, kept_row])
team = LiteLLM_TeamTable(
team_id="team-doomed",
team_alias="doomed-team",
members_with_roles=[],
metadata={},
model_max_budget={},
model_spend={},
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
async def cascading_delete_data(team_id_list, table_name):
jwt_table.cascade(("hashed-doomed-key",))
return {"deleted_keys": 1}
mock_prisma_client.delete_data = AsyncMock(side_effect=cascading_delete_data)
mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock()
mock_prisma_client.db.litellm_deletedverificationtoken.create_many = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[LiteLLM_VerificationToken(token="hashed-doomed-key", team_id="team-doomed")]
)
mock_prisma_client.db.litellm_jwtkeymapping = jwt_table
mock_prisma_client.db.execute_raw = AsyncMock()
mock_prisma_client.db.litellm_teammembership.delete_many = AsyncMock()
mock_tx = AsyncMock()
mock_tx.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_tx_cm = MagicMock()
mock_tx_cm.__aenter__ = AsyncMock(return_value=mock_tx)
mock_tx_cm.__aexit__ = AsyncMock(return_value=False)
mock_prisma_client.db.tx = MagicMock(return_value=mock_tx_cm)
_wire_team_delete_tx(mock_prisma_client)
cache: Final = UserApiKeyCache()
doomed_cache_keys: Final = _seed_jwt_mapping_cache(cache, doomed_rows)
(kept_cache_key,) = _seed_jwt_mapping_cache(cache, (kept_row,))
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
monkeypatch.setattr("litellm.proxy.proxy_server.create_audit_log_for_update", AsyncMock())
monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")
await delete_team(
data=DeleteTeamRequest(team_ids=["team-doomed"]),
http_request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(
user_id="admin-user", api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN.value
),
litellm_changed_by="admin-user",
)
assert all(cache.get_cache(key=cache_key) is None for cache_key in doomed_cache_keys)
assert cache.get_cache(key=kept_cache_key) == "hashed-unrelated-key"
assert jwt_table.rows == [kept_row]
@pytest.mark.asyncio
async def test_new_team_negative_max_budget():
"""