From 5aba2e33d3fa8420d0f0f18ce33e08430bdff9e8 Mon Sep 17 00:00:00 2001 From: Devin Date: Tue, 11 Aug 2026 13:50:05 +0000 Subject: [PATCH] chore(typing): clear basedpyright Any errors in google genai adapter and response polling --- .../google_genai/adapters/transformation.py | 683 +++++++++--------- .../response_polling/background_streaming.py | 395 +++++----- .../proxy/response_polling/polling_handler.py | 93 ++- 3 files changed, 614 insertions(+), 557 deletions(-) diff --git a/litellm/google_genai/adapters/transformation.py b/litellm/google_genai/adapters/transformation.py index 4f127f476c3..2c9d32707de 100644 --- a/litellm/google_genai/adapters/transformation.py +++ b/litellm/google_genai/adapters/transformation.py @@ -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, - } diff --git a/litellm/proxy/response_polling/background_streaming.py b/litellm/proxy/response_polling/background_streaming.py index 383ada5a1bc..49042b2eb74 100644 --- a/litellm/proxy/response_polling/background_streaming.py +++ b/litellm/proxy/response_polling/background_streaming.py @@ -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), ) diff --git a/litellm/proxy/response_polling/polling_handler.py b/litellm/proxy/response_polling/polling_handler.py index 3dfb67efb50..3449b6e7aef 100644 --- a/litellm/proxy/response_polling/polling_handler.py +++ b/litellm/proxy/response_polling/polling_handler.py @@ -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: