mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
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:
parent
5457f48290
commit
70ddc7e492
3 changed files with 265 additions and 65 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue