/streamGenerateContent - non-gemini model support (#12647)

* fix(google_genai/adapters/transformation.py): enable calling non-googlegenai models via streaming

Fixes https://github.com/BerriAI/litellm/issues/12562

* test(test_openai.py): add unit test asserting streaming works as expected
This commit is contained in:
Krish Dholakia 2025-07-18 10:56:29 -07:00 • committed by GitHub
parent 7c49197f29
commit 60c7537cc7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 153 additions and 78 deletions

View file

@ -12,7 +12,7 @@ GOOGLE_GENAI_ADAPTER = GoogleGenAIAdapter()
class GenerateContentToCompletionHandler:
"""Handler for transforming generate_content calls to completion format when provider config is None"""
@staticmethod
def _prepare_completion_kwargs(
model: str,
@ -23,21 +23,23 @@ class GenerateContentToCompletionHandler:
extra_kwargs: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""Prepare kwargs for litellm.completion/acompletion"""
# Transform generate_content request to completion format
completion_request = GOOGLE_GENAI_ADAPTER.translate_generate_content_to_completion(
model=model,
contents=contents,
config=config,
litellm_params=litellm_params,
**(extra_kwargs or {})
completion_request = (
GOOGLE_GENAI_ADAPTER.translate_generate_content_to_completion(
model=model,
contents=contents,
config=config,
litellm_params=litellm_params,
**(extra_kwargs or {}),
)
)
completion_kwargs: Dict[str, Any] = dict(completion_request)
if stream:
completion_kwargs["stream"] = stream
return completion_kwargs
@staticmethod
@ -50,34 +52,40 @@ class GenerateContentToCompletionHandler:
**kwargs,
) -> Union[Dict[str, Any], AsyncIterator[bytes]]:
"""Handle generate_content call asynchronously using completion adapter"""
completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs(
model=model,
contents=contents,
config=config,
stream=stream,
litellm_params=litellm_params,
extra_kwargs=kwargs,
completion_kwargs = (
GenerateContentToCompletionHandler._prepare_completion_kwargs(
model=model,
contents=contents,
config=config,
stream=stream,
litellm_params=litellm_params,
extra_kwargs=kwargs,
)
)
try:
completion_response = await litellm.acompletion(**completion_kwargs)
if stream:
# Transform streaming completion response to generate_content format
transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
completion_response
transformed_stream = (
GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
completion_response
)
)
if transformed_stream is not None:
return transformed_stream
raise ValueError("Failed to transform streaming response")
else:
# Transform completion response back to generate_content format
generate_content_response = GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
cast(ModelResponse, completion_response)
generate_content_response = (
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
cast(ModelResponse, completion_response)
)
)
return generate_content_response
except Exception as e:
raise ValueError(
f"Error calling litellm.acompletion for generate_content: {str(e)}"
@ -92,9 +100,13 @@ class GenerateContentToCompletionHandler:
stream: bool = False,
_is_async: bool = False,
**kwargs,
) -> Union[Dict[str, Any], AsyncIterator[bytes], Coroutine[Any, Any, Union[Dict[str, Any], AsyncIterator[bytes]]]]:
) -> Union[
Dict[str, Any],
AsyncIterator[bytes],
Coroutine[Any, Any, Union[Dict[str, Any], AsyncIterator[bytes]]],
]:
"""Handle generate_content call using completion adapter"""
if _is_async:
return GenerateContentToCompletionHandler.async_generate_content_handler(
model=model,
@ -104,35 +116,41 @@ class GenerateContentToCompletionHandler:
litellm_params=litellm_params,
**kwargs,
)
completion_kwargs = GenerateContentToCompletionHandler._prepare_completion_kwargs(
model=model,
contents=contents,
config=config,
stream=stream,
litellm_params=litellm_params,
extra_kwargs=kwargs,
completion_kwargs = (
GenerateContentToCompletionHandler._prepare_completion_kwargs(
model=model,
contents=contents,
config=config,
stream=stream,
litellm_params=litellm_params,
extra_kwargs=kwargs,
)
)
try:
completion_response = litellm.completion(**completion_kwargs)
if stream:
# Transform streaming completion response to generate_content format
transformed_stream = GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
completion_response
transformed_stream = (
GOOGLE_GENAI_ADAPTER.translate_completion_output_params_streaming(
completion_response
)
)
if transformed_stream is not None:
return transformed_stream
raise ValueError("Failed to transform streaming response")
else:
# Transform completion response back to generate_content format
generate_content_response = GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
cast(ModelResponse, completion_response)
generate_content_response = (
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
cast(ModelResponse, completion_response)
)
)
return generate_content_response
except Exception as e:
raise ValueError(
f"Error calling litellm.completion for generate_content: {str(e)}"
)
)

View file

@ -18,6 +18,7 @@ from litellm.types.utils import (
AdapterCompletionStreamWrapper,
Choices,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
@ -35,6 +36,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
def __init__(self, completion_stream: Any):
self.sent_first_chunk = False
self.accumulated_tool_calls = {}
super().__init__(completion_stream)
def __next__(self):
try:
@ -78,7 +80,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
"""
Convert Google GenAI streaming chunks to Server-Sent Events format.
"""
for chunk in self:
for chunk in self.completion_stream:
if isinstance(chunk, dict):
payload = f"data: {json.dumps(chunk)}\n\n"
yield payload.encode()
@ -89,12 +91,25 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
"""
Async version of google_genai_sse_wrapper.
"""
async for chunk in self:
from litellm.types.utils import ModelResponseStream
async for chunk in self.completion_stream:
if isinstance(chunk, dict):
payload = f"data: {json.dumps(chunk)}\n\n"
yield payload.encode()
elif isinstance(chunk, ModelResponseStream):
# Transform OpenAI streaming chunk to Google GenAI format
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
chunk, self
)
if isinstance(transformed_chunk, dict): # Only return non-empty chunks
payload = f"data: {json.dumps(transformed_chunk)}\n\n"
yield payload.encode()
else:
raise ValueError(f"Invalid chunk 1: {chunk}")
else:
yield chunk
raise ValueError(f"Invalid chunk 2: {chunk}")
class GoogleGenAIAdapter:
@ -184,7 +199,7 @@ class GoogleGenAIAdapter:
)
if tool_choice:
completion_request["tool_choice"] = tool_choice
#########################################################
# forward any litellm specific params
#########################################################
@ -192,15 +207,15 @@ class GoogleGenAIAdapter:
if litellm_params:
completion_request_dict = self._add_generic_litellm_params_to_request(
completion_request_dict=completion_request_dict,
litellm_params=litellm_params
litellm_params=litellm_params,
)
return completion_request_dict
def _add_generic_litellm_params_to_request(
self,
completion_request_dict: Dict[str, Any],
litellm_params: Optional[GenericLiteLLMParams] = None
self,
completion_request_dict: Dict[str, Any],
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> dict:
"""Add generic litellm params to request. e.g add api_base, api_key, api_version, etc.
@ -420,7 +435,9 @@ class GoogleGenAIAdapter:
return generate_content_response
def translate_streaming_completion_to_generate_content(
self, response: ModelResponse, wrapper: GoogleGenAIStreamWrapper
self,
response: Union[ModelResponse, ModelResponseStream],
wrapper: GoogleGenAIStreamWrapper,
) -> Dict[str, Any]:
"""
Transform streaming litellm completion chunk to Google GenAI generate_content format

View file

@ -384,13 +384,15 @@ async def agenerate_content_stream(
# Check if we should use the adapter (when provider config is None)
if setup_result.generate_content_provider_config is None:
# Use the adapter to convert to completion format
return await GenerateContentToCompletionHandler.async_generate_content_handler(
model=model,
contents=contents, # type: ignore
config=setup_result.generate_content_config_dict,
litellm_params=setup_result.litellm_params,
stream=True,
**kwargs
return (
await GenerateContentToCompletionHandler.async_generate_content_handler(
model=model,
contents=contents, # type: ignore
config=setup_result.generate_content_config_dict,
litellm_params=setup_result.litellm_params,
stream=True,
**kwargs,
)
)
# Call the handler with async enabled and streaming

View file

@ -1,21 +1,13 @@
model_list:
- model_name: gpt-3.5-turbo
- model_name: openai/gpt-3.5-turbo
litellm_params:
model: gpt-3.5-turbo
api_key: os.environ/OPENAI_API_KEY
model: openai/gpt-3.5-turbo
model_info:
version: 2
guardrails:
- guardrail_name: "guardrails_ai-guard"
litellm_params:
guardrail: guardrails_ai
guard_name: "pii_detect" # 👈 Guardrail AI guard name
mode: "logging_only"
api_base: os.environ/GUARDRAILS_AI_API_BASE # 👈 Guardrails AI API Base. Defaults to "http://0.0.0.0:8000"
default_on: true
router_settings:
model_group_alias: {"gemini-2.5-pro": "openai/gpt-3.5-turbo"}
# Tried adding the below settings as well
litellm_settings:
default_internal_user_params:
teams:
- team_id: "team_id_1"
max_budget_in_team: 100
user_role: "user"
stream: false

View file

@ -117,6 +117,7 @@ async def create_streaming_response(
first_chunk_value = await generator.__anext__()
if first_chunk_value is not None:
try:
error_code_from_chunk = await _parse_event_data_for_error(
first_chunk_value
@ -130,6 +131,7 @@ async def create_streaming_response(
verbose_proxy_logger.debug(f"Error parsing first chunk value: {e}")
except StopAsyncIteration:
# Generator was empty. Default status
async def empty_gen() -> AsyncGenerator[str, None]:
if False:
@ -142,6 +144,7 @@ async def create_streaming_response(
status_code=default_status_code,
)
except Exception as e:
# Unexpected error consuming first chunk.
verbose_proxy_logger.exception(
f"Error consuming first chunk from generator: {e}"
@ -454,6 +457,7 @@ class ProxyBaseLLMRequestProcessing:
) or self._is_streaming_response(
response
): # use generate_responses to stream responses
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=logging_obj.litellm_call_id,
@ -492,6 +496,7 @@ class ProxyBaseLLMRequestProcessing:
headers=custom_headers,
)
else:
selected_data_generator = select_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,

View file

@ -2548,7 +2548,6 @@ class Router:
self.total_calls[model_name] += 1
### get custom
response = original_generic_function(
**{
**data,
@ -2585,6 +2584,7 @@ class Router:
verbose_router_logger.info(
f"ageneric_api_call_with_fallbacks(model={model_name})\033[32m 200 OK\033[0m"
)
return response
except Exception as e:
verbose_router_logger.info(

View file

@ -544,6 +544,47 @@ async def test_openai_codex(sync_mode):
assert response.choices[0].message.content is not None
@pytest.mark.asyncio
async def test_openai_via_gemini_streaming_bridge():
"""
Test that the openai via gemini streaming bridge works correctly
"""
from litellm import Router
from litellm.types.utils import ModelResponseStream
router = Router(
model_list=[
{
"model_name": "openai/gpt-3.5-turbo",
"litellm_params": {
"model": "openai/gpt-3.5-turbo",
},
"model_info": {
"version": 2,
},
}
],
model_group_alias={"gemini-2.5-pro": "openai/gpt-3.5-turbo"},
)
response = await router.agenerate_content_stream(
model="openai/gpt-3.5-turbo",
contents=[
{
"parts": [{"text": "Write a long story about space exploration"}],
"role": "user",
}
],
generationConfig={"maxOutputTokens": 500},
)
printed_chunks = []
async for chunk in response:
print("chunk: ", chunk)
printed_chunks.append(chunk)
assert not isinstance(chunk, ModelResponseStream)
assert len(printed_chunks) > 0
def test_openai_deepresearch_model_bridge():
"""