diff --git a/litellm/constants.py b/litellm/constants.py index 89996a36f96..e7ba1f6b07f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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" diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index f41a19e12b3..5004a48623f 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 616ca638a80..61304f0d919 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index fcbaf7c8d8c..9cb3dace305 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -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, diff --git a/tests/integration/observability/test_passthrough_upstream_error_chaos.py b/tests/integration/observability/test_passthrough_upstream_error_chaos.py index 70035030ebc..d4a9bf4fcf0 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_chaos.py +++ b/tests/integration/observability/test_passthrough_upstream_error_chaos.py @@ -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( diff --git a/tests/integration/observability/test_passthrough_upstream_error_visibility.py b/tests/integration/observability/test_passthrough_upstream_error_visibility.py index c302d36ad88..9f4f83ee005 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_visibility.py +++ b/tests/integration/observability/test_passthrough_upstream_error_visibility.py @@ -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 diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 7ec0491c824..1604970796a 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 3cbbe140f2c..6feb37e9867 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/unit/test_constants.py b/tests/unit/test_constants.py index d13cc2b0e88..12e473f68a4 100644 --- a/tests/unit/test_constants.py +++ b/tests/unit/test_constants.py @@ -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)