mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(passthrough): raise_for_status in _async_streaming to propagate Azure 429s
This commit is contained in:
parent
6b7d767637
commit
f8afa04e5a
4 changed files with 239 additions and 5 deletions
|
|
@ -289,7 +289,7 @@ def llm_passthrough_route(
|
|||
request = client.client.build_request(
|
||||
method=method,
|
||||
url=updated_url,
|
||||
content=signed_json_body,
|
||||
content=signed_json_body if signed_json_body is not None else content,
|
||||
data=data if signed_json_body is None else None,
|
||||
files=files,
|
||||
json=json if signed_json_body is None else None,
|
||||
|
|
@ -412,6 +412,7 @@ async def _async_streaming(
|
|||
):
|
||||
try:
|
||||
iter_response = await response
|
||||
iter_response.raise_for_status()
|
||||
raw_bytes: List[bytes] = []
|
||||
|
||||
async for chunk in iter_response.aiter_bytes(): # type: ignore
|
||||
|
|
|
|||
0
tests/test_litellm/passthrough/__init__.py
Normal file
0
tests/test_litellm/passthrough/__init__.py
Normal file
|
|
@ -0,0 +1,122 @@
|
|||
"""
|
||||
Tests for error propagation in _async_streaming passthrough routes.
|
||||
|
||||
Verifies that HTTP 4xx/5xx errors from upstream (e.g. Azure 429 rate limits)
|
||||
raise exceptions instead of being silently forwarded as raw bytes under HTTP 200.
|
||||
|
||||
See: litellm/passthrough/main.py _async_streaming()
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_mock_response(status_code: int, body: bytes, headers: dict = None): # type: ignore[assignment]
|
||||
mock = MagicMock(spec=httpx.Response)
|
||||
mock.status_code = status_code
|
||||
mock.headers = httpx.Headers(headers or {"content-type": "text/event-stream"})
|
||||
|
||||
def _raise_for_status():
|
||||
if status_code >= 400:
|
||||
request = httpx.Request("POST", "https://azure.example.com/openai/responses")
|
||||
real_response = httpx.Response(
|
||||
status_code=status_code,
|
||||
content=body,
|
||||
request=request,
|
||||
headers=headers or {},
|
||||
)
|
||||
raise httpx.HTTPStatusError(
|
||||
message=f"{status_code} Error",
|
||||
request=request,
|
||||
response=real_response,
|
||||
)
|
||||
|
||||
mock.raise_for_status = _raise_for_status
|
||||
|
||||
async def _aiter_bytes():
|
||||
yield body
|
||||
|
||||
mock.aiter_bytes = _aiter_bytes
|
||||
return mock
|
||||
|
||||
|
||||
def _make_mock_logging_obj():
|
||||
mock = MagicMock()
|
||||
mock.async_flush_passthrough_collected_chunks = AsyncMock()
|
||||
return mock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_429_raises():
|
||||
"""429 from upstream should raise HTTPStatusError, not yield error bytes."""
|
||||
from litellm.passthrough.main import _async_streaming
|
||||
|
||||
error_body = json.dumps(
|
||||
{"error": {"code": "429", "message": "Rate limit exceeded."}}
|
||||
).encode()
|
||||
mock_response = _make_mock_response(429, error_body)
|
||||
|
||||
async def response_coro():
|
||||
return mock_response
|
||||
|
||||
chunks = []
|
||||
with pytest.raises(httpx.HTTPStatusError) as exc_info:
|
||||
async for chunk in _async_streaming(
|
||||
response=response_coro(),
|
||||
litellm_logging_obj=_make_mock_logging_obj(),
|
||||
provider_config=MagicMock(),
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert exc_info.value.response.status_code == 429
|
||||
assert len(chunks) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_500_raises():
|
||||
"""500 from upstream should also raise, not yield error bytes."""
|
||||
from litellm.passthrough.main import _async_streaming
|
||||
|
||||
error_body = json.dumps(
|
||||
{"error": {"code": "500", "message": "Internal server error"}}
|
||||
).encode()
|
||||
mock_response = _make_mock_response(500, error_body)
|
||||
|
||||
async def response_coro():
|
||||
return mock_response
|
||||
|
||||
with pytest.raises(httpx.HTTPStatusError) as exc_info:
|
||||
async for _ in _async_streaming(
|
||||
response=response_coro(),
|
||||
litellm_logging_obj=_make_mock_logging_obj(),
|
||||
provider_config=MagicMock(),
|
||||
):
|
||||
pass
|
||||
|
||||
assert exc_info.value.response.status_code == 500
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_streaming_200_yields_chunks():
|
||||
"""Successful 200 streaming responses should continue to work normally."""
|
||||
from litellm.passthrough.main import _async_streaming
|
||||
|
||||
sse_data = b'data: {"type":"response.created"}\n\ndata: [DONE]\n\n'
|
||||
mock_response = _make_mock_response(200, sse_data)
|
||||
|
||||
async def response_coro():
|
||||
return mock_response
|
||||
|
||||
chunks = []
|
||||
async for chunk in _async_streaming(
|
||||
response=response_coro(),
|
||||
litellm_logging_obj=_make_mock_logging_obj(),
|
||||
provider_config=MagicMock(),
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert len(chunks) == 1
|
||||
assert b"response.created" in chunks[0]
|
||||
|
|
@ -1,13 +1,13 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../..")
|
||||
|
|
@ -17,7 +17,7 @@ sys.path.insert(
|
|||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import litellm
|
||||
from litellm.passthrough.main import llm_passthrough_route
|
||||
from litellm.passthrough.main import allm_passthrough_route, llm_passthrough_route
|
||||
|
||||
|
||||
def test_llm_passthrough_route():
|
||||
|
|
@ -507,4 +507,115 @@ def test_azure_with_custom_api_base_and_key():
|
|||
json_body = call_args.kwargs["json"]
|
||||
assert json_body["model"] == "gpt-4.1"
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.status_code == 200 # type: ignore[union-attr]
|
||||
|
||||
|
||||
def _make_429_streaming_response() -> MagicMock:
|
||||
"""Build a mock httpx.Response that looks like a streaming 429 from Azure."""
|
||||
error_body = json.dumps(
|
||||
{"error": {"code": "429", "message": "Rate limit exceeded. Retry after 10 seconds."}}
|
||||
).encode()
|
||||
|
||||
mock = MagicMock(spec=httpx.Response)
|
||||
mock.status_code = 429
|
||||
mock.headers = httpx.Headers({"content-type": "application/json"})
|
||||
|
||||
def _raise_for_status():
|
||||
request = httpx.Request(
|
||||
"POST",
|
||||
"https://my-azure.openai.azure.com/openai/deployments/gpt-4/responses",
|
||||
)
|
||||
raise httpx.HTTPStatusError(
|
||||
message="429 Too Many Requests",
|
||||
request=request,
|
||||
response=httpx.Response(
|
||||
status_code=429,
|
||||
content=error_body,
|
||||
request=request,
|
||||
),
|
||||
)
|
||||
|
||||
mock.raise_for_status = _raise_for_status
|
||||
|
||||
async def _aiter_bytes():
|
||||
yield error_body
|
||||
|
||||
mock.aiter_bytes = _aiter_bytes
|
||||
return mock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_allm_passthrough_route_429_streaming_raises():
|
||||
"""
|
||||
Regression test: Azure 429 during streaming must raise HTTPStatusError,
|
||||
not be silently forwarded as raw bytes under HTTP 200.
|
||||
|
||||
Before the fix, _async_streaming() would yield the 429 error JSON as
|
||||
chunks and allm_passthrough_route returned an async generator. The
|
||||
caller (azure_proxy_route) wrapped it in StreamingResponse(status_code=200),
|
||||
so the client saw HTTP 200 + unparseable SSE body → silent task_complete(null).
|
||||
|
||||
After the fix, raise_for_status() fires inside _async_streaming() before
|
||||
any chunks are yielded, so the exception propagates all the way up.
|
||||
"""
|
||||
mock_provider_config = MagicMock()
|
||||
mock_provider_config.get_complete_url.return_value = (
|
||||
httpx.URL(
|
||||
"https://my-azure.openai.azure.com/openai/deployments/gpt-4/responses"
|
||||
),
|
||||
"https://my-azure.openai.azure.com",
|
||||
)
|
||||
mock_provider_config.get_api_key.return_value = "fake-azure-key"
|
||||
mock_provider_config.validate_environment.return_value = {"api-key": "fake-azure-key"}
|
||||
mock_provider_config.sign_request.return_value = ({"api-key": "fake-azure-key"}, None)
|
||||
mock_provider_config.is_streaming_request.return_value = True
|
||||
|
||||
mock_429_response = _make_429_streaming_response()
|
||||
|
||||
async_client = AsyncHTTPHandler()
|
||||
mock_send = AsyncMock(return_value=mock_429_response)
|
||||
mock_build_request = MagicMock(return_value=MagicMock())
|
||||
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.update_environment_variables = MagicMock()
|
||||
mock_logging_obj.async_flush_passthrough_collected_chunks = AsyncMock()
|
||||
|
||||
with patch(
|
||||
"litellm.utils.ProviderConfigManager.get_provider_passthrough_config",
|
||||
return_value=mock_provider_config,
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.get_litellm_params.get_litellm_params",
|
||||
return_value={},
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
|
||||
return_value=(
|
||||
"gpt-4",
|
||||
"azure",
|
||||
"fake-azure-key",
|
||||
"https://my-azure.openai.azure.com",
|
||||
),
|
||||
), patch.object(
|
||||
async_client.client, "send", mock_send
|
||||
), patch.object(
|
||||
async_client.client, "build_request", mock_build_request
|
||||
):
|
||||
result = await allm_passthrough_route(
|
||||
model="azure/gpt-4",
|
||||
endpoint="openai/deployments/gpt-4/responses",
|
||||
method="POST",
|
||||
custom_llm_provider="azure",
|
||||
api_base="https://my-azure.openai.azure.com",
|
||||
api_key="fake-azure-key",
|
||||
json={"model": "gpt-4", "input": "hello", "stream": True},
|
||||
client=async_client,
|
||||
litellm_logging_obj=mock_logging_obj,
|
||||
)
|
||||
|
||||
# result is an async generator — consuming it must raise, not silently yield error bytes
|
||||
chunks = []
|
||||
with pytest.raises(httpx.HTTPStatusError) as exc_info:
|
||||
async for chunk in result: # type: ignore[union-attr]
|
||||
chunks.append(chunk)
|
||||
|
||||
assert exc_info.value.response.status_code == 429
|
||||
assert len(chunks) == 0, "No chunks should be yielded before the 429 raises"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue