fix(auth): drop the membership write-epoch and background Redis replicate, load membership lazily

Move the cache-miss marker out of litellm.constants into auth_checks (CodeQL cyclic import) and stop
logging user_id/team_id in the lookup failure (CodeQL log injection).

Write the membership row through DualCache synchronously again instead of a background Redis task
guarded by a bounded write-epoch map: the epoch was sampled after the Prisma read, so an invalidate
that raced the read could be cached as current, and eviction of the epoch entry could let an old
Redis write land. The synchronous write keeps invalidate_team_member_spend_state authoritative.

Lookup failures return None again (fail-open like main) instead of 503, and the load is skipped on
routes that neither resolve a model nor run budget checks.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-14 19:06:29 +00:00
parent f4e21430c0
commit b00bb35563
3 changed files with 106 additions and 306 deletions

View file

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

View file

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

View file

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