diff --git a/litellm/google_genai/streaming_iterator.py b/litellm/google_genai/streaming_iterator.py index e2fac6a615b..a49e43e7bdc 100644 --- a/litellm/google_genai/streaming_iterator.py +++ b/litellm/google_genai/streaming_iterator.py @@ -75,6 +75,9 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: self.collected_chunks: list[bytes] = [] self.model = model self.custom_llm_provider = custom_llm_provider + self.endpoint_type: Final = ( + EndpointType.GEMINI if custom_llm_provider == litellm.LlmProviders.GEMINI.value else EndpointType.VERTEX_AI + ) self._hidden_params: dict[str, Any] = hidden_params or {} async def _handle_async_streaming_logging( @@ -86,18 +89,13 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: ) end_time: Final = datetime.now() - endpoint_type: Final = ( - EndpointType.GEMINI - if self.custom_llm_provider == litellm.LlmProviders.GEMINI.value - else EndpointType.VERTEX_AI - ) asyncio.create_task( PassThroughStreamingHandler._route_streaming_logging_to_handler( litellm_logging_obj=self.litellm_logging_obj, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, url_route="/v1/generateContent", request_body=self.request_body or {}, - endpoint_type=endpoint_type, + endpoint_type=self.endpoint_type, start_time=self.start_time, raw_bytes=self.collected_chunks, end_time=end_time, diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index c4dcd086629..5ad41b00890 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -248,7 +248,7 @@ class PassThroughStreamingHandler: kwargs = vertex_passthrough_logging_handler_result["kwargs"] elif endpoint_type == EndpointType.GEMINI: gemini_passthrough_logging_handler_result: Final = ( - GeminiPassthroughLoggingHandler._handle_logging_gemini_collected_chunks( + GeminiPassthroughLoggingHandler._handle_logging_gemini_collected_chunks( # pyright: ignore[reportPrivateUsage] # mirrors sibling handler dispatch litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, @@ -260,8 +260,12 @@ class PassThroughStreamingHandler: model=model, ) ) - standard_logging_response_object = gemini_passthrough_logging_handler_result["result"] - kwargs = gemini_passthrough_logging_handler_result["kwargs"] + standard_logging_response_object = ( # rebind-ok: branch bind in shared if/elif dispatch + gemini_passthrough_logging_handler_result["result"] + ) + kwargs = ( # rebind-ok: branch bind in shared if/elif dispatch + gemini_passthrough_logging_handler_result["kwargs"] + ) elif endpoint_type == EndpointType.OPENAI: openai_passthrough_logging_handler_result: Final = ( OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks( diff --git a/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py b/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py index 91058767730..e8ec2848233 100644 --- a/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py +++ b/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py @@ -1,6 +1,5 @@ -import asyncio import json -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import MagicMock import pytest @@ -12,24 +11,25 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType -@pytest.mark.asyncio @pytest.mark.parametrize( "custom_llm_provider, expected_endpoint_type", [("gemini", EndpointType.GEMINI), ("vertex_ai", EndpointType.VERTEX_AI)], ) -async def test_streaming_logging_routes_to_the_provider_that_served_the_request( - custom_llm_provider, expected_endpoint_type +@pytest.mark.parametrize( + "iterator_cls", + [ + AsyncGoogleGenAIGenerateContentStreamingIterator, + GoogleGenAIGenerateContentStreamingIterator, + ], +) +def test_streaming_logging_targets_the_provider_that_served_the_request( + iterator_cls: type, + custom_llm_provider: str, + expected_endpoint_type: EndpointType, ): """Routing every google stream through the vertex handler bills gemini/* at vertex_ai/ rates.""" - mock_response = MagicMock() - - async def _aiter_lines(): - yield 'data: {"candidates": []}' - - mock_response.aiter_lines = _aiter_lines - - iterator = AsyncGoogleGenAIGenerateContentStreamingIterator( - response=mock_response, + iterator = iterator_cls( + response=MagicMock(), model="gemini-3.1-flash-image", logging_obj=MagicMock(spec=LiteLLMLoggingObj), generate_content_provider_config=MagicMock(), @@ -37,15 +37,7 @@ async def test_streaming_logging_routes_to_the_provider_that_served_the_request( custom_llm_provider=custom_llm_provider, ) - with patch( - "litellm.proxy.pass_through_endpoints.streaming_handler.PassThroughStreamingHandler._route_streaming_logging_to_handler", - new=AsyncMock(), - ) as mock_route: - async for _ in iterator: - pass - - await asyncio.sleep(0) - assert mock_route.call_args.kwargs["endpoint_type"] == expected_endpoint_type + assert iterator.endpoint_type is expected_endpoint_type def _large_inline_data_event() -> str: @@ -91,9 +83,7 @@ async def test_async_streaming_iterator_yields_complete_sse_events(): assert chunk.startswith(b"data: ") assert chunk.endswith(b"\n\n") assert ( - json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][0][ - "inlineData" - ]["mimeType"] + json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][0]["inlineData"]["mimeType"] == "image/jpeg" ) @@ -114,9 +104,9 @@ def test_sync_streaming_iterator_yields_complete_sse_events(): chunk = next(iterator) assert chunk.startswith(b"data: ") assert chunk.endswith(b"\n\n") - assert json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][ - 0 - ]["inlineData"]["data"].startswith("A") + assert json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][0]["inlineData"][ + "data" + ].startswith("A") @pytest.mark.asyncio