mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
734fd29e00
commit
b2e2a38bc0
6 changed files with 582 additions and 68 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue