From 2101c860c25546250fa55c5a67653093a9d28873 Mon Sep 17 00:00:00 2001 From: Stewart Park <388348+stewartpark@users.noreply.github.com> Date: Sat, 26 Sep 2026 21:26:38 -0700 Subject: [PATCH] fix(vertex_ai): make Gemma fake streams work with traced Responses (#43147) * test(vertex_ai): reproduce traced Gemma Responses stream failure * fix(vertex_ai): wrap Gemma fake streams for Responses tracing * test(vertex_ai): cover Gemma traced streams and usage options * test(vertex_ai): inject gemma test deps and assert hidden usage accounting Replace class-level patches in the Vertex AI shard test with the provider's documented dependency-injection seams (httpx.MockTransport client + credential cache), and pin the default/omit-usage trace behavior: LiteLLM still accounts all tokens; ddtrace's metric is absent by design, asserted rather than silent. Mutation-checked: commenting out CustomStreamWrapper chunk accumulation turns the new assertions red; restoring them turns green. * test(vertex_ai): drop explanatory comment from usage-option assertions --- .../vertex_gemma_models/transformation.py | 30 ++++- .../test_vertex_gemma_transformation.py | 93 ++++++++++++++ .../test_vertex_gemma_transformation.py | 117 +++++++++++++++++- 3 files changed, 229 insertions(+), 11 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index ea97f0a0a9a..33922e38674 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -28,8 +28,8 @@ from litellm.types.utils import ModelResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer - from litellm.llms.base_llm.base_model_iterator import MockResponseIterator def parse_vertex_gemma_container_error(predictions: object) -> VertexGemmaContainerError | None: @@ -73,7 +73,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): self, model_response: ModelResponse, stream: bool, - ) -> "ModelResponse | MockResponseIterator": + model: str, + logging_obj: "LiteLLMLoggingObj", + ) -> "ModelResponse | CustomStreamWrapper": """ Helper method to return fake stream iterator if streaming is requested. @@ -82,12 +84,18 @@ class VertexGemmaConfig(OpenAIGPTConfig): stream: Whether streaming was requested Returns: - MockResponseIterator if stream=True, otherwise the model_response + CustomStreamWrapper if stream=True, otherwise the model_response """ if stream: + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator - return MockResponseIterator(model_response=model_response) + return CustomStreamWrapper( + completion_stream=MockResponseIterator(model_response=model_response), + model=model, + custom_llm_provider="vertex_ai", + logging_obj=logging_obj, + ) return model_response def transform_request( @@ -373,7 +381,12 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) async def _async_completion( self, @@ -463,4 +476,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py new file mode 100644 index 00000000000..294d26b2e58 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -0,0 +1,93 @@ +import json +from collections.abc import AsyncIterator +from types import SimpleNamespace +from typing import Any, cast + +import httpx +import pytest + +import litellm +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.main import vertex_gemma_chat_completion +from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse + +_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" +_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}] +_FAKE_CREDENTIALS = "gemma-test-credentials" + + +def _vertex_response(): + return { + "predictions": { + "id": "chatcmpl-stream-test", + "created": 1759863903, + "model": "google/gemma-3-12b-it", + "object": "chat.completion", + "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}], + "usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15}, + } + } + + +@pytest.fixture(autouse=True) +def _cached_access_token(): + """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" + cache = vertex_gemma_chat_completion._credentials_project_mapping + key = (_FAKE_CREDENTIALS, "test") + cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") + yield + cache.pop(key, None) + + +def test_sync_gemma_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + stream = litellm.completion( + model="vertex_ai/gemma/test-model", + messages=_MESSAGES, + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.Client(transport=httpx.MockTransport(handle)), + ) + + assert isinstance(stream, CustomStreamWrapper) + chunks = list(stream) + + assert "stream" not in captured["body"]["instances"][0] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_async_gemma_responses_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), + ) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + + assert "stream" not in captured["body"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage is not None + assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index e5ca31833ce..97f4f290958 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -5,11 +5,18 @@ Maps to: litellm/llms/vertex_ai/vertex_gemma_models/transformation.py """ import json +from collections.abc import AsyncIterator +from typing import cast from unittest.mock import AsyncMock, Mock, patch import pytest import litellm +from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIStreamingResponse, +) @pytest.fixture(autouse=True) @@ -439,8 +446,9 @@ class TestVertexGemmaCompletion: Verifies: 1. Request body does NOT include 'stream' parameter (model doesn't support it) - 2. Response returns a MockResponseIterator that yields chunks + 2. Response wraps a MockResponseIterator and yields chunks """ + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator # Mock Vertex response @@ -502,8 +510,8 @@ class TestVertexGemmaCompletion: vertex_location="us-central1", ) - # Verify the response is a MockResponseIterator - assert isinstance(response, MockResponseIterator), f"Expected MockResponseIterator, got {type(response)}" + assert isinstance(response, CustomStreamWrapper) + assert isinstance(response.completion_stream, MockResponseIterator) # Verify the request sent to Vertex does NOT include 'stream' call_args = mock_client.post.call_args @@ -520,8 +528,9 @@ class TestVertexGemmaCompletion: async for chunk in response: chunks.append(chunk) - # Should get exactly one chunk (fake streaming) - assert len(chunks) == 1, f"Expected 1 chunk from fake stream, got {len(chunks)}" + assert len(chunks) == 2 + assert chunks[1].choices[0].finish_reason == "stop" + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) # Verify the chunk has the expected content chunk = chunks[0] @@ -529,6 +538,104 @@ class TestVertexGemmaCompletion: assert len(chunk.choices) > 0 assert chunk.choices[0].delta.content == "Streaming test response" + @pytest.mark.asyncio + async def test_aresponses_streams_vertex_gemma_with_llm_tracing(self): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + from ddtrace.llmobs._integrations.base_stream_handler import TracedAsyncStream + + from litellm.responses.litellm_completion_transformation.streaming_iterator import ( + LiteLLMCompletionStreamingIterator, + ) + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock() + client.post = AsyncMock(return_value=reply) + + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + bridge = cast(LiteLLMCompletionStreamingIterator, response) + traced_stream = bridge.litellm_custom_stream_wrapper + assert isinstance(traced_stream, TracedAsyncStream) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + span = traced_stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + finally: + unpatch_litellm() + + assert "stream" not in client.post.call_args.kwargs["json"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage.total_tokens == 114 + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream_options", [None, {"include_usage": False}, {"include_usage": True}]) + async def test_acompletion_stream_respects_usage_option_with_llm_tracing(self, stream_options): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock(post=AsyncMock(return_value=reply)) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + stream = await litellm.acompletion( + model="vertex_ai/gemma/test-model", + messages=[{"role": "user", "content": "Reply exactly READY"}], + stream=True, + **({"stream_options": stream_options} if stream_options is not None else {}), + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + chunks = [chunk async for chunk in stream] + span = stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + finally: + unpatch_litellm() + + assert len(chunks) == (3 if stream_options and stream_options["include_usage"] else 2) + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + if stream_options and stream_options["include_usage"]: + assert chunks[-1].choices[0].delta.content is None + assert chunks[-1].usage.total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + else: + from litellm.litellm_core_utils.streaming_handler import calculate_total_usage + + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) + assert calculate_total_usage(chunks=stream.chunks).total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") is None + @pytest.mark.asyncio async def test_acompletion_filters_stream_and_stream_options(self): """