mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
chore(typing): clear basedpyright Any errors in google genai adapter and response polling
This commit is contained in:
parent
b0fac57fe4
commit
5aba2e33d3
3 changed files with 614 additions and 557 deletions
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue