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
This commit is contained in:
Shivi Jain 2026-09-14 21:28:08 +05:30 committed by yassin
parent 5457f48290
commit 70ddc7e492
3 changed files with 265 additions and 65 deletions

View file

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

View file

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

View file

@ -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():
"""