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 20:50:28 +05:30 committed by yassin
parent 6e9e08475a
commit 5457f48290
2 changed files with 369 additions and 71 deletions

View file

@ -327,6 +327,9 @@ _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]] = {}
all_routes: Final = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
@ -873,6 +876,21 @@ 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,
@ -936,6 +954,8 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
)
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
@ -987,6 +1007,8 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
)
# Run before apply_key_tags_pre_auth injects key metadata.tags into request_body.
@ -1096,6 +1118,8 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
),
_check_end_user_budget(end_user_obj=end_user_object, route=route)
if end_user_object is not None and end_user_object.litellm_budget_table is not None
@ -2141,7 +2165,142 @@ 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."""
if cached is None:
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 _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:
"""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)
async def _replicate_to_redis() -> None:
try:
await user_api_key_cache.async_set_cache(key=key, value=value, **kwargs)
except Exception:
return
asyncio.create_task(_replicate_to_redis())
@log_db_metrics
async def _fetch_team_membership_from_db(
user_id: str,
team_id: str,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
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."""
_ = 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}},
include={"litellm_budget_table": True},
)
_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,
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,
model_type=LiteLLM_TeamMembership,
)
return membership
async def _load_team_membership_on_cache_miss(
user_id: str,
team_id: str,
cache_key: str,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging | None,
) -> LiteLLM_TeamMembership | None:
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)
# #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,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception:
verbose_proxy_logger.exception(
"Error getting team membership for user_id: %s, team_id: %s",
user_id,
team_id,
)
return None
async def get_team_membership(
user_id: str,
team_id: str,
@ -2155,54 +2314,43 @@ async def get_team_membership(
Do a isolated check for team membership vs. doing a combined key + team + user + team-membership check, as key might come in frequently for different users/teams. Larger call will slowdown query time. This way we get to cache the constant (key/team/user info) and only update based on the changing value (team membership).
"""
from litellm.proxy._types import LiteLLM_TeamMembership
if prisma_client is None:
raise Exception("No db connected")
if user_id is None or team_id is None:
return None
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
# check if in cache
cached: Final[object] = await user_api_key_cache.async_get_cache(key=_key)
if cached == NO_TEAM_MEMBERSHIP_SENTINEL:
return None
cached_membership_obj: Final = CacheCodec.deserialize(cached, model_type=LiteLLM_TeamMembership)
if cached_membership_obj is not None:
return cached_membership_obj
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)
# else, check db
try:
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
include={"litellm_budget_table": True},
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
if prisma_client is None:
raise Exception("No db connected")
task: Final = asyncio.ensure_future(
_load_team_membership_on_cache_miss(
user_id=user_id,
team_id=team_id,
cache_key=_key,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if response is None:
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
_response: Final = LiteLLM_TeamMembership.model_validate(response.dict())
await user_api_key_cache.async_set_cache(
key=_key,
value=_response,
model_type=LiteLLM_TeamMembership,
)
return _response
except Exception:
verbose_proxy_logger.exception(
"Error getting team membership for user_id: %s, team_id: %s",
user_id,
team_id,
)
return None
)
_team_membership_inflight[_key] = task
task.add_done_callback(lambda _t, k=_key: _team_membership_inflight.pop(k, None))
return await task
def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool:
@ -4122,18 +4270,21 @@ async def _team_member_granted_models(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> Sequence[str]:
"""The member's own ``allowed_models`` scope; empty when the member is not narrowed below the team."""
if team_object is None or valid_token.user_id is None:
return ()
team_membership: Final = 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,
)
if not team_membership_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,
)
return () if team_membership is None else _member_allowed_models(team_membership)
@ -4169,6 +4320,8 @@ async def _granted_model_lists(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[Sequence[str], ...]:
"""One model allowlist per level that participates in authorizing the request."""
return (
@ -4180,6 +4333,8 @@ async def _granted_model_lists(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
),
project_object.models if project_object is not None else (),
await _org_granted_models(
@ -4274,6 +4429,8 @@ async def collect_matched_model_access_groups(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[str, ...]:
"""
The budgeted model access groups that authorized this request, sorted and deduplicated.
@ -4319,6 +4476,8 @@ async def collect_matched_model_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
)
for granted_model in granted_models
)
@ -4334,6 +4493,8 @@ async def stamp_matched_model_access_groups(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> tuple[str, ...]:
"""Record the groups that authorized this request on its auth object, for the post-call spend
writer and the reservation counters, and hand them back for the budget check."""
@ -4350,6 +4511,8 @@ async def stamp_matched_model_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
team_membership=team_membership,
team_membership_loaded=team_membership_loaded,
)
except Exception as e: # noqa: BLE001 # fail-safe: attribution is spend telemetry, it must never break auth
verbose_proxy_logger.debug("model access group attribution failed: %s", e)
@ -5152,6 +5315,8 @@ async def _check_team_member_budget(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
):
"""Check if team member is over their max budget within the team."""
if (
@ -5160,23 +5325,25 @@ async def _check_team_member_budget(
and valid_token is not None
and valid_token.user_id is not None
):
team_membership: Final = 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,
)
if not team_membership_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,
)
loaded_membership = team_membership
# Per-member override wins; otherwise fall back to the team-level
# default configured via team.metadata["team_member_budget_id"].
team_member_budget: float | None = None
if (
team_membership is not None
and team_membership.litellm_budget_table is not None
and team_membership.litellm_budget_table.max_budget is not None
loaded_membership is not None
and loaded_membership.litellm_budget_table is not None
and loaded_membership.litellm_budget_table.max_budget is not None
):
team_member_budget = team_membership.litellm_budget_table.max_budget
team_member_budget = loaded_membership.litellm_budget_table.max_budget
else:
default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id")
if isinstance(default_budget_id, str):
@ -5195,7 +5362,7 @@ async def _check_team_member_budget(
team_member_budget = default_budget.max_budget
if team_member_budget is not None:
team_member_spend = (team_membership.spend if team_membership is not None else 0.0) or 0.0
team_member_spend = (loaded_membership.spend if loaded_membership is not None else 0.0) or 0.0
# Read from cross-pod counter (Redis-first) if available
from litellm.proxy.proxy_server import get_current_spend
@ -5224,6 +5391,8 @@ async def _check_team_member_model_access(
prisma_client: Optional["PrismaClient"],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
) -> None:
"""
Check if a team member's per-member model scope allows access to the requested model.
@ -5234,22 +5403,24 @@ async def _check_team_member_model_access(
if valid_token.user_id is None or team_object.team_id is None:
return
team_membership: Final = 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,
)
if not team_membership_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,
)
loaded_membership = team_membership
if (
team_membership is None
or team_membership.litellm_budget_table is None
or not team_membership.litellm_budget_table.allowed_models
loaded_membership is None
or loaded_membership.litellm_budget_table is None
or not loaded_membership.litellm_budget_table.allowed_models
):
return # no per-member restriction — inherit team-level check
member_allowed_models: Final[list[str]] = team_membership.litellm_budget_table.allowed_models
member_allowed_models: Final[list[str]] = loaded_membership.litellm_budget_table.allowed_models
try:
_can_object_call_model(
model=model,

View file

@ -6454,6 +6454,133 @@ async def test_get_team_membership_reads_sentinel_as_no_membership_not_a_model()
mock_prisma_client.db.litellm_teammembership.find_unique.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_team_membership_coalesces_parallel_db_fetches():
"""Concurrent misses for the same member must share one Prisma round-trip."""
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-parallel", "team_id": "t-parallel", "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-parallel",
team_id="t-parallel",
prisma_client=mock_prisma_client,
user_api_key_cache=cache,
)
first = asyncio.create_task(_load())
second = asyncio.create_task(_load())
await started.wait()
await asyncio.sleep(0)
release.set()
results = await asyncio.gather(first, second)
assert results[0] is not None and results[1] is not None
assert results[0].user_id == "u-parallel"
assert results[1].user_id == "u-parallel"
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."""
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
team = LiteLLM_TeamTable(team_id="t-once")
token = UserAPIKeyAuth(token="k-once", user_id="u-once", team_id="t-once", models=["gpt-4o-mini"])
membership = MagicMock()
membership.litellm_budget_table = None
membership.spend = 0.0
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,
return_value=membership,
) as load_membership,
patch("litellm.proxy.proxy_server.get_current_spend", new_callable=AsyncMock, return_value=0.0),
):
result = await common_checks(
request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]},
team_object=team,
user_object=LiteLLM_UserTable(user_id="u-once"),
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 result is True
assert load_membership.await_count == 1
@pytest.mark.asyncio
async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sentinel():
"""