diff --git a/litellm/constants.py b/litellm/constants.py index 3ff80d8b7dd..7b40f432446 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1772,6 +1772,8 @@ RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS: Final = max(1, int(os.getenv("RESPONSES_S RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL: Final = float(os.getenv("RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL", "0.2")) SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)) PROXY_DB_LOOKUP_MAX_CONCURRENCY: Final = max(1, int(os.getenv("PROXY_DB_LOOKUP_MAX_CONCURRENCY", "25"))) +PROXY_DB_LOOKUP_DEADLINE_SECONDS: Final = max(0.1, float(os.getenv("PROXY_DB_LOOKUP_DEADLINE_SECONDS", "10"))) +PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS: Final = max(0.0, float(os.getenv("PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS", "30"))) DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597)) RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500"))) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index f5e98b40ca7..67950e603c0 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -15,11 +15,11 @@ import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeAlias +from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias from fastapi import HTTPException, Request, status from pydantic import BaseModel, TypeAdapter -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack import litellm from litellm._logging import verbose_proxy_logger @@ -110,7 +110,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( team_membership_auth_cache_key, team_membership_reservation_cache_key, ) -from litellm.proxy.db.db_lookup_gate import db_lookup_gate +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup, db_lookup_gate from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( TOOL_CAPABLE_CALL_TYPES, @@ -223,30 +223,79 @@ class _PrismaTableHolder(Protocol[RowT_co]): def table(self) -> _PrismaAuthTable[RowT_co]: ... -def _dictable_table(repo: _PrismaTableHolder[_PrismaDictableRow]) -> _PrismaAuthTable[_PrismaDictableRow]: - return repo.table +class _FindOneKwargs(TypedDict): + where: ReadOnly[Required[Mapping[str, object]]] + include: ReadOnly[NotRequired[Mapping[str, object] | None]] + + +class _FindManyKwargs(TypedDict): + where: ReadOnly[NotRequired[Mapping[str, object] | None]] + include: ReadOnly[NotRequired[Mapping[str, object] | None]] + take: ReadOnly[NotRequired[int | None]] + + +class _DeadlineBoundedTable(Generic[RowT_co]): + """Every read on the wrapped table fails with ``DBLookupDeadlineExceeded`` once + ``PROXY_DB_LOOKUP_DEADLINE_SECONDS`` passes, so a stalled database fails the + request fast instead of parking it in the pod until it fills its memory.""" + + __slots__ = ("_lookup", "_table") + + def __init__(self, table: _PrismaAuthTable[RowT_co], lookup: str) -> None: + self._table: Final = table + self._lookup: Final = lookup + + async def find_unique( + self, + **kwargs: Unpack[_FindOneKwargs], # kwargs-ok: typed pass-through that forwards exactly what the caller passed + ) -> RowT_co | None: + return await bounded_db_lookup(self._table.find_unique(**kwargs), name=self._lookup) + + async def find_first( + self, + **kwargs: Unpack[_FindOneKwargs], # kwargs-ok: typed pass-through that forwards exactly what the caller passed + ) -> RowT_co | None: + return await bounded_db_lookup(self._table.find_first(**kwargs), name=self._lookup) + + async def find_many( + self, + **kwargs: Unpack[_FindManyKwargs], # kwargs-ok: typed pass-through that forwards exactly what the caller passed + ) -> Sequence[RowT_co]: + return await bounded_db_lookup(self._table.find_many(**kwargs), name=self._lookup) + + async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> RowT_co | None: + return await self._table.update(where=where, data=data) + + async def create(self, *, data: Mapping[str, object], include: Mapping[str, object] | None = None) -> RowT_co: + return await self._table.create(data=data, include=include) + + +def _dictable_table(repo: _PrismaTableHolder[_PrismaDictableRow], lookup: str) -> _PrismaAuthTable[_PrismaDictableRow]: + return _DeadlineBoundedTable(repo.table, lookup) def _jwt_key_mapping_table( repo: _PrismaTableHolder[_PrismaJWTKeyMappingRow], ) -> _PrismaAuthTable[_PrismaJWTKeyMappingRow]: - return repo.table + return _DeadlineBoundedTable(repo.table, "jwt_key_mapping") -def _model_dump_table(repo: _PrismaTableHolder[_PrismaModelDumpRow]) -> _PrismaAuthTable[_PrismaModelDumpRow]: - return repo.table +def _model_dump_table( + repo: _PrismaTableHolder[_PrismaModelDumpRow], lookup: str +) -> _PrismaAuthTable[_PrismaModelDumpRow]: + return _DeadlineBoundedTable(repo.table, lookup) def _team_table(repo: _PrismaTableHolder[_PrismaTeamRow]) -> _PrismaAuthTable[_PrismaTeamRow]: - return repo.table + return _DeadlineBoundedTable(repo.table, "team") def _vector_store_table(repo: _PrismaTableHolder[_PrismaVectorStoreRow]) -> _PrismaAuthTable[_PrismaVectorStoreRow]: - return repo.table + return _DeadlineBoundedTable(repo.table, "vector_store") def _user_table(repo: _PrismaTableHolder[_PrismaUserRow]) -> _PrismaAuthTable[_PrismaUserRow]: - return repo.table + return _DeadlineBoundedTable(repo.table, "user") class _VectorStorePermissionsRow(Protocol): @@ -257,7 +306,7 @@ class _VectorStorePermissionsRow(Protocol): def _object_permission_table( repo: _PrismaTableHolder[_VectorStorePermissionsRow], ) -> _PrismaAuthTable[_VectorStorePermissionsRow]: - return repo.table + return _DeadlineBoundedTable(repo.table, "object_permission") class _PrismaTagRow(Protocol): @@ -1422,7 +1471,7 @@ async def get_default_end_user_budget( # Fetch from database try: - budget_record: Final = await _dictable_table(BudgetRepository(prisma_client)).find_unique( + budget_record: Final = await _dictable_table(BudgetRepository(prisma_client), "budget").find_unique( where={"budget_id": default_budget_id} # mutable-ok: prisma where clause ) @@ -1483,7 +1532,7 @@ async def get_team_member_default_budget( return cached_budget try: - budget_record: Final = await _dictable_table(BudgetRepository(prisma_client)).find_unique( + budget_record: Final = await _dictable_table(BudgetRepository(prisma_client), "budget").find_unique( where={"budget_id": budget_id} ) except Exception: @@ -1877,7 +1926,7 @@ async def get_end_user_object( # Fetch from database try: - response: Final = await _dictable_table(EndUserRepository(prisma_client)).find_unique( + response: Final = await _dictable_table(EndUserRepository(prisma_client), "end_user").find_unique( where={"user_id": end_user_id}, include={"litellm_budget_table": True, "object_permission": True}, ) @@ -2286,7 +2335,7 @@ async def _fetch_team_membership_from_db( proxy_logging_obj: ProxyLogging | None = None, ) -> LiteLLM_TeamMembership | None: _ = parent_otel_span, proxy_logging_obj - response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique( + response: Final = await _dictable_table(TeamMembershipRepository(prisma_client), "team_membership").find_unique( where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, include={"litellm_budget_table": True}, ) @@ -3290,7 +3339,7 @@ async def get_access_object( # Not in cache - fetch from DB try: - response: Final = await _dictable_table(AccessGroupRepository(prisma_client)).find_unique( + response: Final = await _dictable_table(AccessGroupRepository(prisma_client), "access_group").find_unique( where={"access_group_id": access_group_id} ) @@ -3472,7 +3521,7 @@ async def get_org_object_by_alias( # Query database by organization_alias try: - orgs = await _model_dump_table(OrganizationRepository(prisma_client)).find_many( + orgs = await _model_dump_table(OrganizationRepository(prisma_client), "organization").find_many( where={"organization_alias": org_alias} ) @@ -3650,10 +3699,32 @@ async def _fetch_key_object_from_db_with_reconnect( prisma_client: PrismaClient, parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging | None, + deadline_seconds: float | None = None, ) -> BaseModel | None: """ Fetch key object from DB and retry once if a DB connection error can be healed. + The gate wait, the query, the reconnect, and the retry share one deadline, so a + stalled database fails the request with ``DBLookupDeadlineExceeded`` instead of + parking it. """ + return await bounded_db_lookup( + _fetch_key_object_from_db_unbounded( + hashed_token=hashed_token, + prisma_client=prisma_client, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ), + name="key", + deadline_seconds=deadline_seconds, + ) + + +async def _fetch_key_object_from_db_unbounded( + hashed_token: str, + prisma_client: PrismaClient, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging | None, +) -> BaseModel | None: async with db_lookup_gate.current(): try: return await prisma_client.get_data( @@ -3874,9 +3945,9 @@ async def get_object_permission( # else, check db try: - response: Final = await _dictable_table(ObjectPermissionRepository(prisma_client)).find_unique( - where={"object_permission_id": object_permission_id} - ) + response: Final = await _dictable_table( + ObjectPermissionRepository(prisma_client), "object_permission" + ).find_unique(where={"object_permission_id": object_permission_id}) if response is None: return None @@ -4008,7 +4079,9 @@ async def get_org_object( if include_budget_table: query_kwargs["include"] = {"litellm_budget_table": True} - response: Final = await _model_dump_table(OrganizationRepository(prisma_client)).find_unique(**query_kwargs) + response: Final = await _model_dump_table(OrganizationRepository(prisma_client), "organization").find_unique( + **query_kwargs + ) except Exception: # An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed # missing row, and relabelling it as "doesn't exist" made every caller unable to tell them @@ -5948,7 +6021,7 @@ async def get_project_object( return deserialized_project # Fetch from DB - project_row: Final = await _model_dump_table(ProjectRepository(prisma_client)).find_unique( + project_row: Final = await _model_dump_table(ProjectRepository(prisma_client), "project").find_unique( where={"project_id": project_id}, include={"litellm_budget_table": True}, ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a6c0792a86f..ae95e94dd2d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -120,6 +120,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, team_membership_auth_cache_key, ) +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.spend_tracking.carried_budget_state import carry_team_and_user_budget_state @@ -735,8 +736,9 @@ async def _fetch_global_spend_with_event_coordination( """ async def _load_global_spend() -> float | None: - proxy_budget_row: Final = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": LITELLM_PROXY_BUDGET_NAME} + proxy_budget_row: Final = await bounded_db_lookup( + prisma_client.db.litellm_usertable.find_unique(where={"user_id": LITELLM_PROXY_BUDGET_NAME}), + name="proxy_budget", ) return float(proxy_budget_row.spend) if proxy_budget_row is not None else None diff --git a/litellm/proxy/db/db_lookup_gate.py b/litellm/proxy/db/db_lookup_gate.py index 2fd427687bd..d2aefde929f 100644 --- a/litellm/proxy/db/db_lookup_gate.py +++ b/litellm/proxy/db/db_lookup_gate.py @@ -1,7 +1,14 @@ import asyncio -from typing import Final +import time +from collections.abc import Awaitable, Callable +from typing import Final, TypeVar -from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY +from litellm.constants import ( + PROXY_DB_LOOKUP_DEADLINE_SECONDS, + PROXY_DB_LOOKUP_MAX_CONCURRENCY, +) + +LookupT = TypeVar("LookupT") class LoopBoundSemaphore: @@ -20,4 +27,64 @@ class LoopBoundSemaphore: return self._semaphore +class DBLookupDeadlineExceeded(asyncio.TimeoutError): + def __init__(self, lookup: str, deadline_seconds: float) -> None: + super().__init__(f"{lookup} lookup did not answer within {deadline_seconds:g}s") + self.lookup: Final = lookup + self.deadline_seconds: Final = deadline_seconds + + +class DBLookupStallTracker: + __slots__ = ("_clock", "_last_hit") + + def __init__(self, clock: Callable[[], float] = time.monotonic) -> None: + self._clock: Final = clock + self._last_hit: float | None = None + + def record_hit(self) -> None: + self._last_hit = self._clock() + + def clear(self) -> None: + self._last_hit = None + + def stalled_within(self, window_seconds: float) -> bool: + if self._last_hit is None: + return False + return self._clock() - self._last_hit < window_seconds + + db_lookup_gate: Final = LoopBoundSemaphore(PROXY_DB_LOOKUP_MAX_CONCURRENCY) +db_lookup_stall_tracker: Final = DBLookupStallTracker() + + +def _consume_abandoned_lookup(task: asyncio.Future[LookupT]) -> None: + if not task.cancelled(): + task.exception() + + +async def bounded_db_lookup( + lookup: Awaitable[LookupT], + *, + name: str, + deadline_seconds: float | None = None, + tracker: DBLookupStallTracker = db_lookup_stall_tracker, +) -> LookupT: + timeout: Final = PROXY_DB_LOOKUP_DEADLINE_SECONDS if deadline_seconds is None else deadline_seconds + task: Final = asyncio.ensure_future(lookup) + try: + done, _ = await asyncio.wait({task}, timeout=timeout) + except asyncio.CancelledError: + task.cancel() + raise + if task not in done: + task.cancel() + task.add_done_callback(_consume_abandoned_lookup) + tracker.record_hit() + raise DBLookupDeadlineExceeded(name, timeout) + try: + return task.result() + except DBLookupDeadlineExceeded: + raise + except asyncio.TimeoutError as e: + tracker.record_hit() + raise DBLookupDeadlineExceeded(name, timeout) from e diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index 38d6fb9b99d..022ca2efc9e 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -11,6 +11,7 @@ from litellm.proxy._types import ( ProxyErrorTypes, ProxyException, ) +from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded from litellm.secret_managers.main import str_to_bool # Bounds the __cause__/__context__ walk in find_database_service_unavailable_error_in_chain. @@ -104,7 +105,7 @@ class PrismaDBExceptionHandler: """ import prisma.engine.errors - if isinstance(e, DB_CONNECTION_ERROR_TYPES): + if isinstance(e, (*DB_CONNECTION_ERROR_TYPES, DBLookupDeadlineExceeded)): return True if isinstance(e, _exception_types(prisma.engine.errors.EngineConnectionError)): return True diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 2dd028454d6..f8e102d2682 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -23,7 +23,7 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import Litellm_EntityType -from litellm.proxy.db.db_lookup_gate import db_lookup_gate +from litellm.proxy.db.db_lookup_gate import bounded_db_lookup, db_lookup_gate from litellm.proxy.spend_tracking.spend_counter_batch import read_batched_spend_counter, record_spend_counter_value from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.project_repository import ProjectRepository @@ -134,36 +134,9 @@ class SpendCounterReseed: if SpendCounterReseed._is_key_or_team_window_counter(counter_key): return None try: - async with db_lookup_gate.current(): - if counter_key.startswith("spend:key:"): - token: Final = counter_key[len("spend:key:") :] - row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token}) - elif counter_key.startswith("spend:team_member:"): - suffix: Final = counter_key[len("spend:team_member:") :] - if ":" not in suffix: - return None - user_id, team_id = suffix.rsplit(":", 1) - row = await TeamMembershipRepository(prisma_client).table.find_unique( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}} - ) - elif counter_key.startswith("spend:team:"): - team_id = counter_key[len("spend:team:") :] - row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) - elif counter_key.startswith("spend:user:"): - user_id = counter_key[len("spend:user:") :] - row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) - elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"): - return None - elif counter_key.startswith("spend:org:"): - org_id: Final = counter_key[len("spend:org:") :] - row = await OrganizationRepository(prisma_client).table.find_unique( - where={"organization_id": org_id} - ) - elif counter_key.startswith("spend:project:"): - project_id: Final = counter_key[len("spend:project:") :] - row = await ProjectRepository(prisma_client).table.find_unique(where={"project_id": project_id}) - else: - return None + row: Final = await bounded_db_lookup( + SpendCounterReseed._counter_row(prisma_client, counter_key), name="spend_counter" + ) except Exception: verbose_proxy_logger.exception("SpendCounterReseed.from_db: failed for %s", counter_key) return None @@ -171,13 +144,47 @@ class SpendCounterReseed: return None return float(getattr(row, "spend", 0.0) or 0.0) + @staticmethod + async def _counter_row(prisma_client: "PrismaClient", counter_key: str) -> object | None: + async with db_lookup_gate.current(): + if counter_key.startswith("spend:key:"): + token: Final = counter_key[len("spend:key:") :] + return await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token}) + if counter_key.startswith("spend:team_member:"): + suffix: Final = counter_key[len("spend:team_member:") :] + if ":" not in suffix: + return None + user_id, team_id = suffix.rsplit(":", 1) + return await TeamMembershipRepository(prisma_client).table.find_unique( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}} + ) + if counter_key.startswith("spend:team:"): + return await TeamRepository(prisma_client).table.find_unique( + where={"team_id": counter_key[len("spend:team:") :]} + ) + if counter_key.startswith("spend:user:"): + return await UserRepository(prisma_client).table.find_unique( + where={"user_id": counter_key[len("spend:user:") :]} + ) + if counter_key.startswith("spend:org:"): + return await OrganizationRepository(prisma_client).table.find_unique( + where={"organization_id": counter_key[len("spend:org:") :]} + ) + if counter_key.startswith("spend:project:"): + return await ProjectRepository(prisma_client).table.find_unique( + where={"project_id": counter_key[len("spend:project:") :]} + ) + return None + @staticmethod async def end_user_from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None: if prisma_client is None or not counter_key.startswith(END_USER_COUNTER_PREFIX): return None where: Final[LiteLLM_EndUserTableWhereUniqueInput] = {"user_id": counter_key[len(END_USER_COUNTER_PREFIX) :]} try: - row: Final = await EndUserRepository(prisma_client).table.find_unique(where=where) + row: Final = await bounded_db_lookup( + EndUserRepository(prisma_client).table.find_unique(where=where), name="end_user_spend" + ) except Exception: # noqa: BLE001 # a failed floor read falls back to the cached spend, like from_db verbose_proxy_logger.exception("SpendCounterReseed.end_user_from_db: failed for %s", counter_key) return None diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 64fd59bbe44..f8801e65c82 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -16,7 +16,7 @@ from typing_extensions import ReadOnly import litellm from litellm._logging import verbose_logger, verbose_proxy_logger -from litellm.constants import HEALTH_CHECK_TIMEOUT_SECONDS +from litellm.constants import HEALTH_CHECK_TIMEOUT_SECONDS, PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS from litellm.integrations.SlackAlerting.ms_teams import ( MS_TEAMS_ALERT_HEADERS, build_ms_teams_payload, @@ -44,6 +44,7 @@ from litellm.proxy.auth.auth_utils import ( ) from litellm.proxy.auth.model_checks import get_key_models from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.db.health_check_latest import ( LatestHealthCheckRow, @@ -1723,7 +1724,7 @@ async def _get_health_readiness_details( # check DB if prisma_client is not None: # if db passed in, check if it's connected - db_health_status: Final = await _db_health_readiness_check() + db_status: Final = _readiness_db_status(await _db_health_readiness_check()) # A configured DB that is not reachable means the worker cannot # serve requests that depend on persisted state (keys, budgets, # spend logs). Return 503 so orchestrators take this pod out of @@ -1733,13 +1734,13 @@ async def _get_health_readiness_details( # report the DB state through the body instead. if ( response is not None - and db_health_status["status"] != "connected" + and db_status != "connected" and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable() ): response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE return { "status": "healthy", - "db": db_health_status["status"], + "db": db_status, "cache": cache_type, "litellm_version": version, "success_callbacks": success_callback_names, @@ -1816,24 +1817,32 @@ def _authorize_drain_request(request: Request) -> None: ) +def _readiness_db_status(db_health_status: DBHealthCache) -> str: + """A pod whose pre-request lookups hit their deadline inside the stall window + reports "stalled" even though the ping succeeds: the ping is a fresh + connection, the stalled lookups are the ones requests actually wait on.""" + if db_health_status["status"] != "connected": + return db_health_status["status"] + if db_lookup_stall_tracker.stalled_within(PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS): + return "stalled" + return "connected" + + async def _resolve_public_readiness_db(response: Response) -> str: """ Return the db status string for the public probe and flip the response to - 503 when a configured DB is unreachable. Mirrors the legacy values: - "Not connected" (no DB configured), "connected", "disconnected". + 503 when a configured DB is unreachable or stalled. Mirrors the legacy values: + "Not connected" (no DB configured), "connected", "disconnected", plus "stalled". """ from litellm.proxy.proxy_server import prisma_client if prisma_client is None: return "Not connected" - db_health_status: Final = await _db_health_readiness_check() - if ( - db_health_status["status"] != "connected" - and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable() - ): + db_status: Final = _readiness_db_status(await _db_health_readiness_check()) + if db_status != "connected" and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable(): response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE - return db_health_status["status"] + return db_status @router.get( diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 5dae9e8bb10..4dfea5cd472 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -23,6 +23,7 @@ from litellm.proxy.auth.auth_checks import ( log_db_metrics, ) from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded from litellm.proxy.db.db_spend_update_writer import ( DBSpendUpdateWriter, debitable_model_access_groups, @@ -186,8 +187,8 @@ class _ProxyDBLogger(CustomLogger): ) _metadata["error_information"] = _error_information - _metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( - metadata=_metadata, + _metadata = await _ProxyDBLogger._enrich_failure_metadata_unless_db_stalled( + metadata=_metadata, original_exception=original_exception ) existing_metadata: Final[dict] = request_data.get("metadata", None) or {} @@ -472,6 +473,12 @@ class _ProxyDBLogger(CustomLogger): spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) + @staticmethod + async def _enrich_failure_metadata_unless_db_stalled(metadata: dict, original_exception: Exception) -> dict: + if isinstance(original_exception, DBLookupDeadlineExceeded): + return metadata + return await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) + @staticmethod async def _enrich_failure_metadata_with_key_info(metadata: dict, resolve_missing_key_identity: bool = True) -> dict: """ diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index aa4bb2e49d3..30f5abdbb98 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1,7 +1,7 @@ import asyncio import json import time -from collections.abc import Mapping +from collections.abc import Iterator, Mapping from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal, Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -646,6 +646,114 @@ async def test_fetch_key_object_from_db_bounds_in_flight_prisma_requests(): assert prisma.max_in_flight == PROXY_DB_LOOKUP_MAX_CONCURRENCY +@pytest.fixture +def _clear_db_lookup_stall() -> Iterator[None]: + from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + + db_lookup_stall_tracker.clear() + yield + db_lookup_stall_tracker.clear() + + +class _StalledPrisma: + def __init__(self) -> None: + self.attempt_db_reconnect = AsyncMock(return_value=True) + self.db = MagicMock() + self.db.litellm_teamtable.find_unique = AsyncMock(side_effect=_stall_forever) + self.db.litellm_teamtable.update = AsyncMock(side_effect=_answer_slowly) + + async def get_data(self, token: str, table_name: str, parent_otel_span: None, proxy_logging_obj: None) -> None: + await _stall_forever() + + +async def _stall_forever(**kwargs: object) -> None: + await asyncio.Event().wait() + + +async def _answer_slowly(**kwargs: object) -> Mapping[str, object]: + await asyncio.sleep(0.15) + return {"team_id": "slow-write"} + + +@pytest.mark.asyncio +async def test_fetch_key_object_from_db_fails_a_stalled_burst_within_the_deadline_without_reconnecting( + _clear_db_lookup_stall, +): + """The incident: a stalled database parked every request in the pod with liveness + and readiness green until it OOMed. Every lookup in a burst larger than the gate, + the ones queued behind it included, must fail within one deadline, must not try to + reconnect (the transport is fine, the query is slow), and must leave every gate slot + free for the next burst.""" + from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded + + prisma: Final = _StalledPrisma() + burst: Final = PROXY_DB_LOOKUP_MAX_CONCURRENCY * 3 + started: Final = time.monotonic() + + results: Final = await asyncio.gather( + *( + _fetch_key_object_from_db_with_reconnect( + hashed_token=f"hashed-token-{i}", + prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient + parent_otel_span=None, + proxy_logging_obj=None, + deadline_seconds=0.2, + ) + for i in range(burst) + ), + return_exceptions=True, + ) + elapsed: Final = time.monotonic() - started + + assert len(results) == burst + assert all(isinstance(result, DBLookupDeadlineExceeded) for result in results) + assert all(PrismaDBExceptionHandler.is_database_service_unavailable_error(result) for result in results) + assert elapsed < 3 + prisma.attempt_db_reconnect.assert_not_awaited() + + recovered: Final = _InFlightCountingPrisma() + after: Final = await asyncio.wait_for( + asyncio.gather( + *( + _fetch_key_object_from_db_with_reconnect( + hashed_token=f"after-{i}", + prisma_client=recovered, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient + parent_otel_span=None, + proxy_logging_obj=None, + ) + for i in range(PROXY_DB_LOOKUP_MAX_CONCURRENCY) + ) + ), + timeout=5, + ) + assert {r.token for r in after if r is not None} == {f"after-{i}" for i in range(PROXY_DB_LOOKUP_MAX_CONCURRENCY)} + + +@pytest.mark.asyncio +async def test_team_lookup_fails_at_the_db_lookup_deadline_while_writes_stay_unbounded(_clear_db_lookup_stall): + """Team, user, budget, and membership reads share the key lookup's deadline through + the typed table wrappers; writes do not, since a slow write must land rather than + fail the request that already passed auth.""" + from litellm.proxy.auth.auth_checks import _team_table + from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded + from litellm.repositories.table_repositories import TeamRepository + + prisma: Final = _StalledPrisma() + with patch( # test-quality-ok: lowers the module-level lookup deadline so the stalled-read test finishes fast + "litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05 + ): + started: Final = time.monotonic() + with pytest.raises(DBLookupDeadlineExceeded, match=r"team lookup did not answer within 0\.05s"): + await _get_team_db_check(team_id="stalled-team", prisma_client=prisma) # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient + assert time.monotonic() - started < 2 + + written: Final = await _team_table(TeamRepository(prisma)).update( + where={"team_id": "slow-write"}, data={"spend": 1.0} + ) + + assert written == {"team_id": "slow-write"} + + def _fake_redis_cache(): fake_redis = MagicMock() fake_redis.async_get_cache = AsyncMock(return_value=None) @@ -6184,7 +6292,9 @@ async def test_get_org_object_for_request_serves_last_known_org_through_db_outag proxy_logging_obj=None, ) - with patch("litellm.proxy.proxy_server.general_settings", {}): # test-quality-ok: the outage fallback reads this module global; no dependency injection seam exists + with patch( + "litellm.proxy.proxy_server.general_settings", {} + ): # test-quality-ok: the outage fallback reads this module global; no dependency injection seam exists warm = await _lookup() assert warm is not None and warm.organization_alias == "platform-org" await user_api_key_cache.async_delete_cache("org_id:org-1:with_budget") @@ -8933,20 +9043,34 @@ async def test_access_group_model_fallback_uses_the_injected_database(channel: s reader: Final = AsyncMock(return_value=group) client: Final = MagicMock(db=MagicMock(litellm_accessgrouptable=MagicMock(find_unique=reader))) with ( - patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: [TQ008] prove reads stay on the injected connection - patch("litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache()), # test-quality-ok: [TQ008] isolate the process cache + patch( + "litellm.proxy.proxy_server.prisma_client", None + ), # test-quality-ok: [TQ008] prove reads stay on the injected connection + patch( + "litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache() + ), # test-quality-ok: [TQ008] isolate the process cache ): if channel == "team": - assert await can_team_access_model( - model="allowed", team_object=LiteLLM_TeamTable(team_id="team-a", models=["other"], access_group_ids=["group-a"]), - llm_router=None, prisma_client=client, - ) is True + assert ( + await can_team_access_model( + model="allowed", + team_object=LiteLLM_TeamTable(team_id="team-a", models=["other"], access_group_ids=["group-a"]), + llm_router=None, + prisma_client=client, + ) + is True + ) else: - assert await can_key_call_model( - model="allowed", llm_model_list=None, - valid_token=UserAPIKeyAuth(models=["other"], access_group_ids=["group-a"]), - llm_router=None, prisma_client=client, - ) is True + assert ( + await can_key_call_model( + model="allowed", + llm_model_list=None, + valid_token=UserAPIKeyAuth(models=["other"], access_group_ids=["group-a"]), + llm_router=None, + prisma_client=client, + ) + is True + ) reader.assert_awaited_once_with(where={"access_group_id": "group-a"}) @@ -8967,6 +9091,7 @@ def test_jwt_team_role_reaches_the_gateway_token_endpoint_by_default(): litellm_proxy_roles=LiteLLM_JWTAuth(team_allowed_routes=[]), ) + def test_route_skips_budget_checks_marks_only_spend_free_routes() -> None: assert route_skips_budget_checks(route="/v1/models") is True assert route_skips_budget_checks(route="/spend/logs") is True @@ -9083,7 +9208,9 @@ async def test_team_member_budget_check_temp_budget_increase_extends_cap(): return fallback_spend with ( - patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), # test-quality-ok: [TQ008] no seam on the cross-pod spend counter + patch( + "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend + ), # test-quality-ok: [TQ008] no seam on the cross-pod spend counter patch( # test-quality-ok: [TQ008] isolates the check from the DB fetch "litellm.proxy.auth.auth_checks.get_team_membership", new_callable=AsyncMock, @@ -9111,7 +9238,9 @@ async def test_team_member_budget_check_temp_budget_increase_extends_cap(): ), ) with ( - patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), # test-quality-ok: [TQ008] no seam on the cross-pod spend counter + patch( + "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend + ), # test-quality-ok: [TQ008] no seam on the cross-pod spend counter patch( # test-quality-ok: [TQ008] isolates the check from the DB fetch "litellm.proxy.auth.auth_checks.get_team_membership", new_callable=AsyncMock, @@ -9173,7 +9302,9 @@ async def test_team_member_budget_check_adds_temp_increase_to_live_team_default( return fallback_spend with ( - patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), # test-quality-ok: [TQ008] no seam on the cross-pod spend counter + patch( + "litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend + ), # test-quality-ok: [TQ008] no seam on the cross-pod spend counter patch( # test-quality-ok: [TQ008] isolates the check from the DB fetch "litellm.proxy.auth.auth_checks.get_team_membership", new_callable=AsyncMock, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index f03abe8f124..de669449f85 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -4,6 +4,7 @@ import logging import os import subprocess import sys +import time from collections.abc import Mapping from contextlib import contextmanager from datetime import datetime, timedelta, timezone @@ -186,11 +187,7 @@ async def test_disable_budget_reservation_does_not_log_per_request(caplog): general_settings={"disable_budget_reservation": True}, ) - records = [ - record - for record in caplog.records - if "disable_budget_reservation is enabled" in record.message - ] + records = [record for record in caplog.records if "disable_budget_reservation is enabled" in record.message] assert records == [] assert user_api_key_auth_obj.budget_reservation is None @@ -234,9 +231,7 @@ async def test_budget_reservation_runs_when_not_disabled(): ({}, False), ], ) -async def test_fail_closed_budget_enforcement_reaches_reservation( - general_settings, expected_flag -): +async def test_fail_closed_budget_enforcement_reaches_reservation(general_settings, expected_flag): """#33923: the strict flag must be threaded into reserve_budget_for_request so a failed reservation write can reject instead of failing open.""" user_api_key_auth_obj = UserAPIKeyAuth(token="test_token") @@ -259,10 +254,7 @@ async def test_fail_closed_budget_enforcement_reaches_reservation( general_settings=general_settings, ) - assert ( - mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"] - is expected_flag - ) + assert mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"] is expected_flag @pytest.mark.asyncio @@ -274,9 +266,7 @@ async def test_fail_closed_budget_enforcement_reaches_reservation( ({}, False), ], ) -async def test_apply_user_budget_to_team_keys_reaches_reservation( - general_settings, expected_flag -): +async def test_apply_user_budget_to_team_keys_reaches_reservation(general_settings, expected_flag): """The opt-in lives in general_settings but is consumed inside _get_budget_counters, so it has to be threaded through reserve_budget_for_request or the reservation path keeps exempting team keys while the read path enforces.""" @@ -300,9 +290,7 @@ async def test_apply_user_budget_to_team_keys_reaches_reservation( general_settings=general_settings, ) - assert ( - mock_reserve.await_args.kwargs["apply_user_budget_to_team_keys"] is expected_flag - ) + assert mock_reserve.await_args.kwargs["apply_user_budget_to_team_keys"] is expected_flag @pytest.mark.asyncio @@ -402,9 +390,7 @@ async def test_custom_auth_honors_key_level_model_access_restriction_allowed_wit "litellm.proxy.auth.user_api_key_auth.can_key_call_model", new_callable=AsyncMock, ) as mock_can_key, - patch( - "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock - ), + patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock), patch( "litellm.proxy.proxy_server.general_settings", {"custom_auth_run_common_checks": True}, @@ -435,9 +421,7 @@ async def test_custom_auth_enforces_key_model_access_from_file_route_header_with "litellm.proxy.auth.user_api_key_auth.can_key_call_model", new_callable=AsyncMock, ) as mock_can_key, - patch( - "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock - ), + patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock), patch( "litellm.proxy.proxy_server.general_settings", {"custom_auth_run_common_checks": True}, @@ -468,9 +452,7 @@ async def test_custom_auth_honors_key_level_model_access_restriction_denied_with "litellm.proxy.auth.user_api_key_auth.can_key_call_model", new_callable=AsyncMock, ) as mock_can_key, - patch( - "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock - ), + patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock), patch( "litellm.proxy.proxy_server.general_settings", {"custom_auth_run_common_checks": True}, @@ -506,9 +488,7 @@ def _proxy_server_attrs_for_custom_auth(*, user_custom_auth): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) return { @@ -770,9 +750,7 @@ async def test_enterprise_custom_auth_runs_post_custom_auth_checks_when_opt_in() litellm.enable_post_custom_auth_checks = original_flag -def _assert_get_api_key_with_custom_litellm_key_header( - custom_litellm_key_header, api_key, passed_in_key -): +def _assert_get_api_key_with_custom_litellm_key_header(custom_litellm_key_header, api_key, passed_in_key): assert get_api_key( custom_litellm_key_header=custom_litellm_key_header, api_key=None, @@ -829,9 +807,7 @@ def _assert_get_api_key_with_custom_litellm_key_header( ("App:LiteLLM", None, False, False), ], ) -def test_routing_selector_matches_claim_parametrized( - selector_value, claim_value, expected, split_space_delimited -): +def test_routing_selector_matches_claim_parametrized(selector_value, claim_value, expected, split_space_delimited): assert ( _routing_selector_matches_claim( selector_value=selector_value, @@ -925,10 +901,7 @@ def test_routing_selector_matches_claim_parametrized( ], ) def test_matches_routing_override_parametrized(override, token_claims, expected): - assert ( - _matches_routing_override(token_claims=token_claims, override=override) - is expected - ) + assert _matches_routing_override(token_claims=token_claims, override=override) is expected def test_get_api_key_with_custom_litellm_key_header_bearer_prefix(): @@ -1007,12 +980,9 @@ def test_team_metadata_with_tags_flows_through_jwt_auth(): ) # Verify team_metadata is set - assert ( - user_api_key_auth.team_metadata is not None - ), "team_metadata should be populated" + assert user_api_key_auth.team_metadata is not None, "team_metadata should be populated" assert user_api_key_auth.team_metadata == team_object.metadata, ( - f"team_metadata not correctly mapped. " - f"Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}" + f"team_metadata not correctly mapped. Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}" ) # Specifically verify tags are present @@ -1051,9 +1021,7 @@ def test_route_checks_is_llm_api_route(): ] for route in openai_routes: - assert RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" # Test Anthropic routes anthropic_routes = [ @@ -1062,9 +1030,7 @@ def test_route_checks_is_llm_api_route(): ] for route in anthropic_routes: - assert RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" # Test passthrough routes (this is the key improvement over the old route checking) passthrough_routes = [ @@ -1084,9 +1050,7 @@ def test_route_checks_is_llm_api_route(): ] for route in passthrough_routes: - assert RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" # Test MCP routes mcp_routes = [ @@ -1096,9 +1060,7 @@ def test_route_checks_is_llm_api_route(): ] for route in mcp_routes: - assert RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" # Test LiteLLM native RAG routes rag_routes = [ @@ -1108,9 +1070,7 @@ def test_route_checks_is_llm_api_route(): "/v1/rag/query", ] for route in rag_routes: - assert RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" # Test routes with placeholders placeholder_routes = [ @@ -1125,9 +1085,7 @@ def test_route_checks_is_llm_api_route(): ] for route in placeholder_routes: - assert RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" # Test Azure OpenAI routes azure_routes = [ @@ -1138,9 +1096,7 @@ def test_route_checks_is_llm_api_route(): ] for route in azure_routes: - assert RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" # Test non-LLM routes (should return False) non_llm_routes = [ @@ -1159,9 +1115,7 @@ def test_route_checks_is_llm_api_route(): ] for route in non_llm_routes: - assert not RouteChecks.is_llm_api_route( - route=route - ), f"Route {route} should NOT be identified as LLM API route" + assert not RouteChecks.is_llm_api_route(route=route), f"Route {route} should NOT be identified as LLM API route" # Test invalid inputs invalid_inputs = [ @@ -1173,9 +1127,9 @@ def test_route_checks_is_llm_api_route(): ] for invalid_input in invalid_inputs: - assert not RouteChecks.is_llm_api_route( - route=invalid_input - ), f"Invalid input {invalid_input} should return False" + assert not RouteChecks.is_llm_api_route(route=invalid_input), ( + f"Invalid input {invalid_input} should return False" + ) @pytest.mark.asyncio @@ -1222,9 +1176,7 @@ async def test_proxy_admin_expired_key_from_cache(): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() # Mock post_call_failure_hook as async function returning None (no transformation) mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) @@ -1261,9 +1213,7 @@ async def test_proxy_admin_expired_key_from_cache(): "jwt_handler": None, "litellm_proxy_admin_name": "admin", } - _original_values = { - attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set - } + _original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set} try: for attr, val in _attrs_to_set.items(): setattr(_proxy_server_mod, attr, val) @@ -1287,36 +1237,30 @@ async def test_proxy_admin_expired_key_from_cache(): ) # Verify that ProxyException was raised with expired_key type - assert hasattr( - exc_info.value, "type" - ), "Exception should have 'type' attribute" - assert ( - exc_info.value.type == ProxyErrorTypes.expired_key - ), f"Expected expired_key error type, got {exc_info.value.type}" + assert hasattr(exc_info.value, "type"), "Exception should have 'type' attribute" + assert exc_info.value.type == ProxyErrorTypes.expired_key, ( + f"Expected expired_key error type, got {exc_info.value.type}" + ) assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED - assert "Expired Key" in str( - exc_info.value.message - ), f"Exception message should mention 'Expired Key', got: {exc_info.value.message}" + assert "Expired Key" in str(exc_info.value.message), ( + f"Exception message should mention 'Expired Key', got: {exc_info.value.message}" + ) # Verify that the param field does NOT leak the full API key (Issue #18731) # The param should be abbreviated like "sk-...XXXX" not the full plaintext key - assert ( - exc_info.value.param is not None - ), "Exception should have 'param' attribute" + assert exc_info.value.param is not None, "Exception should have 'param' attribute" assert exc_info.value.param != api_key, ( f"SECURITY: Full API key should NOT be in param field! " f"Got: {exc_info.value.param}, Expected abbreviated format like 'sk-...XXXX'" ) - assert exc_info.value.param.startswith( - "sk-..." - ), f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}" + assert exc_info.value.param.startswith("sk-..."), ( + f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}" + ) # Verify that cache deletion was called mock_delete_cache.assert_called_once() call_args = mock_delete_cache.call_args - assert ( - call_args[1]["hashed_token"] == hashed_key - ), "Cache deletion should be called with the hashed key" + assert call_args[1]["hashed_token"] == hashed_key, "Cache deletion should be called with the hashed key" finally: # Restore all module-level attributes so subsequent tests are not affected for attr, val in _original_values.items(): @@ -1354,9 +1298,7 @@ async def test_scim_deactivated_user_key_is_rejected(): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) mock_prisma_client = MagicMock() @@ -1377,9 +1319,7 @@ async def test_scim_deactivated_user_key_is_rejected(): "jwt_handler": None, "litellm_proxy_admin_name": "admin", } - _original_values = { - attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set - } + _original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set} try: for attr, val in _attrs_to_set.items(): setattr(_proxy_server_mod, attr, val) @@ -1446,9 +1386,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) import litellm.proxy.proxy_server as _proxy_server_mod @@ -1467,9 +1405,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): "jwt_handler": None, "litellm_proxy_admin_name": "admin", } - _original_values = { - attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set - } + _original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set} try: for attr, val in _attrs_to_set.items(): setattr(_proxy_server_mod, attr, val) @@ -1521,9 +1457,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker(): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) import litellm.proxy.proxy_server as _proxy_server_mod @@ -1542,9 +1476,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker(): "jwt_handler": None, "litellm_proxy_admin_name": "admin", } - _original_values = { - attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set - } + _original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set} try: for attr, val in _attrs_to_set.items(): setattr(_proxy_server_mod, attr, val) @@ -1597,9 +1529,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) mock_prisma_client = MagicMock() @@ -1620,9 +1550,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): "jwt_handler": None, "litellm_proxy_admin_name": "admin", } - _original_values = { - attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set - } + _original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set} try: for attr, val in _attrs_to_set.items(): setattr(_proxy_server_mod, attr, val) @@ -2153,7 +2081,10 @@ async def test_auto_register_first_request_propagates_user_email(active: bool) - patch("litellm.proxy.proxy_server.master_key", "sk-master"), patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), - patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock(return_value=None))), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + ), patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), patch( "litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key", @@ -2218,7 +2149,9 @@ async def test_auto_register_stamps_new_key_with_jwt_agent_id(): plaintext = "sk-auto-registered-agent" token_hash = hash_token(plaintext) persisted_principal = IdentityStore._principal_from_key( - UserAPIKeyAuth(token=token_hash, user_id="validated-user", team_id="validated-team", agent_id="canonical-agent-id"), + UserAPIKeyAuth( + token=token_hash, user_id="validated-user", team_id="validated-team", agent_id="canonical-agent-id" + ), auth_method=AuthMethod.API_KEY, credential_ref=CredentialRef(token_id=token_hash), ) @@ -2433,10 +2366,7 @@ class TestJWTOAuth2Coexistence: def test_is_jwt_detects_jwt_tokens(self): """JWT tokens have 3 dot-separated parts.""" assert JWTHandler.is_jwt("header.payload.signature") is True - assert ( - JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") - is True - ) + assert JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") is True def test_is_jwt_rejects_opaque_tokens(self): """Opaque OAuth2 tokens do not have 3 dot-separated parts.""" @@ -2545,10 +2475,7 @@ class TestJWTOAuth2Coexistence: assert exc_info.value.type == ProxyErrorTypes.auth_error assert exc_info.value.code == "403" - assert ( - "Oauth2 token validation is only available for premium users" - in exc_info.value.message - ) + assert "Oauth2 token validation is only available for premium users" in exc_info.value.message mock_oauth2.assert_not_called() @pytest.mark.asyncio @@ -2740,9 +2667,7 @@ class TestJWTOAuth2Coexistence: assert mock_auto_register.call_args.kwargs["team_id"] == "validated-team" assert mock_auto_register.call_args.kwargs["user_id"] == "validated-user" assert mock_auto_register.call_args.kwargs["org_id"] == "validated-org" - assert ( - mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user" - ) + assert mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user" assert result.org_id == "validated-org" assert result.user_email == "validated@example.com" @@ -2820,10 +2745,7 @@ class TestJWTOAuth2Coexistence: assert result.user_id == "mapped-user" assert result.user_email == "mapped@example.com" - assert ( - mock_get_user_object.call_args_list[0].kwargs["user_email"] - == "mapped@example.com" - ) + assert mock_get_user_object.call_args_list[0].kwargs["user_email"] == "mapped@example.com" @pytest.mark.asyncio async def test_mapped_virtual_key_does_not_backfill_mismatched_owner(self): @@ -2899,8 +2821,7 @@ class TestJWTOAuth2Coexistence: assert result.user_id == "other-owner" assert result.user_email is None assert all( - call.kwargs.get("user_email") != "principal@example.com" - for call in mock_get_user_object.call_args_list + call.kwargs.get("user_email") != "principal@example.com" for call in mock_get_user_object.call_args_list ) @pytest.mark.asyncio @@ -3705,9 +3626,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls(): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) import litellm.proxy.proxy_server as _proxy_server_mod @@ -3839,9 +3758,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( - AsyncMock() - ) + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) import litellm.proxy.proxy_server as _proxy_server_mod @@ -3891,9 +3808,9 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): request_data={}, ) - assert result.team_metadata == { - "guardrails": ["test-guardrail-333"] - }, f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" + assert result.team_metadata == {"guardrails": ["test-guardrail-333"]}, ( + f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" + ) finally: for k, v in _originals.items(): @@ -4218,9 +4135,7 @@ async def test_auth_flow_fallback_team_object_permission_none_when_unreadable(): # --------------------------------------------------------------------------- -def _proxy_attrs_for_centralized_checks( - user_custom_auth=None, flag=False, master_key="sk-test-master" -): +def _proxy_attrs_for_centralized_checks(user_custom_auth=None, flag=False, master_key="sk-test-master"): """Build the minimal proxy_server module attributes that _run_centralized_common_checks reads. @@ -4430,9 +4345,7 @@ async def _run_centralized_checks_with_key_end_user_budget( request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") attrs = { - **_proxy_attrs_for_centralized_checks( - user_custom_auth=AsyncMock() if custom_auth else None, flag=custom_auth - ), + **_proxy_attrs_for_centralized_checks(user_custom_auth=AsyncMock() if custom_auth else None, flag=custom_auth), "prisma_client": prisma_client, "user_api_key_cache": user_api_key_cache if user_api_key_cache is not None else DualCache(), "proxy_logging_obj": proxy_logging_obj, @@ -4623,7 +4536,9 @@ async def test_centralized_common_checks_enforces_team_model_max_budget_from_the for k, v in attrs.items(): setattr(_proxy_server_mod, k, v) with ( - patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock), # test-quality-ok: stubs the sibling check so only the team model-budget gate is under test + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock + ), # test-quality-ok: stubs the sibling check so only the team model-budget gate is under test patch( # test-quality-ok: stubs the budget reservation so only the team model-budget gate is under test "litellm.proxy.auth.user_api_key_auth._reserve_budget_after_common_checks", new_callable=AsyncMock, @@ -4656,9 +4571,7 @@ async def test_centralized_common_checks_skipped_for_custom_auth_without_flag(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - attrs = _proxy_attrs_for_centralized_checks( - user_custom_auth=AsyncMock(), flag=False - ) + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=AsyncMock(), flag=False) originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} try: for k, v in attrs.items(): @@ -5063,9 +4976,7 @@ async def test_centralized_common_checks_reserves_request_end_user_budget(): "applied_adjustment": 0.0, } ] - assert counter_cache.in_memory_cache.get_cache( - key="spend:end_user:alice" - ) == pytest.approx(0.6) + assert counter_cache.in_memory_cache.get_cache(key="spend:end_user:alice") == pytest.approx(0.6) @pytest.mark.asyncio @@ -5080,9 +4991,7 @@ async def test_centralized_common_checks_short_circuits_when_master_key_unset(): from litellm.proxy._types import LitellmUserRoles - token = UserAPIKeyAuth( - api_key="sk-test", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER - ) + token = UserAPIKeyAuth(api_key="sk-test", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER) request = Request(scope={"type": "http"}) request._url = URL(url="/get/config/callbacks") @@ -5883,9 +5792,7 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on request._url = URL(url="/chat/completions") request._body = json.dumps({"user": "alice", "model": "gpt-4o"}).encode() - fetched_team = LiteLLM_TeamTableCachedObj( - team_id="t1", max_budget=20.0, models=["gpt-4o"] - ) + fetched_team = LiteLLM_TeamTableCachedObj(team_id="t1", max_budget=20.0, models=["gpt-4o"]) fetched_end_user = LiteLLM_EndUserTable(user_id="alice", blocked=False, spend=1.0) fetched_project = LiteLLM_ProjectTableCachedObj( project_id="proj-1", @@ -6014,10 +5921,46 @@ async def test_centralized_common_checks_backfills_org_id_from_team(key_org_id, ("org-pinned", None, None, "preset", None, "success", False, False, "org-pinned", "preset", (None, None, None)), ("org-view", None, None, None, 3, "success", False, False, "org-view", None, (None, None, 3)), ("org-missing", None, None, None, None, "missing", False, False, "org-missing", None, (None, None, None)), - ("org-db-failure-allowed", None, None, None, None, "db_failure", True, False, "org-db-failure-allowed", None, (None, None, None)), - ("org-db-failure-denied", None, None, None, None, "db_failure", False, True, "org-db-failure-denied", None, (None, None, None)), + ( + "org-db-failure-allowed", + None, + None, + None, + None, + "db_failure", + True, + False, + "org-db-failure-allowed", + None, + (None, None, None), + ), + ( + "org-db-failure-denied", + None, + None, + None, + None, + "db_failure", + False, + True, + "org-db-failure-denied", + None, + (None, None, None), + ), ("org-bad-row", None, None, None, None, "bad_row", False, False, "org-bad-row", None, (None, None, None)), - ("org-nobudget", None, None, None, None, "no_budget", False, False, "org-nobudget", "acme-org", (None, None, None)), + ( + "org-nobudget", + None, + None, + None, + None, + "no_budget", + False, + False, + "org-nobudget", + "acme-org", + (None, None, None), + ), ], ) async def test_centralized_common_checks_inherits_org_identity( @@ -6326,9 +6269,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it(): } ) request._url = URL(url="/chat/completions") - request._body = json.dumps( - {"model": "gpt-4o", "user": "alice@example.com"} - ).encode() + request._body = json.dumps({"model": "gpt-4o", "user": "alice@example.com"}).encode() attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} @@ -6372,9 +6313,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() import litellm.proxy.proxy_server as _proxy_server_mod - builder_token = UserAPIKeyAuth( - api_key="sk-test", user_id="u1", end_user_id="builder-resolved-id" - ) + builder_token = UserAPIKeyAuth(api_key="sk-test", user_id="u1", end_user_id="builder-resolved-id") request = Request( scope={ @@ -6384,9 +6323,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() } ) request._url = URL(url="/chat/completions") - request._body = json.dumps( - {"model": "gpt-4o", "user": "different-id-from-body"} - ).encode() + request._body = json.dumps({"model": "gpt-4o", "user": "different-id-from-body"}).encode() attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} @@ -6717,6 +6654,83 @@ async def _run_builder_with_key_lookup(get_key_object_mock): setattr(_proxy_server_mod, k, v) +class _StalledKeyLookupPrisma: + """A database whose connection answers the readiness ping but whose key lookups + never return, which is what the incident's locked table looked like.""" + + def __init__(self) -> None: + self.health_check = AsyncMock(return_value=True) + self.attempt_db_reconnect = AsyncMock(return_value=True) + self.db = MagicMock() + + async def get_data(self, token: str, table_name: str, parent_otel_span: None, proxy_logging_obj: None) -> None: + await asyncio.Event().wait() + + +@pytest.mark.asyncio +async def test_burst_against_a_stalled_db_fails_fast_with_503_and_turns_readiness_red(): + """The incident, end to end: N requests into a proxy whose database stalls used to + park in the pod with readiness green until it OOMed. Now every one of them fails + within the lookup deadline as a 503, and the next readiness probe takes the pod out + of rotation.""" + import httpx + from fastapi import Depends, FastAPI + + import litellm.proxy.health_endpoints._health_endpoints as health_endpoints + import litellm.proxy.proxy_server as _proxy_server_mod + from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + + app = FastAPI() + + @app.post("/chat/completions", dependencies=[Depends(user_api_key_auth)]) + async def chat_completions() -> Mapping[str, bool]: + return {"served": True} + + app.include_router(health_endpoints.router) + app.add_exception_handler(ProxyException, _proxy_server_mod.openai_exception_handler) + + attrs = {**_proxy_attrs_for_db_lookup(), "prisma_client": _StalledKeyLookupPrisma()} + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + health_endpoints.db_health_cache = {"status": "unknown", "last_updated": datetime.now() - timedelta(seconds=60)} + db_lookup_stall_tracker.clear() + burst = 60 + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with ( + patch( # test-quality-ok: lowers the module-level lookup deadline so the stalled burst finishes fast + "litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.2 + ), + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + ): + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://t") as client: + started = time.monotonic() + responses = await asyncio.gather( + *( + client.post( + "/chat/completions", + json={"model": "gpt-5.5", "messages": [{"role": "user", "content": "hi"}]}, + headers={"Authorization": f"Bearer sk-stalled-{i}"}, + ) + for i in range(burst) + ) + ) + elapsed = time.monotonic() - started + readiness = await client.get("/health/readiness") + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + db_lookup_stall_tracker.clear() + + assert len(responses) == burst + assert {r.status_code for r in responses} == {status.HTTP_503_SERVICE_UNAVAILABLE} + assert {r.json()["error"]["type"] for r in responses} == {ProxyErrorTypes.no_db_connection.value} + assert all("temporarily unreachable" in r.json()["error"]["message"] for r in responses) + assert elapsed < 5 + assert readiness.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert readiness.json()["db"] == "stalled" + + @pytest.mark.asyncio async def test_builder_returns_503_when_db_lookup_raises_infra_error(): """End-to-end: a DB infrastructure failure during the key lookup must @@ -6784,9 +6798,7 @@ def _mint_cli_session_token(monkeypatch, *, user_id="cli-admin"): models=["gpt-3.5-turbo"], max_budget=100.0, ) - return ExperimentalUIJWTToken.get_cli_jwt_auth_token( - user_info, team_id="cli-team", team_alias="cli-team-alias" - ) + return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info, team_id="cli-team", team_alias="cli-team-alias") @pytest.mark.asyncio @@ -6836,7 +6848,7 @@ async def test_random_non_sk_token_is_rejected(monkeypatch): patch("litellm.proxy.proxy_server.master_key", "sk-master"), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), ): - with pytest.raises(Exception, match='LiteLLM Virtual Key expected\\.') as exc_info: + with pytest.raises(Exception, match="LiteLLM Virtual Key expected\\.") as exc_info: await user_api_key_auth( request=mock_request, api_key="Bearer not-a-real-token", @@ -6915,9 +6927,7 @@ async def test_non_admin_cli_session_token_reaches_production_auth_path(monkeypa user_role=LitellmUserRoles.INTERNAL_USER.value, models=[], ) - cli_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( - user_info, team_id="team-abc", team_alias="my-team" - ) + cli_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info, team_id="team-abc", team_alias="my-team") import litellm.proxy.proxy_server as _proxy_server_mod from fastapi import Request @@ -7224,7 +7234,7 @@ async def test_real_jwt_still_requires_license_when_jwt_auth_enabled(monkeypatch patch("litellm.proxy.proxy_server.master_key", "sk-master"), patch("litellm.proxy.proxy_server.prisma_client", None), ): - with pytest.raises(Exception, match='JWT Auth is an enterprise only feature\\. You must be a') as exc_info: + with pytest.raises(Exception, match="JWT Auth is an enterprise only feature\\. You must be a") as exc_info: await user_api_key_auth( request=mock_request, api_key=f"Bearer {jwt_token}", @@ -7263,13 +7273,9 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): metadata={"model_rpm_limit": {"gpt-5.4-mini": 3}}, last_refreshed_at=1000.0, ) - await key_cache.async_set_cache( - key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth - ) + await key_cache.async_set_cache(key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth) - fetch_from_db = AsyncMock( - side_effect=AssertionError("cache-hit auth must not touch the DB") - ) + fetch_from_db = AsyncMock(side_effect=AssertionError("cache-hit auth must not touch the DB")) proxy_logging_obj = MagicMock() proxy_logging_obj.internal_usage_cache = MagicMock() @@ -7316,9 +7322,7 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): assert result.token == hashed_key fetch_from_db.assert_not_called() - cached_after = await key_cache.async_get_cache( - key=hashed_key, model_type=UserAPIKeyAuth - ) + cached_after = await key_cache.async_get_cache(key=hashed_key, model_type=UserAPIKeyAuth) assert cached_after is not None assert cached_after.last_refreshed_at == 1000.0 assert cached_after.metadata == {"model_rpm_limit": {"gpt-5.4-mini": 3}} @@ -7394,7 +7398,9 @@ class TestJWTAuthUserEmail: assert result.user_email == "resolved@example.com" @pytest.mark.asyncio - @pytest.mark.parametrize("route", ["/mcp-rest/tools/list", "/mcp-rest/tools/call", "/v1/chat/completions", "/user/info"]) + @pytest.mark.parametrize( + "route", ["/mcp-rest/tools/list", "/mcp-rest/tools/call", "/v1/chat/completions", "/user/info"] + ) @pytest.mark.parametrize("active", [False, True, None, "false", 0]) @pytest.mark.parametrize("is_admin", [False, True]) async def test_jwt_auth_rejects_deactivated_user( @@ -7469,9 +7475,7 @@ class TestCheckKeyModelBudgetWithFallback: @pytest.mark.asyncio async def test_within_budget_does_not_reroute(self): - valid_token = UserAPIKeyAuth( - token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]} - ) + valid_token = UserAPIKeyAuth(token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}) limiter = AsyncMock() limiter.is_key_within_model_budget.return_value = True request_data = {"model": "gpt-4o"} @@ -7496,9 +7500,7 @@ class TestCheckKeyModelBudgetWithFallback: budget_fallbacks={"gpt-4o": ["gpt-4o-mini", "claude-haiku"]}, ) limiter = AsyncMock() - limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError( - current_cost=10, max_budget=5 - ) + limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5) limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini" request_data = {"model": "gpt-4o"} request = self._make_request() @@ -7512,9 +7514,7 @@ class TestCheckKeyModelBudgetWithFallback: ) assert request_data["model"] == "gpt-4o-mini" - limiter.get_fallback_model_within_budget.assert_awaited_once_with( - user_api_key_dict=valid_token, model="gpt-4o" - ) + limiter.get_fallback_model_within_budget.assert_awaited_once_with(user_api_key_dict=valid_token, model="gpt-4o") # the rerouted model must be visible to a later, separate # `_read_request_body` call on the same `request` (route handlers # re-parse the body from this cache instead of reusing the dict). @@ -7523,9 +7523,7 @@ class TestCheckKeyModelBudgetWithFallback: @pytest.mark.asyncio async def test_raises_when_every_fallback_also_exceeded(self): - valid_token = UserAPIKeyAuth( - token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]} - ) + valid_token = UserAPIKeyAuth(token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}) limiter = AsyncMock() original_error = litellm.BudgetExceededError(current_cost=10, max_budget=5) limiter.is_key_within_model_budget.side_effect = original_error @@ -7595,9 +7593,7 @@ class TestCheckKeyModelBudgetWithFallback: budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}, ) limiter = AsyncMock() - limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError( - current_cost=10, max_budget=5 - ) + limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5) limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini" request_data = {"model": "gpt-4o"} request = self._make_request() @@ -7665,9 +7661,7 @@ class TestCheckKeyModelBudgetWithFallback: budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}, ) limiter = AsyncMock() - limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError( - current_cost=10, max_budget=5 - ) + limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5) limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini" request_data = {"model": "gpt-4o"} request = self._make_request() @@ -7747,9 +7741,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row(): ) assert result == 42.5 - prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with( - where={"user_id": "litellm-proxy-budget"} - ) + prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with(where={"user_id": "litellm-proxy-budget"}) @pytest.mark.asyncio @@ -8124,9 +8116,7 @@ async def test_jwt_shaped_key_error_names_enable_jwt_auth_when_disabled(): Prometheus invalid-key filter and the admin UI both substring-match it. Keys that are not JWT-shaped must not pick up the hint. """ - jwt_error = await _proxy_exception_for_key( - "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9.c2lnbmF0dXJl", {}, True - ) + jwt_error = await _proxy_exception_for_key("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9.c2lnbmF0dXJl", {}, True) assert jwt_error.code == "401" assert "enable_jwt_auth" in jwt_error.message @@ -8136,9 +8126,7 @@ async def test_jwt_shaped_key_error_names_enable_jwt_auth_when_disabled(): assert "is a JWT" not in jwt_error.message opaque_error = await _proxy_exception_for_key("not-a-jwt-at-all", {}, True) - two_segment_error = await _proxy_exception_for_key( - "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9", {}, True - ) + two_segment_error = await _proxy_exception_for_key("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9", {}, True) assert "enable_jwt_auth" not in opaque_error.message assert "enable_jwt_auth" not in two_segment_error.message @@ -8167,9 +8155,7 @@ class TestLitellmReceivedAtStamping: on OTEL being configured to see a true request-arrival timestamp.""" def test_stamped_even_when_otel_is_not_configured(self, monkeypatch): - monkeypatch.setattr( - "litellm.proxy.proxy_server.open_telemetry_logger", None - ) + monkeypatch.setattr("litellm.proxy.proxy_server.open_telemetry_logger", None) request = MagicMock() request.state = SimpleNamespace() @@ -8201,7 +8187,7 @@ class TestLitellmReceivedAtStamping: _RECORDING_DDTRACE = dedent( - ''' + """ import functools import inspect @@ -8250,11 +8236,11 @@ _RECORDING_DDTRACE = dedent( tracer = _Tracer() - ''' + """ ) _DDTRACE_AUTH_PROBE = dedent( - ''' + """ import asyncio import json @@ -8293,7 +8279,7 @@ _DDTRACE_AUTH_PROBE = dedent( asyncio.run(main()) - ''' + """ ) @@ -8440,25 +8426,43 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a @pytest.mark.asyncio -@pytest.mark.parametrize("route", ["/v1/messages", "/messages", "/v1/chat/completions", "/chat/completions", "/v1/responses", "/responses"]) +@pytest.mark.parametrize( + "route", ["/v1/messages", "/messages", "/v1/chat/completions", "/chat/completions", "/v1/responses", "/responses"] +) async def test_claude_view_normalizes_before_model_access(monkeypatch, route): from starlette.requests import Request from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access source = "foo[1m]" encoded = "claude-router-" + source.encode().hex() + "[1m]" - router = litellm.Router(model_list=[{"model_name": source, "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}}]) + router = litellm.Router( + model_list=[{"model_name": source, "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}}] + ) monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) data = {"model": encoded, "messages": [{"role": "user", "content": "hi"}]} request = Request({"type": "http", "method": "POST", "path": route, "headers": [], "query_string": b""}) token = UserAPIKeyAuth(models=[source]) - await _enforce_key_and_fallback_model_access(valid_token=token, request_data=data, route=route, request=request, llm_model_list=router.model_list, llm_router=router) + await _enforce_key_and_fallback_model_access( + valid_token=token, + request_data=data, + route=route, + request=request, + llm_model_list=router.model_list, + llm_router=router, + ) assert data["model"] == source assert (await request.json())["model"] == source assert json.loads(await request.body())["model"] == source assert request.scope["parsed_body"][1]["model"] == source with pytest.raises(ProxyException): - await _enforce_key_and_fallback_model_access(valid_token=UserAPIKeyAuth(models=["other"]), request_data=data, route=route, request=request, llm_model_list=router.model_list, llm_router=router) + await _enforce_key_and_fallback_model_access( + valid_token=UserAPIKeyAuth(models=["other"]), + request_data=data, + route=route, + request=request, + llm_model_list=router.model_list, + llm_router=router, + ) @pytest.mark.asyncio @@ -8470,10 +8474,18 @@ async def test_claude_view_never_reinterprets_explicit_names(monkeypatch, layer) encoded = "claude-router-666f6f" names = ("foo", "other", encoded) if layer == "literal" else ("foo", "other") alias = {encoded: "other"} - router = litellm.Router(model_list=[{"model_name": name, "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}} for name in names], model_group_alias=alias if layer == "router" else None) + router = litellm.Router( + model_list=[ + {"model_name": name, "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}} for name in names + ], + model_group_alias=alias if layer == "router" else None, + ) monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) monkeypatch.setattr(litellm, "model_alias_map", alias if layer == "global" else {}) - token = UserAPIKeyAuth(aliases=alias if layer == "key" else {}, router_settings={"model_group_alias": alias} if layer == "hierarchical" else None) + token = UserAPIKeyAuth( + aliases=alias if layer == "key" else {}, + router_settings={"model_group_alias": alias} if layer == "hierarchical" else None, + ) data = {"model": encoded} request = Request({"type": "http", "method": "POST", "path": "/v1/messages", "headers": [], "query_string": b""}) await _normalize_claude_model(data, token, request, "/v1/messages") diff --git a/tests/test_litellm/proxy/db/test_db_lookup_gate.py b/tests/test_litellm/proxy/db/test_db_lookup_gate.py new file mode 100644 index 00000000000..68903170840 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_db_lookup_gate.py @@ -0,0 +1,111 @@ +import asyncio +import time +from typing import Final + +import pytest + +from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded, DBLookupStallTracker, bounded_db_lookup + + +async def _never_answers() -> None: + await asyncio.Event().wait() + + +class _FakeClock: + def __init__(self) -> None: + self.now = 1000.0 + + def __call__(self) -> float: + return self.now + + +@pytest.mark.asyncio +async def test_bounded_db_lookup_fails_a_stalled_lookup_at_the_deadline_and_records_the_hit(): + tracker: Final = DBLookupStallTracker() + started: Final = time.monotonic() + + with pytest.raises(DBLookupDeadlineExceeded) as exc_info: + await bounded_db_lookup(_never_answers(), name="team", deadline_seconds=0.05, tracker=tracker) + + assert time.monotonic() - started < 2 + assert exc_info.value.lookup == "team" + assert exc_info.value.deadline_seconds == 0.05 + assert str(exc_info.value) == "team lookup did not answer within 0.05s" + assert isinstance(exc_info.value, asyncio.TimeoutError) + assert tracker.stalled_within(30) is True + + +@pytest.mark.asyncio +async def test_bounded_db_lookup_returns_a_prompt_answer_without_recording_a_stall(): + tracker: Final = DBLookupStallTracker() + + async def answers() -> str: + return "row" + + assert await bounded_db_lookup(answers(), name="key", deadline_seconds=0.05, tracker=tracker) == "row" + assert tracker.stalled_within(30) is False + + +@pytest.mark.asyncio +async def test_bounded_db_lookup_fails_a_whole_stalled_burst_within_one_deadline(): + tracker: Final = DBLookupStallTracker() + burst: Final = 200 + started: Final = time.monotonic() + + results: Final = await asyncio.gather( + *( + bounded_db_lookup(_never_answers(), name=f"key-{i}", deadline_seconds=0.1, tracker=tracker) + for i in range(burst) + ), + return_exceptions=True, + ) + + assert time.monotonic() - started < 2 + assert len(results) == burst + assert all(isinstance(result, DBLookupDeadlineExceeded) for result in results) + assert tracker.stalled_within(30) is True + + +@pytest.mark.asyncio +async def test_bounded_db_lookup_fails_at_the_deadline_even_when_the_lookup_absorbs_the_cancel(): + tracker: Final = DBLookupStallTracker() + absorbed: Final = asyncio.Event() + let_go: Final = asyncio.Event() + + async def absorbs_the_cancel() -> str: + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + absorbed.set() + await let_go.wait() + return "late row" + + started: Final = time.monotonic() + with pytest.raises(DBLookupDeadlineExceeded): + await asyncio.wait_for( + bounded_db_lookup(absorbs_the_cancel(), name="key", deadline_seconds=0.05, tracker=tracker), + timeout=2, + ) + + assert time.monotonic() - started < 1 + assert tracker.stalled_within(30) is True + await asyncio.wait_for(absorbed.wait(), timeout=1) + let_go.set() + await asyncio.sleep(0) + + +def test_stall_tracker_reports_a_stall_only_inside_the_window(): + clock: Final = _FakeClock() + tracker: Final = DBLookupStallTracker(clock=clock) + + assert tracker.stalled_within(30) is False + tracker.record_hit() + assert tracker.stalled_within(30) is True + assert tracker.stalled_within(0) is False + clock.now += 29.9 + assert tracker.stalled_within(30) is True + clock.now += 0.2 + assert tracker.stalled_within(30) is False + tracker.record_hit() + tracker.clear() + assert tracker.stalled_within(30) is False diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index 613ca847115..09f4d294ad0 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -774,3 +774,18 @@ def test_connection_error_answers_when_prisma_is_mocked_after_import(): with patch.dict(sys.modules, {"prisma": MagicMock()}): assert PrismaDBExceptionHandler.is_database_connection_error(Exception("x")) is False assert PrismaDBExceptionHandler.is_database_connection_error(httpx.ConnectError("refused")) is True + + +def test_db_lookup_deadline_is_a_connection_and_unavailability_error_but_never_a_transport_error(): + """A lookup that hit its deadline fails the request as a 503 and counts as a + DB outage for ``allow_requests_on_db_unavailable``, but it must not be read + as a broken transport: that would send every parked request into + ``attempt_db_reconnect`` and turn a slow database into a reconnect storm.""" + from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded + + deadline: Final = DBLookupDeadlineExceeded("key", 10.0) + + assert PrismaDBExceptionHandler.is_database_connection_error(deadline) is True + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(deadline) is True + assert PrismaDBExceptionHandler.is_database_transport_error(deadline) is False + assert "temporarily unreachable" in PrismaDBExceptionHandler.database_unavailable_message(deadline) diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index ff0b67d426b..ab931277313 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -12,12 +12,14 @@ from collections.abc import Mapping from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final +from unittest.mock import AsyncMock import pytest from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY +from litellm.proxy.db.db_lookup_gate import LoopBoundSemaphore, db_lookup_stall_tracker from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed WINDOW_START = datetime(2026, 8, 1, tzinfo=timezone.utc) @@ -445,6 +447,31 @@ async def test_from_db_returns_none_for_a_missing_project_row(): assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:project:proj-1") is None +@pytest.mark.asyncio +async def test_from_db_deadline_covers_the_wait_for_a_gate_slot(monkeypatch: pytest.MonkeyPatch) -> None: + """A saturated gate must fail the lookup at the deadline instead of parking + the request on a gate slot outside the bounded window.""" + gate: Final = LoopBoundSemaphore(1) + monkeypatch.setattr("litellm.proxy.db.spend_counter_reseed.db_lookup_gate", gate) + monkeypatch.setattr("litellm.proxy.db.db_lookup_gate.PROXY_DB_LOOKUP_DEADLINE_SECONDS", 0.05) + find_unique: Final = AsyncMock() + prisma: Final = SimpleNamespace( + db=SimpleNamespace(litellm_verificationtoken=SimpleNamespace(find_unique=find_unique)) + ) + db_lookup_stall_tracker.clear() + try: + async with gate.current(): + result: Final = await asyncio.wait_for( + SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:key:abc"), + timeout=1.0, + ) + assert result is None + assert db_lookup_stall_tracker.stalled_within(60.0) + find_unique.assert_not_called() + finally: + db_lookup_stall_tracker.clear() + + @pytest.mark.asyncio async def test_from_db_still_never_reads_the_end_user_row(): """A cold end-user counter keeps seeding from the cached end-user object the auth diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index 761cd0685f2..ee4c468a460 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -2703,6 +2703,134 @@ async def test_health_readiness_details_returns_200_when_db_down_and_allow_reque assert result["db"] == "disconnected" +@pytest.fixture +def _clear_db_lookup_stall() -> Iterator[None]: + from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + + db_lookup_stall_tracker.clear() + yield + db_lookup_stall_tracker.clear() + + +def _connected_prisma() -> MagicMock: + mock_prisma = MagicMock() + mock_prisma.health_check = AsyncMock(return_value=True) + return mock_prisma + + +def _forget_db_health_cache() -> None: + _health_endpoints_module.db_health_cache = { + "status": "unknown", + "last_updated": datetime.now() - timedelta(seconds=60), + } + + +@pytest.mark.asyncio +async def test_health_readiness_returns_503_stalled_after_a_db_lookup_deadline_hit(_clear_db_lookup_stall): + """The incident's readiness stayed green while every request sat parked on the + database: the probe's own ping is a fresh connection that answers fine. A lookup + that hit its deadline inside the stall window must take the pod out of rotation.""" + from fastapi import Response + + from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + from litellm.proxy.health_endpoints._health_endpoints import health_readiness + + _forget_db_health_cache() + db_lookup_stall_tracker.record_hit() + + response = Response() + with patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam + "litellm.proxy.proxy_server.prisma_client", _connected_prisma() + ): + result = await health_readiness(response=response) + + assert response.status_code == 503 + assert result == {"status": "healthy", "db": "stalled"} + + +@pytest.mark.asyncio +async def test_health_readiness_details_returns_503_stalled_after_a_db_lookup_deadline_hit(_clear_db_lookup_stall): + from fastapi import Response + + from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + from litellm.proxy.health_endpoints._health_endpoints import _get_health_readiness_details + + _forget_db_health_cache() + db_lookup_stall_tracker.record_hit() + + response = Response() + with patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam + "litellm.proxy.proxy_server.prisma_client", _connected_prisma() + ): + result = await _get_health_readiness_details(response=response) + + assert response.status_code == 503 + assert result["db"] == "stalled" + + +@pytest.mark.asyncio +async def test_health_readiness_stays_200_with_stalled_body_when_requests_are_allowed_on_db_unavailable( + _clear_db_lookup_stall, +): + """The fail-open deployment keeps serving through a stalled database, so the pod + must stay in rotation and report the stall through the body, exactly as it does + for a disconnected one.""" + from fastapi import Response + + from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + from litellm.proxy.health_endpoints._health_endpoints import health_readiness + + _forget_db_health_cache() + db_lookup_stall_tracker.record_hit() + + response = Response() + with ( + patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam + "litellm.proxy.proxy_server.prisma_client", _connected_prisma() + ), + patch.dict( # test-quality-ok: the fail-open flag lives in the proxy-global general_settings; no injection seam + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": True}, + ), + ): + result = await health_readiness(response=response) + + assert response.status_code == 200 + assert result == {"status": "healthy", "db": "stalled"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("hit_recorded", [False, True]) +async def test_health_readiness_reports_connected_without_a_stall_inside_the_window( + _clear_db_lookup_stall, hit_recorded: bool +): + """No deadline hit, or a window of 0 (the opt-out), keeps the ordinary connected + answer, so a healthy pod never leaves rotation over the stall check.""" + from fastapi import Response + + from litellm.proxy.db.db_lookup_gate import db_lookup_stall_tracker + from litellm.proxy.health_endpoints._health_endpoints import health_readiness + + _forget_db_health_cache() + if hit_recorded: + db_lookup_stall_tracker.record_hit() + + response = Response() + with ( + patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam + "litellm.proxy.proxy_server.prisma_client", _connected_prisma() + ), + patch( # test-quality-ok: lowers the module-level stall window to its opt-out value for the recorded-hit case + "litellm.proxy.health_endpoints._health_endpoints.PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS", + 0.0 if hit_recorded else 30.0, + ), + ): + result = await health_readiness(response=response) + + assert response.status_code == 200 + assert result == {"status": "healthy", "db": "connected"} + + @pytest.mark.asyncio async def test_db_health_readiness_check_bounds_hung_health_check(): """ diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 0e9c336a9eb..b5e594db701 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -5,12 +5,14 @@ from datetime import datetime from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.proxy._types import SpendLogsPayload, UserAPIKeyAuth from litellm.proxy.collector import SpendEventConsumer +from litellm.proxy.db.db_lookup_gate import DBLookupDeadlineExceeded from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.spend_log_tool_index import response_tool_call_names from litellm.proxy.hooks.proxy_track_cost_callback import ( @@ -1679,6 +1681,92 @@ async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): assert metadata["user_api_key_team_alias"] == "my-team-alias" +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_skips_the_key_lookup_when_the_failure_is_a_db_stall(): + logger = _ProxyDBLogger() + user_api_key_dict = UserAPIKeyAuth(api_key="hashed_key") + request_data = { + "model": "gpt-5.6", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {}, + "litellm_params": {}, + } + + with ( + patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + ) as mock_get_key_object, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + ) as mock_get_team_object, + ): + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=DBLookupDeadlineExceeded("key", 10.0), + user_api_key_dict=user_api_key_dict, + ) + + mock_get_key_object.assert_not_called() + mock_get_team_object.assert_not_called() + mock_update_database.assert_called_once() + metadata = mock_update_database.call_args[1]["kwargs"]["litellm_params"]["metadata"] + assert metadata["status"] == "failure" + assert metadata["user_api_key"] == "hashed_key" + assert metadata["user_api_key_alias"] is None + + +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_still_enriches_metadata_for_a_non_stall_failure(): + """Only a DBLookupDeadlineExceeded skips the key lookup; a transport error + from the provider call must still resolve the key's alias for the failure row.""" + logger = _ProxyDBLogger() + user_api_key_dict = UserAPIKeyAuth(api_key="hashed_key") + request_data = { + "model": "gpt-5.6", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {}, + "litellm_params": {}, + } + + mock_key_obj = MagicMock() + mock_key_obj.key_alias = "my-key-alias" + mock_key_obj.user_id = "my-user-id" + mock_key_obj.team_id = "my-team-id" + mock_key_obj.org_id = None + mock_key_obj.project_id = None + + with ( + patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ) as mock_get_key_object, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + ), + ): + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=httpx.ConnectError("boom"), + user_api_key_dict=user_api_key_dict, + ) + + mock_get_key_object.assert_called_once() + metadata = mock_update_database.call_args[1]["kwargs"]["litellm_params"]["metadata"] + assert metadata["user_api_key_alias"] == "my-key-alias" + + @pytest.mark.asyncio async def test_async_post_call_failure_hook_enriches_missing_team_alias(): """ @@ -2035,9 +2123,15 @@ async def test_track_cost_callback_keeps_guardrail_cost_on_cache_hit(): } with ( - patch("litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock) as mock_increment, # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam - patch("litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock), # test-quality-ok: same function-body import, no injection seam - patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging, # test-quality-ok: same function-body import, no injection seam + patch( + "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock + ) as mock_increment, # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam + patch( + "litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock + ), # test-quality-ok: same function-body import, no injection seam + patch( + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_proxy_logging, # test-quality-ok: same function-body import, no injection seam ): mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()