diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 43d840130b6..22644bcd207 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2212,13 +2212,14 @@ async def _fetch_team_membership_from_db( value=NO_TEAM_MEMBERSHIP_SENTINEL, ttl=get_management_object_ttl(user_api_key_cache), ) - return None - - await user_api_key_cache.async_set_cache( - key=_key, - value=membership, - model_type=LiteLLM_TeamMembership, - ) + else: + await user_api_key_cache.async_set_cache( + key=_key, + value=membership, + model_type=LiteLLM_TeamMembership, + ) + if _team_membership_inflight.get(_key) is not asyncio.current_task(): + await user_api_key_cache.async_delete_cache(key=_key) return membership diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 986d53b80b0..61f0675aa1e 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -6543,6 +6543,38 @@ async def test_get_team_membership_invalidation_mid_flight_discards_stale_load() assert cached is not None and cached.budget_id == "budget-new" +@pytest.mark.asyncio +async def test_get_team_membership_invalidation_during_cache_write_evicts_stale_entry(): + from litellm.proxy.auth.auth_checks import get_team_membership, invalidate_team_member_spend_state + from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key + + write_started = asyncio.Event() + release_write = asyncio.Event() + + class _SlowWriteCache(UserApiKeyCache): + async def async_set_cache(self, key, value, local_only=False, **kwargs): + write_started.set() + await release_write.wait() + return await super().async_set_cache(key, value, local_only=local_only, **kwargs) + + row = MagicMock() + row.dict = lambda: {"user_id": "u-w", "team_id": "t-w", "spend": 1.0, "budget_id": "budget-old"} + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=row) + cache = _SlowWriteCache() + + stale = asyncio.create_task( + get_team_membership(user_id="u-w", team_id="t-w", prisma_client=mock_prisma_client, user_api_key_cache=cache) + ) + await asyncio.wait_for(write_started.wait(), timeout=2) + await invalidate_team_member_spend_state(user_id="u-w", team_id="t-w", user_api_key_cache=cache) + release_write.set() + stale_result = await stale + + assert stale_result is not None and stale_result.budget_id == "budget-old" + assert await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-w", team_id="t-w")) is None + + @pytest.mark.asyncio async def test_common_checks_calls_get_team_membership_once_per_request(): from fastapi import Request