From 70ddc7e4928b19fe2d73cc97ff707e0919a8bceb Mon Sep 17 00:00:00 2001 From: Shivi Jain Date: Mon, 14 Sep 2026 21:28:08 +0530 Subject: [PATCH] fix(auth): load team membership once per request and skip prisma on an L1 hit common_checks was querying get_team_membership twice, and DualCache awaited Redis SET on the auth path, so LRU eviction plus a hung Redis write showed up as two postgres spans --- litellm/constants.py | 7 + litellm/proxy/auth/auth_checks.py | 163 +++++++++++------- .../proxy/auth/test_auth_checks.py | 160 +++++++++++++++++ 3 files changed, 265 insertions(+), 65 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 5751e6e46af..5c7a02d0743 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2039,3 +2039,10 @@ BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit" # Shared read-only empty mapping, for defaulting optional Mapping parameters without # constructing a fresh mutable dict at each call site. EMPTY_MAPPING: Final = MappingProxyType({}) + + +class TeamMembershipCacheMiss: + __slots__ = () + + +TEAM_MEMBERSHIP_CACHE_MISS: Final = TeamMembershipCacheMiss() diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index abcade6b1c8..0baa79fb95a 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -35,6 +35,8 @@ from litellm.constants import ( MODEL_ACCESS_GROUP_REGISTRY_MAX_SIZE, REGISTRY_ERROR_NEGATIVE_CACHE_TTL, TAG_REGISTRY_MAX_SIZE, + TEAM_MEMBERSHIP_CACHE_MISS, + TeamMembershipCacheMiss, ) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider @@ -327,12 +329,31 @@ _safe_json_loads_obj: Final = _typed_json_loads(safe_json_loads) last_db_access_time: Final = LimitedSizeOrderedDict(max_size=100) db_cache_expiry: Final = DEFAULT_IN_MEMORY_TTL # refresh every 5s -_TEAM_MEMBERSHIP_CACHE_MISS: Final = object() -_team_membership_inflight: dict[str, asyncio.Task[LiteLLM_TeamMembership | None]] = {} +_TEAM_MEMBERSHIP_INFLIGHT_MAX: Final = 10000 +_team_membership_inflight: Final = LimitedSizeOrderedDict(max_size=_TEAM_MEMBERSHIP_INFLIGHT_MAX) +_team_membership_write_epoch: Final = LimitedSizeOrderedDict(max_size=_TEAM_MEMBERSHIP_INFLIGHT_MAX) all_routes: Final = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value +def _membership_write_epoch(key: str) -> int: + cached: Final[object] = _team_membership_write_epoch.get(key, 0) + return cached if isinstance(cached, int) else 0 + + +def _bump_membership_write_epoch(key: str) -> None: + _team_membership_write_epoch[key] = _membership_write_epoch(key) + 1 + + +def _membership_from_shared_load(result: object) -> LiteLLM_TeamMembership | None: + if result is None or isinstance(result, LiteLLM_TeamMembership): + return result + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Failed to load team membership", + ) + + def _log_budget_lookup_failure(entity: str, error: Exception) -> None: """ Log a warning when budget lookup fails; cache will not be populated. @@ -2165,38 +2186,39 @@ async def get_tag_object( return tag_objects.get(tag_name) -def _debug_team_membership_log(hypothesis_id: str, message: str, data: dict[str, object]) -> None: - # #region agent log - try: - import json as _json - - with open("/Users/shijain/genai-apps/genai-proxy/.cursor/debug-86534f.log", "a", encoding="utf-8") as _f: - _f.write( - _json.dumps( - { - "sessionId": "86534f", - "timestamp": int(time.time() * 1000), - "location": "auth_checks.py:get_team_membership", - "message": message, - "hypothesisId": hypothesis_id, - "data": data, - } - ) - + "\n" - ) - except Exception: - pass - # #endregion - - -def _membership_from_cached_payload(cached: object) -> LiteLLM_TeamMembership | None | object: - """Decode a DualCache payload. ``_TEAM_MEMBERSHIP_CACHE_MISS`` means try the next tier.""" +def _membership_from_cached_payload( + cached: object, +) -> LiteLLM_TeamMembership | None | TeamMembershipCacheMiss: if cached is None: - return _TEAM_MEMBERSHIP_CACHE_MISS + return TEAM_MEMBERSHIP_CACHE_MISS if cached == NO_TEAM_MEMBERSHIP_SENTINEL: return None cached_membership: Final = CacheCodec.deserialize(cached, model_type=LiteLLM_TeamMembership) - return cached_membership if cached_membership is not None else _TEAM_MEMBERSHIP_CACHE_MISS + return cached_membership if cached_membership is not None else TEAM_MEMBERSHIP_CACHE_MISS + + +async def _set_team_membership_cache_entry( + user_api_key_cache: UserApiKeyCache, + key: str, + value: object, + *, + local_only: bool, + model_type: type[LiteLLM_TeamMembership] | None, + ttl: float | None, +) -> None: + match (model_type is not None, ttl is not None): + case (False, False): + await user_api_key_cache.async_set_cache(key=key, value=value, local_only=local_only) + case (False, True): + await user_api_key_cache.async_set_cache(key=key, value=value, local_only=local_only, ttl=ttl) + case (True, False): + await user_api_key_cache.async_set_cache( + key=key, value=value, local_only=local_only, model_type=model_type + ) + case (True, True): + await user_api_key_cache.async_set_cache( + key=key, value=value, local_only=local_only, model_type=model_type, ttl=ttl + ) async def _populate_team_membership_cache( @@ -2207,17 +2229,33 @@ async def _populate_team_membership_cache( model_type: type[LiteLLM_TeamMembership] | None = None, ttl: float | None = None, ) -> None: - """Await in-memory write; replicate to Redis off the auth await path.""" - kwargs: dict[str, object] = {} - if model_type is not None: - kwargs["model_type"] = model_type - if ttl is not None: - kwargs["ttl"] = ttl - await user_api_key_cache.async_set_cache(key=key, value=value, local_only=True, **kwargs) + write_epoch: Final = _membership_write_epoch(key) + await _set_team_membership_cache_entry( + user_api_key_cache, + key, + value, + local_only=True, + model_type=model_type, + ttl=ttl, + ) + if _membership_write_epoch(key) != write_epoch: + await user_api_key_cache.async_delete_cache(key) + return async def _replicate_to_redis() -> None: try: - await user_api_key_cache.async_set_cache(key=key, value=value, **kwargs) + if _membership_write_epoch(key) != write_epoch: + return + await _set_team_membership_cache_entry( + user_api_key_cache, + key, + value, + local_only=False, + model_type=model_type, + ttl=ttl, + ) + if _membership_write_epoch(key) != write_epoch: + await user_api_key_cache.async_delete_cache(key) except Exception: return @@ -2271,19 +2309,9 @@ async def _load_team_membership_on_cache_miss( try: redis_cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key) redis_membership: Final = _membership_from_cached_payload(redis_cached) - if redis_membership is not _TEAM_MEMBERSHIP_CACHE_MISS: - # #region agent log - _debug_team_membership_log( - "H3", - "membership redis hit after l1 miss", - {"prisma": False, "source": "redis"}, - ) - # #endregion - return cast(LiteLLM_TeamMembership | None, redis_membership) + if not isinstance(redis_membership, TeamMembershipCacheMiss): + return redis_membership - # #region agent log - _debug_team_membership_log("H1", "membership prisma fetch", {"prisma": True}) - # #endregion return await _fetch_team_membership_from_db( user_id=user_id, team_id=team_id, @@ -2292,13 +2320,18 @@ async def _load_team_membership_on_cache_miss( parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) - except Exception: + except HTTPException: + raise + except Exception as e: verbose_proxy_logger.exception( "Error getting team membership for user_id: %s, team_id: %s", user_id, team_id, ) - return None + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Failed to load team membership", + ) from e async def get_team_membership( @@ -2321,18 +2354,12 @@ async def get_team_membership( l1_cached: Final[object] = await user_api_key_cache.async_get_cache(key=_key, local_only=True) l1_membership: Final = _membership_from_cached_payload(l1_cached) - if l1_membership is not _TEAM_MEMBERSHIP_CACHE_MISS: - # #region agent log - _debug_team_membership_log("H4", "membership l1 hit", {"prisma": False, "source": "l1"}) - # #endregion - return cast(LiteLLM_TeamMembership | None, l1_membership) + if not isinstance(l1_membership, TeamMembershipCacheMiss): + return l1_membership - inflight: Final = _team_membership_inflight.get(_key) - if inflight is not None: - # #region agent log - _debug_team_membership_log("H5", "membership coalesced waiter", {"prisma": False, "coalesced": True}) - # #endregion - return await inflight + inflight: Final[object] = _team_membership_inflight.get(_key) + if isinstance(inflight, asyncio.Task): + return _membership_from_shared_load(await asyncio.shield(inflight)) if prisma_client is None: raise Exception("No db connected") @@ -2349,8 +2376,13 @@ async def get_team_membership( ) ) _team_membership_inflight[_key] = task - task.add_done_callback(lambda _t, k=_key: _team_membership_inflight.pop(k, None)) - return await task + + def _clear_inflight(_done: object) -> None: + if _team_membership_inflight.get(_key) is task: + _team_membership_inflight.pop(_key, None) + + task.add_done_callback(_clear_inflight) + return _membership_from_shared_load(await asyncio.shield(task)) def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool: @@ -2861,6 +2893,7 @@ async def invalidate_team_member_spend_state( ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS, ) + _bump_membership_write_epoch(team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)) await evict_and_broadcast( cache_keys=( team_membership_auth_cache_key(team_id=team_id, user_id=user_id), diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 755fe9efe2c..d16c75a4a3b 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -6581,6 +6581,166 @@ async def test_common_checks_calls_get_team_membership_once_per_request(): assert load_membership.await_count == 1 +@pytest.mark.asyncio +async def test_get_team_membership_db_error_raises_503_not_none(): + """A Prisma failure must fail closed as 503, not look like a missing membership row.""" + from fastapi import HTTPException + + from litellm.proxy.auth.auth_checks import get_team_membership + from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=RuntimeError("db down")) + cache = UserApiKeyCache() + + with pytest.raises(HTTPException) as exc: + await get_team_membership( + user_id="u-fail", + team_id="t-fail", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) + + cached = await cache.async_get_cache( + key=team_membership_reservation_cache_key(user_id="u-fail", team_id="t-fail") + ) + assert exc.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert cached is None + + +@pytest.mark.asyncio +async def test_common_checks_does_not_skip_member_limits_when_membership_lookup_fails(): + """common_checks must not mark membership loaded-absent after a lookup error.""" + from fastapi import HTTPException, Request + + from litellm.proxy.auth.auth_checks import common_checks + + team = LiteLLM_TeamTable(team_id="t-fail-closed") + token = UserAPIKeyAuth( + token="k-fail-closed", + user_id="u-fail-closed", + team_id="t-fail-closed", + models=["gpt-4o-mini"], + ) + lookup_error = HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Failed to load team membership", + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache()), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + side_effect=lookup_error, + ), + patch("litellm.proxy.proxy_server.get_current_spend", new_callable=AsyncMock, return_value=0.0), + ): + with pytest.raises(HTTPException) as exc: + await common_checks( + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + team_object=team, + user_object=LiteLLM_UserTable(user_id="u-fail-closed"), + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route="/chat/completions", + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=MagicMock(spec=Request), + ) + + assert exc.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + + +@pytest.mark.asyncio +async def test_get_team_membership_waiter_cancel_does_not_cancel_shared_load(): + """Cancelling one coalesced waiter must not cancel the shared Prisma load.""" + from litellm.proxy.auth.auth_checks import get_team_membership + + started = asyncio.Event() + release = asyncio.Event() + membership_row = MagicMock() + membership_row.dict = lambda: {"user_id": "u-shield", "team_id": "t-shield", "spend": 1.0} + + async def _slow_find_unique(*args, **kwargs): + started.set() + await release.wait() + return membership_row + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_slow_find_unique) + cache = UserApiKeyCache() + + async def _load(): + return await get_team_membership( + user_id="u-shield", + team_id="t-shield", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) + + owner = asyncio.create_task(_load()) + await started.wait() + waiter = asyncio.create_task(_load()) + await asyncio.sleep(0) + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + release.set() + result = await owner + + assert result is not None + assert result.user_id == "u-shield" + mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_stale_membership_redis_replicate_does_not_restore_after_invalidate(): + """A delayed DualCache Redis SET must not resurrect membership after invalidation.""" + 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 + + membership_row = MagicMock() + membership_row.dict = lambda: {"user_id": "u-stale", "team_id": "t-stale", "spend": 9.0} + + hang_redis_set = asyncio.Event() + + async def _hanging_redis_set(*args, **kwargs): + await hang_redis_set.wait() + + redis_cache = MagicMock() + redis_cache.async_get_cache = AsyncMock(return_value=None) + redis_cache.async_set_cache = AsyncMock(side_effect=_hanging_redis_set) + redis_cache.async_delete_cache = AsyncMock() + redis_cache.delete_cache = MagicMock() + + cache = UserApiKeyCache(redis_cache=redis_cache) + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) + + loaded = await get_team_membership( + user_id="u-stale", + team_id="t-stale", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) + cache_key = team_membership_reservation_cache_key(user_id="u-stale", team_id="t-stale") + await invalidate_team_member_spend_state(user_id="u-stale", team_id="t-stale", user_api_key_cache=cache) + after_invalidate = await cache.async_get_cache(key=cache_key, local_only=True) + + hang_redis_set.set() + await asyncio.sleep(0) + await asyncio.sleep(0) + after_replicate = await cache.async_get_cache(key=cache_key, local_only=True) + + assert loaded is not None + assert after_invalidate is None + assert after_replicate is None + + @pytest.mark.asyncio async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sentinel(): """