mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
829cba1bf1
commit
2101c860c2
3 changed files with 229 additions and 11 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue