mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(proxy): fail parked DB lookups at a deadline and flip readiness while they stall (#42654)
* 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:
parent
b8f3ba03b3
commit
e135a199ad
15 changed files with 1016 additions and 330 deletions
|
|
@ -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")))
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
111
tests/test_litellm/proxy/db/test_db_lookup_gate.py
Normal file
111
tests/test_litellm/proxy/db/test_db_lookup_gate.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue