From c58a0bb124698a900934ab3e5aa6c0b2c0cb6b22 Mon Sep 17 00:00:00 2001 From: Adam Holmberg Date: Sat, 24 May 2025 14:04:46 -0500 Subject: [PATCH] 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 --- .../proxy/anthropic_endpoints/endpoints.py | 11 +- litellm/proxy/common_request_processing.py | 89 +++++++- litellm/proxy/proxy_server.py | 13 +- .../proxy/test_common_request_processing.py | 212 ++++++++++++++++++ 4 files changed, 313 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index a84c1e84ab0..024a56fc5b6 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -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)) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 6226c2fa540..7ba3efb4811 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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, ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5ea29b2ac1d..df40577e1a0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 ### diff --git a/tests/litellm/proxy/test_common_request_processing.py b/tests/litellm/proxy/test_common_request_processing.py index 8e795f8b3b9..299bdda917b 100644 --- a/tests/litellm/proxy/test_common_request_processing.py +++ b/tests/litellm/proxy/test_common_request_processing.py @@ -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", + ]