fix(proxy): fail parked DB lookups at a deadline and flip readiness while they stall (#42654)
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

* fix(proxy): fail parked DB lookups at a deadline and flip readiness while they stall

Under a load burst with a slow authentication database every request parked inside
the pod with no deadline while /health/readiness kept answering 200 (its own ping
gets a fresh connection), so the load balancer kept sending traffic until the pod hit
its memory limit, and the parked requests completed against the provider minutes
after every client had hung up

Every pre-request read (key, team, user, end user, budget, membership, organization,
object permission, jwt mapping, project, proxy budget, spend counter reseed) now runs
under one deadline, PROXY_DB_LOOKUP_DEADLINE_SECONDS (default 10 s). A lookup that
hits it fails the request with the existing 503 "authentication database is
temporarily unreachable" answer, honours allow_requests_on_db_unavailable, and never
triggers the transport reconnect (the transport is fine, the query is slow), which is
what turned the repro's stall into "too many clients". Writes stay unbounded

A deadline hit marks the pod stalled for PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS
(default 30 s, 0 disables), during which /health/readiness answers 503 with
"db": "stalled" behind the same fail-open gate, so the pod leaves rotation before it
fills its memory. The existing litellm_in_flight_requests gauge already exposes the
parked set on /metrics

The deadline is enforced on the wall clock: bounded_db_lookup waits on the lookup
task with asyncio.wait and raises DBLookupDeadlineExceeded when the deadline passes
even if the lookup absorbs its cancellation, where asyncio.wait_for on 3.12+ would
sit on the cancelled task for as long as it takes

The failure spend-log row no longer re-runs the key and team lookups when the
failure itself is a database connection or deadline error, so a request that hit
the deadline is answered after one deadline instead of two

* fix(proxy): bound the spend counter gate wait and narrow the stalled lookup shortcut

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep the global spend lookup on the prisma client handle

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-24 10:09:49 -05:00 • committed by GitHub
parent b8f3ba03b3
commit e135a199ad
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 1016 additions and 330 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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