mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #32133 from BerriAI/litellm_passthrough_error_normalisation
fix(proxy): return upstream error bodies unchanged in passthrough
This commit is contained in:
commit
a3a3201e12
6 changed files with 494 additions and 32 deletions
|
|
@ -676,6 +676,75 @@ def _carry_guardrail_logging_info(request_data: dict, guardrail_data: Optional[d
|
|||
metadata.setdefault("standard_logging_guardrail_information", list(entries))
|
||||
|
||||
|
||||
def _build_passthrough_failure_request_payload(
|
||||
parsed_body: Optional[dict],
|
||||
kwargs: Optional[dict],
|
||||
logging_obj: Optional[LiteLLMLoggingObj],
|
||||
custom_llm_provider: Optional[str],
|
||||
) -> dict:
|
||||
"""Build the ``request_data`` dict passed to ``post_call_failure_hook``.
|
||||
|
||||
Shared by the outer exception handler (LiteLLM-internal failures) and
|
||||
upstream HTTP error logging, so both failure paths report the same shape
|
||||
of request data (model, custom_llm_provider, litellm_logging_obj, ...).
|
||||
"""
|
||||
request_payload: dict = dict(parsed_body or {})
|
||||
if kwargs:
|
||||
request_payload.update(kwargs)
|
||||
if logging_obj is not None:
|
||||
request_payload["litellm_logging_obj"] = logging_obj
|
||||
if "model" not in request_payload and parsed_body and isinstance(parsed_body, dict):
|
||||
request_payload["model"] = parsed_body.get("model", "")
|
||||
if "custom_llm_provider" not in request_payload and custom_llm_provider:
|
||||
request_payload["custom_llm_provider"] = custom_llm_provider
|
||||
return request_payload
|
||||
|
||||
|
||||
async def _log_passthrough_upstream_failure(
|
||||
response: httpx.Response,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_payload: dict,
|
||||
) -> None:
|
||||
"""Fire LiteLLM-side failure hooks (spend tracking, alerting callbacks) for
|
||||
an upstream 4xx/5xx passthrough response.
|
||||
|
||||
Passthrough must return the upstream status/body/headers to the client
|
||||
unchanged, so this never raises or transforms the response - it only
|
||||
mirrors the monitoring side effect that ``post_call_failure_hook`` would
|
||||
have received had the error originated inside LiteLLM.
|
||||
"""
|
||||
if response.status_code < 400:
|
||||
return
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError:
|
||||
# Reported as an HTTPException, not the raw httpx error: ProxyLogging's
|
||||
# alerting path only excludes HTTPException/ProxyException from its
|
||||
# "High" severity llm_exceptions alert, treating everything else as an
|
||||
# operational LLM-API failure. An upstream 4xx/5xx returned unchanged
|
||||
# to the client is a user-facing error like any other, not something
|
||||
# ops needs paged for, so it must be excluded the same way auth and
|
||||
# rate-limit errors already are.
|
||||
synthetic_exception = HTTPException(
|
||||
status_code=response.status_code,
|
||||
detail=f"Upstream passthrough request failed with status {response.status_code}",
|
||||
)
|
||||
try:
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=synthetic_exception,
|
||||
request_data=request_payload,
|
||||
traceback_str=traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG),
|
||||
)
|
||||
except Exception: # noqa: BLE001 - a failing logging callback must never break the passthrough response
|
||||
verbose_proxy_logger.warning(
|
||||
"pass_through_endpoint: post_call_failure_hook raised for upstream error",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
from litellm.passthrough.timeout_utils import (
|
||||
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, # noqa: F401 - re-exported for backward compat
|
||||
resolve_llm_passthrough_timeout, # noqa: F401 - re-exported for backward compat
|
||||
|
|
@ -1053,10 +1122,16 @@ async def pass_through_request(
|
|||
|
||||
response = await async_client.send(req, stream=stream)
|
||||
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise HTTPException(status_code=e.response.status_code, detail=await e.response.aread())
|
||||
await _log_passthrough_upstream_failure(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_payload=_build_passthrough_failure_request_payload(
|
||||
parsed_body=_parsed_body,
|
||||
kwargs=kwargs,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
|
||||
# Call response headers hook for streaming pass-through
|
||||
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
|
||||
|
|
@ -1112,10 +1187,16 @@ async def pass_through_request(
|
|||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["stream"] = True
|
||||
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise HTTPException(status_code=e.response.status_code, detail=await e.response.aread())
|
||||
await _log_passthrough_upstream_failure(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_payload=_build_passthrough_failure_request_payload(
|
||||
parsed_body=_parsed_body,
|
||||
kwargs=kwargs,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
|
||||
# Call response headers hook for detected streaming pass-through
|
||||
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
|
||||
|
|
@ -1145,20 +1226,29 @@ async def pass_through_request(
|
|||
status_code=response.status_code,
|
||||
)
|
||||
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise HTTPException(status_code=e.response.status_code, detail=e.response.text)
|
||||
|
||||
if response.status_code >= 300:
|
||||
raise HTTPException(status_code=response.status_code, detail=response.text)
|
||||
|
||||
content = await response.aread()
|
||||
|
||||
## POST-CALL GUARDRAILS ##
|
||||
# Guardrails and managed-id rewriting only apply to successful upstream
|
||||
# responses; response_body itself is parsed unconditionally so the
|
||||
# failure-hook log payload below still reflects upstream error bodies.
|
||||
_content_modified = False
|
||||
response_body: Optional[dict] = get_response_body(response)
|
||||
if response_body is not None and guardrails_to_run:
|
||||
|
||||
failure_request_payload = _build_passthrough_failure_request_payload(
|
||||
parsed_body=_parsed_body,
|
||||
kwargs=kwargs,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
failure_request_payload["response_body"] = response_body
|
||||
await _log_passthrough_upstream_failure(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_payload=failure_request_payload,
|
||||
)
|
||||
|
||||
if response.status_code < 400 and response_body is not None and guardrails_to_run:
|
||||
# Build an enriched data dict: _parsed_body has been stripped of
|
||||
# `metadata` by both pre_call_hook and _init_kwargs_for_pass_through_endpoint,
|
||||
# so we re-attach the configured guardrails here so should_run_guardrail
|
||||
|
|
@ -1249,23 +1339,28 @@ async def pass_through_request(
|
|||
)
|
||||
|
||||
## LOG SUCCESS
|
||||
# Upstream errors are already logged via _log_passthrough_upstream_failure
|
||||
# above; the success handler has no status-code awareness of its own; so
|
||||
# calling it here for a 4xx/5xx would double-log the same request as both
|
||||
# a failure and a success (corrupting spend tracking).
|
||||
passthrough_logging_payload["response_body"] = response_body
|
||||
end_time = datetime.now()
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
|
||||
httpx_response=response,
|
||||
response_body=response_body,
|
||||
url_route=str(url),
|
||||
result="",
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
logging_obj=logging_obj,
|
||||
cache_hit=False,
|
||||
request_body=_parsed_body or {},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
if response.status_code < 400:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
|
||||
httpx_response=response,
|
||||
response_body=response_body,
|
||||
url_route=str(url),
|
||||
result="",
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
logging_obj=logging_obj,
|
||||
cache_hit=False,
|
||||
request_body=_parsed_body or {},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
## CUSTOM HEADERS - `x-litellm-*`
|
||||
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
|
|
|
|||
|
|
@ -89,7 +89,11 @@ class PassThroughStreamingHandler:
|
|||
# GeneratorExit (raised on client disconnect) is not caught by
|
||||
# `except Exception`; the finally block ensures partial usage
|
||||
# still gets logged for spend tracking. See LIT-2642.
|
||||
if not logging_scheduled and raw_bytes:
|
||||
# Upstream 4xx/5xx responses are already logged as a failure by
|
||||
# the caller before this generator starts (see
|
||||
# _log_passthrough_upstream_failure); logging them again here as
|
||||
# a success would double-log the same request.
|
||||
if not logging_scheduled and raw_bytes and response.status_code < 400:
|
||||
logging_scheduled = True
|
||||
try:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ async def test_chunk_processor_yields_raw_bytes(endpoint_type, url_route):
|
|||
"""
|
||||
# Mock inputs
|
||||
response = AsyncMock(spec=httpx.Response)
|
||||
response.status_code = 200
|
||||
raw_chunks = [
|
||||
b'{"id": "1", "content": "Hello"}',
|
||||
b'{"id": "2", "content": "World"}',
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ async def test_vertex_ai_anthropic_streaming_cost_injection_enabled():
|
|||
try:
|
||||
# Mock response with Anthropic SSE format chunks
|
||||
response = AsyncMock(spec=httpx.Response)
|
||||
response.status_code = 200
|
||||
|
||||
# Create chunks with message_delta event containing usage
|
||||
chunks_with_usage = [
|
||||
|
|
@ -120,6 +121,7 @@ async def test_vertex_ai_anthropic_streaming_cost_injection_disabled():
|
|||
try:
|
||||
# Mock response with Anthropic SSE format chunks
|
||||
response = AsyncMock(spec=httpx.Response)
|
||||
response.status_code = 200
|
||||
|
||||
chunks_with_usage = [
|
||||
b'data: {"type": "message_delta", "usage": {"input_tokens": 10, "output_tokens": 5}}\n\n',
|
||||
|
|
@ -178,6 +180,7 @@ async def test_vertex_ai_anthropic_streaming_cost_injection_no_usage_chunk():
|
|||
|
||||
try:
|
||||
response = AsyncMock(spec=httpx.Response)
|
||||
response.status_code = 200
|
||||
|
||||
# Chunks without usage (should not be modified)
|
||||
chunks_without_usage = [
|
||||
|
|
@ -233,6 +236,7 @@ async def test_vertex_ai_anthropic_streaming_model_extraction():
|
|||
|
||||
try:
|
||||
response = AsyncMock(spec=httpx.Response)
|
||||
response.status_code = 200
|
||||
|
||||
chunks = [
|
||||
b'data: {"type": "message_delta", "usage": {"input_tokens": 10, "output_tokens": 5}}\n\n',
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
|
@ -3619,3 +3620,325 @@ async def test_non_guardrail_exception_still_logs_with_traceback():
|
|||
assert (
|
||||
logger.warning.call_count == 0
|
||||
), "a genuine failure must not be downgraded to WARNING"
|
||||
|
||||
|
||||
# Regression: generic config-based passthrough (`pass_through_request`) used to
|
||||
# call `response.raise_for_status()` on upstream errors and re-raise as an
|
||||
# `HTTPException`, which the outer `except` block then reshaped into a
|
||||
# `ProxyException` (`{"error": {"message": "<stringified upstream body>", ...}}`).
|
||||
# Upstream error responses must reach the client byte-for-byte, with the
|
||||
# original status code, exactly like success responses already do.
|
||||
_UPSTREAM_ERROR_BODY = {
|
||||
"error": "Permission denied",
|
||||
"error_code": "ACCESS_DENIED",
|
||||
"request_id": "req_mock_403",
|
||||
"trace_id": "trace_mock_403",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_non_streaming_upstream_error_returned_unchanged():
|
||||
upstream_content = json.dumps(_UPSTREAM_ERROR_BODY).encode("utf-8")
|
||||
upstream_response = httpx.Response(
|
||||
status_code=403,
|
||||
headers={"content-type": "application/json"},
|
||||
content=upstream_content,
|
||||
request=httpx.Request("POST", "http://target-api.com/api/denied"),
|
||||
)
|
||||
|
||||
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.ProxyBaseLLMRequestProcessing"
|
||||
) as mock_processing:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_processing.get_custom_headers.return_value = {}
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.request = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
|
||||
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
response = await pass_through_request(
|
||||
request=mock_request,
|
||||
target="http://target-api.com/api/denied",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert response.status_code == 403
|
||||
body = json.loads(response.body)
|
||||
# Exact dict equality proves the upstream body was forwarded verbatim,
|
||||
# not stringified into a ProxyException's `error.message` field.
|
||||
assert body == _UPSTREAM_ERROR_BODY
|
||||
assert set(body.keys()) != {"error"} or not isinstance(body["error"], dict)
|
||||
|
||||
# Regression: the success handler has no status-code awareness, so it must
|
||||
# never be called for a 4xx/5xx upstream response - otherwise the same
|
||||
# request gets recorded as both a failure and a success in SpendLogs.
|
||||
mock_success_handler.assert_not_called()
|
||||
|
||||
# Regression: post_call_failure_hook (spend-tracking, alerting callbacks)
|
||||
# must still fire for upstream errors even though the client-facing
|
||||
# response is unchanged and no ProxyException is raised.
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
failure_call_kwargs = mock_proxy_logging.post_call_failure_hook.call_args.kwargs
|
||||
# Must be reported as HTTPException, not the raw httpx error: ProxyLogging's
|
||||
# alerting only excludes HTTPException/ProxyException from its "High"
|
||||
# severity llm_exceptions alert, so a raw HTTPStatusError here would page
|
||||
# ops for every routine upstream 4xx returned through passthrough.
|
||||
assert isinstance(failure_call_kwargs["original_exception"], HTTPException)
|
||||
assert failure_call_kwargs["original_exception"].status_code == 403
|
||||
|
||||
# Regression: the failure-hook log payload's response_body must reflect
|
||||
# the upstream error JSON, not None, so downstream spend-tracking/logging
|
||||
# integrations can see what the upstream actually returned.
|
||||
assert failure_call_kwargs["request_data"]["response_body"] == _UPSTREAM_ERROR_BODY
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_upstream_error_failure_hook_exception_is_swallowed():
|
||||
"""
|
||||
A broken failure-hook callback (e.g. a misconfigured alerting integration)
|
||||
must never take down the passthrough response - the upstream error body
|
||||
must still reach the client unchanged, and the callback's exception must
|
||||
only be logged, not raised.
|
||||
"""
|
||||
upstream_content = json.dumps(_UPSTREAM_ERROR_BODY).encode("utf-8")
|
||||
upstream_response = httpx.Response(
|
||||
status_code=403,
|
||||
headers={"content-type": "application/json"},
|
||||
content=upstream_content,
|
||||
request=httpx.Request("POST", "http://target-api.com/api/denied"),
|
||||
)
|
||||
|
||||
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.ProxyBaseLLMRequestProcessing"
|
||||
) as mock_processing:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(
|
||||
side_effect=RuntimeError("alerting integration misconfigured")
|
||||
)
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_processing.get_custom_headers.return_value = {}
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.request = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
|
||||
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
response = await pass_through_request(
|
||||
request=mock_request,
|
||||
target="http://target-api.com/api/denied",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
assert response.status_code == 403
|
||||
assert json.loads(response.body) == _UPSTREAM_ERROR_BODY
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_streaming_upstream_error_returned_unchanged():
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
upstream_content = json.dumps(_UPSTREAM_ERROR_BODY).encode("utf-8")
|
||||
upstream_response = httpx.Response(
|
||||
status_code=403,
|
||||
headers={"content-type": "application/json"},
|
||||
content=upstream_content,
|
||||
request=httpx.Request("GET", "http://target-api.com/api/stream-denied"),
|
||||
)
|
||||
|
||||
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(
|
||||
return_value=None
|
||||
)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client = 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)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "GET"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/stream-denied"
|
||||
mock_request.body = AsyncMock(return_value=b"")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
response = await pass_through_request(
|
||||
request=mock_request,
|
||||
target="http://target-api.com/api/stream-denied",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 403
|
||||
|
||||
streamed_chunks = [chunk async for chunk in response.body_iterator]
|
||||
await asyncio.sleep(0)
|
||||
streamed_bytes = b"".join(
|
||||
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8")
|
||||
for chunk in streamed_chunks
|
||||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
assert json.loads(streamed_bytes) == _UPSTREAM_ERROR_BODY
|
||||
|
||||
# Regression: chunk_processor's end-of-stream success logging has no
|
||||
# status-code awareness, so it must never fire for a 4xx/5xx upstream
|
||||
# response - otherwise the same request gets recorded as both a failure
|
||||
# (via the hook below) and a success in SpendLogs.
|
||||
mock_success_handler.assert_not_called()
|
||||
|
||||
# Regression: post_call_failure_hook must still fire for streaming
|
||||
# upstream errors, mirroring the non-streaming behavior, and must also
|
||||
# report an HTTPException (not the raw httpx error) to avoid triggering
|
||||
# a "High" severity llm_exceptions alert for a routine upstream 4xx.
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
failure_call_kwargs = mock_proxy_logging.post_call_failure_hook.call_args.kwargs
|
||||
assert isinstance(failure_call_kwargs["original_exception"], HTTPException)
|
||||
assert failure_call_kwargs["original_exception"].status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_non_streaming_success_unchanged():
|
||||
"""Success (2xx) passthrough behavior must remain unchanged by the error fix."""
|
||||
upstream_success_body = {"status": "ok", "message": "mock upstream success"}
|
||||
upstream_content = json.dumps(upstream_success_body).encode("utf-8")
|
||||
upstream_response = httpx.Response(
|
||||
status_code=200,
|
||||
headers={"content-type": "application/json"},
|
||||
content=upstream_content,
|
||||
request=httpx.Request("GET", "http://target-api.com/api/success"),
|
||||
)
|
||||
|
||||
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.ProxyBaseLLMRequestProcessing"
|
||||
) as mock_processing:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_processing.get_custom_headers.return_value = {}
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.request = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "GET"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
|
||||
mock_request.body = AsyncMock(return_value=b"")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
response = await pass_through_request(
|
||||
request=mock_request,
|
||||
target="http://target-api.com/api/success",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert json.loads(response.body) == upstream_success_body
|
||||
# Regression guard: the failure hook must only fire for upstream errors,
|
||||
# never for a successful upstream response.
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
# ...and the success handler must still fire exactly once for a 2xx,
|
||||
# proving the status_code gate doesn't also swallow real successes.
|
||||
mock_success_handler.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_internal_failure_still_raises_proxy_exception():
|
||||
"""
|
||||
Internal proxy failures (e.g. a hook raising before any upstream request is
|
||||
made) must still surface as ProxyException, distinct from upstream
|
||||
passthrough errors which are now returned unchanged.
|
||||
"""
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(
|
||||
side_effect=RuntimeError("auth backend unavailable")
|
||||
)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "GET"
|
||||
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
|
||||
mock_request.body = AsyncMock(return_value=b"")
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await pass_through_request(
|
||||
request=mock_request,
|
||||
target="http://target-api.com/api/success",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
assert int(exc_info.value.code) == 500
|
||||
assert "auth backend unavailable" in exc_info.value.message
|
||||
|
|
|
|||
|
|
@ -93,6 +93,41 @@ async def test_chunk_processor_logs_on_client_disconnect():
|
|||
assert call_kwargs["raw_bytes"] == [chunks[0]]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_error():
|
||||
"""A 4xx/5xx upstream response is already logged as a failure by the caller
|
||||
before this generator starts; scheduling success logging here too would
|
||||
double-log the same request in SpendLogs."""
|
||||
chunks = [b'{"error": "denied"}']
|
||||
response = _make_streaming_response(chunks)
|
||||
response.status_code = 403
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_passthrough_handler = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
PassThroughStreamingHandler,
|
||||
"_route_streaming_logging_to_handler",
|
||||
new=AsyncMock(),
|
||||
) as mock_route:
|
||||
received = []
|
||||
async for chunk in PassThroughStreamingHandler.chunk_processor(
|
||||
response=response,
|
||||
request_body={"model": "claude-3-haiku"},
|
||||
litellm_logging_obj=mock_logging_obj,
|
||||
endpoint_type=EndpointType.GENERIC,
|
||||
start_time=datetime.now(),
|
||||
passthrough_success_handler_obj=mock_passthrough_handler,
|
||||
url_route="/bedrock/model/claude/invoke-with-response-stream",
|
||||
):
|
||||
received.append(chunk)
|
||||
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert received == chunks
|
||||
mock_route.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_does_not_schedule_logging_when_no_chunks():
|
||||
response = _make_streaming_response([])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue