From b04c9830b95160d326bc684ea8f215a4e48564e8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 10:12:01 +0000 Subject: [PATCH] test(passthrough): cover round 5 dispatch, bounded teardown, and bounded redaction Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_passthrough_upstream_error_chaos.py | 291 ++++++++++++++++++ .../test_pass_through_endpoints.py | 142 +++++++++ .../proxy/test_common_request_processing.py | 21 ++ 3 files changed, 454 insertions(+) diff --git a/tests/integration/observability/test_passthrough_upstream_error_chaos.py b/tests/integration/observability/test_passthrough_upstream_error_chaos.py index 0376ed5459c..e424e8597dc 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_chaos.py +++ b/tests/integration/observability/test_passthrough_upstream_error_chaos.py @@ -3,6 +3,8 @@ import json import re import signal import threading +import time +import uuid from pathlib import Path from typing import Final @@ -461,3 +463,292 @@ async def test_passthrough_error_report_survives_sigterm_after_full_response( owned.process.wait(timeout=60) await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_done").exists(), bool, 30) _single_spend_row(call_id) + + +_BUDGET_MARKER_FAILURE_HOOK: Final = """ +import asyncio +from pathlib import Path + +from litellm.integrations.custom_logger import CustomLogger + + +class BudgetMarkerFailureHook(CustomLogger): + async def async_post_call_failure_hook( + self, request_data, original_exception, user_api_key_dict, traceback_str=None + ): + call_id = (request_data or {{}}).get("litellm_call_id") or "unknown" + Path({directory!r}, f"started-{{call_id}}").touch() + await asyncio.sleep(12) + Path({directory!r}, f"done-{{call_id}}").touch() + + +instance = BudgetMarkerFailureHook() +""" + +_BUDGET_CHUNK: Final = b"e" * 8000 + + +async def test_passthrough_sigterm_graceful_window_reports_dispatched_at_preview_budget( + gateway: Gateway, tmp_path: Path +) -> None: + """The upstream sends 8000 bytes then stalls for 60 s: the failure report must + dispatch when the preview budget is crossed, so a graceful shutdown can drain + the 12 s hooks it would otherwise find still buffered behind the gate.""" + gate: Final = threading.Event() + markers: Final = tmp_path / "budget-markers" + markers.mkdir() + + def respond(request: Request) -> Reply: + if "streamGenerateContent" in request.target: + return Reply( + status=500, + content_type="text/event-stream", + chunks=(_BUDGET_CHUNK, b"data: tail\n\n"), + gate_after_first=gate, + gate_timeout_seconds=120, + ) + return Reply(status=200, body=json.dumps({"ok": True}).encode()) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update({"callbacks": ["budget_hook.instance"]}) + (tmp_path / "budget_hook.py").write_text(_BUDGET_MARKER_FAILURE_HOOK.format(directory=str(markers))) + path: Final = tmp_path / "chaos-budget-sigterm.yaml" + path.write_text(yaml.safe_dump(config)) + 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=20) as owned: + + async def hold_stream() -> str | None: + async with httpx.AsyncClient( + base_url=str(owned.gateway.client.base_url), timeout=httpx.Timeout(30, connect=5), trust_env=False + ) as client: + try: + async with client.stream( + "POST", + "/gemini/v1beta/models/nope-9:streamGenerateContent", + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers={ + "Authorization": f"Bearer {owned.gateway.key}", + "x-goog-api-key": owned.gateway.key, + }, + ) as response: + return response.headers["x-litellm-call-id"] + except (httpx.TransportError, httpx.TimeoutException): + return None + + try: + streams: Final = asyncio.gather(*(hold_stream() for _ in range(20)), return_exceptions=True) + await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 20, 60) + owned.process.send_signal(signal.SIGTERM) + owned.process.wait(timeout=90) + await streams + finally: + gate.set() + done_markers: Final = tuple(markers.glob("done-*")) + assert len(done_markers) == 20, sorted(p.name for p in markers.iterdir()) + call_ids: Final = [call_id for call_id in await streams if isinstance(call_id, str)] + assert len(call_ids) == 20, call_ids + for call_id in call_ids: + _single_spend_row(call_id) + + +_SLOW_HEADERS_HOOK: Final = """ +import asyncio +from pathlib import Path + +from litellm.integrations.custom_logger import CustomLogger + + +class SlowHeadersHook(CustomLogger): + async def async_post_call_response_headers_hook( + self, data, user_api_key_dict, response, request_headers, **kwargs + ): + await asyncio.sleep(15) + return {{}} + + async def async_post_call_failure_hook( + self, request_data, original_exception, user_api_key_dict, traceback_str=None + ): + call_id = (request_data or {{}}).get("litellm_call_id") or "unknown" + Path({directory!r}, f"failure-reported-{{call_id}}").touch() + + +instance = SlowHeadersHook() +""" + + +async def test_passthrough_sigterm_during_headers_hook_still_dispatches_the_report( + gateway: Gateway, tmp_path: Path +) -> None: + """SIGTERM cancelling the request task mid-headers-hook is an exit between + building the relay and returning its StreamingResponse: the relay must + dispatch its report there instead of losing the failure.""" + markers: Final = tmp_path / "headers-markers" + markers.mkdir() + marker: Final = uuid.uuid4().hex + + def respond(request: Request) -> Reply: + if "streamGenerateContent" in request.target: + return Reply( + status=500, + content_type="text/event-stream", + chunks=(f"data: upstream blew up {marker}\n\n".encode(), b"data: [DONE]\n\n"), + delay_before_headers_seconds=1, + ) + return Reply(status=200, body=json.dumps({"ok": True}).encode()) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update({"callbacks": ["slow_headers_hook.instance"]}) + (tmp_path / "slow_headers_hook.py").write_text(_SLOW_HEADERS_HOOK.format(directory=str(markers))) + path: Final = tmp_path / "chaos-headers-hook-sigterm.yaml" + path.write_text(yaml.safe_dump(config)) + 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: + + async def stream_error() -> None: + async with httpx.AsyncClient( + base_url=str(owned.gateway.client.base_url), timeout=httpx.Timeout(60, connect=5), trust_env=False + ) as client: + try: + await client.post( + "/gemini/v1beta/models/nope-9:streamGenerateContent", + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers={ + "Authorization": f"Bearer {owned.gateway.key}", + "x-goog-api-key": owned.gateway.key, + }, + ) + except (httpx.TransportError, httpx.TimeoutException): + pass + + request_task: Final = asyncio.create_task(stream_error()) + await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 1, 30) + owned.process.send_signal(signal.SIGTERM) + owned.process.wait(timeout=30) + await request_task + reported: Final = await asyncio.to_thread( + eventually, + lambda: tuple(markers.glob("failure-reported-*")), + lambda paths: len(paths) == 1, + 30, + ) + call_id: Final = reported[0].name.removeprefix("failure-reported-") + _single_spend_row(call_id) + error_information: Final = _error_information(call_id) + assert error_information["error_code"] == "500", error_information + + +_PARKING_HEADERS_HOOK_FAILURE_ONLY: Final = """ +import asyncio +from pathlib import Path + +from litellm.integrations.custom_logger import CustomLogger + + +class NeverReturningFailureHook(CustomLogger): + async def async_post_call_failure_hook( + self, request_data, original_exception, user_api_key_dict, traceback_str=None + ): + Path({marker!r}).touch() + await asyncio.Event().wait() + + +instance = NeverReturningFailureHook() +""" + + +async def test_passthrough_sigterm_exits_when_late_relay_background_never_returns( + gateway: Gateway, tmp_path: Path +) -> None: + """Keepalive replays the late response's background task inside aclose: a + failure hook that never returns must be abandoned after the bound so SIGTERM + still gets the process out.""" + hook_started: Final = tmp_path / "late-relay-hook-started" + + def respond(request: Request) -> Reply: + if "streamGenerateContent" in request.target: + return Reply( + status=500, + content_type="text/event-stream", + chunks=(b"data: err\n\n", b"data: [DONE]\n\n"), + delay_before_headers_seconds=3, + ) + return Reply(status=200, body=json.dumps({"ok": True}).encode()) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update({"callbacks": ["park_hook.instance"], "sse_keepalive_ping_interval_seconds": 0.2}) + (tmp_path / "park_hook.py").write_text(_PARKING_HEADERS_HOOK_FAILURE_ONLY.format(marker=str(hook_started))) + path: Final = tmp_path / "chaos-late-relay-bound.yaml" + path.write_text(yaml.safe_dump(config)) + 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: + async with httpx.AsyncClient( + base_url=str(owned.gateway.client.base_url), + timeout=httpx.Timeout(10, read=1.5, connect=5), + trust_env=False, + ) as client: + try: + async with client.stream( + "POST", + "/gemini/v1beta/models/nope-9:streamGenerateContent", + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers={ + "Authorization": f"Bearer {owned.gateway.key}", + "x-goog-api-key": owned.gateway.key, + }, + ) as response: + await response.aread() + except (httpx.TransportError, httpx.TimeoutException): + pass + await asyncio.to_thread(eventually, lambda: hook_started.exists(), bool, 60) + owned.process.send_signal(signal.SIGTERM) + owned.process.wait(timeout=40) + log_text: Final = owned.log.read_text() + assert "relayed response background task still running after 10s" in log_text, log_text[-4000:] + assert "upstream error reports still running after 10s shutdown wait" in log_text, log_text[-4000:] + + +_BIG_ERROR_BODY: Final = ( + '{"error":{"code":500,"message":"upstream blew up with key sk-leak0leak0leak0leak0 then ' + + "x" * (5 * 1024 * 1024) + + '","status":"INTERNAL"}}' +).encode() + + +def test_passthrough_buffered_error_body_redaction_stays_bounded(gateway: Gateway, tmp_path: Path) -> None: + """A 5 MiB buffered error body must not be fully decoded and redacted: the + request returns fast and the spend row still carries the redacted prefix.""" + + def respond(request: Request) -> Reply: + if "generateContent" in request.target: + return Reply(status=500, body=_BIG_ERROR_BODY) + return Reply(status=200, body=json.dumps({"ok": True}).encode()) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + path: Final = tmp_path / "chaos-big-error-body.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: + started: Final = time.perf_counter() + response: Final = owned.gateway.request( + "POST", + "/gemini/v1beta/models/nope-9:generateContent", + _GENERATE_CONTENT, + headers={"x-goog-api-key": owned.gateway.key}, + ) + elapsed: Final = time.perf_counter() - started + assert response.status_code == 500, response.status_code + call_id: Final = response.headers["x-litellm-call-id"] + _single_spend_row(call_id) + error_information: Final = _error_information(call_id) + assert "REDACTED" in str(error_information["error_message"]), error_information + assert elapsed < 0.5, elapsed 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 ef56cf057c7..cac18ded4ac 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 @@ -34,6 +34,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( HttpPassThroughEndpointHelpers, InitPassThroughEndpointHelpers, _PreviewReportingStream, + _passthrough_upstream_failure_reporter, _registered_pass_through_routes, _truncate_upstream_error_body, _with_trace_context, @@ -8189,3 +8190,144 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token() metadata = kwargs["litellm_params"]["metadata"] assert metadata["user_api_key"] == "cli-session-alice" assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice" + + +@pytest.mark.asyncio +async def test_preview_stream_dispatches_the_report_as_soon_as_the_preview_budget_is_crossed(): + """A body that crosses the 4 KiB preview budget reports immediately instead of + waiting for EOF, so a held-open upstream cannot stall the failure hook.""" + chunks: Final = tuple(b"d" * 1000 for _ in range(5)) + hold: Final = asyncio.Event() + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=_UpstreamErrorBodyStreamHeld(chunks, hold), + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + reported: list[bytes] = [] + report_started: Final = asyncio.Event() + release_report: Final = asyncio.Event() + + async def report(preview: bytes) -> None: + reported.append(preview) + report_started.set() + await release_report.wait() + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=MagicMock(), + ) + received: list[bytes] = [] + + async def consume() -> None: + async for chunk in relay.__aiter__(): + received.append(chunk) + + consumer: Final = asyncio.create_task(consume()) + await asyncio.wait_for(report_started.wait(), timeout=5) + named: Final = [t for t in asyncio.all_tasks() if t.get_name() == "passthrough-upstream-error-report"] + assert len(named) == 1, [t.get_name() for t in asyncio.all_tasks()] + assert reported == [b"d" * 5000], reported + release_report.set() + hold.set() + await consumer + assert b"".join(received) == b"d" * 5000, received + + +@pytest.mark.asyncio +async def test_pass_through_request_dispatches_the_report_when_the_headers_hook_fails(): + """An exception between building the relay and returning the StreamingResponse + (here, a raising post_call_response_headers_hook) must still start the report.""" + 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"), + ) + + 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() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock( + side_effect=RuntimeError("headers hook blew up") + ) + 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) + + with pytest.raises(ProxyException, match="headers hook blew up"): + 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, + ) + + for _ in range(10): + await asyncio.sleep(0) + upstream_exceptions: Final = [ + call.kwargs["original_exception"] + for call in mock_proxy_logging.post_call_failure_hook.call_args_list + if "Upstream passthrough request failed" in str(getattr(call.kwargs["original_exception"], "detail", "")) + ] + assert len(upstream_exceptions) == 1, mock_proxy_logging.post_call_failure_hook.call_args_list + assert upstream_exceptions[0].status_code == 500, upstream_exceptions[0] + await upstream_response.aclose() + + +@pytest.mark.asyncio +async def test_passthrough_upstream_failure_reporter_redacts_only_the_bounded_prefix(): + """A huge upstream error body must not make the report decode and redact all of + it: only MAX+MARGIN chars reach redact_secrets, and content before the cut still + lands in the failure detail redacted.""" + marker_key: Final = "sk-" + "leak0" * 8 + preview: Final = b"A" * 10 + marker_key.encode() + b"z" * (5 * 1024 * 1024) + redacted_inputs: list[str] = [] + + def recording_redact(text: str) -> str: + redacted_inputs.append(text) + return text.replace(marker_key, "REDACTED-KEY") + + logged_details: list[str] = [] + + def recording_warning(fmt, *args, **kwargs): + if str(fmt).startswith("pass_through_endpoint: upstream"): + logged_details.append(str(fmt % args)) + + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "application/json"}, + request=httpx.Request("POST", "http://target-api.com/v1/chat/completions"), + content=b"{}", + ) + logging_obj: Final = MagicMock() + logging_obj.model_call_details = {} + proxy_logging: Final = MagicMock() + proxy_logging.post_call_failure_hook = AsyncMock() + report: Final = _passthrough_upstream_failure_reporter( + response=upstream_response, + user_api_key_dict=MagicMock(), + request_payload={}, + logging_obj=logging_obj, + proxy_logging=proxy_logging, + log_warning=recording_warning, + redact=recording_redact, + ) + await report(preview) + max_plus_margin: Final = 4096 + 256 + assert len(redacted_inputs) == 1, redacted_inputs + assert len(redacted_inputs[0]) <= max_plus_margin, len(redacted_inputs[0]) + assert len(logged_details) == 1, logged_details + assert "REDACTED-KEY" in logged_details[0], logged_details[0] + assert marker_key not in logged_details[0], logged_details[0] diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index f5f7f7e2a26..f50ae112239 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -10262,3 +10262,24 @@ async def test_aclose_late_response_runs_background_task_for_non_streaming_respo produced: Final = StarletteResponse(content=b"{}", background=BackgroundTask(mark)) await _aclose_late_response(produced) assert ran == [True] + + +@pytest.mark.asyncio +async def test_aclose_late_response_bounds_a_never_returning_background_task(caplog): + import logging + + from starlette.background import BackgroundTask + from starlette.responses import StreamingResponse + + from litellm.proxy.common_request_processing import _aclose_late_response + + async def body(): + yield b"x" + + async def parks() -> None: + await asyncio.Event().wait() + + produced: Final = StreamingResponse(body(), background=BackgroundTask(parks)) + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await asyncio.wait_for(_aclose_late_response(produced, background_wait_seconds=0.05), timeout=1) + assert "relayed response background task still running after 0s" in caplog.text, caplog.text