mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-22 00:31:44 +00:00
chore: merge litellm_internal_staging
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
1075b8c073
22 changed files with 1430 additions and 59 deletions
7
.github/ci-coverage-allowlist.yml
vendored
7
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -136,13 +136,6 @@ test_paths:
|
|||
- tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py
|
||||
|
||||
dockerfiles:
|
||||
- reason: >-
|
||||
The componentized images the microservices chart deploys are built by no job; wiring both into
|
||||
the scan workflow costs a full image build each and is deferred to a change that prices the
|
||||
whole set
|
||||
paths:
|
||||
- backend/Dockerfile
|
||||
- gateway/Dockerfile
|
||||
- reason: >-
|
||||
The dashboard container is a static Next.js export served by nginx, and the dashboard build
|
||||
and lint workflows already exercise that output, so building the image adds no signal about it
|
||||
|
|
|
|||
65
.github/workflows/image-scan.yml
vendored
65
.github/workflows/image-scan.yml
vendored
|
|
@ -12,6 +12,11 @@ on:
|
|||
- docker/Dockerfile.non_root
|
||||
- migrations/Dockerfile
|
||||
- migrations/run.py
|
||||
- gateway/Dockerfile
|
||||
- gateway/main.py
|
||||
- backend/Dockerfile
|
||||
- backend/main.py
|
||||
- docker/component_entrypoint.sh
|
||||
- litellm-proxy-extras/**
|
||||
- tests/proxy_migration_tests/**
|
||||
- uv.lock
|
||||
|
|
@ -147,3 +152,63 @@ jobs:
|
|||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v
|
||||
|
||||
gateway-image:
|
||||
name: gateway-image
|
||||
runs-on: ubuntu-latest
|
||||
if: >-
|
||||
github.event_name != 'pull_request' ||
|
||||
github.event.pull_request.head.repo.full_name == github.repository
|
||||
timeout-minutes: 30
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Build gateway image
|
||||
run: docker build -f gateway/Dockerfile -t litellm-gateway-scan:${{ github.sha }} .
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Verify the gateway serves offline as a non-root uid
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-gateway-scan:${{ github.sha }}
|
||||
LITELLM_COMPONENT_PORT: "4000"
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py -v
|
||||
|
||||
backend-image:
|
||||
name: backend-image
|
||||
runs-on: ubuntu-latest
|
||||
if: >-
|
||||
github.event_name != 'pull_request' ||
|
||||
github.event.pull_request.head.repo.full_name == github.repository
|
||||
timeout-minutes: 30
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Build backend image
|
||||
run: docker build -f backend/Dockerfile -t litellm-backend-scan:${{ github.sha }} .
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Verify the backend serves offline as a non-root uid
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-backend-scan:${{ github.sha }}
|
||||
LITELLM_COMPONENT_PORT: "4001"
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py -v
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ import asyncio
|
|||
import datetime
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Coroutine
|
||||
from http.cookiejar import DefaultCookiePolicy
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -80,8 +79,6 @@ from litellm.a2a_protocol.exceptions import A2ALocalhostURLError
|
|||
# Use our custom resolver instead of the default A2A SDK resolver
|
||||
A2ACardResolver: Final = LiteLLMA2ACardResolver
|
||||
|
||||
_BLOCK_ALL_COOKIES: Final = DefaultCookiePolicy(allowed_domains=())
|
||||
|
||||
|
||||
def _set_usage_on_logging_obj(
|
||||
kwargs: dict[str, Any],
|
||||
|
|
@ -770,7 +767,6 @@ async def create_a2a_client(
|
|||
params={"timeout": timeout},
|
||||
)
|
||||
httpx_client: Final = _async_handler.client
|
||||
httpx_client.cookies.jar.set_policy(_BLOCK_ALL_COOKIES)
|
||||
if extra_headers:
|
||||
verbose_proxy_logger.debug("A2A client created with extra_headers=%s", list(extra_headers.keys()))
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_non
|
|||
DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
|
||||
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
|
||||
ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5))
|
||||
ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS: Final = 2000
|
||||
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
||||
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
|
||||
|
|
|
|||
|
|
@ -8,11 +8,12 @@ import sys
|
|||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from http.cookiejar import CookieJar, DefaultCookiePolicy
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn, Optional
|
||||
|
||||
import certifi
|
||||
import httpx
|
||||
from aiohttp import ClientSession, TCPConnector
|
||||
from aiohttp import ClientSession, DummyCookieJar, TCPConnector
|
||||
from httpx import USE_CLIENT_DEFAULT, AsyncHTTPTransport, HTTPTransport
|
||||
from httpx._types import RequestFiles
|
||||
|
||||
|
|
@ -144,6 +145,15 @@ def _handler_may_close_client(client_refcount: int, owns_client: bool) -> bool:
|
|||
return owns_client and client_refcount <= _CLIENT_REFCOUNT_WHEN_HANDLER_IS_SOLE_REFERRER
|
||||
|
||||
|
||||
def blocked_cookie_jar() -> CookieJar:
|
||||
"""A jar that stores no response cookie and sends none, for httpx clients.
|
||||
|
||||
LiteLLM's outbound clients are pooled and shared by every caller, so a cookie one
|
||||
upstream sets would be replayed to every other upstream on a matching domain.
|
||||
"""
|
||||
return CookieJar(policy=DefaultCookiePolicy(allowed_domains=()))
|
||||
|
||||
|
||||
_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS: Final = 5.0
|
||||
_STREAMING_ERROR_BODY_READ_EXECUTOR: Final = concurrent.futures.ThreadPoolExecutor(
|
||||
max_workers=50,
|
||||
|
|
@ -587,6 +597,7 @@ class AsyncHTTPHandler:
|
|||
verify=ssl_config,
|
||||
cert=cert,
|
||||
headers=default_headers,
|
||||
cookies=blocked_cookie_jar(),
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
||||
|
|
@ -1063,6 +1074,7 @@ class AsyncHTTPHandler:
|
|||
def session_factory() -> ClientSession:
|
||||
return ClientSession(
|
||||
connector=TCPConnector(**transport_connector_kwargs),
|
||||
cookie_jar=DummyCookieJar(),
|
||||
trust_env=trust_env,
|
||||
)
|
||||
|
||||
|
|
@ -1132,6 +1144,7 @@ class HTTPHandler:
|
|||
verify=ssl_config,
|
||||
cert=cert,
|
||||
headers=default_headers,
|
||||
cookies=blocked_cookie_jar(),
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -486,6 +486,38 @@ class PrismaWrapper:
|
|||
os.environ[self._db_url_env_var] = _db_url
|
||||
return _db_url
|
||||
|
||||
@property
|
||||
def engine_generation(self) -> int:
|
||||
"""How many query-engine replacements have completed on this wrapper.
|
||||
|
||||
Bumped under `_reconnection_lock` only after a replacement engine has
|
||||
connected, so a change across an await proves a *successful* planned
|
||||
replacement happened in between — a replacement that failed (a real
|
||||
outage) leaves it untouched.
|
||||
"""
|
||||
return self._engine_generation
|
||||
|
||||
async def _reconnection_settled(self) -> None:
|
||||
async with self._reconnection_lock:
|
||||
pass
|
||||
|
||||
async def wait_for_planned_engine_replacement(self, timeout_seconds: float) -> None:
|
||||
"""Wait, bounded, for an in-flight planned engine replacement to finish.
|
||||
|
||||
Both replacement paths (`recreate_prisma_client` and
|
||||
`_safe_refresh_token`) hold `_reconnection_lock` across their whole
|
||||
kill/connect window, so re-acquiring it means the replacement has
|
||||
settled one way or the other. Gives up silently on timeout: a caller
|
||||
that stopped waiting must treat the replacement as not completed and
|
||||
consult `engine_generation` rather than assume success.
|
||||
"""
|
||||
if timeout_seconds <= 0 or not self._reconnection_lock.locked():
|
||||
return
|
||||
try:
|
||||
await asyncio.wait_for(self._reconnection_settled(), timeout=timeout_seconds)
|
||||
except asyncio.TimeoutError:
|
||||
return
|
||||
|
||||
async def recreate_prisma_client(
|
||||
self,
|
||||
new_db_url: str,
|
||||
|
|
|
|||
|
|
@ -871,7 +871,7 @@ async def proxy_shutdown_event():
|
|||
async def _initialize_shared_aiohttp_session():
|
||||
"""Initialize shared aiohttp session for connection reuse with connection limits."""
|
||||
try:
|
||||
from aiohttp import ClientSession, TCPConnector
|
||||
from aiohttp import ClientSession, DummyCookieJar, TCPConnector
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
_build_aiohttp_keepalive_socket_factory,
|
||||
|
|
@ -892,7 +892,7 @@ async def _initialize_shared_aiohttp_session():
|
|||
connector_kwargs["socket_factory"] = socket_factory
|
||||
|
||||
connector: Final = TCPConnector(**connector_kwargs)
|
||||
session: Final = ClientSession(connector=connector)
|
||||
session: Final = ClientSession(connector=connector, cookie_jar=DummyCookieJar())
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"SESSION REUSE: Created shared aiohttp session for connection pooling (ID: %s, limit=%s, limit_per_host=%s)",
|
||||
|
|
|
|||
|
|
@ -3000,6 +3000,13 @@ class PrismaClient:
|
|||
] = [] # mutable-ok: drained queue, mirrors tool_usage_transactions
|
||||
_autorouter_turn_transactions_lock = asyncio.Lock()
|
||||
|
||||
# How long a health probe failure waits for an in-flight planned engine
|
||||
# replacement to settle before deciding whether to report itself. Generous
|
||||
# against a replacement that takes well under a second, and far short of the
|
||||
# reconnect budget an outage-hung `connect()` runs under, so a real outage
|
||||
# is never waited out.
|
||||
PLANNED_ENGINE_REPLACEMENT_SETTLE_SECONDS: ClassVar[float] = 5.0
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
database_url: str,
|
||||
|
|
@ -4970,6 +4977,101 @@ class PrismaClient:
|
|||
else:
|
||||
verbose_proxy_logger.debug("Prisma DB health watchdog observed non-DB error: %s", e)
|
||||
|
||||
def _probe_target_wrapper(self) -> PrismaWrapper:
|
||||
"""The Prisma wrapper a `SELECT 1` health probe actually reaches.
|
||||
|
||||
`health_check()` issues `query_raw`, which `RoutingPrismaWrapper` sends
|
||||
to the reader unless the reader is degraded. The writer's engine state
|
||||
therefore says nothing about a probe that failed against the reader, so
|
||||
the gate has to follow the same routing rule the probe did.
|
||||
"""
|
||||
if isinstance(self.db, RoutingPrismaWrapper):
|
||||
return self.db.writer if self.db.reader_unavailable else self.db.reader
|
||||
return self.db
|
||||
|
||||
async def _run_health_probe(self, wrapper: PrismaWrapper) -> object:
|
||||
"""Issue the `SELECT 1` a health check is made of, against `wrapper`.
|
||||
|
||||
Takes the wrapper rather than re-reading `self.db`, because routing is
|
||||
re-resolved on every attribute access: a reader that recovers between
|
||||
the caller picking its target and the query going out would send the
|
||||
probe to a different engine than the one whose generation the caller is
|
||||
about to check, and attribute the failure to the wrong replacement.
|
||||
"""
|
||||
sql_query: Final = "SELECT 1"
|
||||
response: Final = await wrapper.query_raw(sql_query)
|
||||
return response
|
||||
|
||||
async def _probe_answers_now(self, wrapper: PrismaWrapper) -> bool:
|
||||
try:
|
||||
await self._run_health_probe(wrapper)
|
||||
except Exception as probe_error: # noqa: BLE001 # any failure means the database is not answering
|
||||
verbose_proxy_logger.debug("Prisma health_check() confirmation probe failed: %s", probe_error)
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _planned_engine_replacement_absorbed(
|
||||
self,
|
||||
e: Exception,
|
||||
wrapper: PrismaWrapper,
|
||||
generation_before: int,
|
||||
) -> bool:
|
||||
"""True iff `e` is a connection-class probe failure that a completed
|
||||
planned query-engine replacement explains.
|
||||
|
||||
Planned replacements (RDS IAM token refresh, guarded reconnect) kill the
|
||||
running query engine and spawn a new one. A `SELECT 1` probe that races
|
||||
that sub-second window fails with a transport error against the engine's
|
||||
local HTTP port even though nothing is wrong with the database, and
|
||||
reporting it drives a false-positive `db_exceptions` alert on every
|
||||
replacement.
|
||||
|
||||
Two things must both hold, because neither is sufficient alone. The
|
||||
engine generation must have moved, which says a replacement completed
|
||||
rather than merely being attempted: reconnect attempts during a real
|
||||
outage hold the same lock for tens of seconds, so gating on an in-flight
|
||||
replacement would swallow most of an outage's alerts. And a fresh probe
|
||||
must succeed, because `Prisma.connect()` polls the query engine's own
|
||||
`/status` endpoint rather than round-tripping to the database, so a
|
||||
future engine that binds before it validates its connection pool would
|
||||
let the generation advance with the database still unreachable.
|
||||
|
||||
Waiting for an in-flight replacement to settle is what makes the
|
||||
generation check meaningful, since the generation has not moved yet at
|
||||
the instant the probe fails. The wait is generous against a replacement
|
||||
that takes well under a second and short enough that an outage-hung
|
||||
reconnect is not waited out; a replacement that has not settled by then
|
||||
reports rather than stays silent.
|
||||
"""
|
||||
if not PrismaDBExceptionHandler.is_database_connection_error(e):
|
||||
return False
|
||||
await wrapper.wait_for_planned_engine_replacement(self.PLANNED_ENGINE_REPLACEMENT_SETTLE_SECONDS)
|
||||
if wrapper.engine_generation == generation_before:
|
||||
return False
|
||||
return await self._probe_answers_now(wrapper)
|
||||
|
||||
async def _report_health_check_failure(
|
||||
self,
|
||||
e: Exception,
|
||||
duration: float,
|
||||
traceback_str: str,
|
||||
wrapper: PrismaWrapper,
|
||||
generation_before: int,
|
||||
) -> None:
|
||||
if await self._planned_engine_replacement_absorbed(e, wrapper, generation_before):
|
||||
verbose_proxy_logger.info(
|
||||
"Prisma health_check() connection error raced a planned query-engine replacement; "
|
||||
"not reporting it as a DB exception: %s",
|
||||
e,
|
||||
)
|
||||
return
|
||||
await self.proxy_logging_obj.failure_handler(
|
||||
original_exception=e,
|
||||
duration=duration,
|
||||
call_type="health_check",
|
||||
traceback_str=traceback_str,
|
||||
)
|
||||
|
||||
@backoff.on_exception(
|
||||
backoff.expo,
|
||||
Exception,
|
||||
|
|
@ -4982,13 +5084,10 @@ class PrismaClient:
|
|||
Health check endpoint for the prisma client
|
||||
"""
|
||||
start_time: Final = time.time()
|
||||
probe_wrapper: Final = self._probe_target_wrapper()
|
||||
generation_before: Final = probe_wrapper.engine_generation
|
||||
try:
|
||||
sql_query: Final = "SELECT 1"
|
||||
|
||||
# Execute the raw query
|
||||
# The asterisk before `user_id_list` unpacks the list into separate arguments
|
||||
response: Final = await self.db.query_raw(sql_query)
|
||||
return response
|
||||
return await self._run_health_probe(probe_wrapper)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
|
|
@ -4998,11 +5097,12 @@ class PrismaClient:
|
|||
end_time: Final = time.time()
|
||||
_duration: Final = end_time - start_time
|
||||
asyncio.create_task(
|
||||
self.proxy_logging_obj.failure_handler(
|
||||
original_exception=e,
|
||||
self._report_health_check_failure(
|
||||
e=e,
|
||||
duration=_duration,
|
||||
call_type="health_check",
|
||||
traceback_str=error_traceback,
|
||||
wrapper=probe_wrapper,
|
||||
generation_before=generation_before,
|
||||
)
|
||||
)
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -109,6 +109,7 @@ from litellm.router_utils.common_utils import (
|
|||
filter_team_based_models,
|
||||
filter_web_search_deployments,
|
||||
resolve_model_group_alias,
|
||||
truncate_fallback_error_detail,
|
||||
)
|
||||
from litellm.router_utils.cooldown_cache import CooldownCache
|
||||
from litellm.router_utils.cooldown_handlers import (
|
||||
|
|
@ -342,6 +343,12 @@ def _replay_live_router_model_cost() -> None:
|
|||
set_live_deployment_replay(_replay_live_router_model_cost)
|
||||
|
||||
|
||||
# Kwargs that log_retry must not copy into a retry breadcrumb. The breadcrumbs reach spend
|
||||
# logs and logging callbacks, and these carry either the request payload or router-internal
|
||||
# walk state rather than anything that identifies the failed attempt.
|
||||
RETRY_BREADCRUMB_EXCLUDED_KWARGS: Final = frozenset(("messages", "original_function", "attempted_targets"))
|
||||
|
||||
|
||||
class Router:
|
||||
model_names: set = set()
|
||||
cache_responses: bool | None = False
|
||||
|
|
@ -6361,17 +6368,16 @@ class Router:
|
|||
return response
|
||||
except Exception as new_exception:
|
||||
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
fallback_failure_exception_str = redact_string(str(new_exception))
|
||||
fallback_failure_exception_str = truncate_fallback_error_detail(redact_string(str(new_exception)))
|
||||
cooldown_info: Final = await _async_get_cooldown_deployments_with_debug_info(
|
||||
litellm_router_instance=self,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
verbose_router_logger.error(
|
||||
"litellm.router.py::async_function_with_fallbacks() - "
|
||||
"Error occurred while trying to do fallbacks - %s\n%s\n"
|
||||
"Error occurred while trying to do fallbacks - %s\n"
|
||||
"Debug Information:\nCooldown Deployments=%s",
|
||||
fallback_failure_exception_str,
|
||||
redact_string(traceback.format_exc()),
|
||||
cooldown_info,
|
||||
)
|
||||
|
||||
|
|
@ -7162,7 +7168,7 @@ class Router:
|
|||
k,
|
||||
v,
|
||||
) in kwargs.items(): # log everything in kwargs except the old previous_models value - prevent nesting
|
||||
if k not in [_metadata_var, "messages", "original_function"]:
|
||||
if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS:
|
||||
previous_model[k] = v
|
||||
elif k == _metadata_var and isinstance(v, dict):
|
||||
previous_model[_metadata_var] = {}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ if TYPE_CHECKING:
|
|||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.types.router import CredentialLiteLLMParams
|
||||
|
||||
|
|
@ -43,6 +44,22 @@ def resolve_model_group_alias(model_group_alias: object, model: str) -> str | No
|
|||
return target
|
||||
|
||||
|
||||
def truncate_fallback_error_detail(detail: str) -> str:
|
||||
"""
|
||||
Bound a fallback failure detail before it is logged or appended to an exception message.
|
||||
|
||||
Each level of the fallback walk records the failure of the level below it, so an
|
||||
untruncated detail carries every nested failure with it and grows superlinearly with
|
||||
the number of attempted model groups. One deterministic pre-network failure walked
|
||||
through a small fallback graph is enough to turn that into hundreds of megabytes of
|
||||
output on the event-loop thread, which starves the process that produced it.
|
||||
"""
|
||||
if len(detail) <= ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS:
|
||||
return detail
|
||||
dropped: Final = len(detail) - ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS
|
||||
return f"{detail[:ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS]}... [truncated {dropped} characters]"
|
||||
|
||||
|
||||
def get_litellm_params_sensitive_credential_hash(litellm_params: dict) -> str:
|
||||
"""
|
||||
Hash of the credential params, used for mapping the file id to the right model
|
||||
|
|
|
|||
|
|
@ -1,3 +1,6 @@
|
|||
import hashlib
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
|
|
@ -19,6 +22,52 @@ else:
|
|||
LitellmRouter = Any
|
||||
|
||||
|
||||
def fallback_attempt_key(fallback_target: object) -> str | None:
|
||||
"""
|
||||
Identity of one fallback attempt, so the same attempt is never made twice per request.
|
||||
|
||||
A bare model group name and a `{"model": name}` entry describe the same attempt. An
|
||||
entry carrying anything else describes a different one and keeps its own identity: a
|
||||
client-side fallback list overrides request params such as `messages`, and the router
|
||||
re-targets the group that just failed by attaching `_target_order` or
|
||||
`_excluded_deployment_ids` to select a different set of deployments inside it. The
|
||||
payload is hashed rather than kept, so a large `messages` override does not make the
|
||||
request hold a second copy of itself.
|
||||
|
||||
Returns None for a shape with no usable identity, which is never skipped.
|
||||
"""
|
||||
if isinstance(fallback_target, str):
|
||||
return fallback_target
|
||||
if not isinstance(fallback_target, dict):
|
||||
return None
|
||||
model: Final = fallback_target.get("model")
|
||||
if tuple(fallback_target) == ("model",) and isinstance(model, str):
|
||||
return model
|
||||
serialized: Final = json.dumps(fallback_target, sort_keys=True, default=str)
|
||||
return hashlib.sha256(serialized.encode()).hexdigest()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class AttemptedFallbackTargets:
|
||||
"""
|
||||
The fallback attempts a single request has already made.
|
||||
|
||||
One instance is created on the first fallback hop and shared by reference for the rest
|
||||
of the walk, so an attempt made in one branch is not repeated in a sibling branch.
|
||||
Without it the walk enumerates paths rather than attempts: a fallback graph containing
|
||||
a cycle retries one deterministic failure once per path through the cycle, and a
|
||||
client-side fallback list is re-walked at every level of the recursion.
|
||||
"""
|
||||
|
||||
keys: frozenset[str] = frozenset()
|
||||
|
||||
def __contains__(self, key: str) -> bool:
|
||||
return key in self.keys
|
||||
|
||||
def record(self, key: str) -> None:
|
||||
self.keys = self.keys | frozenset((key,))
|
||||
|
||||
|
||||
def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool:
|
||||
"""
|
||||
Handles wildcard routing scenario
|
||||
|
|
@ -106,7 +155,14 @@ async def run_async_fallback(
|
|||
fallback_model_group: List[str] of fallback model groups. example: ["gpt-4", "gpt-3.5-turbo"]
|
||||
original_model_group: The original model group. example: "gpt-3.5-turbo"
|
||||
original_exception: The original exception.
|
||||
**kwargs: Keyword arguments.
|
||||
**kwargs: Keyword arguments. `attempted_targets` carries the fallback attempts
|
||||
already made for this request, created on the first hop and shared by reference
|
||||
for the rest of the walk. A target already in it is skipped, so neither a
|
||||
fallback graph that loops back on itself nor a client-side fallback list
|
||||
re-walked at each level can repeat an attempt that has already failed. Identity
|
||||
comes from `fallback_attempt_key`, so an entry that overrides request params or
|
||||
re-targets the failed group with a different deployment selection stays distinct
|
||||
from a bare name.
|
||||
|
||||
Returns:
|
||||
The response from the successful fallback model group.
|
||||
|
|
@ -120,10 +176,27 @@ async def run_async_fallback(
|
|||
|
||||
error_from_fallbacks = original_exception
|
||||
fallback_errors = (get_fallback_error_info(original_exception),)
|
||||
# Read out of kwargs and narrowed here rather than declared as a parameter: every caller
|
||||
# reaches this function by spreading a loosely-typed kwargs dict, so a declared parameter
|
||||
# would carry an annotation that no call site can actually be checked against.
|
||||
carried_targets: Final = kwargs.get("attempted_targets")
|
||||
attempted: Final = (
|
||||
carried_targets if isinstance(carried_targets, AttemptedFallbackTargets) else AttemptedFallbackTargets()
|
||||
)
|
||||
attempted.record(original_model_group)
|
||||
|
||||
for mg in fallback_model_group:
|
||||
if mg == original_model_group:
|
||||
continue
|
||||
attempt_key = fallback_attempt_key(mg)
|
||||
if attempt_key is not None:
|
||||
if attempt_key in attempted:
|
||||
verbose_router_logger.info(
|
||||
"Skipping fallback to model_group = %s, already attempted for this request",
|
||||
mask_sensitive_structure(mg),
|
||||
)
|
||||
continue
|
||||
attempted.record(attempt_key)
|
||||
try:
|
||||
# LOGGING
|
||||
kwargs = litellm_router.log_retry(kwargs=kwargs, e=original_exception)
|
||||
|
|
@ -138,6 +211,7 @@ async def run_async_fallback(
|
|||
fallback_depth = fallback_depth + 1
|
||||
kwargs["fallback_depth"] = fallback_depth
|
||||
kwargs["max_fallbacks"] = max_fallbacks
|
||||
kwargs["attempted_targets"] = attempted
|
||||
if include_fallback_errors:
|
||||
kwargs["include_fallback_errors"] = include_fallback_errors
|
||||
response = await litellm_router.async_function_with_fallbacks(*args, **kwargs)
|
||||
|
|
|
|||
|
|
@ -3479,6 +3479,7 @@ all_litellm_params = (
|
|||
"user_continue_message",
|
||||
"fallback_depth",
|
||||
"max_fallbacks",
|
||||
"attempted_targets",
|
||||
"max_budget",
|
||||
"budget_duration",
|
||||
"use_in_pass_through",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,187 @@
|
|||
"""Image-level regression net for the prisma bake in the componentized images.
|
||||
|
||||
The gateway and backend serve requests; they never shell out to the Prisma CLI
|
||||
(``PrismaManager.setup_database`` is reachable only from ``proxy_cli.py``, which
|
||||
uvicorn'ing ``gateway.main:app`` bypasses). What they do need is the generated
|
||||
client's baked query engine, and prisma-python resolves those baked paths
|
||||
eagerly, with an existence check that propagates EACCES rather than skipping the
|
||||
candidate. An engine baked under a build-time ``HOME`` is therefore unreadable to
|
||||
any other runtime uid, and the process dies during startup before
|
||||
``PRISMA_QUERY_ENGINE_BINARY`` is ever consulted.
|
||||
|
||||
That is what an OpenShift ``restricted-v2`` namespace produces: the image
|
||||
``USER`` is ignored and an arbitrary uid in GID 0 is assigned instead. The
|
||||
symptom is not a degraded proxy, it is a proxy that does not serve at all.
|
||||
|
||||
Booting the image the way that deployment does, and requiring it to answer a
|
||||
request with a live database connection, is what catches the whole class:
|
||||
a boot as the default uid, or one that reaches the internet, passes even when
|
||||
the bake is unusable everywhere it actually ships.
|
||||
|
||||
Gated on LITELLM_IMAGE (the tag of the image to exercise) so it is skipped in
|
||||
the normal unit-test run and exercised only where an image has been built (the
|
||||
image-scan workflow). Requires a working docker CLI.
|
||||
"""
|
||||
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import os
|
||||
import pytest
|
||||
|
||||
IMAGE = os.getenv("LITELLM_IMAGE")
|
||||
POSTGRES_IMAGE = os.getenv("LITELLM_TEST_POSTGRES_IMAGE", "postgres:16-alpine")
|
||||
CURL_IMAGE = os.getenv("LITELLM_TEST_CURL_IMAGE", "curlimages/curl:8.11.1")
|
||||
COMPONENT_PORT = os.getenv("LITELLM_COMPONENT_PORT", "4000")
|
||||
NON_ROOT_UID = "12345:0"
|
||||
STARTUP_TIMEOUT_SECONDS = int(os.getenv("LITELLM_COMPONENT_STARTUP_TIMEOUT", "180"))
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(IMAGE is None, reason="requires a built image (set LITELLM_IMAGE)"),
|
||||
pytest.mark.skipif(shutil.which("docker") is None, reason="requires the docker CLI"),
|
||||
]
|
||||
|
||||
|
||||
def _docker(*args: str, check: bool = True) -> subprocess.CompletedProcess:
|
||||
return subprocess.run(
|
||||
["docker", *args], capture_output=True, text=True, check=check
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def offline_stack():
|
||||
"""A component container and a fresh Postgres on a network with no egress.
|
||||
|
||||
NON_ROOT_UID is an arbitrary uid in GID 0, the shape OpenShift restricted-v2
|
||||
assigns. Postgres and curl are pulled while egress still exists, because the
|
||||
``--internal`` network below has none: that is what makes a prisma engine
|
||||
download (binaries.prisma.sh / npm) fail rather than mask a bake that is not
|
||||
self-contained.
|
||||
|
||||
The container runs with DISABLE_SCHEMA_UPDATE, since applying the schema is
|
||||
the migration job's responsibility in this topology and needs the Prisma CLI
|
||||
these images deliberately omit, and with LITELLM_LOCAL_MODEL_COST_MAP, or the
|
||||
proxy spends the whole startup budget timing out on a cost-map fetch over the
|
||||
network it does not have.
|
||||
|
||||
Yields (network_name, component_container). Both are torn down afterwards.
|
||||
"""
|
||||
run_id = f"componentserve-{uuid.uuid4().hex[:8]}"
|
||||
network = f"{run_id}-net"
|
||||
pg = f"{run_id}-pg"
|
||||
component = f"{run_id}-app"
|
||||
|
||||
_docker("pull", "--quiet", POSTGRES_IMAGE)
|
||||
_docker("pull", "--quiet", CURL_IMAGE)
|
||||
_docker("network", "create", "--internal", network)
|
||||
try:
|
||||
_docker(
|
||||
"run", "-d", "--name", pg, "--network", network,
|
||||
"-e", "POSTGRES_PASSWORD=pw", "-e", "POSTGRES_DB=litellm",
|
||||
POSTGRES_IMAGE,
|
||||
)
|
||||
_wait_until_postgres_ready(pg)
|
||||
assert IMAGE is not None
|
||||
_docker(
|
||||
"run", "-d", "--name", component, "--network", network,
|
||||
"--user", NON_ROOT_UID,
|
||||
"-e", f"DATABASE_URL=postgresql://postgres:pw@{pg}:5432/litellm",
|
||||
"-e", "LITELLM_MASTER_KEY=sk-component-serve-test",
|
||||
"-e", "DISABLE_SCHEMA_UPDATE=true",
|
||||
"-e", "LITELLM_LOCAL_MODEL_COST_MAP=True",
|
||||
IMAGE,
|
||||
)
|
||||
yield network, component
|
||||
finally:
|
||||
_docker("logs", component, check=False)
|
||||
_docker("rm", "-f", component, check=False)
|
||||
_docker("rm", "-f", pg, check=False)
|
||||
_docker("network", "rm", network, check=False)
|
||||
|
||||
|
||||
def _wait_until_postgres_ready(pg: str, attempts: int = 60) -> None:
|
||||
for _ in range(attempts):
|
||||
running = _docker(
|
||||
"ps", "--filter", f"name={pg}", "--filter", "status=running",
|
||||
"--format", "{{.Names}}", check=False,
|
||||
).stdout
|
||||
if pg not in running:
|
||||
logs = _docker("logs", pg, check=False)
|
||||
pytest.fail(f"postgres container is not running:\n{logs.stdout}\n{logs.stderr}")
|
||||
ready = _docker(
|
||||
"exec", pg, "pg_isready", "-U", "postgres", "-d", "litellm", check=False
|
||||
)
|
||||
if ready.returncode == 0:
|
||||
return
|
||||
time.sleep(1)
|
||||
pytest.fail(f"postgres never became ready after {attempts}s")
|
||||
|
||||
|
||||
def _container_logs(container: str) -> str:
|
||||
logs = _docker("logs", container, check=False)
|
||||
return f"stdout:\n{logs.stdout}\nstderr:\n{logs.stderr}"
|
||||
|
||||
|
||||
def _is_running(container: str) -> bool:
|
||||
return bool(
|
||||
_docker(
|
||||
"ps", "--filter", f"name={container}", "--filter", "status=running",
|
||||
"--format", "{{.Names}}", check=False,
|
||||
).stdout.strip()
|
||||
)
|
||||
|
||||
|
||||
def _readiness(network: str, component: str) -> subprocess.CompletedProcess:
|
||||
"""Ask the component for its readiness, from a peer on the same egress-less network."""
|
||||
return _docker(
|
||||
"run", "--rm", "--network", network, CURL_IMAGE,
|
||||
"--silent", "--max-time", "10",
|
||||
f"http://{component}:{COMPONENT_PORT}/health/readiness",
|
||||
check=False,
|
||||
)
|
||||
|
||||
|
||||
def test_component_serves_offline_as_non_root_uid(offline_stack):
|
||||
"""The component answers a request with a live DB connection, offline, as an arbitrary uid.
|
||||
|
||||
On the pre-fix image this never gets a response: the engine baked under
|
||||
/home/nonroot (mode 0700, owned by uid 65532) raises
|
||||
``PermissionError: .../query-engine-linux-...`` out of pathlib and uvicorn
|
||||
reports ``Application startup failed. Exiting.``. A bake at the fixed,
|
||||
world-readable /opt/prisma is what lets any uid start the client.
|
||||
|
||||
`db: connected` is the load-bearing part of the assertion: it means the
|
||||
query engine binary was found, executed, and reached Postgres. A liveness
|
||||
probe alone would pass on an image whose engine never resolved.
|
||||
"""
|
||||
network, component = offline_stack
|
||||
|
||||
deadline = time.time() + STARTUP_TIMEOUT_SECONDS
|
||||
probe = None
|
||||
while time.time() < deadline:
|
||||
if not _is_running(component):
|
||||
pytest.fail(
|
||||
f"the component exited during startup as uid {NON_ROOT_UID} with no egress. "
|
||||
"The prisma bake is not readable to a uid other than the one that built it, "
|
||||
"so the proxy does not serve at all.\n"
|
||||
f"{_container_logs(component)}"
|
||||
)
|
||||
probe = _readiness(network, component)
|
||||
if probe.returncode == 0 and probe.stdout.strip():
|
||||
break
|
||||
time.sleep(2)
|
||||
|
||||
assert probe is not None and probe.returncode == 0 and probe.stdout.strip(), (
|
||||
f"/health/readiness never answered within {STARTUP_TIMEOUT_SECONDS}s as uid "
|
||||
f"{NON_ROOT_UID} with no egress.\n{_container_logs(component)}"
|
||||
)
|
||||
|
||||
payload = json.loads(probe.stdout)
|
||||
assert payload.get("db") == "connected", (
|
||||
f"the component answered but its database is {payload.get('db')!r}, so the baked "
|
||||
f"query engine did not resolve as uid {NON_ROOT_UID}.\nresponse: {probe.stdout}\n"
|
||||
f"{_container_logs(component)}"
|
||||
)
|
||||
|
|
@ -174,21 +174,15 @@ _RPC_REPLY = {
|
|||
|
||||
_AGENT_A_HEADERS = {"x-agent-token": "token-for-a", "x-tenant": "tenant-a"}
|
||||
_AGENT_B_HEADERS = {"x-agent-token": "token-for-b", "x-tenant": "tenant-b"}
|
||||
_UPSTREAM_SESSION_COOKIE = "a2a_session=only-agent-a-may-hold-this; Path=/"
|
||||
|
||||
|
||||
class _RequestRecorder:
|
||||
"""Records the headers httpx put on the wire, per outbound request.
|
||||
"""Records the headers httpx put on the wire, per outbound request."""
|
||||
|
||||
``cookie_from_tenant`` makes that tenant's agent answer with a Set-Cookie, standing in
|
||||
for an upstream that issues a session cookie.
|
||||
"""
|
||||
|
||||
def __init__(self, cookie_from_tenant: str | None = None):
|
||||
def __init__(self):
|
||||
self.card_requests = []
|
||||
self.rpc_requests = []
|
||||
self.client = None
|
||||
self.cookie_from_tenant = cookie_from_tenant
|
||||
|
||||
def __call__(self, request: httpx.Request) -> httpx.Response:
|
||||
headers = {k.lower(): v for k, v in request.headers.items()}
|
||||
|
|
@ -196,8 +190,6 @@ class _RequestRecorder:
|
|||
self.card_requests.append(headers)
|
||||
return httpx.Response(200, json=_AGENT_CARD)
|
||||
self.rpc_requests.append(headers)
|
||||
if self.cookie_from_tenant is not None and headers.get("x-tenant") == self.cookie_from_tenant:
|
||||
return httpx.Response(200, json=_RPC_REPLY, headers={"set-cookie": _UPSTREAM_SESSION_COOKIE})
|
||||
return httpx.Response(200, json=_RPC_REPLY)
|
||||
|
||||
|
||||
|
|
@ -205,15 +197,14 @@ def _a2a_client_cache_key(timeout: float) -> str:
|
|||
return "async_httpx_client" + f"timeout_{timeout}" + httpxSpecialProvider.A2AProvider
|
||||
|
||||
|
||||
async def _seed_shared_a2a_client(cookie_from_tenant: str | None = None) -> _RequestRecorder:
|
||||
async def _seed_shared_a2a_client() -> _RequestRecorder:
|
||||
"""Put the one A2A client the cache will hand out behind a mock transport.
|
||||
|
||||
Seeding has to happen on the test's own event loop, because the client cache keys on
|
||||
it. The injected client is a real httpx.AsyncClient, so the merge of per-request
|
||||
headers over client defaults, and httpx's own cookie handling, which is what these
|
||||
tests are about, stay real.
|
||||
headers over client defaults, which is what these tests are about, stays real.
|
||||
"""
|
||||
recorder = _RequestRecorder(cookie_from_tenant=cookie_from_tenant)
|
||||
recorder = _RequestRecorder()
|
||||
handler = AsyncHTTPHandler(timeout=DEFAULT_A2A_AGENT_TIMEOUT)
|
||||
owned_client = handler.client
|
||||
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(recorder))
|
||||
|
|
@ -333,17 +324,21 @@ async def test_agent_card_fetch_carries_the_callers_headers(isolated_client_cach
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_agents_session_cookie_never_reaches_another_agent(isolated_client_cache):
|
||||
"""One pooled client is also one httpx cookie jar. httpx stores every Set-Cookie on the
|
||||
client and replays it on any later request to a matching domain, so an agent's session
|
||||
cookie would ride along on a different agent's call to the same host."""
|
||||
recorder = await _seed_shared_a2a_client(cookie_from_tenant="tenant-a")
|
||||
async def test_the_pooled_a2a_client_arrives_with_cookie_persistence_disabled(isolated_client_cache):
|
||||
"""create_a2a_client takes its client from the shared builder rather than building one,
|
||||
and the builder is what refuses to persist cookies. This pins the join between those
|
||||
two facts, so the A2A path cannot quietly start acquiring a client that keeps a jar.
|
||||
|
||||
client_a = await create_a2a_client(base_url="http://127.0.0.1:9", extra_headers=_AGENT_A_HEADERS)
|
||||
await _send_message(client_a, _send_request("a"))
|
||||
client_b = await create_a2a_client(base_url="http://127.0.0.1:9", extra_headers=_AGENT_B_HEADERS)
|
||||
await _send_message(client_b, _send_request("b"))
|
||||
test_callers_with_different_headers_reuse_one_pooled_client pins the other half, that
|
||||
create_a2a_client hands back exactly this cached client."""
|
||||
handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.A2AProvider,
|
||||
params={"timeout": DEFAULT_A2A_AGENT_TIMEOUT},
|
||||
)
|
||||
request = httpx.Request("GET", "https://agent-a.example.com/")
|
||||
handler.client.cookies.extract_cookies(
|
||||
httpx.Response(200, headers={"set-cookie": "SESSION=only-agent-a-may-hold-this"}, request=request)
|
||||
)
|
||||
|
||||
assert dict(recorder.client.cookies) == {}, "the shared client kept an agent's session cookie"
|
||||
assert "cookie" not in recorder.card_requests[-1]
|
||||
assert "cookie" not in recorder.rpc_requests[-1]
|
||||
assert dict(handler.client.cookies) == {}, "the pooled A2A client kept an upstream's cookie"
|
||||
await handler.close()
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ def test_create_aiohttp_transport_sets_enable_cleanup_closed_when_needed(monkeyp
|
|||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")):
|
||||
transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport(
|
||||
shared_session=None
|
||||
)
|
||||
|
|
@ -36,7 +36,7 @@ def test_create_aiohttp_transport_omits_enable_cleanup_closed_when_not_needed(
|
|||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")):
|
||||
transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport(
|
||||
shared_session=None
|
||||
)
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ def test_socket_factory_omitted_when_disabled(monkeypatch):
|
|||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")):
|
||||
_invoke_connector_factory(http_handler_module)
|
||||
|
||||
assert mock_tcp_connector.call_count >= 1
|
||||
|
|
@ -55,7 +55,7 @@ def test_socket_factory_attached_when_enabled(monkeypatch):
|
|||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")):
|
||||
_invoke_connector_factory(http_handler_module)
|
||||
|
||||
assert mock_tcp_connector.call_count >= 1
|
||||
|
|
@ -77,7 +77,7 @@ def test_socket_factory_skipped_on_old_aiohttp(monkeypatch):
|
|||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")):
|
||||
_invoke_connector_factory(http_handler_module)
|
||||
|
||||
assert mock_tcp_connector.call_count >= 1
|
||||
|
|
|
|||
|
|
@ -1171,3 +1171,77 @@ async def test_client_handed_out_by_async_cache_survives_eviction_and_collection
|
|||
assert not consumer_client.is_closed
|
||||
|
||||
await consumer_client.aclose()
|
||||
|
||||
|
||||
_SET_COOKIE = "SESSION=upstream-a-secret; Path=/"
|
||||
|
||||
|
||||
def _cookie_recorder():
|
||||
"""A transport that hands out a Set-Cookie once, and records what comes back."""
|
||||
seen = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen.append(request.headers.get("cookie"))
|
||||
if request.url.path == "/set":
|
||||
return httpx.Response(200, headers={"set-cookie": _SET_COOKIE})
|
||||
return httpx.Response(200)
|
||||
|
||||
return handler, seen
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_client_never_replays_one_upstreams_cookie_to_another():
|
||||
"""LiteLLM's async clients are pooled and shared by every caller, so a cookie one
|
||||
upstream sets would be attached to every later request on a matching domain, reaching
|
||||
a different tenant's upstream. The client must persist no response cookie."""
|
||||
handler, seen = _cookie_recorder()
|
||||
http_handler = AsyncHTTPHandler()
|
||||
client = http_handler.client
|
||||
client._transport = httpx.MockTransport(handler)
|
||||
|
||||
await client.get("https://upstream-a.example.com/set")
|
||||
await client.get("https://upstream-b.example.com/rpc")
|
||||
await client.aclose()
|
||||
|
||||
assert dict(client.cookies) == {}, "the shared client stored an upstream's cookie"
|
||||
assert seen == [None, None]
|
||||
|
||||
|
||||
def test_sync_client_never_replays_one_upstreams_cookie_to_another():
|
||||
"""Same invariant on the sync client, which is pooled the same way."""
|
||||
handler, seen = _cookie_recorder()
|
||||
http_handler = HTTPHandler()
|
||||
client = http_handler.client
|
||||
client._transport = httpx.MockTransport(handler)
|
||||
|
||||
client.get("https://upstream-a.example.com/set")
|
||||
client.get("https://upstream-b.example.com/rpc")
|
||||
client.close()
|
||||
|
||||
assert dict(client.cookies) == {}
|
||||
assert seen == [None, None]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aiohttp_session_never_replays_one_upstreams_cookie_to_another():
|
||||
"""The httpx jar is not the only one. AiohttpTransport is litellm's default transport
|
||||
and the aiohttp ClientSession keeps its own cookie jar, which httpx-level assertions
|
||||
cannot see, so blocking only the httpx jar leaves the leak intact on the real path.
|
||||
|
||||
aiohttp's default jar refuses cookies for IP hosts, so this drives a hostname. An
|
||||
IP-addressed check passes whether or not the session jar is blocked."""
|
||||
from aiohttp import DummyCookieJar
|
||||
from yarl import URL
|
||||
|
||||
http_handler = AsyncHTTPHandler(timeout=61.0)
|
||||
transport = http_handler.client._transport
|
||||
assert isinstance(transport, LiteLLMAiohttpTransport), "aiohttp is no longer the default transport"
|
||||
|
||||
session = transport.client() if callable(transport.client) else transport.client
|
||||
jar = session.cookie_jar
|
||||
assert isinstance(jar, DummyCookieJar)
|
||||
|
||||
jar.update_cookies({"SESSION": "upstream-a-secret"}, URL("https://upstream-a.example.com"))
|
||||
assert len(jar) == 0
|
||||
assert dict(jar.filter_cookies(URL("https://upstream-a.example.com"))) == {}
|
||||
await session.close()
|
||||
|
|
|
|||
|
|
@ -6,13 +6,24 @@ subprocess), and the engine-death watcher / in-flight transport-error
|
|||
retries must not treat that planned restart as a crash and recreate the
|
||||
client a second time.
|
||||
|
||||
A planned restart is also invisible to the database itself, so a ``SELECT 1``
|
||||
health probe that races the kill/connect window fails with a transport error
|
||||
against the engine's local HTTP port. That failure must not be reported as a
|
||||
``db_exceptions`` DB failure, or every IAM refresh cycle raises a false alarm.
|
||||
|
||||
Symbols pinned here:
|
||||
- ``PrismaWrapper._expected_engine_deaths``
|
||||
- ``PrismaWrapper._engine_generation``
|
||||
- ``PrismaWrapper.engine_generation``
|
||||
- ``PrismaWrapper.wait_for_planned_engine_replacement``
|
||||
- ``PrismaWrapper.on_engine_replaced``
|
||||
- ``PrismaWrapper.recreate_prisma_client`` (expected_generation guard)
|
||||
- ``PrismaWrapper._safe_refresh_token`` (refresh coalescing)
|
||||
- ``RoutingPrismaWrapper.recreate_prisma_client`` (guard forwarding)
|
||||
- ``PrismaClient.health_check`` (planned-replacement alert suppression)
|
||||
- ``PrismaClient._probe_target_wrapper``
|
||||
- ``PrismaClient._probe_answers_now``
|
||||
- ``PrismaClient._planned_engine_replacement_absorbed``
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
|
@ -21,22 +32,29 @@ import signal
|
|||
import sys
|
||||
import urllib.parse
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, List
|
||||
from unittest.mock import AsyncMock, MagicMock, call, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from prisma import Prisma as GeneratedPrisma
|
||||
from prisma.engine.errors import EngineConnectionError
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_prisma_binary():
|
||||
"""Mock prisma.Prisma to avoid requiring generated Prisma binaries for unit tests."""
|
||||
mock_module = MagicMock()
|
||||
# Production code isinstance-checks against this, which a bare MagicMock
|
||||
# attribute cannot satisfy.
|
||||
mock_module.engine.errors.EngineConnectionError = EngineConnectionError
|
||||
with patch.dict(sys.modules, {"prisma": mock_module}):
|
||||
yield mock_module
|
||||
|
||||
|
|
@ -49,6 +67,87 @@ def _make_wrapper(engine_pid: int = 111, iam: bool = False) -> PrismaWrapper:
|
|||
return PrismaWrapper(original_prisma=prisma, iam_token_db_auth=iam)
|
||||
|
||||
|
||||
def _make_prisma_client(db: Any) -> PrismaClient:
|
||||
"""A ``PrismaClient`` whose ``db`` is a real wrapper and whose alerting
|
||||
hook is observable."""
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.failure_handler = AsyncMock()
|
||||
client = PrismaClient(
|
||||
database_url="postgresql://user:pass@localhost:5432/db",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
client.db = db
|
||||
client._db_watchdog_reconnect_timeout_seconds = 5.0
|
||||
return client
|
||||
|
||||
|
||||
_real_asyncio_sleep = asyncio.sleep
|
||||
|
||||
|
||||
async def _yield_to_loop(times: int = 10) -> None:
|
||||
"""Let already-scheduled tasks make progress.
|
||||
|
||||
Bound to the real ``asyncio.sleep`` at import time: the tests below patch
|
||||
``asyncio.sleep`` to skip the SIGTERM/SIGKILL grace, and an ``AsyncMock``
|
||||
stand-in never yields to the event loop, which would silently leave every
|
||||
background task un-started and the assertions vacuous.
|
||||
"""
|
||||
for _ in range(times):
|
||||
await _real_asyncio_sleep(0)
|
||||
|
||||
|
||||
async def _await_health_check_reports() -> None:
|
||||
"""Await the fire-and-forget reporting tasks ``health_check()`` scheduled.
|
||||
|
||||
Selected by coroutine qualname rather than by draining every pending task,
|
||||
so an unrelated background task can never make these assertions pass by
|
||||
accident.
|
||||
"""
|
||||
reports = [
|
||||
task
|
||||
for task in asyncio.all_tasks()
|
||||
if getattr(task.get_coro(), "__qualname__", "")
|
||||
== "PrismaClient._report_health_check_failure"
|
||||
]
|
||||
if reports:
|
||||
await asyncio.gather(*reports, return_exceptions=True)
|
||||
|
||||
|
||||
def _fails_then_answers(error: Exception, failures: int = 3) -> Any:
|
||||
"""Raise ``error`` for the first ``failures`` probes, then answer.
|
||||
|
||||
``health_check`` retries up to three times, so this exhausts the retries and
|
||||
still lets the confirmation probe that decides suppression succeed. Without
|
||||
that, a test would report for the wrong reason: the confirmation probe would
|
||||
fail too, masking whether the error type was classified at all.
|
||||
"""
|
||||
seen: List[int] = []
|
||||
|
||||
async def _query_raw(_sql: str) -> Any:
|
||||
seen.append(1)
|
||||
if len(seen) <= failures:
|
||||
raise error
|
||||
return [{"?column?": 1}]
|
||||
|
||||
return _query_raw
|
||||
|
||||
|
||||
def _blocking_replacement(gate: asyncio.Event, fail: bool = False) -> MagicMock:
|
||||
"""A replacement Prisma whose ``connect()`` parks until ``gate`` is set.
|
||||
|
||||
Holds ``_reconnection_lock`` open for as long as the test needs, which is
|
||||
how a health probe is made to fail *while* a planned replacement is in
|
||||
flight rather than after it.
|
||||
"""
|
||||
|
||||
async def _connect(*_: Any, **__: Any) -> None:
|
||||
await gate.wait()
|
||||
if fail:
|
||||
raise ConnectionRefusedError("database is down")
|
||||
|
||||
return MagicMock(connect=AsyncMock(side_effect=_connect))
|
||||
|
||||
|
||||
def _token_db_url(created: datetime, expires_in: int = 900) -> str:
|
||||
"""Build a DATABASE_URL whose password is a parseable RDS IAM token."""
|
||||
token = (
|
||||
|
|
@ -603,3 +702,348 @@ async def test_recreate_caps_expected_engine_deaths_set(mock_prisma_binary):
|
|||
await wrapper.recreate_prisma_client("postgresql://new")
|
||||
|
||||
assert wrapper._expected_engine_deaths == {111}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wait_for_planned_engine_replacement_returns_once_recreate_settles(
|
||||
mock_prisma_binary,
|
||||
):
|
||||
wrapper = _make_wrapper(engine_pid=111)
|
||||
gate = asyncio.Event()
|
||||
mock_prisma_binary.Prisma.return_value = _blocking_replacement(gate)
|
||||
|
||||
with patch("os.kill"), patch("asyncio.sleep", new_callable=AsyncMock):
|
||||
recreate = asyncio.create_task(
|
||||
wrapper.recreate_prisma_client("postgresql://new")
|
||||
)
|
||||
await _yield_to_loop()
|
||||
assert wrapper._reconnection_lock.locked() is True
|
||||
|
||||
waiter = asyncio.create_task(wrapper.wait_for_planned_engine_replacement(5.0))
|
||||
await _yield_to_loop()
|
||||
blocked_while_in_flight = not waiter.done()
|
||||
|
||||
gate.set()
|
||||
await recreate
|
||||
await waiter
|
||||
|
||||
assert {
|
||||
"blocked_while_in_flight": blocked_while_in_flight,
|
||||
"generation": wrapper.engine_generation,
|
||||
} == {"blocked_while_in_flight": True, "generation": 1}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wait_for_planned_engine_replacement_gives_up_at_timeout(
|
||||
mock_prisma_binary,
|
||||
):
|
||||
"""A replacement that never settles must not stall the caller forever; the
|
||||
caller then sees an unchanged generation and reports the failure."""
|
||||
wrapper = _make_wrapper(engine_pid=111)
|
||||
gate = asyncio.Event()
|
||||
mock_prisma_binary.Prisma.return_value = _blocking_replacement(gate)
|
||||
|
||||
with patch("os.kill"), patch("asyncio.sleep", new_callable=AsyncMock):
|
||||
recreate = asyncio.create_task(
|
||||
wrapper.recreate_prisma_client("postgresql://new")
|
||||
)
|
||||
await _yield_to_loop()
|
||||
assert wrapper._reconnection_lock.locked() is True
|
||||
|
||||
await asyncio.wait_for(
|
||||
wrapper.wait_for_planned_engine_replacement(0.05), timeout=5.0
|
||||
)
|
||||
gave_up_with_replacement_still_in_flight = not recreate.done()
|
||||
|
||||
gate.set()
|
||||
await recreate
|
||||
|
||||
assert gave_up_with_replacement_still_in_flight is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_does_not_alert_when_probe_races_a_completed_replacement(
|
||||
mock_prisma_binary,
|
||||
):
|
||||
"""The reported bug: an IAM-refresh engine recreate makes a concurrent
|
||||
readiness probe fail transiently, and that failure was alerting as a DB
|
||||
exception on every refresh cycle.
|
||||
|
||||
The reporting task is drained while the replacement is still in flight,
|
||||
which is when it runs in production; a decision taken at that instant sees
|
||||
an engine generation that has not moved yet.
|
||||
"""
|
||||
wrapper = _make_wrapper(engine_pid=111)
|
||||
client = _make_prisma_client(wrapper)
|
||||
wrapper.query_raw = AsyncMock(
|
||||
side_effect=[
|
||||
httpx.ConnectError("All connection attempts failed"),
|
||||
[{"?column?": 1}],
|
||||
[{"?column?": 1}],
|
||||
]
|
||||
)
|
||||
gate = asyncio.Event()
|
||||
mock_prisma_binary.Prisma.return_value = _blocking_replacement(gate)
|
||||
|
||||
with patch("os.kill"), patch("asyncio.sleep", new_callable=AsyncMock):
|
||||
recreate = asyncio.create_task(
|
||||
wrapper.recreate_prisma_client("postgresql://new")
|
||||
)
|
||||
await _yield_to_loop()
|
||||
assert wrapper._reconnection_lock.locked() is True
|
||||
|
||||
probe_result = await client.health_check()
|
||||
|
||||
drain = asyncio.create_task(_await_health_check_reports())
|
||||
await _yield_to_loop()
|
||||
alerts_while_replacement_in_flight = (
|
||||
client.proxy_logging_obj.failure_handler.await_count
|
||||
)
|
||||
|
||||
gate.set()
|
||||
await recreate
|
||||
await drain
|
||||
|
||||
assert {
|
||||
"probe_result": probe_result,
|
||||
"probe_attempts": wrapper.query_raw.await_count,
|
||||
"alerts_while_in_flight": alerts_while_replacement_in_flight,
|
||||
"alerts": client.proxy_logging_obj.failure_handler.await_count,
|
||||
} == {
|
||||
"probe_result": [{"?column?": 1}],
|
||||
"probe_attempts": 3,
|
||||
"alerts_while_in_flight": 0,
|
||||
"alerts": 0,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_alerts_when_a_completed_replacement_still_cannot_reach_the_database(
|
||||
mock_prisma_binary,
|
||||
):
|
||||
"""``Prisma.connect()`` polls the query engine's own ``/status`` endpoint
|
||||
rather than round-tripping to the database, so a replacement can complete
|
||||
against a database that is still unreachable. The engine generation alone
|
||||
must not be enough to stay silent."""
|
||||
wrapper = _make_wrapper(engine_pid=111)
|
||||
client = _make_prisma_client(wrapper)
|
||||
wrapper.query_raw = AsyncMock(
|
||||
side_effect=httpx.ConnectError("All connection attempts failed")
|
||||
)
|
||||
gate = asyncio.Event()
|
||||
mock_prisma_binary.Prisma.return_value = _blocking_replacement(gate)
|
||||
|
||||
with patch("os.kill"), patch("asyncio.sleep", new_callable=AsyncMock):
|
||||
recreate = asyncio.create_task(
|
||||
wrapper.recreate_prisma_client("postgresql://new")
|
||||
)
|
||||
await _yield_to_loop()
|
||||
assert wrapper._reconnection_lock.locked() is True
|
||||
|
||||
with pytest.raises(httpx.ConnectError):
|
||||
await client.health_check()
|
||||
|
||||
gate.set()
|
||||
await recreate
|
||||
await _await_health_check_reports()
|
||||
|
||||
assert {
|
||||
"replacement_completed": wrapper.engine_generation,
|
||||
"alerted": client.proxy_logging_obj.failure_handler.await_count > 0,
|
||||
} == {"replacement_completed": 1, "alerted": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_alerts_when_the_replacement_never_completes(
|
||||
mock_prisma_binary,
|
||||
):
|
||||
"""A real outage also has a replacement in flight, but it fails, so the
|
||||
engine generation never advances and the probe failure must still alert."""
|
||||
wrapper = _make_wrapper(engine_pid=111)
|
||||
client = _make_prisma_client(wrapper)
|
||||
wrapper.query_raw = AsyncMock(
|
||||
side_effect=httpx.ConnectError("All connection attempts failed")
|
||||
)
|
||||
gate = asyncio.Event()
|
||||
mock_prisma_binary.Prisma.return_value = _blocking_replacement(gate, fail=True)
|
||||
|
||||
with patch("os.kill"), patch("asyncio.sleep", new_callable=AsyncMock):
|
||||
recreate = asyncio.create_task(
|
||||
wrapper.recreate_prisma_client("postgresql://new")
|
||||
)
|
||||
await _yield_to_loop()
|
||||
assert wrapper._reconnection_lock.locked() is True
|
||||
|
||||
with pytest.raises(httpx.ConnectError):
|
||||
await client.health_check()
|
||||
|
||||
gate.set()
|
||||
with pytest.raises(ConnectionRefusedError):
|
||||
await recreate
|
||||
await _await_health_check_reports()
|
||||
|
||||
call_types: List[str] = [
|
||||
c.kwargs["call_type"]
|
||||
for c in client.proxy_logging_obj.failure_handler.await_args_list
|
||||
]
|
||||
assert {
|
||||
"generation": wrapper.engine_generation,
|
||||
"alerted": len(call_types) > 0,
|
||||
"call_types": set(call_types),
|
||||
} == {"generation": 0, "alerted": True, "call_types": {"health_check"}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_alerts_for_non_connection_errors_during_a_replacement(
|
||||
mock_prisma_binary,
|
||||
):
|
||||
"""Suppression is scoped to transport failures. A query the database itself
|
||||
rejected is a real defect and must alert even mid-replacement, and even
|
||||
though the database is plainly reachable a moment later."""
|
||||
wrapper = _make_wrapper(engine_pid=111)
|
||||
client = _make_prisma_client(wrapper)
|
||||
wrapper.query_raw = AsyncMock(side_effect=_fails_then_answers(ValueError("malformed SELECT")))
|
||||
gate = asyncio.Event()
|
||||
mock_prisma_binary.Prisma.return_value = _blocking_replacement(gate)
|
||||
|
||||
with patch("os.kill"), patch("asyncio.sleep", new_callable=AsyncMock):
|
||||
recreate = asyncio.create_task(
|
||||
wrapper.recreate_prisma_client("postgresql://new")
|
||||
)
|
||||
await _yield_to_loop()
|
||||
assert wrapper._reconnection_lock.locked() is True
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await client.health_check()
|
||||
|
||||
gate.set()
|
||||
await recreate
|
||||
await _await_health_check_reports()
|
||||
|
||||
assert {
|
||||
"generation": wrapper.engine_generation,
|
||||
"alerted": client.proxy_logging_obj.failure_handler.await_count > 0,
|
||||
} == {"generation": 1, "alerted": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_alerts_for_a_transient_failure_with_no_engine_replacement(
|
||||
mock_prisma_binary,
|
||||
):
|
||||
"""Suppression is scoped to failures a planned replacement explains. A
|
||||
transport blip that self-heals with no engine replacement at all still
|
||||
alerts, so the gate cannot be widened into silencing every failure whose
|
||||
database happens to answer a moment later."""
|
||||
wrapper = _make_wrapper(engine_pid=111)
|
||||
client = _make_prisma_client(wrapper)
|
||||
wrapper.query_raw = AsyncMock(
|
||||
side_effect=[
|
||||
httpx.ConnectError("All connection attempts failed"),
|
||||
[{"?column?": 1}],
|
||||
[{"?column?": 1}],
|
||||
]
|
||||
)
|
||||
|
||||
probe_result = await client.health_check()
|
||||
await _await_health_check_reports()
|
||||
|
||||
assert {
|
||||
"probe_result": probe_result,
|
||||
"generation": wrapper.engine_generation,
|
||||
"alerted": client.proxy_logging_obj.failure_handler.await_count > 0,
|
||||
} == {"probe_result": [{"?column?": 1}], "generation": 0, "alerted": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_probe_stays_on_its_target_when_reader_availability_flips(
|
||||
mock_prisma_binary, monkeypatch
|
||||
):
|
||||
"""Routing is re-resolved on every attribute access, so a reader that
|
||||
recovers mid-call would otherwise send the probe to a different engine than
|
||||
the one whose generation is being checked, and blame the wrong replacement.
|
||||
The probe follows the wrapper it was handed."""
|
||||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||||
|
||||
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "postgresql://reader")
|
||||
writer = _make_wrapper(engine_pid=111)
|
||||
reader = _make_wrapper(engine_pid=222)
|
||||
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
|
||||
client = _make_prisma_client(routing)
|
||||
writer.query_raw = AsyncMock(return_value=[{"writer": 1}])
|
||||
reader.query_raw = AsyncMock(return_value=[{"reader": 1}])
|
||||
|
||||
routing._reader_unavailable = False
|
||||
target = client._probe_target_wrapper()
|
||||
routing._reader_unavailable = True
|
||||
result = await client._run_health_probe(target)
|
||||
|
||||
assert {
|
||||
"target_is_reader": target is reader,
|
||||
"result": result,
|
||||
"reader_probes": reader.query_raw.await_count,
|
||||
"writer_probes": writer.query_raw.await_count,
|
||||
} == {
|
||||
"target_is_reader": True,
|
||||
"result": [{"reader": 1}],
|
||||
"reader_probes": 1,
|
||||
"writer_probes": 0,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_consults_the_reader_wrapper_under_read_replica_routing(
|
||||
mock_prisma_binary, monkeypatch
|
||||
):
|
||||
"""``query_raw`` is routed to the reader, so a reader-side planned
|
||||
replacement is the one that explains a probe failure."""
|
||||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||||
|
||||
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "postgresql://reader")
|
||||
writer = _make_wrapper(engine_pid=111)
|
||||
reader = _make_wrapper(engine_pid=222)
|
||||
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
|
||||
client = _make_prisma_client(routing)
|
||||
reader.query_raw = AsyncMock(
|
||||
side_effect=[
|
||||
httpx.ConnectError("All connection attempts failed"),
|
||||
[{"?column?": 1}],
|
||||
[{"?column?": 1}],
|
||||
]
|
||||
)
|
||||
gate = asyncio.Event()
|
||||
mock_prisma_binary.Prisma.return_value = _blocking_replacement(gate)
|
||||
|
||||
with patch("os.kill"), patch("asyncio.sleep", new_callable=AsyncMock):
|
||||
recreate = asyncio.create_task(
|
||||
reader.recreate_prisma_client("postgresql://new-reader")
|
||||
)
|
||||
await _yield_to_loop()
|
||||
assert reader._reconnection_lock.locked() is True
|
||||
|
||||
probe_result = await client.health_check()
|
||||
|
||||
drain = asyncio.create_task(_await_health_check_reports())
|
||||
await _yield_to_loop()
|
||||
alerts_while_replacement_in_flight = (
|
||||
client.proxy_logging_obj.failure_handler.await_count
|
||||
)
|
||||
|
||||
gate.set()
|
||||
await recreate
|
||||
await drain
|
||||
|
||||
assert {
|
||||
"probe_target_is_the_reader": client._probe_target_wrapper() is reader,
|
||||
"writer_generation": writer.engine_generation,
|
||||
"reader_generation": reader.engine_generation,
|
||||
"probe_result": probe_result,
|
||||
"alerts_while_in_flight": alerts_while_replacement_in_flight,
|
||||
"alerts": client.proxy_logging_obj.failure_handler.await_count,
|
||||
} == {
|
||||
"probe_target_is_the_reader": True,
|
||||
"writer_generation": 0,
|
||||
"reader_generation": 1,
|
||||
"probe_result": [{"?column?": 1}],
|
||||
"alerts_while_in_flight": 0,
|
||||
"alerts": 0,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1096,6 +1096,7 @@ async def test_prisma_health_check_failure_names_itself_at_operator_visible_leve
|
|||
function and reads as "the check never ran", and reporting it only at debug
|
||||
level hides a database fault behind a flag nobody enables in production."""
|
||||
import logging
|
||||
from functools import partial
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
|
@ -1103,6 +1104,9 @@ async def test_prisma_health_check_failure_names_itself_at_operator_visible_leve
|
|||
client = MagicMock()
|
||||
client.db.query_raw = AsyncMock(side_effect=Exception("connection refused"))
|
||||
client.proxy_logging_obj.failure_handler = AsyncMock()
|
||||
client._probe_target_wrapper = MagicMock(return_value=client.db)
|
||||
client._run_health_probe = partial(PrismaClient._run_health_probe, client)
|
||||
client._report_health_check_failure = AsyncMock()
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
with pytest.raises(Exception, match="connection refused"):
|
||||
|
|
@ -1142,6 +1146,7 @@ async def test_prisma_health_check_failure_redacts_database_credentials(caplog):
|
|||
text can carry a full connection string, so the credential has to be gone
|
||||
from the emitted record."""
|
||||
import logging
|
||||
from functools import partial
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
|
@ -1151,6 +1156,9 @@ async def test_prisma_health_check_failure_redacts_database_credentials(caplog):
|
|||
side_effect=Exception("could not connect to postgresql://admin:hunter2@db.internal:5432/litellm")
|
||||
)
|
||||
client.proxy_logging_obj.failure_handler = AsyncMock()
|
||||
client._probe_target_wrapper = MagicMock(return_value=client.db)
|
||||
client._run_health_probe = partial(PrismaClient._run_health_probe, client)
|
||||
client._report_health_check_failure = AsyncMock()
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
with pytest.raises(Exception):
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ import json
|
|||
import pytest
|
||||
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
AttemptedFallbackTargets,
|
||||
fallback_attempt_key,
|
||||
get_fallback_model_group,
|
||||
run_async_fallback,
|
||||
)
|
||||
|
|
@ -142,6 +144,208 @@ async def test_run_async_fallback_skips_original_model_group():
|
|||
assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
|
||||
|
||||
|
||||
class RecordingFailRouter:
|
||||
def __init__(self):
|
||||
self.attempted_models = []
|
||||
|
||||
def log_retry(self, kwargs, e):
|
||||
return kwargs
|
||||
|
||||
async def async_function_with_fallbacks(self, *args, **kwargs):
|
||||
self.attempted_models.append(kwargs.get("model"))
|
||||
raise RuntimeError("fallback model also failed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_skips_model_group_already_attempted():
|
||||
"""A fallback graph that loops back on itself must not re-attempt a model group that
|
||||
already failed for this request. Every group in a cycle fails identically, so
|
||||
revisiting one multiplies the work and the error output without any chance of
|
||||
succeeding."""
|
||||
router = RecordingFailRouter()
|
||||
|
||||
with pytest.raises(RuntimeError, match="original failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["already-attempted"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
attempted_targets=AttemptedFallbackTargets(frozenset({"already-attempted"})),
|
||||
)
|
||||
|
||||
assert router.attempted_models == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_attempts_a_repeated_target_once():
|
||||
router = RecordingFailRouter()
|
||||
|
||||
with pytest.raises(RuntimeError, match="fallback model also failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["fallback-model", "fallback-model", "other-model"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=5,
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert router.attempted_models == ["fallback-model", "other-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_forwards_attempted_model_groups_to_nested_call():
|
||||
"""The nested call is where the next hop of the walk decides what to skip, so the
|
||||
accumulated set has to reach it, carrying both the group that just failed and the
|
||||
target being attempted."""
|
||||
router = RecordingRouter()
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["fallback-model"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
attempted_targets=AttemptedFallbackTargets(frozenset({"earlier-model"})),
|
||||
)
|
||||
|
||||
assert router.received_kwargs["attempted_targets"].keys == frozenset(
|
||||
{"earlier-model", "primary-model", "fallback-model"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"entry",
|
||||
[
|
||||
{"model": "primary-model", "_target_order": 2},
|
||||
{"model": "primary-model", "_excluded_deployment_ids": ["dep-1"]},
|
||||
],
|
||||
)
|
||||
async def test_run_async_fallback_still_retargets_the_same_group_via_dict_entry(entry):
|
||||
"""Order-based fallback and weighted intra-group failover both re-target the group that
|
||||
just failed, selecting a different set of deployments inside it. Those entries are dicts
|
||||
rather than plain names and must survive a guard that skips already-attempted names."""
|
||||
router = RecordingRouter()
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[entry],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
attempted_targets=AttemptedFallbackTargets(frozenset({"primary-model"})),
|
||||
)
|
||||
|
||||
assert router.received_kwargs["model"] == "primary-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_skips_a_repeated_dict_target():
|
||||
"""A client-side fallback list names its targets with dicts, and that list is re-walked
|
||||
at every level of the recursion, so an entry that carries no request override has to be
|
||||
recognised as the same attempt as the bare name."""
|
||||
router = RecordingFailRouter()
|
||||
|
||||
with pytest.raises(RuntimeError, match="original failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[{"model": "already-attempted"}],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
attempted_targets=AttemptedFallbackTargets(frozenset({"already-attempted"})),
|
||||
)
|
||||
|
||||
assert router.attempted_models == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_attempts_a_repeated_dict_target_once():
|
||||
router = RecordingFailRouter()
|
||||
entry = {"model": "fallback-model", "messages": [{"role": "user", "content": "shorter"}]}
|
||||
|
||||
with pytest.raises(RuntimeError, match="fallback model also failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[entry, entry, {"model": "other-model"}],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=5,
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert router.attempted_models == ["fallback-model", "other-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_keeps_a_request_override_distinct_from_the_bare_name():
|
||||
"""The documented use of the client-side form is to retry a group with different request
|
||||
params, so an entry carrying an override must survive even when the bare name of that
|
||||
same group has already been attempted."""
|
||||
router = RecordingFailRouter()
|
||||
|
||||
with pytest.raises(RuntimeError, match="fallback model also failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[
|
||||
{"model": "already-attempted", "messages": [{"role": "user", "content": "shorter"}]}
|
||||
],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
attempted_targets=AttemptedFallbackTargets(frozenset({"already-attempted"})),
|
||||
)
|
||||
|
||||
assert router.attempted_models == ["already-attempted"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"target, expected",
|
||||
[
|
||||
("group-a", "group-a"),
|
||||
({"model": "group-a"}, "group-a"),
|
||||
(None, None),
|
||||
(["group-a"], None),
|
||||
],
|
||||
)
|
||||
def test_fallback_attempt_key_identity(target, expected):
|
||||
"""A bare name and a `{"model": name}` entry are the same attempt. A shape with no
|
||||
usable identity returns None and is never skipped, so an unrecognised entry keeps
|
||||
today's behaviour rather than being silently dropped."""
|
||||
assert fallback_attempt_key(target) == expected
|
||||
|
||||
|
||||
def test_fallback_attempt_key_gives_a_param_only_entry_its_own_identity():
|
||||
"""An entry with no `model` re-targets the group currently being attempted with
|
||||
different request params, so it is a distinct attempt and still needs an identity."""
|
||||
key = fallback_attempt_key({"messages": [{"role": "user", "content": "shorter"}]})
|
||||
|
||||
assert key is not None
|
||||
assert key != fallback_attempt_key({"messages": [{"role": "user", "content": "other"}]})
|
||||
|
||||
|
||||
def test_fallback_attempt_key_separates_overrides_from_the_bare_name():
|
||||
bare = fallback_attempt_key("group-a")
|
||||
override = fallback_attempt_key({"model": "group-a", "messages": [{"role": "user", "content": "x"}]})
|
||||
other_override = fallback_attempt_key({"model": "group-a", "messages": [{"role": "user", "content": "y"}]})
|
||||
order_retarget = fallback_attempt_key({"model": "group-a", "_target_order": 2})
|
||||
|
||||
assert len({bare, override, other_override, order_retarget}) == 4
|
||||
|
||||
|
||||
def test_fallback_attempt_key_is_stable_across_key_order():
|
||||
assert fallback_attempt_key({"model": "group-a", "_target_order": 2}) == fallback_attempt_key(
|
||||
{"_target_order": 2, "model": "group-a"}
|
||||
)
|
||||
|
||||
|
||||
def test_get_fallback_model_group_does_not_mutate_fallbacks():
|
||||
"""A string fallback must be resolved without mutating the caller's
|
||||
fallbacks list, which is the live router config shared across requests."""
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from unittest.mock import Mock
|
|||
import pytest
|
||||
|
||||
from litellm import Router
|
||||
from litellm.constants import ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.router_utils.common_utils import (
|
||||
_deployment_supports_web_search,
|
||||
|
|
@ -11,6 +12,7 @@ from litellm.router_utils.common_utils import (
|
|||
filter_team_based_models,
|
||||
filter_web_search_deployments,
|
||||
resolve_model_group_alias,
|
||||
truncate_fallback_error_detail,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -558,3 +560,27 @@ class TestResolveModelGroupAlias:
|
|||
assert router._get_model_from_alias("group-a") == "group-b"
|
||||
assert router._get_model_from_alias("group-item") == "group-b"
|
||||
assert router._get_model_from_alias("group-b") is None
|
||||
|
||||
|
||||
class TestTruncateFallbackErrorDetail:
|
||||
def test_short_detail_is_returned_unchanged(self):
|
||||
assert truncate_fallback_error_detail("boom") == "boom"
|
||||
|
||||
def test_detail_at_the_limit_is_returned_unchanged(self):
|
||||
detail = "x" * ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS
|
||||
assert truncate_fallback_error_detail(detail) == detail
|
||||
|
||||
def test_long_detail_is_bounded_and_reports_what_was_dropped(self):
|
||||
detail = "x" * (ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS + 500)
|
||||
|
||||
truncated = truncate_fallback_error_detail(detail)
|
||||
|
||||
assert truncated.startswith("x" * ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS)
|
||||
assert truncated.endswith("... [truncated 500 characters]")
|
||||
assert len(truncated) < len(detail)
|
||||
|
||||
def test_a_megabyte_of_detail_comes_back_small(self):
|
||||
"""The detail is what a fallback level records about the level below it, so it has
|
||||
to stay small enough that a walk over many model groups cannot compound it into an
|
||||
output volume that starves the process."""
|
||||
assert len(truncate_fallback_error_detail("x" * 1_000_000)) < 3_000
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -14,6 +15,7 @@ sys.path.insert(
|
|||
|
||||
import litellm
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata():
|
||||
|
|
@ -7415,3 +7417,136 @@ class TestAutoRouterMaxInputCharsWiring:
|
|||
router = self._router()
|
||||
|
||||
assert self._registered_auto_router(router).max_input_chars == DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
|
||||
|
||||
|
||||
class _LogCapture(logging.Handler):
|
||||
def __init__(self, level):
|
||||
super().__init__(level=level)
|
||||
self._level = level
|
||||
self.messages = []
|
||||
|
||||
def emit(self, record):
|
||||
if record.levelno == self._level:
|
||||
self.messages.append(record.getMessage())
|
||||
|
||||
|
||||
class _FallbackAttemptRecorder(CustomLogger):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.failed_targets = []
|
||||
|
||||
async def log_failure_fallback_event(self, original_model_group, kwargs, original_exception):
|
||||
self.failed_targets.append(kwargs.get("model"))
|
||||
|
||||
|
||||
def _cyclic_fallback_router(num_retries=0):
|
||||
groups = ["group-a", "group-b", "group-c", "group-d"]
|
||||
return litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": group,
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "sk-fake",
|
||||
"mock_response": "litellm.InternalServerError",
|
||||
},
|
||||
}
|
||||
for group in groups
|
||||
],
|
||||
fallbacks=[
|
||||
{"group-a": ["group-b", "group-c"]},
|
||||
{"group-b": ["group-a", "group-c"]},
|
||||
{"group-c": ["group-d"]},
|
||||
{"group-d": ["group-b", "group-a"]},
|
||||
],
|
||||
num_retries=num_retries,
|
||||
)
|
||||
|
||||
|
||||
async def _drive_cyclic_fallback(router, capture, recorder=None, **request_kwargs):
|
||||
router_logger = logging.getLogger("LiteLLM Router")
|
||||
previous_level = router_logger.level
|
||||
router_logger.setLevel(capture.level)
|
||||
router_logger.addHandler(capture)
|
||||
if recorder is not None:
|
||||
litellm.callbacks.append(recorder)
|
||||
try:
|
||||
with pytest.raises(litellm.InternalServerError):
|
||||
await router.acompletion(
|
||||
model="group-a", messages=[{"role": "user", "content": "hi"}], **request_kwargs
|
||||
)
|
||||
finally:
|
||||
router_logger.removeHandler(capture)
|
||||
router_logger.setLevel(previous_level)
|
||||
if recorder is not None:
|
||||
litellm.callbacks.remove(recorder)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cyclic_fallback_graph_does_not_amplify_one_request():
|
||||
"""A fallback graph whose entries loop back on each other is easy to build by accident,
|
||||
and every group in the loop fails identically on a deterministic error, so the walk must
|
||||
not revisit a group and must not re-emit a growing chained traceback at each level. Left
|
||||
unbounded, one request blocks the event loop long enough for health probes to fail."""
|
||||
recorder = _FallbackAttemptRecorder()
|
||||
capture = _LogCapture(logging.ERROR)
|
||||
|
||||
await _drive_cyclic_fallback(_cyclic_fallback_router(), capture, recorder)
|
||||
|
||||
assert sorted(set(recorder.failed_targets)) == ["group-b", "group-c", "group-d"]
|
||||
assert len(recorder.failed_targets) == len(set(recorder.failed_targets))
|
||||
assert not any("Traceback (most recent call last)" in message for message in capture.messages)
|
||||
assert sum(len(message) for message in capture.messages) < 5_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retry_breadcrumbs_do_not_carry_the_walk_state():
|
||||
"""log_retry copies every kwarg into previous_models, which reaches spend logs and
|
||||
logging callbacks. The set of already-attempted groups is router-internal walk state
|
||||
with no diagnostic value there, and it is the one entry that is not a plain scalar.
|
||||
A retry has to be configured for the walk state to reach log_retry at all."""
|
||||
router = _cyclic_fallback_router(num_retries=1)
|
||||
capture = _LogCapture(logging.ERROR)
|
||||
|
||||
await _drive_cyclic_fallback(router, capture)
|
||||
|
||||
assert router.previous_models, "no retry breadcrumbs were recorded"
|
||||
assert any(
|
||||
"fallback_depth" in breadcrumb for breadcrumb in router.previous_models
|
||||
), "no breadcrumb carried router walk state, so this test cannot see the leak"
|
||||
for breadcrumb in router.previous_models:
|
||||
assert "attempted_targets" not in breadcrumb
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_traceback_stays_available_at_debug_level():
|
||||
"""Dropping the stack from the ERROR line is only safe because the fallback path still
|
||||
emits it once per level at DEBUG, which is what an operator needs to diagnose why every
|
||||
fallback failed. This pins that remaining debug traceback."""
|
||||
capture = _LogCapture(logging.DEBUG)
|
||||
|
||||
await _drive_cyclic_fallback(_cyclic_fallback_router(), capture)
|
||||
|
||||
assert any("Traceback (most recent call last)" in message for message in capture.messages)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_failure_detail_from_upstream_is_bounded():
|
||||
"""The detail each level records about the level below it is attacker-influenced, since
|
||||
it carries whatever the upstream error said. It has to be bounded on its own, so a walk
|
||||
over several groups cannot compound one large message into the log or into the message
|
||||
handed back to the caller."""
|
||||
huge_message = "z" * 50_000
|
||||
capture = _LogCapture(logging.ERROR)
|
||||
|
||||
await _drive_cyclic_fallback(
|
||||
_cyclic_fallback_router(),
|
||||
capture,
|
||||
mock_response=litellm.InternalServerError(
|
||||
message=huge_message, llm_provider="openai", model="group-a"
|
||||
),
|
||||
)
|
||||
|
||||
assert capture.messages, "the fallback failure path did not log at ERROR"
|
||||
assert huge_message not in "".join(capture.messages)
|
||||
assert max(len(message) for message in capture.messages) < 5_000
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue