diff --git a/litellm/constants.py b/litellm/constants.py index 5c7a02d0743..5751e6e46af 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2039,10 +2039,3 @@ 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 89b061bdefe..a8f9cc77577 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -35,8 +35,6 @@ 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 @@ -331,27 +329,19 @@ db_cache_expiry: Final = DEFAULT_IN_MEMORY_TTL # refresh every 5s _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) + + +class _TeamMembershipCacheMiss: + __slots__ = () + + +_TEAM_MEMBERSHIP_CACHE_MISS: Final = _TeamMembershipCacheMiss() 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", - ) + return result if isinstance(result, LiteLLM_TeamMembership) else None def _log_budget_lookup_failure(entity: str, error: Exception) -> None: @@ -897,21 +887,6 @@ async def common_checks( """ from litellm.proxy.proxy_server import prisma_client, user_api_key_cache - # One membership read for model-access, access-group attribution, and - # member-budget. Each used to call get_team_membership independently; - # DualCache Redis SET/GET on every miss made those look like two Postgres spans. - loaded_team_membership: LiteLLM_TeamMembership | None = None - team_membership_loaded = False - if team_object is not None and valid_token is not None and valid_token.user_id is not None: - loaded_team_membership = await get_team_membership( - user_id=valid_token.user_id, - team_id=team_object.team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - team_membership_loaded = True - _model: Final[str | list[str] | None] = get_model_from_request( request_data=request_body, route=route, @@ -926,6 +901,22 @@ async def common_checks( and (route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route)) ) + membership_user_id: Final = ( + valid_token.user_id if valid_token is not None and (bool(_model) or not skip_all_budget_checks) else None + ) + team_membership_loaded: Final = team_object is not None and membership_user_id is not None + loaded_team_membership: Final = ( + await get_team_membership( + user_id=membership_user_id, + team_id=team_object.team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if team_object is not None and membership_user_id is not None + else None + ) + unpriced_models: Final = ( _unpriced_models_in_request(model=_model, llm_router=llm_router) if litellm.block_requests_for_models_without_pricing and RouteChecks.is_llm_api_route(route=route) @@ -2188,90 +2179,13 @@ async def get_tag_object( def _membership_from_cached_payload( cached: object, -) -> LiteLLM_TeamMembership | None | TeamMembershipCacheMiss: +) -> 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 - - -async def _set_team_membership_l1( - user_api_key_cache: UserApiKeyCache, - key: str, - value: object, - *, - 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=True) - case (False, True): - await user_api_key_cache.async_set_cache(key=key, value=value, local_only=True, ttl=ttl) - case (True, False): - await user_api_key_cache.async_set_cache(key=key, value=value, local_only=True, model_type=model_type) - case (True, True): - await user_api_key_cache.async_set_cache( - key=key, value=value, local_only=True, model_type=model_type, ttl=ttl - ) - - -async def _replicate_team_membership_to_redis( - user_api_key_cache: UserApiKeyCache, - key: str, - value: object, - *, - model_type: type[LiteLLM_TeamMembership] | None, - ttl: float | None, - write_epoch: int, -) -> None: - redis_cache: Final = user_api_key_cache.redis_cache - if redis_cache is None or _membership_write_epoch(key) != write_epoch: - return - payload: Final[object] = CacheCodec.serialize(value, model_type=model_type) - try: - if ttl is None: - await redis_cache.async_set_cache(key, payload) - else: - await redis_cache.async_set_cache(key, payload, ttl=ttl) - if _membership_write_epoch(key) != write_epoch: - await redis_cache.async_delete_cache(key) - except Exception: - return - - -async def _populate_team_membership_cache( - user_api_key_cache: UserApiKeyCache, - key: str, - value: object, - *, - model_type: type[LiteLLM_TeamMembership] | None = None, - ttl: float | None = None, -) -> None: - write_epoch: Final = _membership_write_epoch(key) - await _set_team_membership_l1( - user_api_key_cache, - key, - value, - model_type=model_type, - ttl=ttl, - ) - if _membership_write_epoch(key) != write_epoch: - user_api_key_cache.in_memory_cache_for(key).delete_cache(key) - return - - asyncio.create_task( - _replicate_team_membership_to_redis( - user_api_key_cache, - key, - value, - model_type=model_type, - ttl=ttl, - write_epoch=write_epoch, - ) - ) + return cached_membership if cached_membership is not None else _TEAM_MEMBERSHIP_CACHE_MISS @log_db_metrics @@ -2283,7 +2197,7 @@ async def _fetch_team_membership_from_db( parent_otel_span: Span | None = None, proxy_logging_obj: ProxyLogging | None = None, ) -> LiteLLM_TeamMembership | None: - """Prisma read + L1 populate. Decorated so cache hits on ``get_team_membership`` are not postgres spans.""" + """Prisma read + cache populate. Decorated so cache hits on ``get_team_membership`` are not postgres spans.""" _ = parent_otel_span, proxy_logging_obj response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique( where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, @@ -2291,19 +2205,17 @@ async def _fetch_team_membership_from_db( ) _key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id) if response is None: - await _populate_team_membership_cache( - user_api_key_cache, - _key, - NO_TEAM_MEMBERSHIP_SENTINEL, + await user_api_key_cache.async_set_cache( + key=_key, + value=NO_TEAM_MEMBERSHIP_SENTINEL, ttl=get_management_object_ttl(user_api_key_cache), ) return None membership: Final = LiteLLM_TeamMembership.model_validate(response.dict()) - await _populate_team_membership_cache( - user_api_key_cache, - _key, - membership, + await user_api_key_cache.async_set_cache( + key=_key, + value=membership, model_type=LiteLLM_TeamMembership, ) return membership @@ -2321,7 +2233,7 @@ 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 not isinstance(redis_membership, TeamMembershipCacheMiss): + if not isinstance(redis_membership, _TeamMembershipCacheMiss): return redis_membership return await _fetch_team_membership_from_db( @@ -2332,18 +2244,9 @@ async def _load_team_membership_on_cache_miss( parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) - 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, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Failed to load team membership", - ) from e + except Exception: + verbose_proxy_logger.exception("Error getting team membership") + return None async def get_team_membership( @@ -2366,7 +2269,7 @@ 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 not isinstance(l1_membership, TeamMembershipCacheMiss): + if not isinstance(l1_membership, _TeamMembershipCacheMiss): return l1_membership inflight: Final[object] = _team_membership_inflight.get(_key) @@ -2375,9 +2278,6 @@ async def get_team_membership( if prisma_client is None: raise Exception("No db connected") - prisma: Final[object] = prisma_client - if isinstance(prisma, str): - return None task: Final = asyncio.ensure_future( _load_team_membership_on_cache_miss( @@ -2908,7 +2808,6 @@ 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 29f6c3861f6..04e41608b26 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -6494,52 +6494,6 @@ 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_returns_before_redis_set_completes(): - """Auth must not wait on DualCache Redis SET; L1 is enough for the next lookup.""" - from litellm.proxy.auth.auth_checks import get_team_membership - - membership_row = MagicMock() - membership_row.dict = lambda: {"user_id": "u-redis", "team_id": "t-redis", "spend": 2.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) - - cache = UserApiKeyCache(redis_cache=redis_cache) - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) - - first = await asyncio.wait_for( - get_team_membership( - user_id="u-redis", - team_id="t-redis", - prisma_client=mock_prisma_client, - user_api_key_cache=cache, - ), - timeout=0.5, - ) - second = await get_team_membership( - user_id="u-redis", - team_id="t-redis", - prisma_client=mock_prisma_client, - user_api_key_cache=cache, - ) - - hang_redis_set.set() - await asyncio.sleep(0) - - assert first is not None and second is not None - assert first.spend == 2.0 - assert second.spend == 2.0 - mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once() - - @pytest.mark.asyncio async def test_common_checks_calls_get_team_membership_once_per_request(): """Model-access, attribution, and member-budget must reuse one membership load.""" @@ -6588,33 +6542,85 @@ async def test_common_checks_calls_get_team_membership_once_per_request(): @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 +async def test_common_checks_skips_membership_load_when_no_check_reads_it(): + """A management route has no model and no budget gate, so the membership row is never loaded.""" + from fastapi import Request + from litellm.proxy.auth.auth_checks import common_checks + + team = LiteLLM_TeamTable(team_id="t-lazy") + token = UserAPIKeyAuth(token="k-lazy", user_id="u-lazy", team_id="t-lazy") + + with ( + patch( # test-quality-ok: common_checks imports prisma_client from proxy_server + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch( # test-quality-ok: common_checks imports user_api_key_cache from proxy_server + "litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache() + ), + patch( # test-quality-ok: counts membership loads; common_checks has no membership seam + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + ) as load_membership, + ): + result = await common_checks( + request_body={}, + team_object=team, + user_object=LiteLLM_UserTable(user_id="u-lazy"), + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route="/key/info", + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=MagicMock(spec=Request), + ) + + assert result is True + load_membership.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_get_team_membership_db_error_returns_none_and_retries_next_call(): + """A Prisma failure reads as no membership, caches nothing, and the next call hits the DB again.""" 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 + membership_row = MagicMock() + membership_row.dict = lambda: {"user_id": "u-fail", "team_id": "t-fail", "spend": 1.0} mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=RuntimeError("db down")) + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock( + side_effect=[RuntimeError("db down"), membership_row] + ) 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, - ) + failed = await get_team_membership( + user_id="u-fail", + team_id="t-fail", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) + cached_after_failure = await cache.async_get_cache( + key=team_membership_reservation_cache_key(user_id="u-fail", team_id="t-fail") + ) + recovered = 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 + assert failed is None + assert cached_after_failure is None + assert recovered is not None + assert recovered.user_id == "u-fail" + assert mock_prisma_client.db.litellm_teammembership.find_unique.await_count == 2 @pytest.mark.asyncio async def test_get_team_membership_string_prisma_client_returns_none(): - """Unit tests stub prisma_client as a string; that is not a lookup failure and must not 503.""" + """Unit tests stub prisma_client as a string; the lookup fails and reads as no membership.""" from litellm.proxy.auth.auth_checks import get_team_membership result = await get_team_membership( @@ -6626,59 +6632,6 @@ async def test_get_team_membership_string_prisma_client_returns_none(): assert result 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( # test-quality-ok: common_checks imports prisma_client from proxy_server - "litellm.proxy.proxy_server.prisma_client", MagicMock() - ), - patch( # test-quality-ok: common_checks imports user_api_key_cache from proxy_server - "litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache() - ), - patch( # test-quality-ok: injects membership lookup failure; common_checks has no seam - "litellm.proxy.auth.auth_checks.get_team_membership", - new_callable=AsyncMock, - side_effect=lookup_error, - ), - patch( # test-quality-ok: common_checks imports get_current_spend locally - "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.""" @@ -6721,51 +6674,6 @@ async def test_get_team_membership_waiter_cancel_does_not_cancel_shared_load(): 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 Redis SET must not rewrite L1 or leave Redis holding 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 - redis_cache.async_delete_cache.assert_awaited() - - @pytest.mark.asyncio async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sentinel(): """