fix(passthrough): stream non-sse passthrough responses instead of buffering in memory (#32386)

* fix(passthrough): stream non-sse passthrough responses instead of buffering in memory

Non-SSE passthrough responses were fully read into proxy memory (content = await response.aread()) before the first byte reached the client. For large non-JSON bodies such as Anthropic batch results jsonl files this ballooned proxy RSS to a multiple of the file size and produced near-total TTFB dead air, letting intermediaries kill the silent connection and truncate the download.

The upstream request is now sent with httpx stream semantics and the buffering decision is made from the response headers: application/json (and +json) bodies plus upstream errors keep the buffered behavior since spend logging, guardrails and managed-id rewriting inspect them, while every other 2xx body is relayed as a StreamingResponse that iterates upstream bytes without accumulating them, preserving status code and headers (including x-litellm-*) and firing the success-handler logging with response_body=None once the stream completes.

* fix(passthrough): log client disconnects mid-stream and derive test client cache key from production code

* test(passthrough): intercept AsyncClient.send in legacy passthrough tests and assert final wire params

* test(passthrough): fail with a clear assert when the passthrough client cache scan misses
This commit is contained in:
Mateo Wang 2026-07-07 20:51:15 -07:00 • committed by GitHub
parent 734fd29e00
commit b2e2a38bc0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 582 additions and 68 deletions

View file

@ -7,7 +7,7 @@ import traceback
from base64 import b64encode
from datetime import datetime
from itertools import groupby
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast
from typing import Any, AsyncGenerator, Dict, List, Mapping, Optional, Tuple, Union, cast
from urllib.parse import urlencode, urlparse
import httpx
@ -389,18 +389,24 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
forward_multipart: bool = False,
) -> httpx.Response:
"""
Handle non-streaming HTTP requests
Handle non-SSE HTTP requests
Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests
Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests.
GET and generic requests are sent with httpx stream semantics so the caller can
decide from the response headers whether to buffer the body (JSON, inspected for
logging/guardrails) or relay it to the client without materializing it in memory
(LIT-4009: large batch results files must not be buffered in proxy RSS).
"""
if request.method == "GET":
response = await async_client.request(
method=request.method,
url=url,
get_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
)
elif HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart:
return await async_client.send(get_request, stream=True)
if HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart:
# Forward multipart via make_multipart_http_request even when _parsed_body is
# non-empty (pass_through_request always injects litellm_logging_obj, etc.).
# forward_multipart is False when custom_body was supplied (JSON body despite
@ -412,16 +418,14 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
headers=headers,
requested_query_params=requested_query_params,
)
else:
# Generic httpx method
response = await async_client.request(
method=request.method,
url=url,
headers=headers,
params=requested_query_params,
json=_parsed_body,
)
return response
generic_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
json=_parsed_body,
)
return await async_client.send(generic_request, stream=True)
@staticmethod
def is_multipart(request: Request) -> bool:
@ -1161,13 +1165,14 @@ async def pass_through_request(
if state_raw_body is not None:
# SigV4-signed callers (Bedrock) require the exact pre-signed bytes
# to be forwarded so the signature/Content-Length stay valid.
response = await async_client.request(
method=request.method,
url=url,
raw_body_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
content=state_raw_body,
)
response = await async_client.send(raw_body_request, stream=True)
else:
response = await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
request=request,
@ -1223,6 +1228,40 @@ async def pass_through_request(
status_code=response.status_code,
)
if not _should_buffer_passthrough_response(response):
relay_custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=litellm_call_id,
model_id=None,
cache_key=None,
api_base=str(url._uri_reference),
)
relay_callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=_parsed_body or {},
user_api_key_dict=user_api_key_dict,
response=response,
request_headers=dict(request.headers),
)
if relay_callback_headers:
relay_custom_headers.update(relay_callback_headers)
return StreamingResponse(
_relay_passthrough_response_bytes(
response=response,
request_body=_parsed_body or {},
url_route=str(url),
start_time=start_time,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
success_handler_kwargs=kwargs,
),
status_code=response.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=response.headers,
custom_headers=relay_custom_headers,
),
)
content = await response.aread()
## POST-CALL GUARDRAILS ##
@ -2211,6 +2250,70 @@ def _is_streaming_response(response: httpx.Response) -> bool:
return False
def _should_buffer_passthrough_response(response: httpx.Response) -> bool:
"""
Decide from the response headers whether the body must be read into memory.
JSON bodies (and upstream errors) stay buffered: spend logging, guardrails and
managed-id rewriting inspect them, and they are small in practice. Everything
else (jsonl batch results, octet-stream files, ...) is relayed to the client
chunk by chunk so a large body is never resident in full (LIT-4009). A missing
content-type is buffered because the body cannot be classified.
"""
if response.status_code >= 400:
return True
media_type = response.headers.get("content-type", "").split(";")[0].strip().lower()
return media_type in ("", "application/json") or media_type.endswith("+json")
async def _relay_passthrough_response_bytes(
response: httpx.Response,
request_body: dict,
url_route: str,
start_time: datetime,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str],
success_handler_kwargs: dict,
) -> AsyncGenerator[bytes, None]:
"""
Yield upstream bytes to the client without accumulating them, then fire the
passthrough success handler with response_body=None (uninspected body). The
finally block also runs on client disconnect (GeneratorExit) so partial
downloads still produce a spend-log row, mirroring chunk_processor; a
disconnect additionally logs a warning with the number of bytes relayed so
partial deliveries are distinguishable from complete ones in proxy logs.
"""
bytes_relayed = 0
upstream_fully_relayed = False
try:
async for chunk in response.aiter_bytes():
bytes_relayed += len(chunk)
yield chunk
upstream_fully_relayed = True
finally:
if not upstream_fully_relayed:
verbose_proxy_logger.warning(
f"Passthrough stream for {url_route} ended before upstream body was fully relayed; "
f"{bytes_relayed} bytes were sent to the client"
)
await response.aclose()
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
httpx_response=response,
response_body=None,
url_route=url_route,
result="",
start_time=start_time,
end_time=datetime.now(),
logging_obj=logging_obj,
cache_hit=False,
request_body=request_body,
custom_llm_provider=custom_llm_provider,
**success_handler_kwargs,
)
)
def _extract_model_from_vertex_ai_setup(setup_response: dict) -> Optional[str]:
"""
Extract the model name from Vertex AI Live setup response.

View file

@ -34,6 +34,18 @@ from .llm_provider_handlers.vertex_passthrough_logging_handler import (
cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler()
def _safe_response_text(httpx_response: httpx.Response) -> str:
"""
Streamed passthrough responses are relayed to the client without being read
into memory, so accessing .text on them raises ResponseNotRead. Their body is
intentionally uninspected; log an empty string instead of failing the row.
"""
try:
return httpx_response.text
except httpx.ResponseNotRead:
return ""
class PassThroughEndpointLogging:
def __init__(self):
self.TRACKED_VERTEX_ROUTES = [
@ -306,7 +318,9 @@ class PassThroughEndpointLogging:
]
kwargs = normalized_llm_passthrough_logging_payload["kwargs"]
if standard_logging_response_object is None:
standard_logging_response_object = StandardPassThroughResponseObject(response=httpx_response.text)
standard_logging_response_object = StandardPassThroughResponseObject(
response=_safe_response_text(httpx_response)
)
kwargs = self._set_cost_per_request(
logging_obj=logging_obj,

View file

@ -22,10 +22,8 @@ from litellm.proxy.proxy_server import initialize_pass_through_endpoints
# Mock the async_client used in the pass_through_request function
async def mock_request(*args, **kwargs):
mock_response = httpx.Response(200, json={"message": "Mocked response"})
mock_response.request = Mock(spec=httpx.Request)
return mock_response
async def mock_request(self, request, **kwargs):
return httpx.Response(200, json={"message": "Mocked response"}, request=request)
def remove_rerank_route(app):
@ -49,8 +47,8 @@ def client():
@pytest.mark.asyncio
async def test_pass_through_endpoint_no_headers(client, monkeypatch):
# Mock the httpx.AsyncClient.request method
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
# Mock the httpx.AsyncClient.send method
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
# Define a pass-through endpoint
@ -79,8 +77,8 @@ async def test_pass_through_endpoint_no_headers(client, monkeypatch):
@pytest.mark.asyncio
async def test_pass_through_endpoint(client, monkeypatch):
# Mock the httpx.AsyncClient.request method
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
# Mock the httpx.AsyncClient.send method
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
# Define a pass-through endpoint
@ -181,7 +179,7 @@ async def test_pass_through_endpoint_rpm_limit(
expected_status_codes,
num_users,
):
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache
@ -285,7 +283,7 @@ async def test_pass_through_endpoint_rpm_limit(
async def test_pass_through_endpoint_sequential_rpm_limit(
client, monkeypatch, auth, rpm_limit, requests_to_make, expected_status_codes
):
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache
@ -504,10 +502,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
captured_requests = []
async def mock_bing_request(*args, **kwargs):
async def mock_bing_request(self, request, **kwargs):
captured_requests.append((args, kwargs))
mock_response = httpx.Response(
captured_requests.append(request)
return httpx.Response(
200,
json={
"_type": "SearchResponse",
@ -518,11 +516,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
"value": [],
},
},
request=request,
)
mock_response.request = Mock(spec=httpx.Request)
return mock_response
monkeypatch.setattr("httpx.AsyncClient.request", mock_bing_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_bing_request)
# Define a pass-through endpoint
pass_through_endpoints = [
@ -555,8 +552,8 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
client.get("/bing/search?q=bob+barker")
client.get("/bing/search-no-merge-params?q=bob+barker")
first_transformed_url = captured_requests[0][1]["url"]
second_transformed_url = captured_requests[1][1]["url"]
first_transformed_url = captured_requests[0].url
second_transformed_url = captured_requests[1].url
# Parse URLs to compare query params order-independently
# Parse first URL
@ -573,7 +570,7 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
"setLang": ["en-US"],
"mkt": ["en-US"],
}
expected_second_params = {"setLang": ["en-US"], "mkt": ["en-US"]}
expected_second_params = {"q": ["bob barker"]}
# Assert the response - compare base URL and params separately
assert (

View file

@ -387,8 +387,9 @@ async def test_pass_through_request_stream_param_no_override(
# Create mocks for the async client
mock_async_client = AsyncMock()
# Mock request to return the non-streaming response
mock_async_client.request.return_value = mock_response
# Mock build_request/send to return the non-streaming response
mock_async_client.build_request = Mock(return_value=Mock())
mock_async_client.send.return_value = mock_response
# Mock get_async_httpx_client to return our mock client
mock_client_obj = Mock()
@ -420,20 +421,19 @@ async def test_pass_through_request_stream_param_no_override(
stream=False, # Should be used since no stream in request body
)
# Verify that build_request was NOT called (no streaming path)
mock_async_client.build_request.assert_not_called()
# Verify that send was NOT called (no streaming path)
mock_async_client.send.assert_not_called()
# Verify that the non-streaming request method WAS called
mock_async_client.request.assert_called_once_with(
method="POST",
url=httpx.URL("https://api.anthropic.com/v1/messages"),
# Non-SSE requests are sent with stream semantics so large bodies can
# be relayed without buffering; the JSON response below is still
# buffered into a plain Response.
mock_async_client.request.assert_not_called()
mock_async_client.build_request.assert_called_once_with(
"POST",
httpx.URL("https://api.anthropic.com/v1/messages"),
headers={"Authorization": "Bearer test-key"},
params={},
json=request_body,
)
mock_async_client.send.assert_called_once()
assert mock_async_client.send.call_args.kwargs.get("stream") is True
# Verify response is a regular Response (not StreamingResponse)
from fastapi.responses import Response, StreamingResponse

View file

@ -1918,7 +1918,8 @@ class TestForwardHeaders:
):
# Setup mock httpx client
mock_client = MagicMock()
mock_client.request = AsyncMock(return_value=mock_httpx_response)
mock_client.build_request = MagicMock(return_value=MagicMock())
mock_client.send = AsyncMock(return_value=mock_httpx_response)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_client
mock_get_client.return_value = mock_client_obj
@ -1942,10 +1943,10 @@ class TestForwardHeaders:
)
# Verify the httpx client was called
assert mock_client.request.called
assert mock_client.send.called
# Get the headers that were sent to the target
call_args = mock_client.request.call_args
call_args = mock_client.build_request.call_args
sent_headers = call_args[1]["headers"]
# Verify user headers were forwarded (except content-length and host)
@ -2019,7 +2020,8 @@ class TestForwardHeaders:
):
# Setup mock httpx client
mock_client = MagicMock()
mock_client.request = AsyncMock(return_value=mock_httpx_response)
mock_client.build_request = MagicMock(return_value=MagicMock())
mock_client.send = AsyncMock(return_value=mock_httpx_response)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_client
mock_get_client.return_value = mock_client_obj
@ -2043,10 +2045,10 @@ class TestForwardHeaders:
)
# Verify the httpx client was called
assert mock_client.request.called
assert mock_client.send.called
# Get the headers that were sent to the target
call_args = mock_client.request.call_args
call_args = mock_client.build_request.call_args
sent_headers = call_args[1]["headers"]
# Verify only custom headers were sent

View file

@ -1,5 +1,6 @@
import asyncio
import json
import logging
import os
import sys
from contextlib import ExitStack
@ -1337,7 +1338,8 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
upstream_response.raise_for_status = MagicMock()
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
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 def _empty_chunks(*args, **kwargs):
@ -1361,7 +1363,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
stream=False,
)
async_client.request.assert_awaited_once()
async_client.send.assert_awaited_once()
mock_chunk_processor.assert_called_once()
logging_obj = mock_chunk_processor.call_args.kwargs[
@ -3046,7 +3048,8 @@ async def test_pass_through_request_non_streaming_uses_content_for_state_raw_bod
)
mock_async_client = AsyncMock()
mock_async_client.request = AsyncMock(return_value=upstream)
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=upstream)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_async_client
@ -3082,10 +3085,12 @@ async def test_pass_through_request_non_streaming_uses_content_for_state_raw_bod
stream=False,
)
mock_async_client.request.assert_called_once()
req_kw = mock_async_client.request.call_args[1]
assert req_kw.get("content") == raw_signed
assert "json" not in req_kw
mock_async_client.build_request.assert_called_once()
build_kw = mock_async_client.build_request.call_args[1]
assert build_kw.get("content") == raw_signed
assert "json" not in build_kw
mock_async_client.send.assert_awaited_once()
assert mock_async_client.send.call_args.kwargs.get("stream") is True
@pytest.mark.asyncio
@ -3826,7 +3831,8 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan
mock_success_handler.return_value = None
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
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)
mock_request = MagicMock(spec=Request)
@ -3913,7 +3919,8 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa
mock_success_handler.return_value = None
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
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)
mock_request = MagicMock(spec=Request)
@ -4043,7 +4050,8 @@ async def test_pass_through_request_non_streaming_success_unchanged():
mock_success_handler.return_value = None
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
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)
mock_request = MagicMock(spec=Request)
@ -4103,3 +4111,393 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio
assert int(exc_info.value.code) == 500
assert "auth backend unavailable" in exc_info.value.message
class _RecordingUpstreamByteStream(httpx.AsyncByteStream):
def __init__(self, chunks):
self._chunks = chunks
self.chunks_served = 0
self.closed = False
async def __aiter__(self):
for chunk in self._chunks:
self.chunks_served += 1
yield chunk
async def aclose(self):
self.closed = True
class _FakeUpstreamTransport(httpx.AsyncBaseTransport):
def __init__(self, status_code, headers, stream):
self._status_code = status_code
self._headers = headers
self._stream = stream
async def handle_async_request(self, request):
return httpx.Response(
status_code=self._status_code,
headers=self._headers,
stream=self._stream,
request=request,
)
def _inject_fake_passthrough_client(transport, timeout):
"""Dependency-inject a fake upstream via the client cache that
get_async_httpx_client resolves passthrough clients from (no monkeypatching
of the HTTP layer). The cache entry is located by calling the production
get_async_httpx_client and identity-scanning the cache for the handler it
returned, so the internal cache-key format is never duplicated here. Must
run inside the test's event loop because cache keys are loop-scoped.
Returns (client, cleanup)."""
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
real_handler = get_async_httpx_client(
httpxSpecialProvider.PassThroughEndpoint,
params={"timeout": resolve_pass_through_request_timeout(timeout)},
)
cache = litellm.in_memory_llm_clients_cache
cache_key = next(
(key for key, cached in cache.cache_dict.items() if cached is real_handler),
None,
)
assert cache_key is not None, (
"PassThroughEndpoint client not found in in_memory_llm_clients_cache; "
"get_async_httpx_client may not be caching this provider."
)
fake_client = httpx.AsyncClient(transport=transport)
cache.cache_dict[cache_key] = SimpleNamespace(client=fake_client)
def _cleanup():
cache.cache_dict.pop(cache_key, None)
return fake_client, _cleanup
def _enter_relay_logging_mocks(stack, parsed_body):
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
mock_proxy_logging = stack.enter_context(
patch("litellm.proxy.proxy_server.proxy_logging_obj")
)
mock_proxy_logging.pre_call_hook = AsyncMock(return_value=parsed_body)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler = stack.enter_context(
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
)
)
mock_success_handler.return_value = None
stack.enter_context(
patch.object(
GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue", new=MagicMock()
)
)
return mock_proxy_logging, mock_success_handler
def _relay_client_request(method="GET"):
mock_request = MagicMock(spec=Request)
mock_request.method = method
mock_request.url = "http://localhost:4000/passthrough-relay/results"
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
return mock_request
@pytest.mark.asyncio
async def test_pass_through_request_relays_non_json_body_without_buffering():
"""
Regression (LIT-4009): non-SSE passthrough responses used to be fully
buffered in proxy memory (content = await response.aread()) before a single
byte reached the client, ballooning proxy RSS to a multiple of the body size
for large non-JSON downloads (e.g. Anthropic batch results .jsonl files) and
producing near-total TTFB dead air that let intermediaries kill the silent
connection mid-download.
A non-JSON 2xx body must be relayed as a StreamingResponse whose chunks are
pulled from the upstream one at a time, with zero chunks consumed before the
handler returns, upstream status/headers plus x-litellm-* headers preserved,
and the success-handler logging fired with response_body=None once the
stream completes. Pre-fix, the handler returned a plain Response after
reading the entire body, so these assertions fail on the old code.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (
b'{"custom_id": "a", "result": {}}\n',
b'{"custom_id": "b", "result": {}}\n',
b'{"custom_id": "c", "result": {}}\n',
)
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={
"content-type": "application/x-jsonl",
"x-upstream-marker": "batch-results",
"content-length": str(sum(len(c) for c in upstream_chunks)),
},
stream=upstream_stream,
),
timeout=311.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=311.0,
)
assert isinstance(response, StreamingResponse)
assert upstream_stream.chunks_served == 0
mock_success_handler.assert_not_called()
iterator = response.body_iterator
first_chunk = await iterator.__anext__()
assert first_chunk == upstream_chunks[0]
assert upstream_stream.chunks_served == 1
remaining = [chunk async for chunk in iterator]
assert b"".join([first_chunk, *remaining]) == b"".join(upstream_chunks)
assert upstream_stream.closed is True
assert response.status_code == 200
assert response.headers["x-upstream-marker"] == "batch-results"
assert "x-litellm-call-id" in response.headers
assert "content-length" not in response.headers
mock_success_handler.assert_called_once()
success_kwargs = mock_success_handler.call_args.kwargs
assert success_kwargs["response_body"] is None
assert (
success_kwargs["url_route"]
== "http://upstream.test/v1/messages/batches/b1/results"
)
finally:
cleanup()
await fake_client.aclose()
@pytest.mark.asyncio
async def test_pass_through_request_json_response_stays_buffered_for_logging():
"""
JSON responses (content-type application/json) must keep the buffered
behavior: spend logging and guardrails inspect the parsed body, so the
handler reads the full upstream body and passes the parsed dict to the
success handler.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (b'{"id": "file-123"', b', "status": "processed"}')
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={"content-type": "application/json"},
stream=upstream_stream,
),
timeout=312.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/files/file-123",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=312.0,
)
assert not isinstance(response, StreamingResponse)
assert response.status_code == 200
assert response.body == b"".join(upstream_chunks)
assert upstream_stream.chunks_served == len(upstream_chunks)
mock_success_handler.assert_called_once()
success_kwargs = mock_success_handler.call_args.kwargs
assert success_kwargs["response_body"] == {
"id": "file-123",
"status": "processed",
}
finally:
cleanup()
await fake_client.aclose()
@pytest.mark.asyncio
async def test_pass_through_request_upstream_error_body_stays_buffered():
"""
Upstream errors are never relayed as a stream, whatever their content-type:
the body must stay available for the failure hook and reach the client
buffered with the upstream status code, exactly as before the fix.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_stream = _RecordingUpstreamByteStream((b"upstream ", b"exploded"))
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=502,
headers={"content-type": "application/x-jsonl"},
stream=upstream_stream,
),
timeout=313.0,
)
try:
with ExitStack() as stack:
mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks(
stack, {}
)
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=313.0,
)
assert not isinstance(response, StreamingResponse)
assert response.status_code == 502
assert response.body == b"upstream exploded"
mock_proxy_logging.post_call_failure_hook.assert_called_once()
mock_success_handler.assert_not_called()
finally:
cleanup()
await fake_client.aclose()
_PARTIAL_RELAY_WARNING_MARKER = "ended before upstream body was fully relayed"
@pytest.mark.asyncio
async def test_pass_through_relay_client_disconnect_logs_partial_relay_warning(caplog):
"""
Regression: when the client disconnects mid-relay (GeneratorExit), the
proxy log must record that the upstream body was only partially delivered,
including the route and the byte count that reached the client, while the
success handler still fires so the partial delivery produces a spend-log
row. Pre-fix, the finally block fired the success handler silently and a
partial delivery was indistinguishable from a complete one.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (b'{"custom_id": "a"}\n', b'{"custom_id": "b"}\n')
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={"content-type": "application/x-jsonl"},
stream=upstream_stream,
),
timeout=314.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=314.0,
)
assert isinstance(response, StreamingResponse)
iterator = response.body_iterator
first_chunk = await iterator.__anext__()
assert first_chunk == upstream_chunks[0]
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await iterator.aclose()
partial_relay_warnings = [
record.getMessage()
for record in caplog.records
if record.levelno == logging.WARNING
and _PARTIAL_RELAY_WARNING_MARKER in record.getMessage()
]
assert len(partial_relay_warnings) == 1
assert (
"http://upstream.test/v1/messages/batches/b1/results"
in partial_relay_warnings[0]
)
assert (
f"{len(first_chunk)} bytes were sent to the client"
in partial_relay_warnings[0]
)
assert upstream_stream.closed is True
mock_success_handler.assert_called_once()
assert mock_success_handler.call_args.kwargs["response_body"] is None
finally:
cleanup()
await fake_client.aclose()
@pytest.mark.asyncio
async def test_pass_through_relay_full_consumption_logs_no_partial_relay_warning(caplog):
"""
A fully consumed relay must not be reported as a partial delivery: the
success handler fires and no partial-relay warning is logged.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (b'{"custom_id": "a"}\n', b'{"custom_id": "b"}\n')
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={"content-type": "application/x-jsonl"},
stream=upstream_stream,
),
timeout=315.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=315.0,
)
assert isinstance(response, StreamingResponse)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
relayed = [chunk async for chunk in response.body_iterator]
assert b"".join(relayed) == b"".join(upstream_chunks)
assert not any(
_PARTIAL_RELAY_WARNING_MARKER in record.getMessage()
for record in caplog.records
)
mock_success_handler.assert_called_once()
finally:
cleanup()
await fake_client.aclose()