From 4ac168b3fcb440a7e0af0e3b919ad6880baf2fd6 Mon Sep 17 00:00:00 2001 From: yassin Date: Mon, 14 Sep 2026 20:13:39 +0000 Subject: [PATCH] fix(auth): drop in-flight membership load on invalidation so it cannot repopulate the cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 8 ++- .../proxy/auth/test_auth_checks.py | 50 +++++++++++++++++++ 2 files changed, 56 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index fa85400f4e0..43d840130b6 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2202,8 +2202,11 @@ async def _fetch_team_membership_from_db( where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, include={"litellm_budget_table": True}, ) + membership: Final = None if response is None else LiteLLM_TeamMembership.model_validate(response.dict()) _key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id) - if response is None: + if _team_membership_inflight.get(_key) is not asyncio.current_task(): + return membership + if membership is None: await user_api_key_cache.async_set_cache( key=_key, value=NO_TEAM_MEMBERSHIP_SENTINEL, @@ -2211,7 +2214,6 @@ async def _fetch_team_membership_from_db( ) return None - membership: Final = LiteLLM_TeamMembership.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=_key, value=membership, @@ -2762,6 +2764,8 @@ async def invalidate_team_member_spend_state( publish_auth_cache_invalidation, ) + _team_membership_inflight.pop(team_membership_reservation_cache_key(user_id=user_id, team_id=team_id), None) + if new_spend is not None: from litellm.proxy.proxy_server import SPEND_DB_FLOOR_CACHE_TTL_SECONDS, spend_counter_cache diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index b843916debf..986d53b80b0 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -6493,6 +6493,56 @@ async def test_get_team_membership_coalesces_parallel_db_fetches(): mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once() +@pytest.mark.asyncio +async def test_get_team_membership_invalidation_mid_flight_discards_stale_load(): + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import get_team_membership, invalidate_team_member_spend_state + from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec + from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key + + started = asyncio.Event() + release_stale = asyncio.Event() + release_fresh = asyncio.Event() + loads = iter((("budget-old", release_stale), ("budget-new", release_fresh))) + + async def _find_unique(*args, **kwargs): + budget_id, release = next(loads) + row = MagicMock() + row.dict = lambda: {"user_id": "u-inv", "team_id": "t-inv", "spend": 1.0, "budget_id": budget_id} + started.set() + await release.wait() + return row + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_find_unique) + cache = UserApiKeyCache() + + async def _load(): + return await get_team_membership( + user_id="u-inv", team_id="t-inv", prisma_client=mock_prisma_client, user_api_key_cache=cache + ) + + stale = asyncio.create_task(_load()) + await started.wait() + await invalidate_team_member_spend_state(user_id="u-inv", team_id="t-inv", user_api_key_cache=cache) + started.clear() + fresh = asyncio.create_task(_load()) + await asyncio.wait_for(started.wait(), timeout=2) + release_fresh.set() + fresh_result = await fresh + release_stale.set() + stale_result = await stale + + assert stale_result is not None and stale_result.budget_id == "budget-old" + assert fresh_result is not None and fresh_result.budget_id == "budget-new" + assert mock_prisma_client.db.litellm_teammembership.find_unique.await_count == 2 + cached = CacheCodec.deserialize( + await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-inv", team_id="t-inv")), + model_type=LiteLLM_TeamMembership, + ) + assert cached is not None and cached.budget_id == "budget-new" + + @pytest.mark.asyncio async def test_common_checks_calls_get_team_membership_once_per_request(): from fastapi import Request