mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
d14af20bbd
commit
c58a0bb124
4 changed files with 313 additions and 12 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue