mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
38586e9684
commit
e79d03e604
2 changed files with 93 additions and 7 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue