fix(auth): evict the membership cache entry when invalidation lands during the cache write

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-14 20:46:09 +00:00
parent 4ac168b3fc
commit 91c964a338
2 changed files with 40 additions and 7 deletions

View file

@ -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

View file

@ -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