From 0f10c0624180e23fb35da1882ddbcbea1c6421cb Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 11 Sep 2026 16:31:32 -0700 Subject: [PATCH] fix(proxy): evict negative membership cache when a member row is created The get_team_membership negative cache stores a NO_TEAM_MEMBERSHIP_SENTINEL for a session-token member with no LiteLLM_TeamMembership row. The two create paths that add a row with a per-member budget, /team/member_add and the /team/update budget backfill, did not evict that sentinel, so the new per-member budget stayed unenforced until the membership cache TTL expired. Add _evict_created_membership_caches and call it from both sites so the budget applies on the next request. --- .../management_endpoints/team_endpoints.py | 46 ++++++++++++++++++- .../test_team_endpoints.py | 46 +++++++++++++++++++ 2 files changed, 91 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index c050368b3fe..4d7ed30852d 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -14,7 +14,7 @@ import copy import json import math import traceback -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from collections.abc import Set as AbstractSet from datetime import datetime, timezone from types import MappingProxyType @@ -1902,6 +1902,39 @@ def validate_team_org_change( return True +def _member_user_ids(members_with_roles: Sequence[dict[str, object]]) -> tuple[str, ...]: + """Extract the string ``user_id`` of each team member, dropping rows without one. + + ``members_with_roles`` is a Prisma-deserialized JSON column, so its ``user_id`` is typed + ``object``; the ``isinstance`` narrows it to the ``str`` ``invalidate_team_member_spend_state`` needs. + """ + return tuple(user_id for member in members_with_roles if isinstance((user_id := member.get("user_id")), str)) + + +async def _evict_created_membership_caches( + user_ids: Iterable[str], + team_id: str, + user_api_key_cache: UserApiKeyCache, +) -> None: + """Evict the ``get_team_membership`` negative-cache sentinel for members whose row was just created. + + A session-token request caches ``NO_TEAM_MEMBERSHIP_SENTINEL`` for a member with no + ``LiteLLM_TeamMembership`` row. When a create path (``/team/member_add`` or the ``/team/update`` + budget backfill) later writes that row with a per-member budget, the stale sentinel keeps the + member's budget unenforced until the membership cache TTL expires, so it must be evicted here. + """ + await asyncio.gather( + *( + invalidate_team_member_spend_state( + user_id=user_id, + team_id=team_id, + user_api_key_cache=user_api_key_cache, + ) + for user_id in user_ids + ) + ) + + @router.post("/team/update", tags=["team management"], dependencies=[Depends(user_api_key_auth)]) @management_endpoint_wrapper async def update_team( @@ -2238,6 +2271,11 @@ async def update_team( team_member_budget_id=_backfill_budget_id, prisma_client=prisma_client, ) + await _evict_created_membership_caches( + user_ids=_member_user_ids(existing_team_row.members_with_roles), + team_id=data.team_id, + user_api_key_cache=user_api_key_cache, + ) elif _team_member_fields_in_request: updated_kv = await TeamMemberBudgetHandler.clear_team_member_budget_fields( team_table=existing_team_row, @@ -3190,6 +3228,12 @@ async def team_member_add( litellm_proxy_admin_name=litellm_proxy_admin_name, ) + await _evict_created_membership_caches( + user_ids=(tm.user_id for tm in updated_team_memberships), + team_id=data.team_id, + user_api_key_cache=user_api_key_cache, + ) + _emit_team_members_metric(complete_team_data) await _create_team_member_add_audit_logs( diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 2f6561046b1..ccc639620cc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -13941,6 +13941,52 @@ async def test_team_member_update_skips_invalidation_when_no_budget_fields_sent( assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 1.5 +@pytest.mark.asyncio +async def test_evict_created_membership_caches_drops_the_negative_sentinel(): + """ + Regression: a membership-create path (/team/member_add, the /team/update budget backfill) must + evict any cached "no membership" sentinel a prior session-token read left, so a per-member budget + attached at create time is enforced on the next request instead of after the membership cache TTL. + Uses a real cache so the assertion is that the sentinel is actually gone, not that a mock was called. + """ + from litellm.proxy.common_utils.user_api_key_cache import ( + NO_TEAM_MEMBERSHIP_SENTINEL, + UserApiKeyCache, + team_membership_reservation_cache_key, + ) + from litellm.proxy.management_endpoints.team_endpoints import _evict_created_membership_caches + + cache = UserApiKeyCache() + kept_key = team_membership_reservation_cache_key(user_id="carol", team_id="team-eviction") + evicted_key = team_membership_reservation_cache_key(user_id="bob", team_id="team-eviction") + await cache.async_set_cache(key=kept_key, value=NO_TEAM_MEMBERSHIP_SENTINEL) + await cache.async_set_cache(key=evicted_key, value=NO_TEAM_MEMBERSHIP_SENTINEL) + + await _evict_created_membership_caches(user_ids=("bob",), team_id="team-eviction", user_api_key_cache=cache) + + assert await cache.async_get_cache(key=evicted_key) is None + assert await cache.async_get_cache(key=kept_key) == NO_TEAM_MEMBERSHIP_SENTINEL + + +def test_member_user_ids_keeps_only_string_user_ids(): + """ + The /team/update backfill feeds Prisma-deserialized member dicts here; a row can be missing + user_id or carry a non-string value. Only real string ids may reach invalidate_team_member_spend_state, + so those get eviction and the malformed rows are dropped rather than crashing the update. + """ + from litellm.proxy.management_endpoints.team_endpoints import _member_user_ids + + members = [ + {"user_id": "alice", "role": "admin"}, + {"role": "user"}, + {"user_id": None, "role": "user"}, + {"user_id": 123, "role": "user"}, + {"user_id": "bob", "role": "user"}, + ] + + assert _member_user_ids(members) == ("alice", "bob") + + def _team_spend_by_user_team(team_id: str, team_alias: str, member: Member, permissions: list[str]) -> MagicMock: team = MagicMock(spec=LiteLLM_TeamTable) team.team_id = team_id