fix(proxy): evict jwt key mapping cache on bulk user and team member deletion

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
ryan 2026-09-18 00:00:17 +00:00
parent 38586e9684
commit e79d03e604
2 changed files with 93 additions and 7 deletions

View file

@ -28,7 +28,7 @@ from litellm.proxy._types import (
MemberDeleteRequest,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import delete_cache_key_objects
from litellm.proxy.auth.auth_checks import delete_cache_key_objects, get_jwt_key_mapping_cache_keys_for_tokens
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks
@ -94,12 +94,20 @@ class _TeamRemoval:
removed: frozenset[str]
matched: frozenset[int]
deleted_key_tokens: tuple[str, ...]
jwt_mapping_cache_keys: tuple[str, ...]
@dataclass(frozen=True, slots=True)
class _UserBatchDeletion:
removals: Mapping[str, _TeamRemoval]
deleted_key_tokens: tuple[str, ...]
jwt_mapping_cache_keys: tuple[str, ...]
@dataclass(frozen=True, slots=True)
class _DeletedKeys:
tokens: tuple[str, ...]
jwt_mapping_cache_keys: tuple[str, ...]
def _team_not_found(team_id: str) -> ManagementProblem:
@ -237,6 +245,10 @@ async def _remove_members_from_team(
if any(_addresses_member(m, r) for m in removed_members) or any(_addresses_user(u, r) for u in stale_rows)
)
keys: Final = await _token_tx_db(tx).find_many(where=_team_users_filter(team_id, cleanup_ids))
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
hashed_tokens=tuple(k.token for k in keys),
prisma_client=prisma_client,
)
if removed_members:
roster_data: Final[_RosterData] = {
@ -265,6 +277,7 @@ async def _remove_members_from_team(
removed=cleanup_ids,
matched=matched,
deleted_key_tokens=tuple(k.token for k in keys),
jwt_mapping_cache_keys=jwt_mapping_cache_keys,
)
@ -322,6 +335,7 @@ async def bulk_remove_team_members(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await evict_and_broadcast(cache_keys=removal.jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
_emit_team_members_metric(removal.team)
matched: Final = frozenset(kept_indexes[j] for j in removal.matched)
@ -368,8 +382,12 @@ async def _delete_user_rows(
user_ids: frozenset[str],
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: str | None,
) -> tuple[str, ...]:
) -> _DeletedKeys:
keys: Final = await _token_tx_db(tx).find_many(where=_in_filter("user_id", user_ids))
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens(
hashed_tokens=tuple(k.token for k in keys),
prisma_client=prisma_client,
)
if keys:
await _persist_deleted_verification_tokens(
keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken
@ -389,7 +407,7 @@ async def _delete_user_rows(
await _org_membership_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _membership_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _user_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
return tuple(k.token for k in keys)
return _DeletedKeys(tokens=tuple(k.token for k in keys), jwt_mapping_cache_keys=jwt_mapping_cache_keys)
async def _delete_users_tx(
@ -423,12 +441,14 @@ async def _delete_users_tx(
for tid in team_ids
}
)
deleted_key_tokens: Final = await _delete_user_rows(
deleted_keys: Final = await _delete_user_rows(
prisma_client, tx, frozenset(u.user_id for u in users), user_api_key_dict, litellm_changed_by
)
return _UserBatchDeletion(
removals=removals,
deleted_key_tokens=deleted_key_tokens + tuple(t for r in removals.values() for t in r.deleted_key_tokens),
deleted_key_tokens=deleted_keys.tokens + tuple(t for r in removals.values() for t in r.deleted_key_tokens),
jwt_mapping_cache_keys=deleted_keys.jwt_mapping_cache_keys
+ tuple(k for r in removals.values() for k in r.jwt_mapping_cache_keys),
)
@ -454,6 +474,7 @@ async def _delete_users(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await evict_and_broadcast(cache_keys=deletion.jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
await evict_and_broadcast(cache_keys=sorted(user_ids), user_api_key_cache=user_api_key_cache)
for removal in deletion.removals.values():
_emit_team_members_metric(removal.team)
@ -534,7 +555,7 @@ async def bulk_delete_users(
litellm_changed_by,
)
if candidates
else _UserBatchDeletion(removals=MappingProxyType({}), deleted_key_tokens=())
else _UserBatchDeletion(removals=MappingProxyType({}), deleted_key_tokens=(), jwt_mapping_cache_keys=())
)
def result(index: int, user_id: str) -> UserDeleteResult:

View file

@ -8,6 +8,7 @@ import pytest
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, 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.list_api.common import ManagementProblem
from litellm.proxy.management_helpers.bulk_user_deletion import bulk_delete_users, bulk_remove_team_members
@ -114,6 +115,7 @@ class _Db:
tokens: Sequence[Mapping[str, object]] = (),
invitations: Sequence[Mapping[str, object]] = (),
org_memberships: Sequence[Mapping[str, object]] = (),
jwt_mappings: Sequence[Mapping[str, object]] = (),
) -> None:
self.litellm_usertable = _UserTable(users)
self.litellm_teamtable = _TeamTable(teams)
@ -122,6 +124,7 @@ class _Db:
self.litellm_deletedverificationtoken = _Rows()
self.litellm_invitationlink = _Rows(invitations)
self.litellm_organizationmembership = _Rows(org_memberships)
self.litellm_jwtkeymapping = _Rows(jwt_mappings)
class _Tx:
@ -163,11 +166,12 @@ class _FakePrisma:
tokens: Sequence[Mapping[str, object]] = (),
invitations: Sequence[Mapping[str, object]] = (),
org_memberships: Sequence[Mapping[str, object]] = (),
jwt_mappings: Sequence[Mapping[str, object]] = (),
on_lock: Callable[[str], None] = lambda _: None,
fail_locks: frozenset[str] = frozenset(),
fail_commit: bool = False,
) -> None:
self.db = _Db(users, teams, memberships, tokens, invitations, org_memberships)
self.db = _Db(users, teams, memberships, tokens, invitations, org_memberships, jwt_mappings)
self._on_lock = on_lock
self._fail_locks = fail_locks
self._fail_commit = fail_commit
@ -212,6 +216,17 @@ def _cache_with(*hashed_tokens: str) -> UserApiKeyCache:
return cache
def _jwt_mapping(token: str, claim_value: str, issuer: str | None = None) -> Mapping[str, object]:
return {"token": token, "jwt_claim_name": "sub", "jwt_claim_value": claim_value, "jwt_issuer": issuer}
def _cache_with_jwt_mapping_keys(*cache_keys: str) -> UserApiKeyCache:
cache = UserApiKeyCache()
for key in cache_keys:
cache.set_cache(key=key, value={"cache_key": key})
return cache
async def _delete(
prisma: _FakePrisma,
user_ids: Sequence[str],
@ -449,6 +464,34 @@ async def test_bulk_delete_evicts_deleted_keys_and_users_from_the_auth_cache():
assert cache.get_cache(key="keep-key") is not None
@pytest.mark.asyncio
async def test_bulk_delete_evicts_jwt_key_mappings_of_the_deleted_users_keys():
issuer: Final = "https://issuer.example"
doomed_global: Final = jwt_key_mapping_cache_key("sub", "alice")
doomed_scoped: Final = jwt_key_mapping_cache_key("sub", "alice", issuer)
kept: Final = jwt_key_mapping_cache_key("sub", "bob")
prisma = _FakePrisma(
users=[_user("u1", "t1"), _user("keep", "t1")],
teams=[_team("t1", "u1", "keep")],
tokens=[
{"token": "team-key", "user_id": "u1", "team_id": "t1"},
{"token": "personal-key", "user_id": "u1"},
{"token": "keep-key", "user_id": "keep", "team_id": "t1"},
],
jwt_mappings=[
_jwt_mapping("personal-key", "alice"),
_jwt_mapping("team-key", "alice", issuer=issuer),
_jwt_mapping("keep-key", "bob"),
],
)
cache = _cache_with_jwt_mapping_keys(doomed_global, doomed_scoped, kept)
await _delete(prisma, ["u1"], cache=cache)
assert cache.get_cache(key=doomed_global) is None and cache.get_cache(key=doomed_scoped) is None
assert cache.get_cache(key=kept) is not None
@pytest.mark.asyncio
async def test_bulk_delete_rejects_non_admin_callers_before_touching_the_db():
prisma = _FakePrisma(users=[_user("u1")])
@ -577,6 +620,28 @@ async def test_bulk_member_delete_evicts_the_removed_team_keys_from_the_auth_cac
assert cache.get_cache(key="keep-key") is not None
@pytest.mark.asyncio
async def test_bulk_member_delete_evicts_jwt_key_mappings_of_the_removed_team_keys():
issuer: Final = "https://issuer.example"
doomed: Final = jwt_key_mapping_cache_key("sub", "alice", issuer)
kept: Final = jwt_key_mapping_cache_key("sub", "bob")
prisma = _FakePrisma(
users=[_user("u1", "t1"), _user("keep", "t1")],
teams=[_team("t1", "u1", "keep")],
tokens=[
{"token": "team-key", "user_id": "u1", "team_id": "t1"},
{"token": "keep-key", "user_id": "keep", "team_id": "t1"},
],
jwt_mappings=[_jwt_mapping("team-key", "alice", issuer=issuer), _jwt_mapping("keep-key", "bob")],
)
cache = _cache_with_jwt_mapping_keys(doomed, kept)
await _remove(prisma, "t1", [{"user_id": "u1"}], cache=cache)
assert cache.get_cache(key=doomed) is None
assert cache.get_cache(key=kept) is not None
@pytest.mark.asyncio
async def test_bulk_member_delete_cleans_a_user_whose_teams_array_still_names_the_team():
prisma = _FakePrisma(users=[_user("stale", "t1")], teams=[_team("t1", "other")], memberships=[("t1", "stale")])