refactor(google_genai): pick the stream logging endpoint type at construction

This commit is contained in:
mateo-berri 2026-08-26 17:18:44 -07:00
parent 057781a187
commit e8ec34c4c8
3 changed files with 30 additions and 38 deletions

View file

@ -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,

View file

@ -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(

View file

@ -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