fix(passthrough): run the upstream error report as a response background task

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-25 22:55:52 +00:00
parent a614d85c92
commit 3f708fa930
9 changed files with 578 additions and 642 deletions

View file

@ -1529,11 +1529,6 @@ CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTER
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS: Final[float | None] = (
max(0.0, float(_raw_drain))
if (_raw_drain := os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS"))
else None
)
# Headers to control callbacks
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"

View file

@ -4,7 +4,6 @@ import copy
import json
import posixpath
import traceback
import weakref
from base64 import b64encode
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
@ -27,6 +26,7 @@ from fastapi import (
status,
)
from fastapi.responses import StreamingResponse
from starlette.background import BackgroundTask
from starlette.datastructures import UploadFile as StarletteUploadFile
from starlette.websockets import WebSocketState
from websockets.asyncio.client import connect
@ -43,7 +43,6 @@ from litellm._uuid import uuid
from litellm.constants import (
MAXIMUM_TRACEBACK_LINES_TO_LOG,
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS,
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
REDACTED_BY_LITELLM,
SESSION_ID_OMITTED_METADATA_KEY,
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
@ -882,99 +881,53 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
upstream: httpx.Response,
report: _ReportPreview,
log_warning: Callable[..., None],
spawn: Callable[[Awaitable[None]], asyncio.Future[None]],
) -> None:
self._upstream: Final = upstream
self._report: Final = report
self._log_warning: Final = log_warning
self._spawn: Final = spawn
self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying
self._dispatched = False
self._pending: asyncio.Future[None] | None = None
self._budget_crossed = False
self._completed = False
self._aborted = False
self._reported = False
def _dispatch_report(self) -> None:
if self._dispatched:
async def report_collected(self) -> None:
if self._reported:
return
self._dispatched = True
self._pending = self._spawn(self._report(b"".join(self._collected)))
async def _drain_pending_report(self) -> None:
pending: Final = self._pending
if pending is not None:
await asyncio.shield(pending)
def _dispatch_disconnect_report(self) -> None:
if not self._completed and not self._dispatched:
self._reported = True
if not self._completed and not self._budget_crossed and not self._aborted:
self._log_warning(
"pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body",
sum(len(part) for part in self._collected),
)
self._dispatch_report()
await self._report(b"".join(self._collected))
async def __aiter__(self) -> AsyncIterator[bytes]:
total = 0 # rebind-ok: running byte count against the preview budget
try:
async for chunk in self._upstream.aiter_bytes():
if not self._dispatched:
if not self._budget_crossed:
self._collected.append(chunk)
total += len(chunk)
if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
self._dispatch_report()
self._budget_crossed = True
yield chunk
self._completed = True
self._dispatch_report()
await self._drain_pending_report()
except httpx.HTTPError as err:
dispatched_before: Final = self._dispatched
self._aborted = True
self._log_warning(
"pass_through_endpoint: upstream error body read failed after %d bytes: %s",
sum(len(part) for part in self._collected),
type(err).__name__,
)
self._dispatch_report()
await self._drain_pending_report()
if dispatched_before:
if self._budget_crossed:
await self.report_collected()
raise
finally:
self._dispatch_disconnect_report()
async def aclose(self) -> None:
self._dispatch_disconnect_report()
await self._upstream.aclose()
_REPORT_TASKS: Final[ # mutable-ok: in-flight report registry per loop, drained at shutdown
weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, set[asyncio.Future[None]]]
] = weakref.WeakKeyDictionary()
def _spawn_report_task(report: Awaitable[None]) -> asyncio.Future[None]:
task: Final = asyncio.ensure_future(report)
registry: Final = _REPORT_TASKS.setdefault(
asyncio.get_running_loop(),
set(), # mutable-ok: per-loop task set, tasks discard themselves on completion
)
registry.add(task)
task.add_done_callback(registry.discard)
return task
async def drain_passthrough_upstream_error_reports(
timeout: float | None = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
log_warning: Callable[..., None] = verbose_proxy_logger.warning,
) -> None:
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
if not pending:
return
_, still_pending = await asyncio.wait(pending, timeout=timeout)
if still_pending:
log_warning(
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending",
len(still_pending),
)
def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers:
return httpx.Headers(
[(name, value) for name, value in headers.raw if name.lower() not in (b"content-encoding", b"content-length")]
@ -1033,14 +986,20 @@ def _passthrough_upstream_failure_reporter(
return report
@dataclass(frozen=True, slots=True)
class _UpstreamRelay:
response: httpx.Response
background: BackgroundTask | None
async def _log_passthrough_upstream_failure(
response: httpx.Response,
user_api_key_dict: UserAPIKeyAuth,
request_payload: dict,
logging_obj: LiteLLMLoggingObj,
) -> httpx.Response:
) -> _UpstreamRelay:
if response.status_code < 400:
return response
return _UpstreamRelay(response=response, background=None)
from litellm.proxy.proxy_server import proxy_logging_obj
log_warning: Final = verbose_proxy_logger.warning
@ -1049,18 +1008,21 @@ async def _log_passthrough_upstream_failure(
)
if response.is_stream_consumed:
await report(response.content)
return response
return httpx.Response(
status_code=response.status_code,
headers=_headers_without_body_framing(response.headers),
stream=_PreviewReportingStream(
upstream=response,
report=report,
log_warning=log_warning,
spawn=_spawn_report_task,
return _UpstreamRelay(response=response, background=None)
stream: Final = _PreviewReportingStream(
upstream=response,
report=report,
log_warning=log_warning,
)
return _UpstreamRelay(
response=httpx.Response(
status_code=response.status_code,
headers=_headers_without_body_framing(response.headers),
stream=stream,
request=response.request,
extensions=response.extensions,
),
request=response.request,
extensions=response.extensions,
background=BackgroundTask(stream.report_collected),
)
@ -1491,7 +1453,7 @@ async def pass_through_request(
headers=response.headers,
)
relay_response: Final = await _log_passthrough_upstream_failure(
relay: Final = await _log_passthrough_upstream_failure(
response=response,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
@ -1506,13 +1468,13 @@ async def pass_through_request(
# Call response headers hook for streaming pass-through
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
headers=relay_response.headers,
headers=relay.response.headers,
litellm_call_id=litellm_call_id,
)
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=_parsed_body or {},
user_api_key_dict=user_api_key_dict,
response=relay_response,
response=relay.response,
request_headers=dict(request.headers),
)
if callback_headers:
@ -1523,7 +1485,7 @@ async def pass_through_request(
stream=_own_streamed_managed_ids(
stream=_relay_reporting_failures(
stream=PassThroughStreamingHandler.chunk_processor(
response=relay_response,
response=relay.response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
@ -1531,7 +1493,7 @@ async def pass_through_request(
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
),
upstream_status=relay_response.status_code,
upstream_status=relay.response.status_code,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
parsed_body=_parsed_body,
@ -1545,10 +1507,11 @@ async def pass_through_request(
user_api_key_dict=user_api_key_dict,
),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=relay_response.headers,
upstream_headers=relay.response.headers,
),
headers=_response_headers,
status_code=relay_response.status_code,
status_code=relay.response.status_code,
background=relay.background,
)
if state_raw_body is not None:
@ -1583,7 +1546,7 @@ async def pass_through_request(
logging_obj.stream = True
logging_obj.model_call_details["stream"] = True
detected_relay_response: Final = await _log_passthrough_upstream_failure(
detected_relay: Final = await _log_passthrough_upstream_failure(
response=response,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
@ -1598,13 +1561,13 @@ async def pass_through_request(
# Call response headers hook for detected streaming pass-through
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
headers=detected_relay_response.headers,
headers=detected_relay.response.headers,
litellm_call_id=litellm_call_id,
)
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=_parsed_body or {},
user_api_key_dict=user_api_key_dict,
response=detected_relay_response,
response=detected_relay.response,
request_headers=dict(request.headers),
)
if callback_headers:
@ -1615,7 +1578,7 @@ async def pass_through_request(
stream=_own_streamed_managed_ids(
stream=_relay_reporting_failures(
stream=PassThroughStreamingHandler.chunk_processor(
response=detected_relay_response,
response=detected_relay.response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
@ -1623,7 +1586,7 @@ async def pass_through_request(
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
),
upstream_status=detected_relay_response.status_code,
upstream_status=detected_relay.response.status_code,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
parsed_body=_parsed_body,
@ -1637,10 +1600,11 @@ async def pass_through_request(
user_api_key_dict=user_api_key_dict,
),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=detected_relay_response.headers,
upstream_headers=detected_relay.response.headers,
),
headers=_response_headers,
status_code=detected_relay_response.status_code,
status_code=detected_relay.response.status_code,
background=detected_relay.background,
)
if not _should_buffer_passthrough_response(response):

View file

@ -733,7 +733,6 @@ from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import (
router as openai_passthrough_router,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
drain_passthrough_upstream_error_reports,
initialize_pass_through_endpoints,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
@ -1097,27 +1096,6 @@ async def _flush_spend_logs_queue_on_shutdown() -> None:
verbose_proxy_logger.exception("Error flushing spend logs queue on shutdown: %s", e)
async def _drain_reports_then_flush_spend(
drain_reports: Callable[[], Awaitable[None]],
drain_spend_events: Callable[[], Awaitable[None]],
stop_scheduler_jobs: Callable[[], Awaitable[None]] | None,
flush_spend_counters: Callable[[], Awaitable[None]],
flush_spend_logs: Callable[[], Awaitable[None]],
) -> None:
try:
await drain_reports()
except Exception as e: # noqa: BLE001 # shutdown must continue when a report drain fails
verbose_proxy_logger.error("Error draining passthrough upstream error reports: %s", e)
await drain_spend_events()
if stop_scheduler_jobs is not None:
try:
await stop_scheduler_jobs()
except Exception as e:
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
await flush_spend_counters()
await flush_spend_logs()
async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = None) -> None:
global prisma_client, master_key, user_custom_auth, user_custom_key_generate, user_custom_key_update
verbose_proxy_logger.info("Shutting down LiteLLM Proxy Server")
@ -1606,18 +1584,18 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
except Exception as e:
verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e)
await _drain_spend_event_producer_on_shutdown()
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
await _drain_reports_then_flush_spend(
drain_reports=drain_passthrough_upstream_error_reports,
drain_spend_events=_drain_spend_event_producer_on_shutdown,
stop_scheduler_jobs=(
partial(stop_in_flight_scheduler_jobs, scheduler, scheduler_executor)
if scheduler is not None and scheduler_executor is not None
else None
),
flush_spend_counters=flush_spend_counters_on_shutdown,
flush_spend_logs=_flush_spend_logs_queue_on_shutdown,
)
if scheduler is not None and scheduler_executor is not None:
try:
await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor)
except Exception as e:
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
await flush_spend_counters_on_shutdown()
await _flush_spend_logs_queue_on_shutdown()
await proxy_config.stop_config_sync_subscriber()
@ -6364,6 +6342,11 @@ class ProxyConfig:
if general_settings is None:
general_settings = {}
if general_settings.get("mcp_advertised_versions") is not None:
from litellm.types.mcp import MCPAdvertisedVersions
TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"])
if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None:
warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1"))
if declared_proxy_ranges(general_settings) is None:

View file

@ -89,6 +89,7 @@ def owned_proxy_process(
config: Path | None = None,
remove_environment: tuple[str, ...] = (),
workers: int = 1,
graceful_shutdown_seconds: int | None = None,
) -> Iterator[OwnedProxy]:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
@ -104,27 +105,43 @@ def owned_proxy_process(
"LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"),
"STORE_MODEL_IN_DB": "True",
**overrides,
**({"CONFIG_FILE_PATH": str(config)} if graceful_shutdown_seconds is not None and config else {}),
}
output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory)))
output.mkdir(parents=True, exist_ok=True)
log_path: Final = output / f"owned-proxy-{uuid.uuid4().hex}.log"
with log_path.open("w") as log:
process: Final = subprocess.Popen(
[
sys.executable,
"-m",
"integration._support.proxy",
"--config",
str(config or "tests/integration/proxy_config.yaml"),
"--host",
"127.0.0.1",
"--port",
str(port),
"--num_workers",
str(workers),
"--use_prisma_db_push",
"--enforce_prisma_migration_check",
],
(
[
sys.executable,
"-m",
"uvicorn",
"litellm.proxy.proxy_server:app",
"--host",
"127.0.0.1",
"--port",
str(port),
"--timeout-graceful-shutdown",
str(graceful_shutdown_seconds),
]
if graceful_shutdown_seconds is not None
else [
sys.executable,
"-m",
"integration._support.proxy",
"--config",
str(config or "tests/integration/proxy_config.yaml"),
"--host",
"127.0.0.1",
"--port",
str(port),
"--num_workers",
str(workers),
"--use_prisma_db_push",
"--enforce_prisma_migration_check",
]
),
cwd=root,
env=environment,
stdout=log,

View file

@ -243,6 +243,73 @@ async def test_passthrough_sigterm_drains_reports_parked_on_a_slow_failure_hook(
_single_spend_row(call_id)
_PARKING_FAILURE_HOOK: Final = """
import asyncio
from litellm.integrations.custom_logger import CustomLogger
class ParkingFailureHook(CustomLogger):
async def async_post_call_failure_hook(
self, request_data, original_exception, user_api_key_dict, traceback_str=None
):
await asyncio.Event().wait()
instance = ParkingFailureHook()
"""
async def test_passthrough_sigterm_with_graceful_timeout_exits_and_flushes_spend(
gateway: Gateway, tmp_path: Path
) -> None:
"""A streamed upstream error whose failure report can never finish must not hold
the proxy open: uvicorn cancels the stuck request task after the graceful
window and the shutdown spend flush still lands rows buffered behind the
3600 s batch interval."""
gate: Final = threading.Event()
def respond(request: Request) -> Reply:
if "streamGenerateContent" in request.target:
return Reply(
status=429, content_type="text/event-stream", chunks=_RATE_LIMITED_FRAMES, gate_after_first=gate
)
return Reply(
status=200,
body=json.dumps(
{
"candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
"usageMetadata": {"promptTokenCount": 3, "candidatesTokenCount": 2, "totalTokenCount": 5},
}
).encode(),
)
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update({"callbacks": ["park_hook.instance"]})
config["general_settings"]["proxy_batch_write_at"] = 3600
(tmp_path / "park_hook.py").write_text(_PARKING_FAILURE_HOOK)
path: Final = tmp_path / "chaos-graceful-sigterm.yaml"
with wire_server(respond) as wire:
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path, graceful_shutdown_seconds=3) as owned:
candidate: Final = owned.gateway
await _first_frame_then_close(str(candidate.client.base_url), candidate.key)
gate.set()
follow_up: Final = candidate.request(
"POST",
"/gemini/v1beta/models/claude-nope-9:generateContent",
_GENERATE_CONTENT,
headers={"x-goog-api-key": candidate.key},
)
assert follow_up.status_code == 200, follow_up.text
call_id: Final = follow_up.headers["x-litellm-call-id"]
await asyncio.sleep(1.5)
owned.process.send_signal(signal.SIGTERM)
owned.process.wait(timeout=20)
_single_spend_row(call_id)
async def _first_frame_then_close(base_url: str, key: str) -> str:
async with httpx.AsyncClient(base_url=base_url, timeout=httpx.Timeout(5, connect=5), trust_env=False) as client:
async with client.stream(

View file

@ -1133,3 +1133,117 @@ def test_llm_endpoint_upstream_quota_429_normalized_error_unchanged(
assert error_information["normalized_error"] == _LLM_429_NORMALIZED, error_information
finally:
delete_scenario(handle)
_TWO_SECOND_FAILURE_HOOK: Final = """
import asyncio
from litellm.integrations.custom_logger import CustomLogger
class TwoSecondFailureHook(CustomLogger):
async def async_post_call_failure_hook(
self, request_data, original_exception, user_api_key_dict, traceback_str=None
):
await asyncio.sleep(2)
instance = TwoSecondFailureHook()
"""
async def test_passthrough_keepalive_pings_never_follow_the_upstream_error_body(
gateway: Gateway, tmp_path: Path
) -> None:
"""Keepalive pings fill idle time while the relay is still producing; once the
upstream error body is fully relayed the response ends. If the failure
report sat on the response path, pings emitted during the slow hook would
land after the body's last frame."""
frames: Final = (b'data: {"error":"one"}\n\n', b'data: {"error":"two"}\n\n')
def respond(request: Request) -> Reply:
return Reply(status=500, content_type="text/event-stream", chunks=frames)
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update(
{"callbacks": ["two_sec_hook.instance"], "sse_keepalive_ping_interval_seconds": 0.2}
)
(tmp_path / "two_sec_hook.py").write_text(_TWO_SECOND_FAILURE_HOOK)
path: Final = tmp_path / "gemini-keepalive-error.yaml"
with wire_server(respond) as wire:
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path) as owned:
candidate: Final = owned.gateway
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
) as response:
assert response.status_code == 500, response.text
call_id: Final = response.headers["x-litellm-call-id"]
streamed: Final = response.read()
assert streamed == b"".join(frames), streamed
_spend_error_information(call_id)
_HEADER_STATE_HOOK: Final = """
from litellm.integrations.custom_logger import CustomLogger
class HeaderStateHook(CustomLogger):
def __init__(self):
self.last_status = None
async def async_post_call_failure_hook(
self, request_data, original_exception, user_api_key_dict, traceback_str=None
):
self.last_status = getattr(original_exception, "status_code", None)
async def async_post_call_response_headers_hook(
self, data, user_api_key_dict, response, request_headers=None, litellm_call_info=None
):
return {"x-failure-for-this-request": str(self.last_status or "none")}
instance = HeaderStateHook()
"""
async def test_passthrough_streamed_error_headers_do_not_carry_failure_hook_state(
gateway: Gateway, tmp_path: Path
) -> None:
"""The streamed error report runs after response headers are sent, so a
callback-derived header sees pre-request state; the failure hook still
records the upstream status in the spend row."""
frames: Final = (b'data: {"error":"quota"}\n\n',)
def respond(request: Request) -> Reply:
return Reply(status=500, content_type="text/event-stream", chunks=frames)
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update({"callbacks": ["header_state_hook.instance"]})
(tmp_path / "header_state_hook.py").write_text(_HEADER_STATE_HOOK)
path: Final = tmp_path / "gemini-header-state.yaml"
with wire_server(respond) as wire:
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path) as owned:
candidate: Final = owned.gateway
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
) as response:
assert response.status_code == 500, response.text
call_id: Final = response.headers["x-litellm-call-id"]
header_value: Final = response.headers["x-failure-for-this-request"]
streamed: Final = response.read()
assert streamed == b"".join(frames), streamed
assert header_value == "none", header_value
error_information: Final = _spend_error_information(call_id)
assert error_information["error_code"] == "500", error_information

View file

@ -1,11 +1,9 @@
import asyncio
import gc
import gzip
import json
import logging
import os
import sys
import time
import zlib
from collections.abc import Callable, Mapping
from contextlib import ExitStack, contextmanager
@ -20,6 +18,7 @@ import pytest
from fastapi import HTTPException, Request, Response, UploadFile
from fastapi.responses import StreamingResponse
from pydantic import TypeAdapter, ValidationError
from starlette.background import BackgroundTask
from starlette.datastructures import FormData, Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
@ -29,19 +28,17 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_REPORT_TASKS,
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS,
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
HttpPassThroughEndpointHelpers,
InitPassThroughEndpointHelpers,
_PreviewReportingStream,
_registered_pass_through_routes,
_spawn_report_task,
_truncate_upstream_error_body,
_with_trace_context,
chat_completion_pass_through_endpoint,
create_pass_through_route,
drain_passthrough_upstream_error_reports,
_log_passthrough_upstream_failure,
initialize_pass_through_endpoints,
pass_through_request,
resolve_llm_passthrough_timeout,
@ -4266,7 +4263,7 @@ async def test_pass_through_request_streaming_upstream_error_body_reaches_client
)
assert streamed_bytes == upstream_content
await _poll(lambda: mock_proxy_logging.post_call_failure_hook.called)
await response.background()
mock_proxy_logging.post_call_failure_hook.assert_called_once()
original_exception: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"]
assert "was not found or your project" in original_exception.detail
@ -4462,42 +4459,30 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[asyncio.Future[None]] = []
served_at_dispatch: list[int] = []
def _recording_spawn(coro):
served_at_dispatch.append(body_stream.served)
enqueued.append(asyncio.ensure_future(coro))
return enqueued[-1]
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=_recording_spawn,
):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 500
@ -4506,12 +4491,12 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == upstream_content
assert body_stream.served == 10
mock_proxy_logging.post_call_failure_hook.assert_not_called()
assert served_at_dispatch == [5], (
"the report is dispatched after the fifth 1024-byte chunk, the first point the preview budget is exceeded"
)
assert len(enqueued) == 1, enqueued
await enqueued[0]
assert response.background is not None
await response.background()
mock_proxy_logging.post_call_failure_hook.assert_called_once()
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
@ -4532,42 +4517,30 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[asyncio.Future[None]] = []
served_at_dispatch: list[int] = []
def _recording_spawn(coro):
served_at_dispatch.append(body_stream.served)
enqueued.append(asyncio.ensure_future(coro))
return enqueued[-1]
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=_recording_spawn,
):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 500
@ -4576,12 +4549,11 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == upstream_content
mock_proxy_logging.post_call_failure_hook.assert_not_called()
assert served_at_dispatch == [1], (
"the report is dispatched after the first raw chunk crosses the preview budget; the second is not pulled first"
)
assert len(enqueued) == 1, enqueued
await enqueued[0]
assert response.background is not None
await response.background()
mock_proxy_logging.post_call_failure_hook.assert_called_once()
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
@ -4589,13 +4561,6 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
)
async def _poll(condition: Callable[[], bool], seconds: float = 5) -> None:
deadline: Final = asyncio.get_running_loop().time() + seconds
while not condition():
assert asyncio.get_running_loop().time() < deadline, "condition not met in time"
await asyncio.sleep(0.01)
class _GatedUpstreamErrorBodyStream(httpx.AsyncByteStream):
def __init__(self, first: bytes, second: bytes) -> None:
self._first: Final = first
@ -4625,8 +4590,6 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[asyncio.Future[None]] = []
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
@ -4634,31 +4597,27 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 429
iterator: Final = response.body_iterator.__aiter__()
assert isinstance(response, StreamingResponse)
assert response.status_code == 429
iterator: Final = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert not body_stream.gate.is_set()
mock_proxy_logging.post_call_failure_hook.assert_not_called()
@ -4670,8 +4629,9 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
assert relayed == first_chunk + second_chunk
await upstream_response.aclose()
assert len(enqueued) == 1, enqueued
await enqueued[0]
mock_proxy_logging.post_call_failure_hook.assert_not_called()
assert response.background is not None
await response.background()
mock_proxy_logging.post_call_failure_hook.assert_called_once()
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
assert (
@ -4680,12 +4640,11 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_client_disconnect_enqueues_failure_report():
async def test_pass_through_request_streaming_upstream_error_client_disconnect_reports_relayed_bytes():
"""
Regression: when the client disconnects mid-relay the response task is
cancelled, so the preview report cannot be awaited inline; it must be
handed to the logging worker, which then fires the failure hook once
with the chunks already relayed.
Regression: when the client disconnects mid-relay the response ends early,
so the report runs as the response background task; it must fire the
failure hook once with only the chunks already relayed.
"""
first_chunk: Final = b'data: {"error":"rate limited"}\n\n'
second_chunk: Final = b"data: [DONE]\n\n"
@ -4697,8 +4656,6 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[asyncio.Future[None]] = []
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
@ -4706,42 +4663,33 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
iterator = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert first_chunk in (first if isinstance(first, bytes) else first.encode())
await iterator.aclose()
del iterator
gc.collect()
for _ in range(40):
if enqueued:
break
await asyncio.sleep(0.05)
assert isinstance(response, StreamingResponse)
iterator = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert first_chunk in (first if isinstance(first, bytes) else first.encode())
await iterator.aclose()
assert len(enqueued) == 1, enqueued
await enqueued[0]
mock_proxy_logging.post_call_failure_hook.assert_not_called()
assert response.background is not None
await response.background()
mock_proxy_logging.post_call_failure_hook.assert_called_once()
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
assert detail == 'Upstream passthrough request failed with status 429: data: {"error":"rate limited"}', detail
@ -4750,9 +4698,9 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_yields_over_budget_chunk_before_report_finishes():
"""
Regression: the failure report is never awaited inside the relay, so a
Regression: the failure report never sits inside the relay, so a
single chunk that crosses the preview budget reaches the client even
while the report coroutine is still running.
while the report would still be running.
"""
first_chunk: Final = b"x" * 6144
release: Final = asyncio.Event()
@ -4764,8 +4712,6 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[asyncio.Future[None]] = []
async def _held_hook(**kwargs):
await release.wait()
@ -4776,43 +4722,40 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_held_hook)
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_held_hook)
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 500
iterator: Final = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert not release.is_set()
assert not enqueued[0].done()
release.set()
relayed: Final = b"".join(
[first]
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
)
assert relayed == first_chunk
assert isinstance(response, StreamingResponse)
assert response.status_code == 500
iterator: Final = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert not release.is_set()
relayed: Final = b"".join(
[first]
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
)
assert relayed == first_chunk
mock_proxy_logging.post_call_failure_hook.assert_not_called()
background: Final = asyncio.ensure_future(response.background())
assert not background.done()
release.set()
await background
assert len(enqueued) == 1, enqueued
await enqueued[0]
mock_proxy_logging.post_call_failure_hook.assert_called_once()
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
assert (
@ -4821,7 +4764,7 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_drains_the_report_before_stream_end():
async def test_pass_through_request_streaming_upstream_error_reports_via_background_task_after_stream_end():
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/plain"},
@ -4858,6 +4801,9 @@ async def test_pass_through_request_streaming_upstream_error_drains_the_report_b
assert response.status_code == 500
drained: Final = [chunk async for chunk in response.body_iterator]
assert drained == [b"x" * 512]
mock_proxy_logging.post_call_failure_hook.assert_not_called()
assert response.background is not None
await response.background()
mock_proxy_logging.post_call_failure_hook.assert_called_once()
@ -4924,11 +4870,7 @@ async def test_pass_through_request_streaming_upstream_error_body_read_failure_k
assert streamed_bytes == b'{"error": "half'
await upstream_response.aclose()
await _poll(
lambda: any(
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
)
)
await response.background()
rendered: Final = [str(args[0]) for args in recorded_warnings]
formats: Final = [args[0] for args in recorded_warnings]
assert any(
@ -4974,40 +4916,30 @@ async def test_pass_through_request_streaming_upstream_abort_after_preview_budge
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[asyncio.Future[None]] = []
def _recording_spawn(coro):
enqueued.append(asyncio.ensure_future(coro))
return enqueued[-1]
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=_recording_spawn,
):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
received: list[bytes] = []
@ -5020,18 +4952,21 @@ async def test_pass_through_request_streaming_upstream_abort_after_preview_budge
await consume_response()
assert b"".join(received) == b"d" * 5000
assert len(enqueued) == 1, enqueued
await enqueued[0]
mock_proxy_logging.post_call_failure_hook.assert_called_once()
expected_body: Final = f"{'d' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
== f"Upstream passthrough request failed with status 500: {expected_body}"
)
assert response.background is not None
await response.background()
mock_proxy_logging.post_call_failure_hook.assert_called_once()
@pytest.mark.asyncio
async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
async def test_preview_report_collected_reports_once_and_warns_on_disconnect():
"""Closing the relay early leaves the preview partial: report_collected logs the
disconnect warning once and reports the collected bytes, and a second call is a no-op."""
upstream_hold: Final = asyncio.Event()
chunks: Final = (b"first",)
upstream_response: Final = httpx.Response(
@ -5040,189 +4975,171 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
release: Final = asyncio.Event()
hook_done: Final = asyncio.Event()
reported: list[bytes] = []
log_warning: Final = MagicMock()
async def report(preview: bytes) -> None:
await release.wait()
hook_done.set()
reported.append(preview)
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=log_warning,
spawn=_spawn_report_task,
)
async def consume() -> None:
async for _ in relay.__aiter__():
pass
iterator: Final = relay.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert first == b"first"
await iterator.aclose()
consumer: Final = asyncio.ensure_future(consume())
for _ in range(20):
await asyncio.sleep(0)
consumer.cancel()
with pytest.raises(asyncio.CancelledError):
await consumer
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
assert len(pending) == 1, pending
assert not pending[0].done()
assert not hook_done.is_set()
await relay.report_collected()
log_warning.assert_called_once_with(
"pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body", 5
)
assert reported == [b"first"], reported
release.set()
await drain_passthrough_upstream_error_reports(timeout=5)
assert hook_done.is_set()
assert pending[0].done()
assert not _REPORT_TASKS.get(asyncio.get_running_loop())
await relay.report_collected()
assert reported == [b"first"], reported
log_warning.assert_called_once()
@pytest.mark.asyncio
async def test_shutdown_drain_returns_with_report_still_pending_on_timeout():
upstream_hold: Final = asyncio.Event()
chunks: Final = (b"first",)
async def test_preview_report_collected_runs_without_disconnect_warning_after_clean_end():
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold),
stream=_UpstreamErrorBodyStream(b"frame"),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
release: Final = asyncio.Event()
reported: list[bytes] = []
log_warning: Final = MagicMock()
async def report(preview: bytes) -> None:
await release.wait()
reported.append(preview)
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=MagicMock(),
spawn=_spawn_report_task,
log_warning=log_warning,
)
async def consume() -> None:
async for _ in relay.__aiter__():
pass
consumer: Final = asyncio.ensure_future(consume())
for _ in range(20):
await asyncio.sleep(0)
consumer.cancel()
with pytest.raises(asyncio.CancelledError):
await consumer
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
assert len(pending) == 1, pending
warnings: Final = MagicMock()
await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings)
warnings.assert_called_once_with(
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1
)
assert not pending[0].done()
pending[0].cancel()
await asyncio.gather(pending[0], return_exceptions=True)
@pytest.mark.asyncio
async def test_parked_reports_do_not_block_the_error_stream():
"""A pile of still-running reports must not sit on the response path: a stream whose
own report finishes immediately delivers its body without waiting on the others."""
hold: Final = asyncio.Event()
async def parked() -> None:
await hold.wait()
spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70))
assert len(_REPORT_TASKS.get(asyncio.get_running_loop(), ())) == 70
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStream(b"d" * 5000),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
hook_done: Final = asyncio.Event()
async def report(preview: bytes) -> None:
hook_done.set()
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=MagicMock(),
spawn=_spawn_report_task,
)
started: Final = time.monotonic()
received: Final = b"".join([chunk async for chunk in relay.__aiter__()])
elapsed: Final = time.monotonic() - started
assert received == b"d" * 5000
assert hook_done.is_set()
assert elapsed < 0.5, elapsed
assert received == b"frame"
assert reported == []
for task in spawned:
task.cancel()
await asyncio.gather(*spawned, return_exceptions=True)
def test_report_registry_is_scoped_to_each_event_loop():
"""Two consecutive asyncio.run calls: each loop's reports register and drain on their
own loop, so a shared module-level primitive bound to the first loop never breaks the second."""
async def run_once() -> None:
loop: Final = asyncio.get_running_loop()
release: Final = asyncio.Event()
async def parked() -> None:
await release.wait()
spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70))
assert len(_REPORT_TASKS.get(loop, ())) == 70
release.set()
await drain_passthrough_upstream_error_reports()
assert all(task.done() for task in spawned)
assert not _REPORT_TASKS.get(loop)
asyncio.run(run_once())
asyncio.run(run_once())
assert all(len(pending) == 0 for pending in _REPORT_TASKS.values())
await relay.report_collected()
assert reported == [b"frame"], reported
log_warning.assert_not_called()
@pytest.mark.asyncio
async def test_shutdown_drain_waits_without_a_timeout():
finished: Final = asyncio.Event()
async def report() -> None:
await asyncio.sleep(0.3)
finished.set()
task: Final = _spawn_report_task(report())
warnings: Final = MagicMock()
await drain_passthrough_upstream_error_reports(timeout=None, log_warning=warnings)
assert task.done()
assert finished.is_set()
warnings.assert_not_called()
@pytest.mark.asyncio
async def test_shutdown_drain_timeout_warns_with_the_pending_count():
hold: Final = asyncio.Event()
async def parked() -> None:
await hold.wait()
task: Final = _spawn_report_task(parked())
warnings: Final = MagicMock()
await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings)
warnings.assert_called_once_with(
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1
async def test_log_passthrough_upstream_failure_returns_background_task_for_unconsumed_error_stream():
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStream(b"frame"),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
task.cancel()
await asyncio.gather(task, return_exceptions=True)
with patch("litellm.proxy.proxy_server.proxy_logging_obj"):
relay: Final = await _log_passthrough_upstream_failure(
response=upstream_response,
user_api_key_dict=MagicMock(),
request_payload={},
logging_obj=MagicMock(),
)
assert isinstance(relay.background, BackgroundTask)
assert relay.background.func == relay.response.stream.report_collected
ok_response: Final = httpx.Response(
status_code=200,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStream(b"ok"),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
ok_relay: Final = await _log_passthrough_upstream_failure(
response=ok_response,
user_api_key_dict=MagicMock(),
request_payload={},
logging_obj=MagicMock(),
)
assert ok_relay.response is ok_response
assert ok_relay.background is None
@pytest.mark.asyncio
async def test_streamed_error_runs_failure_hook_as_background_after_response_headers_hook():
"""The streamed error report runs via the response's background task, so the
failure hook must fire after the response headers hook and after the last
body byte, still exactly once."""
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_ChunkedUpstreamErrorBodyStream((b"data: one\n\n", b"data: two\n\n")),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
calls: list[str] = []
async def _headers_hook(**kwargs):
calls.append("post_call_response_headers_hook")
return None
async def _failure_hook(**kwargs):
calls.append("post_call_failure_hook")
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_failure_hook)
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(side_effect=_headers_hook)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
scope: Final = {
"type": "http",
"asgi": {"version": "3.0", "spec_version": "2.3"},
"http_version": "1.1",
"method": "POST",
"scheme": "http",
"path": "/",
"query_string": b"",
"root_path": "",
"headers": [],
"client": ("127.0.0.1", 1234),
"server": ("127.0.0.1", 80),
}
sent: list = []
async def receive():
await asyncio.sleep(3600)
async def send(message):
sent.append(message)
await response(scope, receive, send)
assert calls == ["post_call_response_headers_hook", "post_call_failure_hook"], calls
body: Final = b"".join(
message.get("body", b"") for message in sent if message["type"] == "http.response.body"
)
assert body == b"data: one\n\ndata: two\n\n"
class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream):
@ -5295,11 +5212,7 @@ async def test_pass_through_request_streaming_upstream_error_gzip_read_failure_r
assert streamed_bytes == plaintext
await upstream_response.aclose()
await _poll(
lambda: any(
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
)
)
await response.background()
rendered: Final = [str(args[0]) for args in recorded_warnings]
assert any(
args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" and plaintext.decode() in str(args[4])
@ -5351,12 +5264,7 @@ async def test_pass_through_request_streaming_upstream_error_gzip_body_decoded_f
)
assert streamed_bytes == upstream_content
await _poll(
lambda: any(
call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
for call in mock_warning.call_args_list
)
)
await response.background()
upstream_warnings: Final = [
call
for call in mock_warning.call_args_list

View file

@ -343,96 +343,6 @@ async def test_flush_spend_counters_on_shutdown_logs_and_swallows_commit_errors(
assert "Error flushing spend counters on shutdown: db gone" in caplog.text
def _recorded_step(calls: List[str], name: str, report_done: "asyncio.Event") -> Callable[[], Awaitable[None]]:
async def _step() -> None:
calls.append(f"{name}:report_done={report_done.is_set()}")
return _step
@pytest.mark.asyncio
async def test_shutdown_flush_runs_after_in_flight_passthrough_error_reports():
"""A passthrough error report callback writes its spend row through the
logging worker, so every spend flush must run only after in-flight reports
have finished: a flush that ran while the report was still pending could
skip its row.
"""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_spawn_report_task,
drain_passthrough_upstream_error_reports,
)
report_done = asyncio.Event()
calls: List[str] = []
async def report() -> None:
await asyncio.sleep(0.05)
report_done.set()
_spawn_report_task(report())
await ps._drain_reports_then_flush_spend(
drain_reports=drain_passthrough_upstream_error_reports,
drain_spend_events=_recorded_step(calls, "drain_spend_events", report_done),
stop_scheduler_jobs=_recorded_step(calls, "stop_scheduler_jobs", report_done),
flush_spend_counters=_recorded_step(calls, "flush_spend_counters", report_done),
flush_spend_logs=_recorded_step(calls, "flush_spend_logs", report_done),
)
assert calls == [
"drain_spend_events:report_done=True",
"stop_scheduler_jobs:report_done=True",
"flush_spend_counters:report_done=True",
"flush_spend_logs:report_done=True",
], calls
@pytest.mark.asyncio
async def test_shutdown_flush_continues_when_report_drain_fails():
report_done = asyncio.Event()
calls: List[str] = []
async def failing_drain() -> None:
raise RuntimeError("drain gone")
await ps._drain_reports_then_flush_spend(
drain_reports=failing_drain,
drain_spend_events=_recorded_step(calls, "drain_spend_events", report_done),
stop_scheduler_jobs=_recorded_step(calls, "stop_scheduler_jobs", report_done),
flush_spend_counters=_recorded_step(calls, "flush_spend_counters", report_done),
flush_spend_logs=_recorded_step(calls, "flush_spend_logs", report_done),
)
assert calls == [
"drain_spend_events:report_done=False",
"stop_scheduler_jobs:report_done=False",
"flush_spend_counters:report_done=False",
"flush_spend_logs:report_done=False",
], calls
@pytest.mark.asyncio
async def test_shutdown_flush_skips_scheduler_step_when_no_scheduler():
report_done = asyncio.Event()
calls: List[str] = []
async def noop_drain() -> None:
report_done.set()
await ps._drain_reports_then_flush_spend(
drain_reports=noop_drain,
drain_spend_events=_recorded_step(calls, "drain_spend_events", report_done),
stop_scheduler_jobs=None,
flush_spend_counters=_recorded_step(calls, "flush_spend_counters", report_done),
flush_spend_logs=_recorded_step(calls, "flush_spend_logs", report_done),
)
assert calls == [
"drain_spend_events:report_done=True",
"flush_spend_counters:report_done=True",
"flush_spend_logs:report_done=True",
], calls
# ---------------------------------------------------------------------------
# _initialize_shared_aiohttp_session
# ---------------------------------------------------------------------------

View file

@ -68,25 +68,3 @@ def _build_constant_env_var_map() -> dict[str, str]:
env_var_map[constant_name] = env_var_name
return env_var_map
@pytest.mark.parametrize(
"env_value, expected",
[
(None, None),
("2.5", 2.5),
("-1", 0.0),
],
)
def test_passthrough_error_report_drain_seconds_env_parsing(monkeypatch, env_value, expected):
"""Unset waits for every report; a value clamps to >= 0 seconds."""
if env_value is None:
monkeypatch.delenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", raising=False)
else:
monkeypatch.setenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", env_value)
try:
reloaded = importlib.reload(constants)
assert reloaded.PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS == expected
finally:
monkeypatch.undo()
importlib.reload(constants)