mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
/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:
parent
7c49197f29
commit
60c7537cc7
7 changed files with 153 additions and 78 deletions
|
|
@ -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)}"
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue