diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5d89d3fa9c..879914e5ac 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6917,6 +6917,15 @@ async def async_data_generator( # noqa: PLR0915 if isinstance(chunk, BaseModel): chunk = _serialize_streaming_chunk(chunk) + elif isinstance(chunk, bytes): + # Some upstream streaming iterators (e.g. AsyncGoogleGenAIGenerateContentStreamingIterator + # for /v1beta/.../streamGenerateContent) yield raw SSE bytes from Gemini. + # Decode to str so the f-string below does not emit a Python b'...' literal, + # and pass already-formatted SSE through unchanged to avoid double "data:" prefix. + chunk = chunk.decode("utf-8", errors="replace") + if chunk.startswith(("data:", "event:", ":")): + yield chunk if chunk.endswith("\n\n") else chunk + "\n\n" + continue elif isinstance(chunk, str) and chunk.startswith("data: "): error_message = chunk break diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 73d5363162..6d10d2a635 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5065,6 +5065,66 @@ async def test_async_data_generator_uses_direct_stream_fast_path_without_callbac mock_response.aclose.assert_awaited_once() +@pytest.mark.asyncio +async def test_async_data_generator_passes_through_google_native_sse_bytes(): + """ + Google-native streamGenerateContent yields raw SSE bytes; they must not be + re-wrapped as data: b'data: {...}'. + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + from litellm.proxy.utils import ProxyLogging + + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_request_data = { + "model": "gemini-2.0-flash", + "messages": [{"role": "user", "content": "test"}], + } + gemini_event = b'data: {"candidates": [{"content": "hi"}]}\n\n' + gemini_event_without_terminator = b'data: {"candidates": [{"content": "there"}]}' + raw_payload = b'{"partial": true}' + + class MockStream: + def __aiter__(self): + return self._stream() + + async def _stream(self): + yield gemini_event + yield gemini_event_without_terminator + yield raw_payload + + async def aclose(self): + pass + + mock_response = MockStream() + mock_response.aclose = AsyncMock() + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.has_streaming_callbacks.return_value = False + mock_proxy_logging_obj.needs_iterator_wrap.return_value = False + mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock() + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): + with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + yielded_data = [] + async for data in async_data_generator( + mock_response, mock_user_api_key_dict, mock_request_data + ): + yielded_data.append(data) + + yielded_text = [ + chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk + for chunk in yielded_data + ] + assert yielded_text[0] == gemini_event.decode("utf-8") + assert yielded_text[1] == gemini_event_without_terminator.decode("utf-8") + "\n\n" + assert yielded_text[2] == f'data: {raw_payload.decode("utf-8")}\n\n' + assert "b'data:" not in "".join(yielded_text) + assert yielded_text[-1] == "data: [DONE]\n\n" + + @pytest.mark.asyncio async def test_async_data_generator_cleanup_on_normal_completion(): """