fix(auth): fail closed when the team membership lookup hits a db outage

This commit is contained in:
mateo-berri 2026-09-19 16:04:13 -07:00
parent d675c1285b
commit 484524b70b
3 changed files with 90 additions and 34 deletions

View file

@ -2296,23 +2296,19 @@ async def _load_team_membership_on_cache_miss(
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 not isinstance(redis_membership, _TeamMembershipCacheMiss):
return redis_membership
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):
return redis_membership
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")
return None
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,
)
async def get_team_membership(

View file

@ -7198,7 +7198,7 @@ async def test_common_checks_skips_membership_load_when_no_check_reads_it():
@pytest.mark.asyncio
async def test_get_team_membership_db_error_returns_none_and_retries_next_call():
async def test_get_team_membership_db_error_surfaces_and_retries_next_call():
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
@ -7210,12 +7210,13 @@ async def test_get_team_membership_db_error_returns_none_and_retries_next_call()
)
cache = UserApiKeyCache()
failed = await get_team_membership(
user_id="u-fail",
team_id="t-fail",
prisma_client=mock_prisma_client,
user_api_key_cache=cache,
)
with pytest.raises(RuntimeError, match="db down"):
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")
)
@ -7226,24 +7227,55 @@ async def test_get_team_membership_db_error_returns_none_and_retries_next_call()
user_api_key_cache=cache,
)
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():
from litellm.proxy.auth.auth_checks import get_team_membership
class _UnreachableMembershipPrisma:
class db:
class litellm_teammembership:
@staticmethod
async def find_unique(where: dict[str, dict[str, str]], include: dict[str, bool]) -> None:
raise httpx.ConnectError("All connection attempts failed")
result = await get_team_membership(
user_id="u-str",
team_id="t-str",
prisma_client="hello-world",
user_api_key_cache=UserApiKeyCache(),
)
assert result is None
def _restricted_member_check_deps() -> dict[str, object]:
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import ProxyLogging
cache = UserApiKeyCache()
return {
"team_object": LiteLLM_TeamTable(team_id="team-outage", models=["claude-sonnet-5"]),
"valid_token": UserAPIKeyAuth(token="hashed-fake", user_id="bob", team_id="team-outage"),
"prisma_client": _UnreachableMembershipPrisma(),
"user_api_key_cache": cache,
"proxy_logging_obj": ProxyLogging(user_api_key_cache=cache),
}
@pytest.mark.asyncio
async def test_check_team_member_model_access_fails_closed_when_the_membership_read_hits_a_db_outage():
"""Regression: with the member's row uncached and the database unreachable, the loader used to swallow the
transport error and return None, which every check reads as "no per-member restriction", so a member
limited to other models got a 200. The outage must surface as the 503 the rest of auth answers with."""
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
from litellm.proxy.auth.auth_exception_handler import _as_proxy_exception
with pytest.raises(httpx.ConnectError) as raised:
await _check_team_member_model_access(
model="claude-sonnet-5", llm_router=None, **_restricted_member_check_deps()
)
surfaced = _as_proxy_exception(raised.value)
assert (surfaced.code, surfaced.type) == ("503", ProxyErrorTypes.no_db_connection)
@pytest.mark.asyncio
async def test_check_team_member_budget_fails_closed_when_the_membership_read_hits_a_db_outage():
with pytest.raises(httpx.ConnectError):
await _check_team_member_budget(user_object=None, **_restricted_member_check_deps())
@pytest.mark.asyncio

View file

@ -1,4 +1,5 @@
from fastapi import HTTPException
import httpx
import pytest
from litellm.proxy._types import (
@ -8,6 +9,7 @@ from litellm.proxy._types import (
ProxyException,
)
from litellm.proxy.auth.auth_checks import TeamNotFoundError, UserNotFoundError
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.auth.resolvers.grants import (
GrantResolver,
LookupDegraded,
@ -172,6 +174,32 @@ async def test_resolve_identity_lets_loader_errors_surface():
await loaders.resolver().resolve_identity(UserLookup(user_id=USER_ID), team_id=None)
class _UnreachableMembershipPrisma:
class db:
class litellm_teammembership:
@staticmethod
async def find_unique(where: dict[str, dict[str, str]], include: dict[str, bool]) -> None:
raise httpx.ConnectError("All connection attempts failed")
async def test_resolve_marks_a_membership_read_that_hits_a_db_outage_as_degraded():
"""Regression: the real membership loader swallowed a database transport error into None, so this outcome
was ResolvedGrants with no membership, never LookupDegraded, and a member's own model or budget limits
silently dropped for the request."""
loaders = _Loaders(user=_user(), team=_team())
resolver = GrantResolver(
_UnreachableMembershipPrisma(),
UserApiKeyCache(),
load_user=loaders.load_user,
load_team=loaders.load_team,
)
outcome = await resolver.resolve(UserLookup(user_id=USER_ID), team_id=TEAM_ID)
assert isinstance(outcome, LookupDegraded)
assert isinstance(outcome.error, httpx.ConnectError)
def test_raise_public_maps_a_deleted_user_to_401():
with pytest.raises(ProxyException) as exc_info:
raise_public(UserGone(user_id=USER_ID))