test_acompletion_fake_streaming

This commit is contained in:
Ishaan Jaffer 2025-10-10 11:28:01 -07:00
parent 7a4c6344ea
commit 5384c2a770
2 changed files with 125 additions and 3 deletions

View file

@ -42,6 +42,26 @@ class VertexGemmaConfig(OpenAIGPTConfig):
"""
return True
def _handle_fake_stream_response(
self,
model_response: ModelResponse,
stream: bool,
) -> Union[ModelResponse, Any]:
"""
Helper method to return fake stream iterator if streaming is requested.
Args:
model_response: The completed model response
stream: Whether streaming was requested
Returns:
MockResponseIterator if stream=True, otherwise the model_response
"""
if stream:
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
return MockResponseIterator(model_response=model_response)
return model_response
def transform_request(
self,
model: str,
@ -167,6 +187,9 @@ class VertexGemmaConfig(OpenAIGPTConfig):
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.utils import convert_to_model_response_object
# Check if streaming is requested (will be faked)
stream = optional_params.get("stream", False)
# Transform the request using parent class methods
request_data = self.transform_request(
model=model,
@ -233,7 +256,8 @@ class VertexGemmaConfig(OpenAIGPTConfig):
additional_args={"complete_input_dict": request_data},
)
return model_response
# Return fake stream iterator if streaming was requested
return self._handle_fake_stream_response(model_response=model_response, stream=stream)
async def _async_completion(
self,
@ -250,9 +274,13 @@ class VertexGemmaConfig(OpenAIGPTConfig):
encoding: Any,
):
"""Asynchronous completion request"""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.utils import LlmProviders
from litellm.utils import convert_to_model_response_object
# Check if streaming is requested (will be faked)
stream = optional_params.get("stream", False)
# Transform the request using parent class async methods
request_data = await self.async_transform_request(
model=model,
@ -321,5 +349,6 @@ class VertexGemmaConfig(OpenAIGPTConfig):
additional_args={"complete_input_dict": request_data},
)
return model_response
# Return fake stream iterator if streaming was requested
return self._handle_fake_stream_response(model_response=model_response, stream=stream)

View file

@ -216,3 +216,96 @@ class TestVertexGemmaCompletion:
# Verify the error message contains the original error
assert "missing 'predictions' field" in str(exc_info.value)
@pytest.mark.asyncio
async def test_acompletion_fake_streaming(self):
"""
Test that streaming requests are faked properly for Vertex AI Gemma models.
Verifies:
1. Request body does NOT include 'stream' parameter (model doesn't support it)
2. Response returns a MockResponseIterator that yields chunks
"""
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
# Mock Vertex response
mock_vertex_response = {
"deployedModelId": "1207280419999999999",
"model": "projects/993702345710/locations/us-central1/models/gemma-3-12b-it-1222199011122",
"modelDisplayName": "gemma-3-12b-it-1222199011122",
"modelVersionId": "1",
"predictions": {
"choices": [
{
"finish_reason": "stop",
"index": 0,
"logprobs": None,
"message": {
"content": "Streaming test response",
"reasoning_content": None,
"role": "assistant",
"tool_calls": [],
},
"stop_reason": None,
}
],
"created": 1759863903,
"id": "chatcmpl-test-stream",
"model": "google/gemma-3-12b-it",
"object": "chat.completion",
"prompt_logprobs": None,
"usage": {
"completion_tokens": 3,
"prompt_tokens": 10,
"prompt_tokens_details": None,
"total_tokens": 13,
},
},
}
with patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
) as mock_get_client:
mock_client = Mock()
mock_response = Mock()
mock_response.status_code = 200
mock_response.json.return_value = mock_vertex_response
mock_client.post = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
# Call litellm.acompletion() with stream=True
response = await litellm.acompletion(
model="vertex_ai/gemma/gemma-3-12b-it-1222199011122",
messages=[{"role": "user", "content": "Test streaming"}],
stream=True,
api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict",
vertex_project="PROJECT_ID",
vertex_location="us-central1",
)
# Verify the response is a MockResponseIterator
assert isinstance(response, MockResponseIterator), f"Expected MockResponseIterator, got {type(response)}"
# Verify the request sent to Vertex does NOT include 'stream'
call_args = mock_client.post.call_args
assert call_args is not None, "HTTP client was not called"
request_data = call_args.kwargs["json"]
instance = request_data["instances"][0]
# Critical: Verify stream parameter is NOT sent to Vertex API
assert "stream" not in instance, "stream parameter should not be sent to Vertex API"
# Verify we can iterate the fake stream and get the response
chunks = []
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)}"
# Verify the chunk has the expected content
chunk = chunks[0]
assert hasattr(chunk, "choices")
assert len(chunk.choices) > 0
assert chunk.choices[0].delta.content == "Streaming test response"