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>
This commit is contained in:
yucheng 2026-09-26 10:12:01 +00:00
parent 18cd44cc5d
commit b04c9830b9
3 changed files with 454 additions and 0 deletions

View file

@ -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

View file

@ -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]

View file

@ -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