fix: detect and return status codes in streaming responses (#10962)

For streaming requests, the remote request is triggered and
first chunk inspected for an error code as emitted by
async_data_generator.

Potential fix for #9035
This commit is contained in:
Adam Holmberg 2025-05-24 14:04:46 -05:00 • committed by GitHub
parent d14af20bbd
commit c58a0bb124
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 313 additions and 12 deletions

View file

@ -8,7 +8,6 @@ import time
import traceback
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from fastapi.responses import StreamingResponse
import litellm
from litellm._logging import verbose_proxy_logger
@ -16,7 +15,10 @@ from litellm.constants import STREAM_SSE_DATA_PREFIX
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
create_streaming_response,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.utils import ProxyLogging
@ -248,9 +250,10 @@ async def anthropic_response( # noqa: PLR0915
proxy_logging_obj=proxy_logging_obj,
)
return StreamingResponse(
selected_data_generator, # type: ignore
return await create_streaming_response(
generator=selected_data_generator,
media_type="text/event-stream",
headers=dict(fastapi_response.headers),
)
verbose_proxy_logger.info("\nResponse from Litellm:\n{}".format(response))

View file

@ -2,9 +2,10 @@ import asyncio
import json
import uuid
from datetime import datetime
from typing import TYPE_CHECKING, Any, Callable, Literal, Optional, Tuple, Union
from typing import TYPE_CHECKING, Any, Callable, Literal, Optional, Tuple, Union, AsyncGenerator
import httpx
import orjson
from fastapi import HTTPException, Request, status
from fastapi.responses import Response, StreamingResponse
@ -30,6 +31,88 @@ else:
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
async def _parse_event_data_for_error(event_line: str) -> Optional[int]:
"""Parses an event line and returns an error code if present, else None."""
if event_line.startswith("data: "):
json_str = event_line[len("data: "):].strip()
if not json_str or json_str == "[DONE]": # handle empty data or [DONE] message
return None
try:
data = orjson.loads(json_str)
if isinstance(data, dict) and "error" in data and isinstance(data["error"], dict):
error_code_raw = data["error"].get("code")
error_code: Optional[int] = None
if isinstance(error_code_raw, int):
error_code = error_code_raw
elif isinstance(error_code_raw, str):
try:
error_code = int(error_code_raw)
except ValueError:
verbose_proxy_logger.warning(f"Error code is a string but not a valid integer: {error_code_raw}")
# Not a valid integer string, treat as if no valid code was found for this check
pass
# Ensure error_code is a valid HTTP status code
if error_code is not None and 100 <= error_code <= 599:
return error_code
elif error_code_raw is not None : # Log if original code was present but not valid
verbose_proxy_logger.warning(f"Error has invalid or non-convertible code: {error_code_raw}")
except (orjson.JSONDecodeError, json.JSONDecodeError):
# not a known error chunk
pass
return None
async def create_streaming_response(
generator: AsyncGenerator[str, None],
media_type: str,
headers: dict,
default_status_code: int = status.HTTP_200_OK,
) -> StreamingResponse:
"""
Creates a StreamingResponse by inspecting the first chunk for an error code.
The entire original generator content is streamed, but the HTTP status code
of the response is set based on the first chunk if it's a recognized error.
"""
first_chunk_value: Optional[str] = None
final_status_code = default_status_code
try:
first_chunk_value = await generator.__anext__()
if first_chunk_value is not None:
error_code_from_chunk = await _parse_event_data_for_error(first_chunk_value)
if error_code_from_chunk is not None:
final_status_code = error_code_from_chunk
verbose_proxy_logger.debug(f"Error detected in first stream chunk. Status code set to: {final_status_code}")
except StopAsyncIteration:
# Generator was empty. Default status
async def empty_gen() -> AsyncGenerator[str, None]:
if False: yield # type: ignore
return StreamingResponse(empty_gen(), media_type=media_type, headers=headers, status_code=default_status_code)
except Exception as e:
# Unexpected error consuming first chunk.
verbose_proxy_logger.error(f"Error consuming first chunk from generator: {e}")
# Fallback to a generic error stream
async def error_gen_message() -> AsyncGenerator[str, None]:
yield f"data: {json.dumps({'error': {'message': 'Error processing stream start', 'code': status.HTTP_500_INTERNAL_SERVER_ERROR}})}\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(error_gen_message(), media_type=media_type, headers=headers, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR)
async def combined_generator() -> AsyncGenerator[str, None]:
if first_chunk_value is not None:
yield first_chunk_value
async for chunk in generator:
yield chunk
return StreamingResponse(
combined_generator(),
media_type=media_type,
headers=headers,
status_code=final_status_code,
)
class ProxyBaseLLMRequestProcessing:
def __init__(self, data: dict):
self.data = data
@ -308,8 +391,8 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict=user_api_key_dict,
request_data=self.data,
)
return StreamingResponse(
selected_data_generator,
return await create_streaming_response(
generator=selected_data_generator,
media_type="text/event-stream",
headers=custom_headers,
)

View file

@ -173,7 +173,7 @@ from litellm.proxy.batches_endpoints.endpoints import router as batches_router
## Import All Misc routes here ##
from litellm.proxy.caching_routes import router as caching_router
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing, create_streaming_response
from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_proxy
from litellm.proxy.common_utils.debug_utils import init_verbose_loggers
from litellm.proxy.common_utils.debug_utils import router as debugging_endpoints_router
@ -3581,6 +3581,7 @@ async def chat_completion( # noqa: PLR0915
return StreamingResponse(
selected_data_generator,
media_type="text/event-stream",
status_code=e.status_code if hasattr(e, "status_code") else status.HTTP_400_BAD_REQUEST,
)
_usage = litellm.Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
_chat_response.usage = _usage # type: ignore
@ -3723,8 +3724,8 @@ async def completion( # noqa: PLR0915
request_data=data,
)
return StreamingResponse(
selected_data_generator,
return await create_streaming_response(
generator=selected_data_generator,
media_type="text/event-stream",
headers=custom_headers,
)
@ -3782,6 +3783,7 @@ async def completion( # noqa: PLR0915
selected_data_generator,
media_type="text/event-stream",
headers={},
status_code=e.status_code if hasattr(e, "status_code") else status.HTTP_400_BAD_REQUEST,
)
else:
_response = litellm.TextCompletionResponse()
@ -5097,13 +5099,14 @@ async def run_thread(
if (
"stream" in data and data["stream"] is True
): # use generate_responses to stream responses
return StreamingResponse(
async_assistants_data_generator(
return await create_streaming_response(
generator=async_assistants_data_generator(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=data,
),
media_type="text/event-stream",
headers={}, # Added empty headers dict, original call missed this argument
)
### ALERTING ###

View file

@ -6,9 +6,13 @@ from unittest.mock import AsyncMock, MagicMock
from fastapi import Request
from litellm.integrations.opentelemetry import UserAPIKeyAuth
from fastapi import status
from fastapi.responses import StreamingResponse
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
ProxyConfig,
_parse_event_data_for_error,
create_streaming_response,
)
from litellm.proxy.utils import ProxyLogging
@ -70,3 +74,211 @@ class TestProxyBaseLLMRequestProcessing:
except ValueError:
pytest.fail("litellm_call_id is not a valid UUID")
assert data_passed["litellm_call_id"] == returned_data["litellm_call_id"]
@pytest.mark.asyncio
class TestCommonRequestProcessingHelpers:
async def consume_stream(self, streaming_response: StreamingResponse) -> list:
content = []
async for chunk_bytes in streaming_response.body_iterator:
content.append(chunk_bytes)
return content
@pytest.mark.parametrize(
"event_line, expected_code",
[
(
'data: {"error": {"code": 400, "message": "bad request"}}',
400,
), # Valid integer code
(
'data: {"error": {"code": "401", "message": "unauthorized"}}',
401,
), # Valid string-integer code
(
'data: {"error": {"code": "invalid_code", "message": "error"}}',
None,
), # Invalid string code
(
'data: {"error": {"code": 99, "message": "too low"}}',
None,
), # Integer code too low
(
'data: {"error": {"code": 600, "message": "too high"}}',
None,
), # Integer code too high
(
'data: {"id": "123", "content": "hello"}',
None,
), # Non-error SSE event
("data: [DONE]", None), # SSE [DONE] event
("data: ", None), # SSE empty data event
(
'data: {"error": {"code": 400',
None,
), # Malformed JSON
("id: 123", None), # Non-SSE event line
(
'data: {"error": {"message": "some error"}}',
None,
), # Error event without 'code' field
(
'data: {"error": {"code": null, "message": "code is null"}}',
None,
), # Error with null code
],
)
async def test_parse_event_data_for_error(self, event_line, expected_code):
assert await _parse_event_data_for_error(event_line) == expected_code
async def test_create_streaming_response_first_chunk_is_error(self):
async def mock_generator():
yield 'data: {"error": {"code": 403, "message": "forbidden"}}\n\n'
yield 'data: {"content": "more data"}\n\n'
yield "data: [DONE]\n\n"
response = await create_streaming_response(
mock_generator(), "text/event-stream", {}
)
assert response.status_code == status.HTTP_403_FORBIDDEN
content = await self.consume_stream(response)
assert content == [
'data: {"error": {"code": 403, "message": "forbidden"}}\n\n',
'data: {"content": "more data"}\n\n',
"data: [DONE]\n\n",
]
async def test_create_streaming_response_first_chunk_not_error(self):
async def mock_generator():
yield 'data: {"content": "first part"}\n\n'
yield 'data: {"content": "second part"}\n\n'
yield "data: [DONE]\n\n"
response = await create_streaming_response(
mock_generator(), "text/event-stream", {}
)
assert response.status_code == status.HTTP_200_OK
content = await self.consume_stream(response)
assert content == [
'data: {"content": "first part"}\n\n',
'data: {"content": "second part"}\n\n',
"data: [DONE]\n\n",
]
async def test_create_streaming_response_empty_generator(self):
async def mock_generator():
if False: # Never yields
yield
# Implicitly raises StopAsyncIteration
response = await create_streaming_response(
mock_generator(), "text/event-stream", {}
)
assert response.status_code == status.HTTP_200_OK
content = await self.consume_stream(response)
assert content == []
async def test_create_streaming_response_generator_raises_stop_async_iteration_immediately(
self,
):
mock_gen = AsyncMock()
mock_gen.__anext__.side_effect = StopAsyncIteration
response = await create_streaming_response(mock_gen, "text/event-stream", {})
assert response.status_code == status.HTTP_200_OK
content = await self.consume_stream(response)
assert content == []
async def test_create_streaming_response_generator_raises_unexpected_exception(
self,
):
mock_gen = AsyncMock()
mock_gen.__anext__.side_effect = ValueError("Test error from generator")
response = await create_streaming_response(mock_gen, "text/event-stream", {})
assert response.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR
content = await self.consume_stream(response)
expected_error_data = {
"error": {
"message": "Error processing stream start",
"code": status.HTTP_500_INTERNAL_SERVER_ERROR,
}
}
assert len(content) == 2
# Use json.dumps to match the formatting in create_streaming_response's exception handler
import json
assert content[0] == f"data: {json.dumps(expected_error_data)}\n\n"
assert content[1] == "data: [DONE]\n\n"
async def test_create_streaming_response_first_chunk_error_string_code(self):
async def mock_generator():
yield 'data: {"error": {"code": "429", "message": "too many requests"}}\n\n'
yield "data: [DONE]\n\n"
response = await create_streaming_response(
mock_generator(), "text/event-stream", {}
)
assert response.status_code == status.HTTP_429_TOO_MANY_REQUESTS
content = await self.consume_stream(response)
assert content == [
'data: {"error": {"code": "429", "message": "too many requests"}}\n\n',
"data: [DONE]\n\n",
]
async def test_create_streaming_response_custom_headers(self):
async def mock_generator():
yield 'data: {"content": "data"}\n\n'
yield "data: [DONE]\n\n"
custom_headers = {"X-Custom-Header": "TestValue"}
response = await create_streaming_response(
mock_generator(), "text/event-stream", custom_headers
)
assert response.headers["x-custom-header"] == "TestValue"
async def test_create_streaming_response_non_default_status_code(self):
async def mock_generator():
yield 'data: {"content": "data"}\n\n'
yield "data: [DONE]\n\n"
response = await create_streaming_response(
mock_generator(),
"text/event-stream",
{},
default_status_code=status.HTTP_201_CREATED,
)
assert response.status_code == status.HTTP_201_CREATED
content = await self.consume_stream(response)
assert content == [
'data: {"content": "data"}\n\n',
"data: [DONE]\n\n",
]
async def test_create_streaming_response_first_chunk_is_done(self):
async def mock_generator():
yield "data: [DONE]\n\n"
response = await create_streaming_response(
mock_generator(), "text/event-stream", {}
)
assert response.status_code == status.HTTP_200_OK # Default status
content = await self.consume_stream(response)
assert content == ["data: [DONE]\n\n"]
async def test_create_streaming_response_first_chunk_is_empty_data(self):
async def mock_generator():
yield "data: \n\n"
yield 'data: {"content": "actual data"}\n\n'
yield "data: [DONE]\n\n"
response = await create_streaming_response(
mock_generator(), "text/event-stream", {}
)
assert response.status_code == status.HTTP_200_OK # Default status
content = await self.consume_stream(response)
assert content == [
"data: \n\n",
'data: {"content": "actual data"}\n\n',
"data: [DONE]\n\n",
]