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
This commit is contained in:
Stewart Park 2026-09-26 21:26:38 -07:00 • committed by GitHub
parent 829cba1bf1
commit 2101c860c2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 229 additions and 11 deletions

View file

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

View file

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

View file

@ -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):
"""