mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
feat(gemini): Add full support for native Gemini API translation
This commit implements a complete, end-to-end fix for the native Gemini API translation feature, allowing requests to be correctly routed to other model providers via `model_group_alias`. The original implementation was broken, causing `systemInstruction` and `tools` to be dropped from requests. This was resolved by refactoring the Gemini endpoint to use a dedicated translation path, similar to the Anthropic adapter. Additionally, this commit hardens the streaming response adapter to correctly handle tool calls generated by the newly-fixed request path. Key improvements to the response handling include: - Replaced the fragile `id`-based tool call tracking with a robust `index`-based accumulation logic. - Fixed a memory leak and improved logging in the stream finalization process. - Prevented empty, non-compliant chunks from being sent to the client during tool call streaming. - Optimized the accumulator to skip and log superfluous empty chunks sent by some models.
This commit is contained in:
parent
7052108d19
commit
b46407fa76
8 changed files with 294 additions and 268 deletions
|
|
@ -1355,6 +1355,7 @@ from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_k
|
|||
|
||||
### PASSTHROUGH ###
|
||||
from .passthrough import allm_passthrough_route, llm_passthrough_route
|
||||
from .google_genai import agenerate_content
|
||||
|
||||
### GLOBAL CONFIG ###
|
||||
global_bitbucket_config: Optional[Dict[str, Any]] = None
|
||||
|
|
|
|||
|
|
@ -72,15 +72,26 @@ class GenerateContentToCompletionHandler:
|
|||
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
|
||||
# Check if completion_response is actually a stream or a ModelResponse
|
||||
# This can happen in error cases or when stream is not properly supported
|
||||
if not hasattr(completion_response, '__aiter__'):
|
||||
# If it's not a stream, treat it as a regular response
|
||||
generate_content_response = (
|
||||
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
)
|
||||
)
|
||||
if transformed_stream is not None:
|
||||
return transformed_stream
|
||||
raise ValueError("Failed to transform streaming response")
|
||||
return generate_content_response
|
||||
else:
|
||||
# Transform streaming completion response to generate_content format
|
||||
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 = (
|
||||
|
|
@ -136,15 +147,26 @@ class GenerateContentToCompletionHandler:
|
|||
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
|
||||
# Check if completion_response is actually a stream or a ModelResponse
|
||||
# This can happen in error cases or when stream is not properly supported
|
||||
if not hasattr(completion_response, '__iter__'):
|
||||
# If it's not a stream, treat it as a regular response
|
||||
generate_content_response = (
|
||||
GOOGLE_GENAI_ADAPTER.translate_completion_to_generate_content(
|
||||
cast(ModelResponse, completion_response)
|
||||
)
|
||||
)
|
||||
)
|
||||
if transformed_stream is not None:
|
||||
return transformed_stream
|
||||
raise ValueError("Failed to transform streaming response")
|
||||
return generate_content_response
|
||||
else:
|
||||
# Transform streaming completion response to generate_content format
|
||||
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 = (
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
import json
|
||||
from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Union, cast
|
||||
|
||||
from litellm import verbose_logger
|
||||
|
||||
from litellm.litellm_core_utils.json_validation_rule import normalize_tool_schema
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
|
|
@ -31,48 +33,106 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
|
||||
sent_first_chunk: bool = False
|
||||
# State tracking for accumulating partial tool calls
|
||||
accumulated_tool_calls: Dict[str, Dict[str, Any]]
|
||||
gccumulated_tool_calls: Dict[str, Dict[str, Any]]
|
||||
|
||||
def __init__(self, completion_stream: Any):
|
||||
self.sent_first_chunk = False
|
||||
self.accumulated_tool_calls = {}
|
||||
self._returned_response = False
|
||||
super().__init__(completion_stream)
|
||||
|
||||
def __next__(self):
|
||||
try:
|
||||
if not hasattr(self.completion_stream, '__iter__'):
|
||||
if self._returned_response:
|
||||
raise StopIteration
|
||||
self._returned_response = True
|
||||
return GoogleGenAIAdapter().translate_completion_to_generate_content(
|
||||
self.completion_stream
|
||||
)
|
||||
|
||||
for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
continue
|
||||
|
||||
# Transform OpenAI streaming chunk to Google GenAI format
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
|
||||
chunk, self
|
||||
)
|
||||
if transformed_chunk: # Only return non-empty chunks
|
||||
if transformed_chunk:
|
||||
return transformed_chunk
|
||||
|
||||
raise StopIteration
|
||||
except StopIteration:
|
||||
raise StopIteration
|
||||
raise
|
||||
except Exception:
|
||||
raise StopIteration
|
||||
|
||||
async def __anext__(self):
|
||||
try:
|
||||
if not hasattr(self.completion_stream, '__aiter__'):
|
||||
if self._returned_response:
|
||||
raise StopAsyncIteration
|
||||
self._returned_response = True
|
||||
return GoogleGenAIAdapter().translate_completion_to_generate_content(
|
||||
self.completion_stream
|
||||
)
|
||||
|
||||
async for chunk in self.completion_stream:
|
||||
if chunk == "None" or chunk is None:
|
||||
continue
|
||||
|
||||
# Transform OpenAI streaming chunk to Google GenAI format
|
||||
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(
|
||||
chunk, self
|
||||
)
|
||||
if transformed_chunk: # Only return non-empty chunks
|
||||
if transformed_chunk:
|
||||
return transformed_chunk
|
||||
|
||||
# After the stream is exhausted, check for any remaining accumulated tool calls
|
||||
if self.accumulated_tool_calls:
|
||||
try:
|
||||
parts = []
|
||||
for (
|
||||
tool_call_index,
|
||||
tool_call_data,
|
||||
) in self.accumulated_tool_calls.items():
|
||||
try:
|
||||
# For tool calls with no arguments, accumulated_args will be "", which is not valid JSON.
|
||||
# We default to an empty JSON object in this case.
|
||||
parsed_args = json.loads(tool_call_data["arguments"] or "{}")
|
||||
function_call_part = {
|
||||
"functionCall": {
|
||||
"name": tool_call_data["name"]
|
||||
or "undefined_tool_name",
|
||||
"args": parsed_args,
|
||||
}
|
||||
}
|
||||
parts.append(function_call_part)
|
||||
except json.JSONDecodeError:
|
||||
# This can happen if the stream is abruptly cut off mid-argument string.
|
||||
verbose_logger.warning(
|
||||
f"Could not parse tool call arguments at end of stream for index {tool_call_index}. "
|
||||
f"Name: {tool_call_data['name']}. "
|
||||
f"Partial args: {tool_call_data['arguments']}"
|
||||
)
|
||||
pass
|
||||
if parts:
|
||||
final_chunk = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
"finishReason": "STOP",
|
||||
"index": 0,
|
||||
"safetyRatings": [],
|
||||
}
|
||||
]
|
||||
}
|
||||
return final_chunk
|
||||
finally:
|
||||
# Ensure the accumulator is always cleared to prevent memory leaks
|
||||
self.accumulated_tool_calls.clear()
|
||||
raise StopAsyncIteration
|
||||
except StopAsyncIteration:
|
||||
raise StopAsyncIteration
|
||||
raise
|
||||
except Exception:
|
||||
raise StopAsyncIteration
|
||||
|
||||
|
|
@ -107,9 +167,14 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
payload = f"data: {json.dumps(transformed_chunk)}\n\n"
|
||||
yield payload.encode()
|
||||
else:
|
||||
raise ValueError(f"Invalid chunk 1: {chunk}")
|
||||
# For empty chunks, continue to next iteration
|
||||
continue
|
||||
else:
|
||||
raise ValueError(f"Invalid chunk 2: {chunk}")
|
||||
# For other chunk types, yield them directly
|
||||
if hasattr(chunk, 'encode'):
|
||||
yield chunk.encode()
|
||||
else:
|
||||
yield str(chunk).encode()
|
||||
|
||||
|
||||
class GoogleGenAIAdapter:
|
||||
|
|
@ -126,6 +191,7 @@ class GoogleGenAIAdapter:
|
|||
litellm_params: Optional[GenericLiteLLMParams] = None,
|
||||
**kwargs,
|
||||
) -> Dict[str, Any]:
|
||||
|
||||
"""
|
||||
Transform generate_content request to litellm completion format
|
||||
|
||||
|
|
@ -133,12 +199,20 @@ class GoogleGenAIAdapter:
|
|||
model: The model name
|
||||
contents: Generate content contents (can be list or single dict)
|
||||
config: Optional config parameters
|
||||
**kwargs: Additional parameters
|
||||
**kwargs: Additional parameters from the original request
|
||||
|
||||
Returns:
|
||||
Dict in OpenAI format
|
||||
"""
|
||||
|
||||
# Extract top-level fields from kwargs
|
||||
system_instruction = kwargs.get("systemInstruction") or kwargs.get(
|
||||
"system_instruction"
|
||||
)
|
||||
tools = kwargs.get("tools")
|
||||
tool_config = kwargs.get("toolConfig") or kwargs.get("tool_config")
|
||||
|
||||
|
||||
# Normalize contents to list format
|
||||
if isinstance(contents, dict):
|
||||
contents_list = [contents]
|
||||
|
|
@ -146,7 +220,10 @@ class GoogleGenAIAdapter:
|
|||
contents_list = contents
|
||||
|
||||
# Transform contents to OpenAI messages format
|
||||
messages = self._transform_contents_to_messages(contents_list)
|
||||
messages = self._transform_contents_to_messages(
|
||||
contents_list, system_instruction=system_instruction
|
||||
)
|
||||
|
||||
|
||||
# Create base request as dict (which is compatible with ChatCompletionRequest)
|
||||
completion_request: ChatCompletionRequest = {
|
||||
|
|
@ -182,20 +259,19 @@ class GoogleGenAIAdapter:
|
|||
completion_request["stop"] = config["stopSequences"]
|
||||
|
||||
# Handle tools transformation
|
||||
if "tools" in kwargs:
|
||||
tools = kwargs["tools"]
|
||||
|
||||
if tools:
|
||||
# Check if tools are already in OpenAI format or Google GenAI format
|
||||
if isinstance(tools, list) and len(tools) > 0:
|
||||
# Tools are in Google GenAI format, transform them
|
||||
openai_tools = self._transform_google_genai_tools_to_openai(tools)
|
||||
|
||||
if openai_tools:
|
||||
completion_request["tools"] = openai_tools
|
||||
|
||||
# Handle tool_config (tool choice)
|
||||
if "tool_config" in kwargs:
|
||||
if tool_config:
|
||||
tool_choice = self._transform_google_genai_tool_config_to_openai(
|
||||
kwargs["tool_config"]
|
||||
tool_config
|
||||
)
|
||||
if tool_choice:
|
||||
completion_request["tool_choice"] = tool_choice
|
||||
|
|
@ -235,7 +311,8 @@ class GoogleGenAIAdapter:
|
|||
return completion_request_dict
|
||||
|
||||
def translate_completion_output_params_streaming(
|
||||
self, completion_stream: Any
|
||||
self,
|
||||
completion_stream: Any,
|
||||
) -> Union[AsyncIterator[bytes], None]:
|
||||
"""Transform streaming completion output to Google GenAI format"""
|
||||
google_genai_wrapper = GoogleGenAIStreamWrapper(
|
||||
|
|
@ -245,7 +322,8 @@ class GoogleGenAIAdapter:
|
|||
return google_genai_wrapper.async_google_genai_sse_wrapper()
|
||||
|
||||
def _transform_google_genai_tools_to_openai(
|
||||
self, tools: List[Dict[str, Any]]
|
||||
self,
|
||||
tools: List[Dict[str, Any]],
|
||||
) -> List[ChatCompletionToolParam]:
|
||||
"""Transform Google GenAI tools to OpenAI tools format"""
|
||||
openai_tools: List[Dict[str, Any]] = []
|
||||
|
|
@ -259,8 +337,10 @@ class GoogleGenAIAdapter:
|
|||
|
||||
if "description" in func_decl:
|
||||
function_chunk["description"] = func_decl["description"]
|
||||
if "parameters" in func_decl:
|
||||
function_chunk["parameters"] = func_decl["parameters"]
|
||||
if "parametersJsonSchema" in func_decl:
|
||||
function_chunk["parameters"] = func_decl[
|
||||
"parametersJsonSchema"
|
||||
]
|
||||
|
||||
openai_tool = {"type": "function", "function": function_chunk}
|
||||
openai_tools.append(openai_tool)
|
||||
|
|
@ -271,7 +351,8 @@ class GoogleGenAIAdapter:
|
|||
return cast(List[ChatCompletionToolParam], normalized_tools)
|
||||
|
||||
def _transform_google_genai_tool_config_to_openai(
|
||||
self, tool_config: Dict[str, Any]
|
||||
self,
|
||||
tool_config: Dict[str, Any],
|
||||
) -> Optional[ChatCompletionToolChoiceValues]:
|
||||
"""Transform Google GenAI tool_config to OpenAI tool_choice"""
|
||||
function_calling_config = tool_config.get("functionCallingConfig", {})
|
||||
|
|
@ -283,11 +364,23 @@ class GoogleGenAIAdapter:
|
|||
return cast(ChatCompletionToolChoiceValues, tool_choice)
|
||||
|
||||
def _transform_contents_to_messages(
|
||||
self, contents: List[Dict[str, Any]]
|
||||
self,
|
||||
contents: List[Dict[str, Any]],
|
||||
system_instruction: Optional[Dict[str, Any]] = None,
|
||||
) -> List[AllMessageValues]:
|
||||
"""Transform Google GenAI contents to OpenAI messages format"""
|
||||
messages: List[AllMessageValues] = []
|
||||
|
||||
# Handle system instruction
|
||||
if system_instruction:
|
||||
system_parts = system_instruction.get("parts", [])
|
||||
if system_parts and "text" in system_parts[0]:
|
||||
messages.append(
|
||||
ChatCompletionUserMessage(
|
||||
role="system", content=system_parts[0]["text"]
|
||||
)
|
||||
)
|
||||
|
||||
for content in contents:
|
||||
role = content.get("role", "user")
|
||||
parts = content.get("parts", [])
|
||||
|
|
@ -364,7 +457,8 @@ class GoogleGenAIAdapter:
|
|||
return messages
|
||||
|
||||
def translate_completion_to_generate_content(
|
||||
self, response: ModelResponse
|
||||
self,
|
||||
response: ModelResponse,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Transform litellm completion response to Google GenAI generate_content format
|
||||
|
|
@ -375,6 +469,8 @@ class GoogleGenAIAdapter:
|
|||
Returns:
|
||||
Dict in Google GenAI generate_content response format
|
||||
"""
|
||||
if isinstance(response, AdapterCompletionStreamWrapper):
|
||||
return self.translate_streaming_completion_to_generate_content(response, wrapper=response)
|
||||
|
||||
# Extract the main response content
|
||||
choice = response.choices[0] if response.choices else None
|
||||
|
|
@ -388,12 +484,6 @@ class GoogleGenAIAdapter:
|
|||
"Invalid completion response: no message found in choice"
|
||||
)
|
||||
parts = self._transform_openai_message_to_google_genai_parts(choice.message)
|
||||
elif isinstance(choice, StreamingChoices):
|
||||
if not choice.delta:
|
||||
raise ValueError(
|
||||
"Invalid completion response: no delta found in streaming choice"
|
||||
)
|
||||
parts = self._transform_openai_delta_to_google_genai_parts(choice.delta)
|
||||
else:
|
||||
# Fallback for generic choice objects
|
||||
message_content = getattr(choice, "message", {}).get(
|
||||
|
|
@ -438,7 +528,8 @@ class GoogleGenAIAdapter:
|
|||
self,
|
||||
response: Union[ModelResponse, ModelResponseStream],
|
||||
wrapper: GoogleGenAIStreamWrapper,
|
||||
) -> Dict[str, Any]:
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
|
||||
"""
|
||||
Transform streaming litellm completion chunk to Google GenAI generate_content format
|
||||
|
||||
|
|
@ -454,7 +545,7 @@ class GoogleGenAIAdapter:
|
|||
choice = response.choices[0] if response.choices else None
|
||||
if not choice:
|
||||
# Return empty chunk if no choices
|
||||
return {}
|
||||
return None
|
||||
|
||||
# Handle streaming choice
|
||||
if isinstance(choice, StreamingChoices):
|
||||
|
|
@ -473,7 +564,7 @@ class GoogleGenAIAdapter:
|
|||
|
||||
# Only create response chunk if we have parts or it's the final chunk
|
||||
if not parts and not finish_reason:
|
||||
return {}
|
||||
return None
|
||||
|
||||
# Create Google GenAI streaming format response
|
||||
streaming_chunk: Dict[str, Any] = {
|
||||
|
|
@ -515,7 +606,8 @@ class GoogleGenAIAdapter:
|
|||
return streaming_chunk
|
||||
|
||||
def _transform_openai_message_to_google_genai_parts(
|
||||
self, message: Any
|
||||
self,
|
||||
message: Any,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Transform OpenAI message to Google GenAI parts format"""
|
||||
parts: List[Dict[str, Any]] = []
|
||||
|
|
@ -537,112 +629,93 @@ class GoogleGenAIAdapter:
|
|||
except json.JSONDecodeError:
|
||||
args = {}
|
||||
|
||||
function_call_part = {
|
||||
"functionCall": {"name": tool_call.function.name, "args": args}
|
||||
}
|
||||
parts.append(function_call_part)
|
||||
|
||||
return parts if parts else [{"text": ""}]
|
||||
|
||||
def _transform_openai_delta_to_google_genai_parts(
|
||||
self, delta: Any
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Transform OpenAI delta to Google GenAI parts format for streaming"""
|
||||
parts: List[Dict[str, Any]] = []
|
||||
|
||||
# Add text content if present
|
||||
if hasattr(delta, "content") and delta.content:
|
||||
parts.append({"text": delta.content})
|
||||
|
||||
# Add tool calls if present (for streaming tool calls)
|
||||
if hasattr(delta, "tool_calls") and delta.tool_calls:
|
||||
for tool_call in delta.tool_calls:
|
||||
if hasattr(tool_call, "function") and tool_call.function:
|
||||
# For streaming, we might get partial function arguments
|
||||
args_str = getattr(tool_call.function, "arguments", "") or ""
|
||||
try:
|
||||
args = json.loads(args_str) if args_str else {}
|
||||
except json.JSONDecodeError:
|
||||
# For partial JSON in streaming, return as text for now
|
||||
args = {"partial": args_str}
|
||||
|
||||
function_call_part = {
|
||||
"functionCall": {
|
||||
"name": getattr(tool_call.function, "name", "") or "",
|
||||
"name": tool_call.function.name or "undefined_tool_name",
|
||||
"args": args,
|
||||
}
|
||||
}
|
||||
parts.append(function_call_part)
|
||||
|
||||
return parts
|
||||
return parts if parts else [{"text": ""}]
|
||||
|
||||
|
||||
def _transform_openai_delta_to_google_genai_parts_with_accumulation(
|
||||
self, delta: Any, wrapper: GoogleGenAIStreamWrapper
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Transform OpenAI delta to Google GenAI parts format with tool call accumulation"""
|
||||
"""Transforms OpenAI delta to Google GenAI parts, accumulating streaming tool calls."""
|
||||
|
||||
# 1. Initialize wrapper state if it doesn't exist
|
||||
if not hasattr(wrapper, "accumulated_tool_calls"):
|
||||
wrapper.accumulated_tool_calls = {}
|
||||
|
||||
parts: List[Dict[str, Any]] = []
|
||||
|
||||
# Add text content if present
|
||||
if hasattr(delta, "content") and delta.content:
|
||||
parts.append({"text": delta.content})
|
||||
|
||||
# Handle tool calls with accumulation for streaming
|
||||
if hasattr(delta, "tool_calls") and delta.tool_calls:
|
||||
for tool_call in delta.tool_calls:
|
||||
if hasattr(tool_call, "function") and tool_call.function:
|
||||
tool_call_id = getattr(tool_call, "id", "") or "call_unknown"
|
||||
function_name = getattr(tool_call.function, "name", "") or ""
|
||||
args_str = getattr(tool_call.function, "arguments", "") or ""
|
||||
# 2. Ensure tool_calls is iterable
|
||||
tool_calls = delta.tool_calls or []
|
||||
|
||||
# Initialize accumulation for this tool call if not exists
|
||||
if tool_call_id not in wrapper.accumulated_tool_calls:
|
||||
wrapper.accumulated_tool_calls[tool_call_id] = {
|
||||
"name": "",
|
||||
"arguments": "",
|
||||
"complete": False,
|
||||
}
|
||||
for tool_call in tool_calls:
|
||||
if not hasattr(tool_call, "function"):
|
||||
continue
|
||||
|
||||
# Accumulate function name if provided
|
||||
if function_name:
|
||||
wrapper.accumulated_tool_calls[tool_call_id][
|
||||
"name"
|
||||
] = function_name
|
||||
# 3. Use `index` as the primary key for accumulation
|
||||
tool_call_index = getattr(tool_call, "index", None)
|
||||
if tool_call_index is None:
|
||||
continue # Index is essential for tracking streaming tool calls
|
||||
|
||||
# Accumulate arguments if provided
|
||||
if args_str:
|
||||
wrapper.accumulated_tool_calls[tool_call_id][
|
||||
"arguments"
|
||||
] += args_str
|
||||
# Initialize accumulator for this index if it's new
|
||||
if tool_call_index not in wrapper.accumulated_tool_calls:
|
||||
wrapper.accumulated_tool_calls[tool_call_index] = {
|
||||
"name": "",
|
||||
"arguments": "",
|
||||
}
|
||||
|
||||
# Try to parse the accumulated arguments as JSON
|
||||
accumulated_args = wrapper.accumulated_tool_calls[tool_call_id][
|
||||
"arguments"
|
||||
]
|
||||
try:
|
||||
if accumulated_args:
|
||||
parsed_args = json.loads(accumulated_args)
|
||||
# JSON is valid, mark as complete and create function call part
|
||||
wrapper.accumulated_tool_calls[tool_call_id][
|
||||
"complete"
|
||||
] = True
|
||||
# Accumulate name and arguments
|
||||
function_name = getattr(tool_call.function, "name", None)
|
||||
args_chunk = getattr(tool_call.function, "arguments", None)
|
||||
|
||||
function_call_part = {
|
||||
"functionCall": {
|
||||
"name": wrapper.accumulated_tool_calls[
|
||||
tool_call_id
|
||||
]["name"],
|
||||
"args": parsed_args,
|
||||
}
|
||||
}
|
||||
parts.append(function_call_part)
|
||||
# Optimization: Skip chunks that have no new data
|
||||
if not function_name and not args_chunk:
|
||||
verbose_logger.debug(
|
||||
f"Skipping empty tool call chunk for index: {tool_call_index}"
|
||||
)
|
||||
continue
|
||||
|
||||
# Clean up completed tool call
|
||||
del wrapper.accumulated_tool_calls[tool_call_id]
|
||||
if function_name:
|
||||
wrapper.accumulated_tool_calls[tool_call_index]["name"] = function_name
|
||||
|
||||
except json.JSONDecodeError:
|
||||
# JSON is still incomplete, continue accumulating
|
||||
# Don't add to parts yet
|
||||
pass
|
||||
if args_chunk:
|
||||
wrapper.accumulated_tool_calls[tool_call_index]["arguments"] += args_chunk
|
||||
|
||||
# Attempt to parse and emit a complete tool call
|
||||
accumulated_data = wrapper.accumulated_tool_calls[tool_call_index]
|
||||
accumulated_name = accumulated_data["name"]
|
||||
accumulated_args = accumulated_data["arguments"]
|
||||
|
||||
# 5. Attempt to parse arguments even if name hasn't arrived.
|
||||
try:
|
||||
# Attempt to parse the accumulated arguments string
|
||||
parsed_args = json.loads(accumulated_args)
|
||||
|
||||
# If parsing succeeds, but we don't have a name yet, wait.
|
||||
# The part will be created by a later chunk that brings the name.
|
||||
if accumulated_name:
|
||||
# If successful, create the part and clean up
|
||||
function_call_part = {
|
||||
"functionCall": {"name": accumulated_name, "args": parsed_args}
|
||||
}
|
||||
parts.append(function_call_part)
|
||||
|
||||
# Remove the completed tool call from the accumulator
|
||||
del wrapper.accumulated_tool_calls[tool_call_index]
|
||||
|
||||
except json.JSONDecodeError:
|
||||
# The JSON for arguments is still incomplete.
|
||||
# We will continue to accumulate and wait for more chunks.
|
||||
pass
|
||||
|
||||
return parts
|
||||
|
||||
|
|
|
|||
|
|
@ -85,7 +85,6 @@ class GenerateContentHelper:
|
|||
contents: GenerateContentContentListUnionDict,
|
||||
config: Optional[GenerateContentConfigDict] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
stream: bool = False,
|
||||
tools: Optional[ToolConfigDict] = None,
|
||||
**kwargs,
|
||||
) -> GenerateContentSetupResult:
|
||||
|
|
@ -97,8 +96,7 @@ class GenerateContentHelper:
|
|||
contents: The content to generate from
|
||||
config: Optional configuration
|
||||
custom_llm_provider: Optional custom LLM provider
|
||||
stream: Whether this is a streaming call
|
||||
local_vars: Local variables from the calling function
|
||||
tools: Optional tools
|
||||
**kwargs: Additional keyword arguments
|
||||
|
||||
Returns:
|
||||
|
|
@ -114,7 +112,7 @@ class GenerateContentHelper:
|
|||
|
||||
## MOCK RESPONSE LOGIC (only for non-streaming)
|
||||
if (
|
||||
not stream
|
||||
not kwargs.get("stream", False)
|
||||
and litellm_params.mock_response
|
||||
and isinstance(litellm_params.mock_response, str)
|
||||
):
|
||||
|
|
@ -289,7 +287,7 @@ def generate_content(
|
|||
"""
|
||||
local_vars = locals()
|
||||
try:
|
||||
_is_async = kwargs.pop("agenerate_content", False) is True
|
||||
_is_async = kwargs.pop("agenerate_content", False)
|
||||
|
||||
# Handle generationConfig parameter from kwargs for backward compatibility
|
||||
if "generationConfig" in kwargs and config is None:
|
||||
|
|
@ -309,7 +307,6 @@ def generate_content(
|
|||
contents=contents,
|
||||
config=config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
stream=False,
|
||||
tools=tools,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -321,7 +318,7 @@ def generate_content(
|
|||
model=model,
|
||||
contents=contents, # type: ignore
|
||||
config=setup_result.generate_content_config_dict,
|
||||
stream=False,
|
||||
tools=tools,
|
||||
_is_async=_is_async,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
**kwargs,
|
||||
|
|
@ -342,7 +339,6 @@ def generate_content(
|
|||
timeout=timeout or request_timeout,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
stream=False,
|
||||
litellm_metadata=kwargs.get("litellm_metadata", {}),
|
||||
)
|
||||
|
||||
|
|
@ -391,15 +387,12 @@ async def agenerate_content_stream(
|
|||
|
||||
# Setup the call
|
||||
setup_result = GenerateContentHelper.setup_generate_content_call(
|
||||
**{
|
||||
"model": model,
|
||||
"contents": contents,
|
||||
"config": config,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"stream": True,
|
||||
"tools": tools,
|
||||
**kwargs,
|
||||
}
|
||||
model=model,
|
||||
contents=contents,
|
||||
config=config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
tools=tools,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Check if we should use the adapter (when provider config is None)
|
||||
|
|
@ -411,7 +404,7 @@ async def agenerate_content_stream(
|
|||
contents=contents, # type: ignore
|
||||
config=setup_result.generate_content_config_dict,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
stream=True,
|
||||
tools=tools,
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
|
|
@ -479,7 +472,6 @@ def generate_content_stream(
|
|||
contents=contents,
|
||||
config=config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
stream=True,
|
||||
tools=tools,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -491,7 +483,6 @@ def generate_content_stream(
|
|||
model=model,
|
||||
contents=contents, # type: ignore
|
||||
config=setup_result.generate_content_config_dict,
|
||||
stream=True,
|
||||
_is_async=_is_async,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
**kwargs,
|
||||
|
|
|
|||
|
|
@ -5139,6 +5139,26 @@ async def aadapter_completion(
|
|||
except Exception as e:
|
||||
raise e
|
||||
|
||||
async def aadapter_generate_content(
|
||||
**kwargs,
|
||||
) -> Union[ModelResponse, CustomStreamWrapper]:
|
||||
from litellm.google_genai.adapters.handler import (
|
||||
GenerateContentToCompletionHandler,
|
||||
)
|
||||
|
||||
custom_llm_provider_params = adapter.translate_generate_content_to_completion(
|
||||
model=model, contents=contents, config=config, **kwargs
|
||||
)
|
||||
|
||||
custom_llm_provider_params["stream"] = stream
|
||||
|
||||
|
||||
if stream:
|
||||
return adapter.translate_completion_output_params_streaming(
|
||||
completion_stream=response
|
||||
)
|
||||
return await handler.async_generate_content_handler(**kwargs, _is_async=True)
|
||||
|
||||
|
||||
def adapter_completion(
|
||||
*, adapter_id: str, **kwargs
|
||||
|
|
|
|||
|
|
@ -379,6 +379,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
user_api_base: Optional[str] = None,
|
||||
version: Optional[str] = None,
|
||||
is_streaming_request: Optional[bool] = False,
|
||||
contents: Optional[list] = None, # Add contents parameter
|
||||
) -> Any:
|
||||
"""
|
||||
Common request processing logic for both chat completions and responses API endpoints
|
||||
|
|
@ -417,6 +418,10 @@ class ProxyBaseLLMRequestProcessing:
|
|||
)
|
||||
)
|
||||
|
||||
# Pass contents if provided
|
||||
if contents:
|
||||
self.data["contents"] = contents
|
||||
|
||||
### ROUTE THE REQUEST ###
|
||||
# Do not change this - it should be a constant time fetch - ALWAYS
|
||||
llm_call = await route_request(
|
||||
|
|
|
|||
|
|
@ -1,8 +1,13 @@
|
|||
from fastapi import APIRouter, Depends, Request, Response
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
create_streaming_response,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.types.llms.vertex_ai import TokenCountDetailsResponse
|
||||
|
||||
router = APIRouter(
|
||||
|
|
@ -18,71 +23,17 @@ async def google_generate_content(
|
|||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Not Implemented, this is a placeholder for the google genai generateContent endpoint.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
_read_request_body,
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
select_data_generator,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
data = await _read_request_body(request=request)
|
||||
if "model" not in data:
|
||||
data["model"] = model_name
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="agenerate_content",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
data["stream"] = False
|
||||
|
||||
# call router
|
||||
response = await llm_router.agenerate_content(**data)
|
||||
return response
|
||||
|
||||
class GoogleAIStudioDataGenerator:
|
||||
"""
|
||||
Ensures SSE data generator is used for Google AI Studio streaming responses
|
||||
|
||||
Thin wrapper around ProxyBaseLLMRequestProcessing.async_sse_data_generator
|
||||
"""
|
||||
@staticmethod
|
||||
def _select_data_generator(response, user_api_key_dict, request_data):
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
return ProxyBaseLLMRequestProcessing.async_sse_data_generator(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
@router.post("/v1beta/models/{model_name}:streamGenerateContent", dependencies=[Depends(user_api_key_auth)])
|
||||
@router.post("/models/{model_name}:streamGenerateContent", dependencies=[Depends(user_api_key_auth)])
|
||||
|
|
@ -92,58 +43,22 @@ async def google_stream_generate_content(
|
|||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Not Implemented, this is a placeholder for the google genai streamGenerateContent endpoint.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
_read_request_body,
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
data = await _read_request_body(request=request)
|
||||
|
||||
if "model" not in data:
|
||||
data["model"] = model_name
|
||||
|
||||
data["stream"] = True # enforce streaming for this endpoint
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="agenerate_content_stream",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=GoogleAIStudioDataGenerator._select_data_generator,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
is_streaming_request=True,
|
||||
)
|
||||
except Exception as e:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
# call router
|
||||
response = await llm_router.agenerate_content(**data)
|
||||
|
||||
# Check if response is an async iterator (streaming response)
|
||||
if hasattr(response, "__aiter__"):
|
||||
return StreamingResponse(response, media_type="text/event-stream")
|
||||
return response
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
@ -171,13 +86,13 @@ async def google_count_tokens(request: Request, model_name: str):
|
|||
}
|
||||
```
|
||||
"""
|
||||
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.proxy_server import token_counter as internal_token_counter
|
||||
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
|
||||
|
||||
data = await _read_request_body(request=request)
|
||||
contents = data.get("contents", [])
|
||||
#Create TokenCountRequest for the internal endpoint
|
||||
# Create TokenCountRequest for the internal endpoint
|
||||
from litellm.proxy._types import TokenCountRequest
|
||||
|
||||
# Translate contents to openai format messages using the adapter
|
||||
|
|
|
|||
|
|
@ -562,15 +562,6 @@ class Router:
|
|||
)
|
||||
else:
|
||||
litellm.failure_callback = [self.deployment_callback_on_failure]
|
||||
verbose_router_logger.debug(
|
||||
f"Intialized router with Routing strategy: {self.routing_strategy}\n\n"
|
||||
f"Routing enable_pre_call_checks: {self.enable_pre_call_checks}\n\n"
|
||||
f"Routing fallbacks: {self.fallbacks}\n\n"
|
||||
f"Routing content fallbacks: {self.content_policy_fallbacks}\n\n"
|
||||
f"Routing context window fallbacks: {self.context_window_fallbacks}\n\n"
|
||||
f"Router Redis Caching={self.cache.redis_cache}\n"
|
||||
)
|
||||
self.service_logger_obj = ServiceLogging()
|
||||
self.routing_strategy_args = routing_strategy_args
|
||||
self.provider_budget_config = provider_budget_config
|
||||
self.router_budget_logger: Optional[RouterBudgetLimiting] = None
|
||||
|
|
@ -774,6 +765,14 @@ class Router:
|
|||
self.aanthropic_messages = self.factory_function(
|
||||
litellm.anthropic_messages, call_type="anthropic_messages"
|
||||
)
|
||||
self.agenerate_content = self.factory_function(
|
||||
litellm.agenerate_content, call_type="agenerate_content"
|
||||
)
|
||||
|
||||
self.aadapter_generate_content = self.factory_function(
|
||||
litellm.aadapter_generate_content, call_type="aadapter_generate_content"
|
||||
)
|
||||
|
||||
self.aresponses = self.factory_function(
|
||||
litellm.aresponses, call_type="aresponses"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue