diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e7a08711eb1..c0ca45cb0da 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1165,12 +1165,32 @@ async def _aclose_late_response(produced: Response) -> None: verbose_proxy_logger.debug("error closing relayed streaming generator: %s", exc) -async def _relay_late_response(produced: Response) -> AsyncGenerator[bytes, None]: +def _late_response_body(produced: object) -> bytes: + """Bytes for one SSE data frame. + + A late payload may be a mapping, raw bytes or text, or a Response whose ``body`` is one of those. + """ + payload: Final = ( + produced if isinstance(produced, (bytes, bytearray, Mapping)) else getattr(produced, "body", produced) + ) + if isinstance(payload, (bytes, bytearray)): + raw: Final = bytes(payload) + return raw or b"{}" + if isinstance(payload, str): + text: Final = payload.encode() + return text or b"{}" + if isinstance(payload, Mapping): + # orjson rejects mappingproxy and other non-dict mappings. + return orjson.dumps(dict(payload)) + return b"{}" + + +async def _relay_late_response(produced: object) -> AsyncGenerator[bytes, None]: """Replay a Response that was built after a keepalive had already opened the wire.""" if not isinstance(produced, StreamingResponse): # The status line is already on the wire, so a non-streaming body, an error # body included, can only reach the client as an SSE frame. - yield b"data: " + (bytes(produced.body) or b"{}") + b"\n\n" + yield b"data: " + _late_response_body(produced) + b"\n\n" yield b"data: [DONE]\n\n" return diff --git a/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py b/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py index 228fd5bcae6..ca447b00538 100644 --- a/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py +++ b/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py @@ -3,7 +3,7 @@ from collections.abc import AsyncGenerator from typing import Final, cast import pytest -from fastapi.responses import StreamingResponse +from fastapi.responses import Response, StreamingResponse from litellm.proxy.common_request_processing import create_response from litellm.types.utils import ModelResponse @@ -408,3 +408,72 @@ def test_ttft_interval_resolves_through_the_deployments_it_could_land_on( deployments, global_interval, expected, why ): assert resolve_ttft_keepalive_interval(deployments, global_interval) == expected, why + + +@pytest.mark.asyncio +async def test_relay_late_response_serializes_a_dict_payload_as_sse(): + from litellm.proxy.common_request_processing import _relay_late_response + + frames: Final = [chunk async for chunk in _relay_late_response({"error": {"message": "late"}})] + + assert frames == [ + b'data: {"error":{"message":"late"}}\n\n', + b"data: [DONE]\n\n", + ] + + +@pytest.mark.asyncio +async def test_relay_late_response_keeps_a_response_body(): + from litellm.proxy.common_request_processing import _relay_late_response + + frames: Final = [chunk async for chunk in _relay_late_response(Response(content=b'{"ok":true}'))] + + assert frames == [ + b'data: {"ok":true}\n\n', + b"data: [DONE]\n\n", + ] + + +@pytest.mark.asyncio +async def test_relay_late_response_serializes_a_dict_body_attribute(): + from litellm.proxy.common_request_processing import _relay_late_response + + class _DictBody: + body = {"ok": True} + + frames: Final = [chunk async for chunk in _relay_late_response(_DictBody())] + + assert frames[0] == b'data: {"ok":true}\n\n' + assert frames[1] == b"data: [DONE]\n\n" + + +@pytest.mark.asyncio +async def test_relay_late_response_encodes_a_string_payload(): + from litellm.proxy.common_request_processing import _relay_late_response + + frames: Final = [chunk async for chunk in _relay_late_response("late")] + + assert frames == [b"data: late\n\n", b"data: [DONE]\n\n"] + + +@pytest.mark.asyncio +async def test_relay_late_response_uses_an_empty_object_for_an_unknown_payload(): + from litellm.proxy.common_request_processing import _relay_late_response + + class _NoBody: + pass + + frames: Final = [chunk async for chunk in _relay_late_response(_NoBody())] + + assert frames == [b"data: {}\n\n", b"data: [DONE]\n\n"] + + +@pytest.mark.asyncio +async def test_relay_late_response_serializes_a_mapping_proxy(): + from types import MappingProxyType + + from litellm.proxy.common_request_processing import _relay_late_response + + frames: Final = [chunk async for chunk in _relay_late_response(MappingProxyType({"ok": True}))] + + assert frames == [b'data: {"ok":true}\n\n', b"data: [DONE]\n\n"]