chore(typing): clear basedpyright Any errors in google genai adapter and response polling

This commit is contained in:
Devin 2026-08-11 13:50:05 +00:00
parent b0fac57fe4
commit 5aba2e33d3
3 changed files with 614 additions and 557 deletions

View file

@ -1,60 +1,195 @@
import json
from collections.abc import AsyncIterator, Iterator
from typing import Any, Final, cast
from collections.abc import AsyncIterable, AsyncIterator, Iterable, Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import Any, Final
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError
from litellm import verbose_logger
from litellm.litellm_core_utils.json_validation_rule import normalize_tool_schema
from litellm.litellm_core_utils.json_validation_rule import normalize_json_schema_types
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionAssistantMessage,
ChatCompletionAssistantToolCall,
ChatCompletionImageObject,
ChatCompletionImageUrlObject,
ChatCompletionRequest,
ChatCompletionSystemMessage,
ChatCompletionTextObject,
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolChoiceStringValues,
ChatCompletionToolChoiceValues,
ChatCompletionToolMessage,
ChatCompletionToolParam,
ChatCompletionToolParamFunctionChunk,
ChatCompletionUserMessage,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
AdapterCompletionStreamWrapper,
Choices,
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Delta,
Message,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
class _InlineData(BaseModel):
model_config = ConfigDict(extra="ignore")
mime_type: str = "image/jpeg"
data: str = ""
class _FunctionCall(BaseModel):
model_config = ConfigDict(extra="ignore")
name: str = "unknown"
args: JsonValue = Field(default_factory=dict)
class _FunctionResponse(BaseModel):
model_config = ConfigDict(extra="ignore")
name: str = "unknown"
response: JsonValue = Field(default_factory=dict)
class _Part(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
text: str | None = None
inline_data: _InlineData | None = Field(default=None, alias="inlineData")
function_call: _FunctionCall | None = Field(default=None, alias="functionCall")
function_response: _FunctionResponse | None = Field(default=None, alias="functionResponse")
class _Content(BaseModel):
model_config = ConfigDict(extra="ignore")
role: str = "user"
parts: tuple[_Part | str, ...] = ()
class _FunctionDeclaration(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
name: str = ""
description: str | None = None
parameters_json_schema: JsonValue = Field(default=None, alias="parametersJsonSchema")
class _Tool(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
function_declarations: tuple[_FunctionDeclaration, ...] = Field(default=(), alias="functionDeclarations")
class _FunctionCallingConfig(BaseModel):
model_config = ConfigDict(extra="ignore")
mode: str = "AUTO"
class _ToolConfig(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
function_calling_config: _FunctionCallingConfig = Field(
default_factory=_FunctionCallingConfig, alias="functionCallingConfig"
)
class _TokenCounts(BaseModel):
model_config = ConfigDict(extra="ignore", from_attributes=True)
prompt_tokens: int = 0
completion_tokens: int = 0
total_tokens: int = 0
class _UsageCarrier(BaseModel):
model_config = ConfigDict(extra="ignore", from_attributes=True)
usage: _TokenCounts | None = None
class _GenerateContentConfig(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
temperature: float | None = None
max_output_tokens: int | None = Field(default=None, alias="maxOutputTokens")
top_p: float | None = Field(default=None, alias="topP")
stop_sequences: list[str] | None = Field(default=None, alias="stopSequences")
_CONTENT_ADAPTER: Final = TypeAdapter(_Content)
_TOOLS_ADAPTER: Final = TypeAdapter(tuple[_Tool, ...])
_TOOL_CONFIG_ADAPTER: Final = TypeAdapter(_ToolConfig)
_CONFIG_ADAPTER: Final = TypeAdapter(_GenerateContentConfig)
_JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
_USAGE_ADAPTER: Final = TypeAdapter(_UsageCarrier)
_CompletionStream = Iterable[object] | AsyncIterable[object] | ModelResponse
_TOOL_CHOICE_BY_MODE: Final[Mapping[str, ChatCompletionToolChoiceStringValues]] = MappingProxyType(
{"AUTO": "auto", "ANY": "required", "NONE": "none"}
)
_EMPTY_USAGE: Final[Mapping[str, int]] = MappingProxyType(
{"promptTokenCount": 0, "candidatesTokenCount": 0, "totalTokenCount": 0}
)
def _litellm_param_values(litellm_params: GenericLiteLLMParams) -> Mapping[str, object]:
return litellm_params.model_dump(exclude_none=True)
def _text_of(parts: Sequence[Mapping[str, object]]) -> str:
texts: Final = (part.get("text") for part in parts)
return "".join(text for text in texts if isinstance(text, str))
def _usage_metadata(response: ModelResponse | ModelResponseStream) -> dict[str, int]:
usage: Final = _USAGE_ADAPTER.validate_python(response).usage
if usage is None:
return dict(_EMPTY_USAGE)
return {
"promptTokenCount": usage.prompt_tokens,
"candidatesTokenCount": usage.completion_tokens,
"totalTokenCount": usage.total_tokens,
}
class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
"""
Wrapper for streaming Google GenAI generate_content responses.
Transforms OpenAI streaming chunks to Google GenAI format.
"""
completion_stream: _CompletionStream
sent_first_chunk: bool = False
# State tracking for accumulating partial tool calls
accumulated_tool_calls: dict[str, dict[str, Any]]
accumulated_tool_calls: dict[int, dict[str, str]]
def __init__(self, completion_stream: Any):
def __init__(self, completion_stream: _CompletionStream) -> None:
self.sent_first_chunk = False
self.accumulated_tool_calls = {}
self._returned_response = False
super().__init__(completion_stream)
def __next__(self):
def __next__(self) -> dict[str, object]:
try:
if not hasattr(self.completion_stream, "__iter__"):
if self._returned_response:
stream: Final = self.completion_stream
if not isinstance(stream, Iterable):
if self._returned_response or not isinstance(stream, ModelResponse):
raise StopIteration
self._returned_response = True
return GoogleGenAIAdapter().translate_completion_to_generate_content(self.completion_stream)
return GoogleGenAIAdapter().translate_completion_to_generate_content(stream)
for chunk in self.completion_stream:
for chunk in stream:
if chunk == "None" or chunk is None:
continue
if not isinstance(chunk, (ModelResponse, ModelResponseStream)):
continue
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(chunk, self)
if transformed_chunk:
@ -66,43 +201,42 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
except Exception:
raise StopIteration
async def __anext__(self):
async def __anext__(self) -> dict[str, object]:
try:
if not hasattr(self.completion_stream, "__aiter__"):
if self._returned_response:
stream: Final = self.completion_stream
if not isinstance(stream, AsyncIterable):
if self._returned_response or not isinstance(stream, ModelResponse):
raise StopAsyncIteration
self._returned_response = True
return GoogleGenAIAdapter().translate_completion_to_generate_content(self.completion_stream)
return GoogleGenAIAdapter().translate_completion_to_generate_content(stream)
async for chunk in self.completion_stream:
async for chunk in stream:
if chunk == "None" or chunk is None:
continue
if not isinstance(chunk, (ModelResponse, ModelResponseStream)):
continue
transformed_chunk = GoogleGenAIAdapter().translate_streaming_completion_to_generate_content(chunk, self)
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: Final = []
parts: Final[list[dict[str, object]]] = []
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 = {
parsed_args = _JSON_ADAPTER.validate_json(tool_call_data["arguments"] or "{}")
function_call_part: dict[str, object] = {
"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.
except ValidationError:
verbose_logger.warning(
"Could not parse tool call arguments at end of stream for index %s. Name: %s. Partial args: %s",
tool_call_index,
@ -110,7 +244,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
tool_call_data["arguments"],
)
if parts:
final_chunk: Final = {
final_chunk: Final[dict[str, object]] = {
"candidates": [
{
"content": {"parts": parts, "role": "model"},
@ -122,7 +256,6 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
}
return final_chunk
finally:
# Ensure the accumulator is always cleared to prevent memory leaks
self.accumulated_tool_calls.clear()
raise StopAsyncIteration
except StopAsyncIteration:
@ -134,39 +267,40 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
"""
Convert Google GenAI streaming chunks to Server-Sent Events format.
"""
for chunk in self.completion_stream:
if isinstance(chunk, dict):
payload = f"data: {json.dumps(chunk)}\n\n"
yield payload.encode()
else:
stream: Final = self.completion_stream
if not isinstance(stream, Iterable):
return
for chunk in stream:
if isinstance(chunk, bytes):
yield chunk
elif isinstance(chunk, str):
yield chunk.encode()
else:
yield f"data: {json.dumps(chunk)}\n\n".encode()
async def async_google_genai_sse_wrapper(self) -> AsyncIterator[bytes]:
"""
Async version of google_genai_sse_wrapper.
"""
from litellm.types.utils import ModelResponseStream
stream: Final = self.completion_stream
if not isinstance(stream, AsyncIterable):
return
async for chunk in self.completion_stream:
if isinstance(chunk, dict):
payload = f"data: {json.dumps(chunk)}\n\n"
yield payload.encode()
async for chunk in stream:
if isinstance(chunk, Mapping):
yield f"data: {json.dumps(chunk)}\n\n".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:
# For empty chunks, continue to next iteration
continue
if transformed_chunk is not None:
yield f"data: {json.dumps(transformed_chunk)}\n\n".encode()
elif isinstance(chunk, str):
yield chunk.encode()
elif isinstance(chunk, bytes):
yield chunk
else:
# For other chunk types, yield them directly
if hasattr(chunk, "encode"):
yield chunk.encode()
else:
yield str(chunk).encode()
yield str(chunk).encode()
class GoogleGenAIAdapter:
@ -178,10 +312,10 @@ class GoogleGenAIAdapter:
def translate_generate_content_to_completion(
self,
model: str,
contents: list[dict[str, Any]] | dict[str, Any],
config: dict[str, Any] | None = None,
contents: Sequence[Mapping[str, object]] | Mapping[str, object],
config: Mapping[str, object] | None = None,
litellm_params: GenericLiteLLMParams | None = None,
**kwargs,
**kwargs: object,
) -> dict[str, Any]:
"""
Transform generate_content request to litellm completion format
@ -196,21 +330,14 @@ class GoogleGenAIAdapter:
Dict in OpenAI format
"""
# Extract top-level fields from kwargs
system_instruction: Final = kwargs.get("systemInstruction") or kwargs.get("system_instruction")
tools: Final = kwargs.get("tools")
tool_config: Final = kwargs.get("toolConfig") or kwargs.get("tool_config")
# Normalize contents to list format
if isinstance(contents, dict):
contents_list = [contents]
else:
contents_list = contents
contents_list: Final = [contents] if isinstance(contents, Mapping) else contents
# Transform contents to OpenAI messages format
messages: Final = self._transform_contents_to_messages(contents_list, system_instruction=system_instruction)
# Create base request as dict (which is compatible with ChatCompletionRequest)
completion_request: Final[ChatCompletionRequest] = {
"model": model,
"messages": messages,
@ -228,32 +355,23 @@ class GoogleGenAIAdapter:
# - tool_choice
#########################################################
# Add config parameters if provided
if config:
# Map common Google GenAI config parameters to OpenAI equivalents
if "temperature" in config:
completion_request["temperature"] = config["temperature"]
if "maxOutputTokens" in config:
completion_request["max_tokens"] = config["maxOutputTokens"]
if "topP" in config:
completion_request["top_p"] = config["topP"]
if "topK" in config:
# OpenAI doesn't have direct topK, but we can pass it as extra
pass
if "stopSequences" in config:
completion_request["stop"] = config["stopSequences"]
generate_content_config: Final = _CONFIG_ADAPTER.validate_python(config)
if generate_content_config.temperature is not None:
completion_request["temperature"] = generate_content_config.temperature
if generate_content_config.max_output_tokens is not None:
completion_request["max_tokens"] = generate_content_config.max_output_tokens
if generate_content_config.top_p is not None:
completion_request["top_p"] = generate_content_config.top_p
if generate_content_config.stop_sequences is not None:
completion_request["stop"] = generate_content_config.stop_sequences
# Handle tools transformation
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: Final = self._transform_google_genai_tools_to_openai(tools)
openai_tools: Final = self._transform_google_genai_tools_to_openai(tools)
if openai_tools:
completion_request["tools"] = openai_tools
if openai_tools:
completion_request["tools"] = openai_tools
# Handle tool_config (tool choice)
if tool_config:
tool_choice: Final = self._transform_google_genai_tool_config_to_openai(tool_config)
if tool_choice:
@ -262,40 +380,39 @@ class GoogleGenAIAdapter:
#########################################################
# forward any litellm specific params
#########################################################
completion_request_dict = dict(completion_request)
if litellm_params:
completion_request_dict = self._add_generic_litellm_params_to_request(
completion_request_dict=completion_request_dict,
litellm_params=litellm_params,
)
completion_request_dict: Final[dict[str, object]] = dict(completion_request)
if litellm_params is None:
return completion_request_dict
return completion_request_dict
return self._add_generic_litellm_params_to_request(
completion_request_dict=completion_request_dict,
litellm_params=litellm_params,
)
def _add_generic_litellm_params_to_request(
self,
completion_request_dict: dict[str, Any],
completion_request_dict: dict[str, object],
litellm_params: GenericLiteLLMParams | None = None,
) -> dict:
) -> dict[str, object]:
"""Add generic litellm params to request. e.g add api_base, api_key, api_version, etc.
Args:
completion_request_dict: Dict[str, Any]
completion_request_dict: dict[str, object]
litellm_params: GenericLiteLLMParams
Returns:
Dict[str, Any]
dict[str, object]
"""
allowed_fields: Final = GenericLiteLLMParams.model_fields.keys()
if litellm_params:
litellm_dict: Final = litellm_params.model_dump(exclude_none=True)
for key, value in litellm_dict.items():
for key, value in _litellm_param_values(litellm_params).items():
if key in allowed_fields:
completion_request_dict[key] = value
return completion_request_dict
def translate_completion_output_params_streaming(
self,
completion_stream: Any,
completion_stream: _CompletionStream,
) -> AsyncIterator[bytes] | None:
"""Transform streaming completion output to Google GenAI format"""
google_genai_wrapper: Final = GoogleGenAIStreamWrapper(completion_stream=completion_stream)
@ -304,164 +421,130 @@ class GoogleGenAIAdapter:
def _transform_google_genai_tools_to_openai(
self,
tools: list[dict[str, Any]],
tools: object,
) -> list[ChatCompletionToolParam]:
"""Transform Google GenAI tools to OpenAI tools format"""
openai_tools: Final[list[dict[str, Any]]] = []
openai_tools: Final[list[ChatCompletionToolParam]] = []
for tool in tools:
if "functionDeclarations" in tool:
for func_decl in tool["functionDeclarations"]:
function_chunk: dict[str, Any] = {
"name": func_decl.get("name", ""),
}
for tool in _TOOLS_ADAPTER.validate_python(tools):
for func_decl in tool.function_declarations:
function_chunk: ChatCompletionToolParamFunctionChunk = {"name": func_decl.name}
if "description" in func_decl:
function_chunk["description"] = func_decl["description"]
if "parametersJsonSchema" in func_decl:
function_chunk["parameters"] = func_decl["parametersJsonSchema"]
if func_decl.description is not None:
function_chunk["description"] = func_decl.description
openai_tool = {"type": "function", "function": function_chunk}
openai_tools.append(openai_tool)
normalized_schema = normalize_json_schema_types(func_decl.parameters_json_schema)
if isinstance(normalized_schema, dict):
function_chunk["parameters"] = normalized_schema
# normalize the tool schemas
normalized_tools: Final = [normalize_tool_schema(tool) for tool in openai_tools]
openai_tools.append(ChatCompletionToolParam(type="function", function=function_chunk))
return cast(list[ChatCompletionToolParam], normalized_tools)
return openai_tools
def _transform_google_genai_tool_config_to_openai(
self,
tool_config: dict[str, Any],
tool_config: object,
) -> ChatCompletionToolChoiceValues | None:
"""Transform Google GenAI tool_config to OpenAI tool_choice"""
function_calling_config: Final = tool_config.get("functionCallingConfig", {})
mode: Final = function_calling_config.get("mode", "AUTO")
mode_mapping: Final = {"AUTO": "auto", "ANY": "required", "NONE": "none"}
tool_choice: Final = mode_mapping.get(mode, "auto")
return cast(ChatCompletionToolChoiceValues, tool_choice)
mode: Final = _TOOL_CONFIG_ADAPTER.validate_python(tool_config).function_calling_config.mode
return _TOOL_CHOICE_BY_MODE.get(mode, "auto")
def _transform_contents_to_messages(
self,
contents: list[dict[str, Any]],
system_instruction: dict[str, Any] | None = None,
contents: Sequence[Mapping[str, object]],
system_instruction: object = None,
) -> list[AllMessageValues]:
"""Transform Google GenAI contents to OpenAI messages format"""
messages: Final[list[AllMessageValues]] = []
# Handle system instruction
if system_instruction:
system_parts: Final = system_instruction.get("parts", [])
if system_parts and "text" in system_parts[0]:
messages.append(ChatCompletionSystemMessage(role="system", content=system_parts[0]["text"]))
if system_instruction is not None:
system_parts: Final = _CONTENT_ADAPTER.validate_python(system_instruction).parts
first_system_part: Final = system_parts[0] if system_parts else None
if isinstance(first_system_part, _Part) and first_system_part.text is not None:
messages.append(ChatCompletionSystemMessage(role="system", content=first_system_part.text))
for content in contents:
role = content.get("role", "user")
parts = content.get("parts", [])
for raw_content in contents:
content = _CONTENT_ADAPTER.validate_python(raw_content)
if role == "user":
# Handle user messages with potential function responses
if content.role == "user":
content_parts: list[ChatCompletionTextObject | ChatCompletionImageObject] = []
tool_messages: list[ChatCompletionToolMessage] = []
for part in parts:
if isinstance(part, dict):
if "text" in part:
content_parts.append(
cast(
ChatCompletionTextObject,
{"type": "text", "text": part["text"]},
)
for part in content.parts:
if isinstance(part, str):
content_parts.append(ChatCompletionTextObject(type="text", text=part))
elif part.text is not None:
content_parts.append(ChatCompletionTextObject(type="text", text=part.text))
elif part.inline_data is not None:
content_parts.append(
ChatCompletionImageObject(
type="image_url",
image_url=ChatCompletionImageUrlObject(
url=f"data:{part.inline_data.mime_type};base64,{part.inline_data.data}"
),
)
elif "inline_data" in part:
# Handle Base64 image data
inline_data = part["inline_data"]
mime_type = inline_data.get("mime_type", "image/jpeg")
data = inline_data.get("data", "")
content_parts.append(
cast(
ChatCompletionImageObject,
{
"type": "image_url",
"image_url": {"url": f"data:{mime_type};base64,{data}"},
},
)
)
elif "functionResponse" in part:
# Transform function response to tool message
func_response = part["functionResponse"]
tool_message = ChatCompletionToolMessage(
)
elif part.function_response is not None:
tool_messages.append(
ChatCompletionToolMessage(
role="tool",
tool_call_id=f"call_{func_response.get('name', 'unknown')}",
content=json.dumps(func_response.get("response", {})),
tool_call_id=f"call_{part.function_response.name}",
content=json.dumps(part.function_response.response),
)
tool_messages.append(tool_message)
elif isinstance(part, str):
content_parts.append(cast(ChatCompletionTextObject, {"type": "text", "text": part}))
)
# Add user message if there's content
if content_parts:
# If only one text part, use simple string format for backward compatibility
if (
len(content_parts) == 1
and isinstance(content_parts[0], dict)
and content_parts[0].get("type") == "text"
):
text_part = cast(ChatCompletionTextObject, content_parts[0])
messages.append(ChatCompletionUserMessage(role="user", content=text_part["text"]))
first_content_part = content_parts[0]
if len(content_parts) == 1 and first_content_part["type"] == "text":
messages.append(ChatCompletionUserMessage(role="user", content=first_content_part["text"]))
else:
# Use multimodal format (array of content parts)
messages.append(ChatCompletionUserMessage(role="user", content=content_parts))
# Add tool messages
messages.extend(tool_messages)
elif role == "model":
# Handle assistant messages with potential function calls
elif content.role == "model":
combined_text = ""
tool_calls: list[ChatCompletionAssistantToolCall] = []
for part in parts:
if isinstance(part, dict):
if "text" in part:
combined_text += part["text"]
elif "functionCall" in part:
# Transform function call to tool call
func_call = part["functionCall"]
tool_call = ChatCompletionAssistantToolCall(
id=f"call_{func_call.get('name', 'unknown')}",
for part in content.parts:
if isinstance(part, str):
combined_text += part
elif part.text is not None:
combined_text += part.text
elif part.function_call is not None:
tool_calls.append(
ChatCompletionAssistantToolCall(
id=f"call_{part.function_call.name}",
type="function",
function=ChatCompletionToolCallFunctionChunk(
name=func_call.get("name", ""),
arguments=json.dumps(func_call.get("args", {})),
name=part.function_call.name,
arguments=json.dumps(part.function_call.args),
),
)
tool_calls.append(tool_call)
elif isinstance(part, str):
combined_text += part
)
# Create assistant message
if tool_calls:
assistant_message = ChatCompletionAssistantMessage(
role="assistant",
content=combined_text if combined_text else None,
tool_calls=tool_calls,
messages.append(
ChatCompletionAssistantMessage(
role="assistant",
content=combined_text if combined_text else None,
tool_calls=tool_calls,
)
)
else:
assistant_message = ChatCompletionAssistantMessage(
role="assistant",
content=combined_text if combined_text else None,
messages.append(
ChatCompletionAssistantMessage(
role="assistant",
content=combined_text if combined_text else None,
)
)
messages.append(assistant_message)
return messages
def translate_completion_to_generate_content(
self,
response: ModelResponse,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Transform litellm completion response to Google GenAI generate_content format
@ -472,49 +555,28 @@ class GoogleGenAIAdapter:
Dict in Google GenAI generate_content response format
"""
# Extract the main response content
choice: Final = response.choices[0] if response.choices else None
if not choice:
raise ValueError("Invalid completion response: no choices found")
# Handle different choice types (Choices vs StreamingChoices)
if isinstance(choice, Choices):
if not choice.message:
raise ValueError("Invalid completion response: no message found in choice")
parts = self._transform_openai_message_to_google_genai_parts(choice.message)
else:
# Fallback for generic choice objects
message_content = getattr(choice, "message", {}).get("content", "") or getattr(choice, "delta", {}).get(
"content", ""
)
parts = [{"text": message_content}] if message_content else []
if not choice.message:
raise ValueError("Invalid completion response: no message found in choice")
# Create Google GenAI format response
generate_content_response: Final[dict[str, Any]] = {
parts: Final = self._transform_openai_message_to_google_genai_parts(choice.message)
generate_content_response: Final[dict[str, object]] = {
"candidates": [
{
"content": {"parts": parts, "role": "model"},
"finishReason": self._map_finish_reason(getattr(choice, "finish_reason", None)),
"finishReason": self._map_finish_reason(choice.finish_reason),
"index": 0,
"safetyRatings": [],
}
],
"usageMetadata": (
self._map_usage(getattr(response, "usage", None))
if hasattr(response, "usage") and getattr(response, "usage", None)
else {
"promptTokenCount": 0,
"candidatesTokenCount": 0,
"totalTokenCount": 0,
}
),
"usageMetadata": _usage_metadata(response),
}
# Add text field for convenience (common in Google GenAI responses)
text_content = ""
for part in parts:
if isinstance(part, dict) and "text" in part:
text_content += part["text"]
text_content: Final = _text_of(parts)
if text_content:
generate_content_response["text"] = text_content
@ -524,7 +586,7 @@ class GoogleGenAIAdapter:
self,
response: ModelResponse | ModelResponseStream,
wrapper: GoogleGenAIStreamWrapper,
) -> dict[str, Any] | None:
) -> dict[str, object] | None:
"""
Transform streaming litellm completion chunk to Google GenAI generate_content format
@ -536,31 +598,22 @@ class GoogleGenAIAdapter:
Dict in Google GenAI streaming generate_content response format
"""
# Extract the main response content from streaming chunk
choice: Final = response.choices[0] if response.choices else None
if not choice:
# Return empty chunk if no choices
return None
# Handle streaming choice
if isinstance(choice, StreamingChoices):
if choice.delta:
parts = self._transform_openai_delta_to_google_genai_parts_with_accumulation(choice.delta, wrapper)
else:
parts = []
finish_reason = getattr(choice, "finish_reason", None)
else:
# Fallback for generic choice objects
message_content: Final = getattr(choice, "delta", {}).get("content", "")
parts = [{"text": message_content}] if message_content else []
finish_reason = getattr(choice, "finish_reason", None)
parts: Final[list[dict[str, object]]] = (
self._transform_openai_delta_to_google_genai_parts_with_accumulation(choice.delta, wrapper)
if isinstance(choice, StreamingChoices) and choice.delta
else []
)
finish_reason: Final = choice.finish_reason
# Only create response chunk if we have parts or it's the final chunk
if not parts and not finish_reason:
return None
# Create Google GenAI streaming format response
streaming_chunk: Final[dict[str, Any]] = {
streaming_chunk: Final[dict[str, object]] = {
"candidates": [
{
"content": {"parts": parts, "role": "model"},
@ -571,24 +624,10 @@ class GoogleGenAIAdapter:
]
}
# Add usage metadata only in the final chunk (when finish_reason is present)
if finish_reason:
usage_metadata: Final = (
self._map_usage(getattr(response, "usage", None))
if hasattr(response, "usage") and getattr(response, "usage", None)
else {
"promptTokenCount": 0,
"candidatesTokenCount": 0,
"totalTokenCount": 0,
}
)
streaming_chunk["usageMetadata"] = usage_metadata
streaming_chunk["usageMetadata"] = _usage_metadata(response)
# Add text field for convenience (common in Google GenAI responses)
text_content = ""
for part in parts:
if isinstance(part, dict) and "text" in part:
text_content += part["text"]
text_content: Final = _text_of(parts)
if text_content:
streaming_chunk["text"] = text_content
@ -596,106 +635,74 @@ class GoogleGenAIAdapter:
def _transform_openai_message_to_google_genai_parts(
self,
message: Any,
) -> list[dict[str, Any]]:
message: Message,
) -> list[dict[str, object]]:
"""Transform OpenAI message to Google GenAI parts format"""
parts: Final[list[dict[str, Any]]] = []
parts: Final[list[dict[str, object]]] = []
# Add text content if present
if hasattr(message, "content") and message.content:
if message.content:
parts.append({"text": message.content})
# Add tool calls if present
if hasattr(message, "tool_calls") and message.tool_calls:
for tool_call in message.tool_calls:
if hasattr(tool_call, "function") and tool_call.function:
try:
args = json.loads(tool_call.function.arguments) if tool_call.function.arguments else {}
except json.JSONDecodeError:
args = {}
for tool_call in message.tool_calls or []:
if not isinstance(tool_call, ChatCompletionMessageToolCall):
continue
function_call_part = {
"functionCall": {
"name": tool_call.function.name or "undefined_tool_name",
"args": args,
}
try:
args = _JSON_ADAPTER.validate_json(tool_call.function.arguments or "{}")
except ValidationError:
args = {}
parts.append(
{
"functionCall": {
"name": tool_call.function.name or "undefined_tool_name",
"args": args,
}
parts.append(function_call_part)
}
)
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]]:
self, delta: Delta, wrapper: GoogleGenAIStreamWrapper
) -> list[dict[str, object]]:
"""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: Final[list[dict[str, object]]] = []
parts: Final[list[dict[str, Any]]] = []
if hasattr(delta, "content") and delta.content:
if delta.content:
parts.append({"text": delta.content})
# 2. Ensure tool_calls is iterable
tool_calls: Final = delta.tool_calls or []
for tool_call in tool_calls:
if not hasattr(tool_call, "function"):
for tool_call in delta.tool_calls or []:
if not isinstance(tool_call, ChatCompletionDeltaToolCall):
continue
# 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
tool_call_index = tool_call.index
accumulated = wrapper.accumulated_tool_calls.setdefault(tool_call_index, {"name": "", "arguments": ""})
# 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": "",
}
function_name = tool_call.function.name
args_chunk = tool_call.function.arguments
# Accumulate name and arguments
function_name = getattr(tool_call.function, "name", None)
args_chunk = getattr(tool_call.function, "arguments", None)
# Optimization: Skip chunks that have no new data
if not function_name and not args_chunk:
verbose_logger.debug("Skipping empty tool call chunk for index: %s", tool_call_index)
continue
if function_name:
wrapper.accumulated_tool_calls[tool_call_index]["name"] = function_name
accumulated["name"] = function_name
if args_chunk:
wrapper.accumulated_tool_calls[tool_call_index]["arguments"] += args_chunk
accumulated["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"]
accumulated_name = accumulated["name"]
# 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)
parsed_args = _JSON_ADAPTER.validate_json(accumulated["arguments"])
except ValidationError:
continue
# 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
if accumulated_name:
parts.append({"functionCall": {"name": accumulated_name, "args": parsed_args}})
del wrapper.accumulated_tool_calls[tool_call_index]
return parts
@ -713,11 +720,3 @@ class GoogleGenAIAdapter:
}
return mapping.get(finish_reason, "STOP")
def _map_usage(self, usage: Any) -> dict[str, int]:
"""Map OpenAI usage to Google GenAI usage format"""
return {
"promptTokenCount": getattr(usage, "prompt_tokens", 0) or 0,
"candidatesTokenCount": getattr(usage, "completion_tokens", 0) or 0,
"totalTokenCount": getattr(usage, "total_tokens", 0) or 0,
}

View file

@ -9,37 +9,191 @@ https://platform.openai.com/docs/api-reference/responses-streaming
"""
import asyncio
import json
from typing import Any, Final, cast
from collections.abc import AsyncIterable, Callable
from typing import TYPE_CHECKING, Final, Literal, Protocol, runtime_checkable
from fastapi import Request, Response
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.response_polling.polling_handler import ResponsePollingHandler
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.types.llms.openai import ResponsesAPIStatus
if TYPE_CHECKING:
from litellm.proxy.proxy_server import ProxyConfig
@runtime_checkable
class StreamingBody(Protocol):
"""Any response exposing a server sent event body, such as a fastapi StreamingResponse"""
body_iterator: AsyncIterable[str | bytes | memoryview]
class ResponsesRequestProcessor(Protocol):
"""Proxy request processor able to run a Responses API request"""
async def base_process_llm_request(
self,
*,
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
route_type: Literal["aresponses"],
proxy_logging_obj: ProxyLogging,
llm_router: Router | None,
general_settings: dict[str, object],
proxy_config: "ProxyConfig",
select_data_generator: Callable[..., object] | None,
model: str | None,
user_model: str | None,
user_temperature: float | None,
user_request_timeout: float | None,
user_max_tokens: int | None,
user_api_base: str | None,
version: str | None,
skip_pre_call_logic: bool,
) -> object: ...
class StreamedResponseSnapshot(BaseModel):
"""Terminal `response` payload of a Responses API stream"""
model_config = ConfigDict(extra="ignore")
status: ResponsesAPIStatus | None = None
error: dict[str, JsonValue] | None = None
incomplete_details: dict[str, JsonValue] | None = None
usage: dict[str, JsonValue] | None = None
reasoning: dict[str, JsonValue] | None = None
tool_choice: JsonValue = None
tools: list[JsonValue] | None = None
model: str | None = None
instructions: str | None = None
temperature: float | None = None
top_p: float | None = None
max_output_tokens: int | None = None
previous_response_id: str | None = None
text: dict[str, JsonValue] | None = None
truncation: str | None = None
parallel_tool_calls: bool | None = None
user: str | None = None
store: bool | None = None
output: list[dict[str, JsonValue]] | None = None
class StreamedEvent(BaseModel):
"""Server sent event of a Responses API stream"""
model_config = ConfigDict(extra="ignore")
type: str = ""
item: dict[str, JsonValue] | None = None
item_id: str | None = None
part: dict[str, JsonValue] | None = None
content_index: int = 0
delta: str = ""
response: StreamedResponseSnapshot | None = None
_EVENT_ADAPTER: Final = TypeAdapter(StreamedEvent)
_EVENT_TO_STATUS: Final[dict[str, ResponsesAPIStatus]] = {
"response.completed": "completed",
"response.failed": "failed",
"response.incomplete": "incomplete",
"response.cancelled": "cancelled",
}
_UPDATE_INTERVAL: Final = 0.150
def _request_processor(data: dict[str, object]) -> ResponsesRequestProcessor:
return ProxyBaseLLMRequestProcessing(data=data)
def _item_content(item: dict[str, JsonValue]) -> list[JsonValue] | None:
content: Final = item.get("content")
return content if isinstance(content, list) else None
def _apply_event(
event: StreamedEvent,
output_items: dict[str, dict[str, JsonValue]],
accumulated_text: dict[tuple[str, int], str],
) -> bool:
"""Apply one streaming event to the accumulated output, reporting whether it changed"""
if event.type == "response.output_item.added" or event.type == "response.output_item.done":
item: Final = event.item or {}
item_id: Final = item.get("id")
if isinstance(item_id, str) and item_id:
output_items[item_id] = item
return True
return False
tracked_item: Final = output_items.get(event.item_id or "")
if tracked_item is None:
return False
content: Final = _item_content(tracked_item)
if event.type == "response.content_part.added":
# Update the output item with new content
part: Final = event.part or {}
if content is None:
tracked_item["content"] = [part]
else:
content.append(part)
return True
if event.type == "response.output_text.delta":
if event.item_id is None:
return False
# Accumulate text delta
# https://platform.openai.com/docs/api-reference/responses-streaming/response-text-delta
key: Final = (event.item_id, event.content_index)
accumulated_text[key] = accumulated_text.get(key, "") + event.delta
if content is not None and event.content_index < len(content):
# Update existing content part with accumulated text
content_part: Final = content[event.content_index]
if isinstance(content_part, dict):
content_part["text"] = accumulated_text[key]
return True
if event.type == "response.content_part.done":
# Update with final content from the event
if content is not None and event.content_index < len(content):
content[event.content_index] = event.part or {}
return True
return False
async def background_streaming_task(
polling_id: str,
data: dict,
data: dict[str, object],
polling_handler: ResponsePollingHandler,
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
general_settings: dict,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_model,
user_temperature,
user_request_timeout,
user_max_tokens,
user_api_base,
version,
):
general_settings: dict[str, object],
llm_router: Router | None,
proxy_config: "ProxyConfig",
proxy_logging_obj: ProxyLogging,
select_data_generator: Callable[..., object] | None,
user_model: str | None,
user_temperature: float | None,
user_request_timeout: float | None,
user_max_tokens: int | None,
user_api_base: str | None,
version: str | None,
) -> None:
"""
Background task to stream response and update cache
@ -64,12 +218,12 @@ async def background_streaming_task(
data.pop("background", None)
# Create processor
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
processor: Final = _request_processor(data)
# Make streaming request.
# Pre-call checks (rate limits, guardrails, budget) were already run
# before polling ID creation, so skip them here to avoid double-counting.
response: Final = await processor.base_process_llm_request(
response: Final[object] = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -91,49 +245,26 @@ async def background_streaming_task(
# Process streaming response following OpenAI events format
# https://platform.openai.com/docs/api-reference/responses-streaming
output_items: Final[dict[str, dict[str, Any]]] = {} # Track output items by ID
accumulated_text: Final = {} # Track accumulated text deltas by (item_id, content_index)
output_items: Final[dict[str, dict[str, JsonValue]]] = {} # Track output items by ID
accumulated_text: Final[dict[tuple[str, int], str]] = {} # Text deltas by (item_id, content_index)
# ResponsesAPIResponse fields to extract from response.completed
usage_data = None
reasoning_data = None
tool_choice_data = None
tools_data = None
model_data = None
instructions_data = None
temperature_data = None
top_p_data = None
max_output_tokens_data = None
previous_response_id_data = None
text_data = None
truncation_data = None
parallel_tool_calls_data = None
user_data = None
store_data = None
incomplete_details_data = None
# ResponsesAPIResponse fields extracted from the terminal event
terminal_response: StreamedResponseSnapshot | None = None
state_dirty = False # Track if state needs to be synced
last_update_time = asyncio.get_event_loop().time()
UPDATE_INTERVAL: Final = 0.150 # 150ms batching interval
# Track the terminal event from the stream (may not be "completed")
terminal_status: ResponsesAPIStatus | None = (
None # Will be set by response.completed/failed/incomplete/cancelled
)
terminal_error = None
_event_to_status: Final = {
"response.completed": "completed",
"response.failed": "failed",
"response.incomplete": "incomplete",
"response.cancelled": "cancelled",
}
async def flush_state_if_needed(force: bool = False) -> None:
"""Flush accumulated state to Redis if interval elapsed or forced"""
nonlocal state_dirty, last_update_time
current_time: Final = asyncio.get_event_loop().time()
if state_dirty and (force or (current_time - last_update_time) >= UPDATE_INTERVAL):
if state_dirty and (force or (current_time - last_update_time) >= _UPDATE_INTERVAL):
# Convert output_items dict to list for update
output_list: Final = list(output_items.values())
await polling_handler.update_state(
@ -144,17 +275,16 @@ async def background_streaming_task(
last_update_time = current_time
# Handle StreamingResponse
if not hasattr(response, "body_iterator"):
if not isinstance(response, StreamingBody):
verbose_proxy_logger.warning(
"background_streaming_task: response for %s has no body_iterator; this may indicate a misconfiguration or provider error",
polling_id,
)
if hasattr(response, "body_iterator"):
async for chunk in response.body_iterator:
if isinstance(response, StreamingBody):
async for raw_chunk in response.body_iterator:
# Parse chunk
if isinstance(chunk, bytes):
chunk = chunk.decode("utf-8")
chunk = raw_chunk.decode("utf-8") if isinstance(raw_chunk, bytes) else raw_chunk
if isinstance(chunk, str) and chunk.startswith("data: "):
chunk_data = chunk[6:].strip()
@ -162,76 +292,10 @@ async def background_streaming_task(
break
try:
event = json.loads(chunk_data)
event_type = event.get("type", "")
event = _EVENT_ADAPTER.validate_json(chunk_data)
# Process different event types based on OpenAI streaming spec
if event_type == "response.output_item.added":
# New output item added
item = event.get("item", {})
item_id = item.get("id")
if item_id:
output_items[item_id] = item
state_dirty = True
elif event_type == "response.content_part.added":
# Content part added to an output item
item_id = event.get("item_id")
content_part = event.get("part", {})
if item_id and item_id in output_items:
# Update the output item with new content
if "content" not in output_items[item_id]:
output_items[item_id]["content"] = []
output_items[item_id]["content"].append(content_part)
state_dirty = True
elif event_type == "response.output_text.delta":
# Text delta - accumulate text content
# https://platform.openai.com/docs/api-reference/responses-streaming/response-text-delta
item_id = event.get("item_id")
content_index = event.get("content_index", 0)
delta = event.get("delta", "")
if item_id and item_id in output_items:
# Accumulate text delta
key = (item_id, content_index)
if key not in accumulated_text:
accumulated_text[key] = ""
accumulated_text[key] += delta
# Update the content in output_items
if "content" in output_items[item_id]:
content_list = output_items[item_id]["content"]
if content_index < len(content_list):
# Update existing content part with accumulated text
if isinstance(content_list[content_index], dict):
content_list[content_index]["text"] = accumulated_text[key]
state_dirty = True
elif event_type == "response.content_part.done":
# Content part completed
item_id = event.get("item_id")
content_part = event.get("part", {})
content_index = event.get("content_index", 0)
if item_id and item_id in output_items:
# Update with final content from event
if "content" in output_items[item_id]:
content_list = output_items[item_id]["content"]
if content_index < len(content_list):
content_list[content_index] = content_part
state_dirty = True
elif event_type == "response.output_item.done":
# Output item completed - use final item data
item = event.get("item", {})
item_id = item.get("id")
if item_id:
output_items[item_id] = item
state_dirty = True
elif event_type == "response.in_progress":
if event.type == "response.in_progress":
# Response is now in progress
# https://platform.openai.com/docs/api-reference/responses-streaming/response-in-progress
await polling_handler.update_state(
@ -239,60 +303,27 @@ async def background_streaming_task(
status="in_progress",
)
elif event_type in (
"response.completed",
"response.failed",
"response.incomplete",
"response.cancelled",
):
elif event.type in _EVENT_TO_STATUS:
# Terminal event - extract all ResponsesAPIResponse fields
# https://platform.openai.com/docs/api-reference/responses-streaming
response_data = event.get("response", {})
terminal_status = cast(
ResponsesAPIStatus,
response_data.get(
"status",
_event_to_status.get(event_type, "completed"),
),
)
# Extract error for failed and incomplete responses
if event_type == "response.failed" or event_type == "response.incomplete":
terminal_error = response_data.get("error")
# Core response fields
usage_data = response_data.get("usage")
reasoning_data = response_data.get("reasoning")
tool_choice_data = response_data.get("tool_choice")
tools_data = response_data.get("tools")
# Additional ResponsesAPIResponse fields
model_data = response_data.get("model")
instructions_data = response_data.get("instructions")
temperature_data = response_data.get("temperature")
top_p_data = response_data.get("top_p")
max_output_tokens_data = response_data.get("max_output_tokens")
previous_response_id_data = response_data.get("previous_response_id")
text_data = response_data.get("text")
truncation_data = response_data.get("truncation")
parallel_tool_calls_data = response_data.get("parallel_tool_calls")
user_data = response_data.get("user")
store_data = response_data.get("store")
incomplete_details_data = response_data.get("incomplete_details")
terminal_response = event.response or StreamedResponseSnapshot()
terminal_status = terminal_response.status or _EVENT_TO_STATUS[event.type]
# Also update output from final response if available
if "output" in response_data:
final_output = response_data.get("output", [])
for item in final_output:
item_id = item.get("id")
if item_id:
output_items[item_id] = item
if terminal_response.output is not None:
for final_item in terminal_response.output:
final_item_id = final_item.get("id")
if isinstance(final_item_id, str) and final_item_id:
output_items[final_item_id] = final_item
state_dirty = True
elif _apply_event(event, output_items, accumulated_text):
state_dirty = True
# Flush state to Redis if interval elapsed
await flush_state_if_needed()
except json.JSONDecodeError as e:
except ValidationError as e:
verbose_proxy_logger.warning("Failed to parse streaming chunk: %s", e)
# Final flush to ensure all accumulated state is saved
@ -300,27 +331,31 @@ async def background_streaming_task(
# Use the terminal status from the stream, default to "completed"
final_status: Final = terminal_status or "completed"
final_response: Final = terminal_response or StreamedResponseSnapshot()
terminal_error: Final = (
final_response.error if final_status == "failed" or final_status == "incomplete" else None
)
await polling_handler.update_state(
polling_id=polling_id,
status=final_status,
usage=usage_data,
usage=final_response.usage,
error=terminal_error,
reasoning=reasoning_data,
tool_choice=tool_choice_data,
tools=tools_data,
model=model_data,
instructions=instructions_data,
temperature=temperature_data,
top_p=top_p_data,
max_output_tokens=max_output_tokens_data,
previous_response_id=previous_response_id_data,
text=text_data,
truncation=truncation_data,
parallel_tool_calls=parallel_tool_calls_data,
user=user_data,
store=store_data,
incomplete_details=incomplete_details_data,
reasoning=final_response.reasoning,
tool_choice=final_response.tool_choice,
tools=final_response.tools,
model=final_response.model,
instructions=final_response.instructions,
temperature=final_response.temperature,
top_p=final_response.top_p,
max_output_tokens=final_response.max_output_tokens,
previous_response_id=final_response.previous_response_id,
text=final_response.text,
truncation=final_response.truncation,
parallel_tool_calls=final_response.parallel_tool_calls,
user=final_response.user,
store=final_response.store,
incomplete_details=final_response.incomplete_details,
)
verbose_proxy_logger.info(
@ -328,7 +363,7 @@ async def background_streaming_task(
polling_id,
final_status,
terminal_error,
incomplete_details_data,
final_response.incomplete_details,
len(output_items),
)

View file

@ -3,14 +3,44 @@ Response Polling Handler for Background Responses with Cache
"""
import json
from collections.abc import Sequence
from datetime import datetime, timezone
from typing import Any, Final
from typing import TYPE_CHECKING, Final, Literal
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid4
from litellm.caching.redis_cache import RedisCache
from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesAPIStatus
if TYPE_CHECKING:
from litellm.router import Router
_STATE_ADAPTER: Final = TypeAdapter(dict[str, JsonValue])
class _DeploymentParams(BaseModel):
model_config = ConfigDict(extra="ignore", from_attributes=True)
custom_llm_provider: str | None = None
model: str | None = None
class _Deployment(BaseModel):
model_config = ConfigDict(extra="ignore", from_attributes=True)
litellm_params: _DeploymentParams = _DeploymentParams()
_DEPLOYMENT_ADAPTER: Final = TypeAdapter(_Deployment)
def _deployment_provider(deployment: object) -> str | None:
litellm_params: Final = _DEPLOYMENT_ADAPTER.validate_python(deployment).litellm_params
dep_model: Final = litellm_params.model or ""
return litellm_params.custom_llm_provider or (dep_model.split("/")[0] if "/" in dep_model else None)
class ResponsePollingHandler:
"""Handles polling-based responses with Redis cache"""
@ -40,7 +70,7 @@ class ResponsePollingHandler:
async def create_initial_state(
self,
polling_id: str,
request_data: dict[str, Any],
request_data: dict[str, JsonValue],
) -> ResponsesAPIResponse:
"""
Create initial state in Redis for a polling request
@ -56,6 +86,7 @@ class ResponsePollingHandler:
ResponsesAPIResponse object following OpenAI spec
"""
created_timestamp: Final = int(datetime.now(timezone.utc).timestamp())
metadata: Final = request_data.get("metadata")
# Create OpenAI-compliant response object
response: Final = ResponsesAPIResponse(
@ -64,7 +95,7 @@ class ResponsePollingHandler:
status="queued", # OpenAI native status
created_at=created_timestamp,
output=[],
metadata=request_data.get("metadata", {}),
metadata=metadata if isinstance(metadata, dict) else {},
usage=None,
)
@ -85,13 +116,13 @@ class ResponsePollingHandler:
self,
polling_id: str,
status: ResponsesAPIStatus | None = None,
usage: dict | None = None,
error: dict | None = None,
incomplete_details: dict | None = None,
reasoning: dict | None = None,
tool_choice: Any | None = None,
tools: list | None = None,
output: list | None = None,
usage: dict[str, JsonValue] | None = None,
error: dict[str, JsonValue] | None = None,
incomplete_details: dict[str, JsonValue] | None = None,
reasoning: dict[str, JsonValue] | None = None,
tool_choice: JsonValue = None,
tools: list[JsonValue] | None = None,
output: Sequence[JsonValue] | None = None,
# Additional ResponsesAPIResponse fields
model: str | None = None,
instructions: str | None = None,
@ -99,7 +130,7 @@ class ResponsePollingHandler:
top_p: float | None = None,
max_output_tokens: int | None = None,
previous_response_id: str | None = None,
text: dict | None = None,
text: dict[str, JsonValue] | None = None,
truncation: str | None = None,
parallel_tool_calls: bool | None = None,
user: str | None = None,
@ -139,13 +170,13 @@ class ResponsePollingHandler:
cache_key: Final = self.get_cache_key(polling_id)
# Get current state
cached_state: Final = await self.redis_cache.async_get_cache(cache_key)
if not cached_state:
cached_state: Final[object] = await self.redis_cache.async_get_cache(cache_key)
if not isinstance(cached_state, str | bytes) or not cached_state:
verbose_proxy_logger.warning("No cached state found for polling_id: %s", polling_id)
return
# Parse existing ResponsesAPIResponse from cache
state: Final = json.loads(cached_state)
state: Final = _STATE_ADAPTER.validate_json(cached_state)
# Update status (using OpenAI native status values)
if status:
@ -153,7 +184,7 @@ class ResponsePollingHandler:
# Replace full output list if provided
if output is not None:
state["output"] = output
state["output"] = list(output)
# Update usage
if usage:
@ -207,21 +238,22 @@ class ResponsePollingHandler:
ttl=self.ttl,
)
output_count: Final = len(state.get("output", []))
final_output: Final = state.get("output")
output_count: Final = len(final_output) if isinstance(final_output, list) else 0
verbose_proxy_logger.debug(
"Updated polling state for %s: status=%s, output_items=%s", polling_id, state["status"], output_count
)
async def get_state(self, polling_id: str) -> dict[str, Any] | None:
async def get_state(self, polling_id: str) -> dict[str, JsonValue] | None:
"""Get current polling state from Redis"""
if not self.redis_cache:
return None
cache_key: Final = self.get_cache_key(polling_id)
cached_state: Final = await self.redis_cache.async_get_cache(cache_key)
cached_state: Final[object] = await self.redis_cache.async_get_cache(cache_key)
if cached_state:
return json.loads(cached_state)
if isinstance(cached_state, str | bytes) and cached_state:
return _STATE_ADAPTER.validate_json(cached_state)
return None
@ -250,11 +282,11 @@ class ResponsePollingHandler:
def should_use_polling_for_request(
background_mode: bool,
polling_via_cache_enabled, # Can be False, "all", or List[str]
redis_cache, # RedisCache or None
polling_via_cache_enabled: bool | Literal["all"] | list[str] | None,
redis_cache: RedisCache | None,
model: str,
llm_router, # Router instance or None
native_background_mode: list[str] | None = None, # List of models that should use native background mode
llm_router: "Router | None",
native_background_mode: list[str] | None = None,
) -> bool:
"""
Determine if polling via cache should be used for a request.
@ -296,18 +328,9 @@ def should_use_polling_for_request(
try:
# Get all deployment indices for this model name
indices: Final = llm_router.model_name_to_deployment_indices.get(model, [])
deployments: Final[list[object]] = llm_router.model_list
for idx in indices:
deployment_dict = llm_router.model_list[idx]
litellm_params = deployment_dict.get("litellm_params", {})
# Check custom_llm_provider first
dep_provider = litellm_params.get("custom_llm_provider")
# Then try to extract from model (e.g., "openai/gpt-5")
if not dep_provider:
dep_model = litellm_params.get("model", "")
if "/" in dep_model:
dep_provider = dep_model.split("/")[0]
dep_provider = _deployment_provider(deployments[idx])
# If ANY deployment's provider matches, enable polling
if dep_provider and dep_provider in polling_via_cache_enabled: