mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(google_genai): pick the stream logging endpoint type at construction
This commit is contained in:
parent
057781a187
commit
e8ec34c4c8
3 changed files with 30 additions and 38 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue