From 16777affe8469413701c13341b3a20afd7898c63 Mon Sep 17 00:00:00 2001 From: Yaniv Israel Date: Wed, 17 Jun 2026 21:17:43 +0300 Subject: [PATCH] fix(lint): black reformat after merge --- .../vertex_and_google_ai_studio_gemini.py | 597 +++++++++++++----- 1 file changed, 447 insertions(+), 150 deletions(-) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 97ed0e46e57..6cf1b938205 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -252,7 +252,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if isinstance(response_format, dict): return response_format - if isinstance(response_format, type) and issubclass(response_format, _BaseModel): + if isinstance(response_format, type) and issubclass( + response_format, _BaseModel + ): schema = response_format.model_json_schema() return { "type": "json_schema", @@ -285,7 +287,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False @staticmethod - def _forward_gemini_function_call_id(model: str, custom_llm_provider: str | None = None) -> bool: + def _forward_gemini_function_call_id( + model: str, custom_llm_provider: str | None = None + ) -> bool: """ Whether to include `id` on function_call / function_response parts. @@ -340,7 +344,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): supported_params.append("thinking") return supported_params - def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> ToolConfig | None: + def map_tool_choice_values( + self, model: str, tool_choice: Union[str, dict] + ) -> ToolConfig | None: if tool_choice == "none": return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE")) elif tool_choice == "required": @@ -350,7 +356,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif isinstance(tool_choice, dict): # only supported for anthropic + mistral models - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html name = tool_choice.get("function", {}).get("name", "") - return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="ANY", allowed_function_names=[name])) + return ToolConfig( + functionCallingConfig=FunctionCallingConfig( + mode="ANY", allowed_function_names=[name] + ) + ) else: raise litellm.utils.UnsupportedParamsError( message="VertexAI doesn't support tool_choice={}. Supported tool_choice values=['auto', 'required', json object]. To drop it from the call, set `litellm.drop_params = True.".format( @@ -398,12 +408,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return search_tool_keys = cls._search_tool_keys() - has_function_declarations = any(isinstance(tool, dict) and tool.get("function_declarations") for tool in tools) + has_function_declarations = any( + isinstance(tool, dict) and tool.get("function_declarations") + for tool in tools + ) if not has_function_declarations: return has_search_tools = any( - isinstance(tool, dict) and any(key in tool for key in search_tool_keys) for tool in tools + isinstance(tool, dict) and any(key in tool for key in search_tool_keys) + for tool in tools ) if not has_search_tools: return @@ -416,7 +430,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "send a request without function calling tools." ) optional_params["tools"] = [ - tool for tool in tools if not (isinstance(tool, dict) and any(key in tool for key in search_tool_keys)) + tool + for tool in tools + if not ( + isinstance(tool, dict) and any(key in tool for key in search_tool_keys) + ) ] def _map_service_tier_param(self, value: str, optional_params: dict) -> None: @@ -460,13 +478,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Transform excluded_predefined_functions to camelCase if "excluded_predefined_functions" in computer_use_config: - transformed_config["excludedPredefinedFunctions"] = computer_use_config["excluded_predefined_functions"] + transformed_config["excludedPredefinedFunctions"] = computer_use_config[ + "excluded_predefined_functions" + ] elif "excludedPredefinedFunctions" in computer_use_config: - transformed_config["excludedPredefinedFunctions"] = computer_use_config["excludedPredefinedFunctions"] + transformed_config["excludedPredefinedFunctions"] = computer_use_config[ + "excludedPredefinedFunctions" + ] return transformed_config - def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> tuple[dict, dict | None]: + def _extract_google_maps_retrieval_config( + self, google_maps_config: dict + ) -> tuple[dict, dict | None]: """ Extract location configuration from googleMaps tool for Vertex AI toolConfig. @@ -499,7 +523,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Remove location fields from tool definition cleaned_config = { - k: v for k, v in google_maps_config.items() if k not in ["latitude", "longitude", "languageCode"] + k: v + for k, v in google_maps_config.items() + if k not in ["latitude", "longitude", "languageCode"] } return cleaned_config, retrieval_config @@ -516,7 +542,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Optional[dict]: The tool value if found, None otherwise """ # Convert camelCase to underscore_case - underscore_name = "".join(["_" + c.lower() if c.isupper() else c for c in tool_name]).lstrip("_") + underscore_name = "".join( + ["_" + c.lower() if c.isupper() else c for c in tool_name] + ).lstrip("_") # Try both camelCase and underscore_case variants if tool.get(tool_name) is not None: @@ -560,8 +588,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): urlContext, ] ) - server_side_tool_invocations = optional_params.get("include_server_side_tool_invocations", False) - if gtool_func_declarations and has_search_tools and not server_side_tool_invocations: + server_side_tool_invocations = optional_params.get( + "include_server_side_tool_invocations", False + ) + if ( + gtool_func_declarations + and has_search_tools + and not server_side_tool_invocations + ): verbose_logger.warning( "Vertex AI does not support mixing function declarations with " "search tools (googleSearch, enterpriseWebSearch, urlContext, " @@ -617,7 +651,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): and _openai_function_object["parameters"] is not None and isinstance(_openai_function_object["parameters"], dict) ): # OPENAI accepts JSON Schema, Google accepts OpenAPI schema. - _openai_function_object["parameters"] = _build_vertex_schema(_openai_function_object["parameters"]) + _openai_function_object["parameters"] = _build_vertex_schema( + _openai_function_object["parameters"] + ) openai_function_object = _openai_function_object @@ -633,43 +669,68 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "web_search", "web_search_preview", ): - verbose_logger.info(f"Gemini: Transforming OpenAI-style '{tool['type']}' tool to googleSearch") + verbose_logger.info( + f"Gemini: Transforming OpenAI-style '{tool['type']}' tool to googleSearch" + ) tool = {VertexToolName.GOOGLE_SEARCH.value: {}} # Handle tools with 'type' field (OpenAI spec compliance) Ignore this field -> https://github.com/BerriAI/litellm/issues/14644#issuecomment-3342061838 elif "type" in tool: tool = {k: tool[k] for k in tool if k != "type"} tool_name = list(tool.keys())[0] if len(tool.keys()) == 1 else None if tool_name and ( - tool_name == "codeExecution" or tool_name == VertexToolName.CODE_EXECUTION.value + tool_name == "codeExecution" + or tool_name == VertexToolName.CODE_EXECUTION.value ): # code_execution maintained for backwards compatibility code_execution = self.get_tool_value(tool, "codeExecution") - elif tool_name and (tool_name == VertexToolName.GOOGLE_SEARCH.value or tool_name == "google_search"): + elif tool_name and ( + tool_name == VertexToolName.GOOGLE_SEARCH.value + or tool_name == "google_search" + ): googleSearch = self.get_tool_value(tool, tool_name) elif tool_name and ( - tool_name == VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value or tool_name == "google_search_retrieval" + tool_name == VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value + or tool_name == "google_search_retrieval" ): googleSearchRetrieval = self.get_tool_value(tool, tool_name) elif tool_name and ( - tool_name == VertexToolName.ENTERPRISE_WEB_SEARCH.value or tool_name == "enterprise_web_search" + tool_name == VertexToolName.ENTERPRISE_WEB_SEARCH.value + or tool_name == "enterprise_web_search" ): enterpriseWebSearch = self.get_tool_value(tool, tool_name) - elif tool_name and (tool_name == VertexToolName.URL_CONTEXT.value or tool_name == "urlContext"): + elif tool_name and ( + tool_name == VertexToolName.URL_CONTEXT.value + or tool_name == "urlContext" + ): urlContext = self.get_tool_value(tool, tool_name) - elif tool_name and (tool_name == VertexToolName.GOOGLE_MAPS.value or tool_name == "google_maps"): - google_maps_value = self.get_tool_value(tool, VertexToolName.GOOGLE_MAPS.value) + elif tool_name and ( + tool_name == VertexToolName.GOOGLE_MAPS.value + or tool_name == "google_maps" + ): + google_maps_value = self.get_tool_value( + tool, VertexToolName.GOOGLE_MAPS.value + ) # Extract and transform location configuration for toolConfig if google_maps_value is not None: ( googleMaps, google_maps_retrieval_config, - ) = self._extract_google_maps_retrieval_config(google_maps_config=google_maps_value) - elif tool_name and (tool_name == VertexToolName.COMPUTER_USE.value or tool_name == "computer_use"): - computer_use_value = self.get_tool_value(tool, VertexToolName.COMPUTER_USE.value) + ) = self._extract_google_maps_retrieval_config( + google_maps_config=google_maps_value + ) + elif tool_name and ( + tool_name == VertexToolName.COMPUTER_USE.value + or tool_name == "computer_use" + ): + computer_use_value = self.get_tool_value( + tool, VertexToolName.COMPUTER_USE.value + ) # Transform Computer Use configuration to Gemini API format if computer_use_value is not None: - computerUse = self._transform_computer_use_config(computer_use_config=computer_use_value) + computerUse = self._transform_computer_use_config( + computer_use_config=computer_use_value + ) else: # Empty config - Gemini will use defaults computerUse = {} @@ -725,11 +786,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _tools_list.append(search_tool) if googleSearchRetrieval is not None: retrieval_tool = Tools() - retrieval_tool[VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value] = googleSearchRetrieval + retrieval_tool[VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value] = ( + googleSearchRetrieval + ) _tools_list.append(retrieval_tool) if enterpriseWebSearch is not None: enterprise_tool = Tools() - enterprise_tool[VertexToolName.ENTERPRISE_WEB_SEARCH.value] = enterpriseWebSearch + enterprise_tool[VertexToolName.ENTERPRISE_WEB_SEARCH.value] = ( + enterpriseWebSearch + ) _tools_list.append(enterprise_tool) if code_execution is not None: code_tool = Tools() @@ -752,7 +817,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if google_maps_retrieval_config is not None: if "toolConfig" not in optional_params: optional_params["toolConfig"] = {} - optional_params["toolConfig"]["retrievalConfig"] = google_maps_retrieval_config + optional_params["toolConfig"][ + "retrievalConfig" + ] = google_maps_retrieval_config return _tools_list @@ -761,13 +828,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if isinstance(old_schema, list): for item in old_schema: if isinstance(item, dict): - item = _build_vertex_schema(parameters=item, add_property_ordering=True) + item = _build_vertex_schema( + parameters=item, add_property_ordering=True + ) elif isinstance(old_schema, dict): - old_schema = _build_vertex_schema(parameters=old_schema, add_property_ordering=True) + old_schema = _build_vertex_schema( + parameters=old_schema, add_property_ordering=True + ) return old_schema - def apply_response_schema_transformation(self, value: dict, optional_params: dict, model: str): + def apply_response_schema_transformation( + self, value: dict, optional_params: dict, model: str + ): new_value = deepcopy(value) # remove 'strict' from json schema (not supported by Gemini) new_value = _remove_strict_from_schema(new_value) @@ -803,13 +876,17 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # - Standard JSON Schema format (lowercase types) # - Supports additionalProperties # - No propertyOrdering needed - optional_params["response_json_schema"] = _build_json_schema(deepcopy(schema)) + optional_params["response_json_schema"] = _build_json_schema( + deepcopy(schema) + ) else: # Use responseSchema (default, backwards compatible) # - OpenAPI-style format (uppercase types) # - No additionalProperties support # - Requires propertyOrdering - optional_params["response_schema"] = self._map_response_schema(value=schema) + optional_params["response_schema"] = self._map_response_schema( + value=schema + ) @staticmethod def _map_reasoning_effort_to_thinking_budget( @@ -824,7 +901,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif model and "gemini-2.5-pro" in model.lower(): budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO elif model and "gemini-2.5-flash" in model.lower(): - budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH + budget = ( + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH + ) else: budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET @@ -877,7 +956,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Check if this is gemini-3-flash which supports MINIMAL thinking level # Covers gemini-3-flash, gemini-3-flash-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview, # gemini-3.5-flash, and any future 3.x-flash variants. - is_gemini3flash = model and ("flash" in model.lower() and "gemini-3" in model.lower()) + is_gemini3flash = model and ( + "flash" in model.lower() and "gemini-3" in model.lower() + ) is_gemini31pro = model and ("gemini-3.1-pro-preview" in model.lower()) if reasoning_effort == "minimal": if is_gemini3flash: @@ -971,14 +1052,20 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): params["includeThoughts"] = True # Follow provider defaults unless explicitly opted into legacy behavior. if litellm.enable_gemini_default_thinking_level_low is True: - is_gemini3flash = "gemini-3" in model.lower() and "flash" in model.lower() - params["thinkingLevel"] = "minimal" if is_gemini3flash else "low" + is_gemini3flash = ( + "gemini-3" in model.lower() and "flash" in model.lower() + ) + params["thinkingLevel"] = ( + "minimal" if is_gemini3flash else "low" + ) else: # Thinking disabled params["includeThoughts"] = False else: # For older Gemini models, use thinkingBudget - if thinking_enabled and not VertexGeminiConfig._is_thinking_budget_zero(thinking_budget): + if thinking_enabled and not VertexGeminiConfig._is_thinking_budget_zero( + thinking_budget + ): params["includeThoughts"] = True if thinking_budget is not None and isinstance(thinking_budget, int): params["thinkingBudget"] = thinking_budget @@ -1085,7 +1172,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): model: str, drop_params: bool, ) -> dict: - self._apply_include_server_side_tool_invocations(non_default_params, optional_params) + self._apply_include_server_side_tool_invocations( + non_default_params, optional_params + ) gemini_sampling_params_warned: bool = False for param, value in non_default_params.items(): if param == "temperature": @@ -1106,7 +1195,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): gemini_sampling_params_warned = True optional_params["temperature"] = value elif param == "top_p": - if VertexGeminiConfig._is_gemini_3_or_newer(model) and not gemini_sampling_params_warned: + if ( + VertexGeminiConfig._is_gemini_3_or_newer(model) + and not gemini_sampling_params_warned + ): verbose_logger.warning( "DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to " f"function for Gemini 3+ ({model}) but are planned for removal in a " @@ -1116,7 +1208,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): gemini_sampling_params_warned = True optional_params["top_p"] = value elif param == "top_k": - if VertexGeminiConfig._is_gemini_3_or_newer(model) and not gemini_sampling_params_warned: + if ( + VertexGeminiConfig._is_gemini_3_or_newer(model) + and not gemini_sampling_params_warned + ): verbose_logger.warning( "DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to " f"function for Gemini 3+ ({model}) but are planned for removal in a " @@ -1141,7 +1236,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif param == "max_tokens" or param == "max_completion_tokens": optional_params["max_output_tokens"] = value elif param == "response_format" and isinstance(value, dict): # type: ignore - self.apply_response_schema_transformation(value=value, optional_params=optional_params, model=model) + self.apply_response_schema_transformation( + value=value, optional_params=optional_params, model=model + ) elif param == "frequency_penalty": if self._supports_penalty_parameters(model): optional_params["frequency_penalty"] = value @@ -1152,11 +1249,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): optional_params["responseLogprobs"] = value elif param == "top_logprobs": optional_params["logprobs"] = value - elif (param == "tools" or param == "functions") and isinstance(value, list) and value: + elif ( + (param == "tools" or param == "functions") + and isinstance(value, list) + and value + ): # Pass optional_params so _map_function can add toolConfig if needed - mapped_tools = self._map_function(value=value, optional_params=optional_params) - optional_params = self._add_tools_to_optional_params(optional_params, mapped_tools) - elif param == "tool_choice" and (isinstance(value, str) or isinstance(value, dict)): + mapped_tools = self._map_function( + value=value, optional_params=optional_params + ) + optional_params = self._add_tools_to_optional_params( + optional_params, mapped_tools + ) + elif param == "tool_choice" and ( + isinstance(value, str) or isinstance(value, dict) + ): _tool_choice_value = self.map_tool_choice_values( model=model, tool_choice=value, # type: ignore @@ -1164,7 +1271,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if _tool_choice_value is not None: optional_params["tool_choice"] = _tool_choice_value elif param == "parallel_tool_calls": - tools_list = non_default_params.get("tools", non_default_params.get("functions")) + tools_list = non_default_params.get( + "tools", non_default_params.get("functions") + ) num_tools = len(tools_list) if isinstance(tools_list, list) else 0 # Gemini does not support parallel_tool_calls=False with multiple # tools. Drop the param instead of failing — Responses API clients @@ -1190,12 +1299,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): param_description="thinking_budget", ) if VertexGeminiConfig._is_gemini_3_or_newer(model): - optional_params["thinkingConfig"] = VertexGeminiConfig._map_reasoning_effort_to_thinking_level( - effort_value, model + optional_params["thinkingConfig"] = ( + VertexGeminiConfig._map_reasoning_effort_to_thinking_level( + effort_value, model + ) ) else: - optional_params["thinkingConfig"] = VertexGeminiConfig._map_reasoning_effort_to_thinking_budget( - effort_value, model + optional_params["thinkingConfig"] = ( + VertexGeminiConfig._map_reasoning_effort_to_thinking_budget( + effort_value, model + ) ) elif param == "thinking": # Validate no conflict with thinking_level @@ -1204,16 +1317,20 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): param_name="thinking", param_description="thinking_budget", ) - optional_params["thinkingConfig"] = VertexGeminiConfig._map_thinking_param( - cast(AnthropicThinkingParam, value), - model=model, + optional_params["thinkingConfig"] = ( + VertexGeminiConfig._map_thinking_param( + cast(AnthropicThinkingParam, value), + model=model, + ) ) elif param == "modalities" and isinstance(value, list): response_modalities = self.map_response_modalities(value) optional_params["responseModalities"] = response_modalities elif param == "web_search_options" and isinstance(value, dict): _tools = self._map_web_search_options(value) - optional_params = self._add_tools_to_optional_params(optional_params, [_tools]) + optional_params = self._add_tools_to_optional_params( + optional_params, [_tools] + ) elif param == "service_tier" and isinstance(value, str): self._map_service_tier_param(value, optional_params) elif param == "include_server_side_tool_invocations" and value is True: @@ -1356,7 +1473,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): """ from litellm.litellm_core_utils.core_helpers import _FINISH_REASON_MAP - return {k: v for k, v in _FINISH_REASON_MAP.items() if k in VertexGeminiConfig._GEMINI_FINISH_REASON_KEYS} + return { + k: v + for k, v in _FINISH_REASON_MAP.items() + if k in VertexGeminiConfig._GEMINI_FINISH_REASON_KEYS + } def translate_exception_str(self, exception_string: str): if ( @@ -1368,7 +1489,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) return exception_string - def get_assistant_content_message(self, parts: list[HttpxPartType]) -> tuple[str | None, str | None]: + def get_assistant_content_message( + self, parts: list[HttpxPartType] + ) -> tuple[str | None, str | None]: content_str: str | None = None reasoning_content_str: str | None = None @@ -1380,7 +1503,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if text_content.startswith("data:audio") and ";base64," in text_content: try: if is_base64_encoded(text_content): - media_type, _ = text_content.split("data:")[1].split(";base64,") + media_type, _ = text_content.split("data:")[1].split( + ";base64," + ) if media_type.startswith("audio/"): continue except (ValueError, IndexError): @@ -1409,7 +1534,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return content_str, reasoning_content_str - def _extract_thinking_blocks_from_parts(self, parts: list[HttpxPartType]) -> list[ChatCompletionThinkingBlock]: + def _extract_thinking_blocks_from_parts( + self, parts: list[HttpxPartType] + ) -> list[ChatCompletionThinkingBlock]: """Extract thinking blocks from parts if present. Per Google's docs (https://ai.google.dev/gemini-api/docs/thinking): @@ -1432,7 +1559,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): thinking_blocks.append(block) return thinking_blocks - def _extract_thought_signatures_from_parts(self, parts: list[HttpxPartType]) -> list[str] | None: + def _extract_thought_signatures_from_parts( + self, parts: list[HttpxPartType] + ) -> list[str] | None: """Extract thoughtSignature values from parts. Per Google's docs, thoughtSignature is returned for multi-turn context preservation @@ -1510,7 +1639,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return invocations if invocations else None - def _extract_image_response_from_parts(self, parts: list[HttpxPartType]) -> list[ImageURLListItem] | None: + def _extract_image_response_from_parts( + self, parts: list[HttpxPartType] + ) -> list[ImageURLListItem] | None: """Extract image response from parts if present""" images: list[ImageURLListItem] = [] for part in parts: @@ -1530,7 +1661,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) return images - def _extract_audio_response_from_parts(self, parts: list[HttpxPartType]) -> ChatCompletionAudioResponse | None: + def _extract_audio_response_from_parts( + self, parts: list[HttpxPartType] + ) -> ChatCompletionAudioResponse | None: """Extract audio response from parts if present""" for part in parts: if "text" in part: @@ -1539,7 +1672,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if text_content.startswith("data:audio") and ";base64," in text_content: try: if is_base64_encoded(text_content): - media_type, audio_data = text_content.split("data:")[1].split(";base64,") + media_type, audio_data = text_content.split("data:")[ + 1 + ].split(";base64,") if media_type.startswith("audio/"): expires_at = int(time.time()) + (24 * 60 * 60) @@ -1562,7 +1697,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): expires_at = int(time.time()) + (24 * 60 * 60) transcript = "" # Gemini doesn't provide transcript - return ChatCompletionAudioResponse(data=data, expires_at=expires_at, transcript=transcript) + return ChatCompletionAudioResponse( + data=data, expires_at=expires_at, transcript=transcript + ) return None @@ -1582,7 +1719,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if "functionCall" in part: _function_chunk: ChatCompletionToolCallFunctionChunk = { "name": part["functionCall"]["name"], - "arguments": json.dumps(part["functionCall"]["args"], ensure_ascii=False), + "arguments": json.dumps( + part["functionCall"]["args"], ensure_ascii=False + ), } # Extract thought signature if present thought_signature = part.get("thoughtSignature") @@ -1596,7 +1735,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if thought_signature: if "provider_specific_fields" not in function_dict: function_dict["provider_specific_fields"] = {} - function_dict["provider_specific_fields"]["thought_signature"] = thought_signature + function_dict["provider_specific_fields"][ + "thought_signature" + ] = thought_signature function = cast(ChatCompletionToolCallFunctionChunk, function_dict) else: _tool_response_chunk: ChatCompletionToolCallChunk = { @@ -1615,8 +1756,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _tool_response_chunk["provider_specific_fields"] = { # type: ignore "thought_signature": thought_signature } - _tool_response_chunk["id"] = _encode_tool_call_id_with_signature( - _tool_response_chunk["id"] or "", thought_signature + _tool_response_chunk["id"] = ( + _encode_tool_call_id_with_signature( + _tool_response_chunk["id"] or "", thought_signature + ) ) _tools.append(_tool_response_chunk) cumulative_tool_call_idx += 1 @@ -1637,11 +1780,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): logprobs_list: list[ChatCompletionTokenLogprob] = [] for index, candidate in enumerate(logprobs_result["chosenCandidates"]): top_logprobs: list[TopLogprob] = [] - if "topCandidates" in logprobs_result and index < len(logprobs_result["topCandidates"]): - top_candidates_for_index = logprobs_result["topCandidates"][index]["candidates"] + if "topCandidates" in logprobs_result and index < len( + logprobs_result["topCandidates"] + ): + top_candidates_for_index = logprobs_result["topCandidates"][index][ + "candidates" + ] for options in top_candidates_for_index: - top_logprobs.append(TopLogprob(token=options["token"], logprob=options["logProbability"])) + top_logprobs.append( + TopLogprob( + token=options["token"], logprob=options["logProbability"] + ) + ) logprobs_list.append( ChatCompletionTokenLogprob( token=candidate["token"], @@ -1676,8 +1827,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ## GET USAGE ## usage = Usage( - prompt_tokens=completion_response["usageMetadata"].get("promptTokenCount", 0), - completion_tokens=completion_response["usageMetadata"].get("candidatesTokenCount", 0), + prompt_tokens=completion_response["usageMetadata"].get( + "promptTokenCount", 0 + ), + completion_tokens=completion_response["usageMetadata"].get( + "candidatesTokenCount", 0 + ), total_tokens=completion_response["usageMetadata"].get("totalTokenCount", 0), ) @@ -1710,8 +1865,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ## GET USAGE ## usage = Usage( - prompt_tokens=completion_response["usageMetadata"].get("promptTokenCount", 0), - completion_tokens=completion_response["usageMetadata"].get("candidatesTokenCount", 0), + prompt_tokens=completion_response["usageMetadata"].get( + "promptTokenCount", 0 + ), + completion_tokens=completion_response["usageMetadata"].get( + "candidatesTokenCount", 0 + ), total_tokens=completion_response["usageMetadata"].get("totalTokenCount", 0), ) @@ -1739,10 +1898,17 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _calculate_usage( - completion_response: Union[GenerateContentResponseBody, BidiGenerateContentServerMessage], + completion_response: Union[ + GenerateContentResponseBody, BidiGenerateContentServerMessage + ], ) -> Usage: - if completion_response is not None and "usageMetadata" not in completion_response: - raise ValueError(f"usageMetadata not found in completion_response. Got={completion_response}") + if ( + completion_response is not None + and "usageMetadata" not in completion_response + ): + raise ValueError( + f"usageMetadata not found in completion_response. Got={completion_response}" + ) cached_tokens: int | None = None # Separate variables for prompt tokens by modality prompt_audio_tokens: int | None = None @@ -1771,11 +1937,17 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): modality = str(detail.get("modality", "")).upper() token_count = _get_token_count(detail) if modality == "TEXT": - response_tokens_details.text_tokens = (response_tokens_details.text_tokens or 0) + token_count + response_tokens_details.text_tokens = ( + response_tokens_details.text_tokens or 0 + ) + token_count elif modality == "AUDIO": - response_tokens_details.audio_tokens = (response_tokens_details.audio_tokens or 0) + token_count + response_tokens_details.audio_tokens = ( + response_tokens_details.audio_tokens or 0 + ) + token_count elif modality == "DOCUMENT": - response_tokens_details.text_tokens = (response_tokens_details.text_tokens or 0) + token_count + response_tokens_details.text_tokens = ( + response_tokens_details.text_tokens or 0 + ) + token_count ######################################################### @@ -1787,15 +1959,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): modality = str(detail.get("modality", "")).upper() token_count = _get_token_count(detail) if modality == "TEXT": - response_tokens_details.text_tokens = (response_tokens_details.text_tokens or 0) + token_count + response_tokens_details.text_tokens = ( + response_tokens_details.text_tokens or 0 + ) + token_count elif modality == "AUDIO": - response_tokens_details.audio_tokens = (response_tokens_details.audio_tokens or 0) + token_count + response_tokens_details.audio_tokens = ( + response_tokens_details.audio_tokens or 0 + ) + token_count elif modality == "IMAGE": - response_tokens_details.image_tokens = (response_tokens_details.image_tokens or 0) + token_count + response_tokens_details.image_tokens = ( + response_tokens_details.image_tokens or 0 + ) + token_count elif modality == "VIDEO": - response_tokens_details.video_tokens = (response_tokens_details.video_tokens or 0) + token_count + response_tokens_details.video_tokens = ( + response_tokens_details.video_tokens or 0 + ) + token_count elif modality == "DOCUMENT": - response_tokens_details.text_tokens = (response_tokens_details.text_tokens or 0) + token_count + response_tokens_details.text_tokens = ( + response_tokens_details.text_tokens or 0 + ) + token_count # Calculate text_tokens if not explicitly provided in candidatesTokensDetails # candidatesTokenCount includes all modalities, so: text = total - (image + audio + video) @@ -1808,7 +1990,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): completion_audio_tokens = response_tokens_details.audio_tokens or 0 completion_video_tokens = response_tokens_details.video_tokens or 0 calculated_text_tokens = ( - candidates_token_count - completion_image_tokens - completion_audio_tokens - completion_video_tokens + candidates_token_count + - completion_image_tokens + - completion_audio_tokens + - completion_video_tokens ) response_tokens_details.text_tokens = calculated_text_tokens ######################################################### @@ -1889,8 +2074,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): video_tokens=prompt_video_tokens, ) - completion_tokens = response_tokens or completion_response["usageMetadata"].get("candidatesTokenCount", 0) - if not VertexGeminiConfig.is_candidate_token_count_inclusive(usage_metadata) and reasoning_tokens: + completion_tokens = response_tokens or completion_response["usageMetadata"].get( + "candidatesTokenCount", 0 + ) + if ( + not VertexGeminiConfig.is_candidate_token_count_inclusive(usage_metadata) + and reasoning_tokens + ): completion_tokens = reasoning_tokens + completion_tokens ## GET USAGE ## usage = Usage( @@ -1971,7 +2161,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None: web_search_requests: int | None = None - if grounding_metadata and isinstance(grounding_metadata, list) and len(grounding_metadata) > 0: + if ( + grounding_metadata + and isinstance(grounding_metadata, list) + and len(grounding_metadata) > 0 + ): for grounding_metadata_item in grounding_metadata: web_search_queries = grounding_metadata_item.get("webSearchQueries") if web_search_queries and web_search_requests: @@ -2085,10 +2279,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) -> None: setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore if grounding_metadata: - model_response._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata + model_response._hidden_params["vertex_ai_grounding_metadata"] = ( + grounding_metadata + ) setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore if url_context_metadata: - model_response._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata + model_response._hidden_params["vertex_ai_url_context_metadata"] = ( + url_context_metadata + ) setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore setattr(model_response, "vertex_ai_safety_results", safety_ratings) # type: ignore if safety_ratings: @@ -2096,7 +2294,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): model_response._hidden_params["vertex_ai_safety_results"] = safety_ratings setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore if citation_metadata: - model_response._hidden_params["vertex_ai_citation_metadata"] = citation_metadata + model_response._hidden_params["vertex_ai_citation_metadata"] = ( + citation_metadata + ) def apply_assembled_streaming_response_metadata( self, @@ -2229,35 +2429,51 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ( content, reasoning_content, - ) = VertexGeminiConfig().get_assistant_content_message(parts=candidate["content"]["parts"]) - - audio_response = VertexGeminiConfig()._extract_audio_response_from_parts( - parts=candidate["content"]["parts"] - ) - image_response = VertexGeminiConfig()._extract_image_response_from_parts( + ) = VertexGeminiConfig().get_assistant_content_message( parts=candidate["content"]["parts"] ) - thinking_blocks = VertexGeminiConfig()._extract_thinking_blocks_from_parts( - parts=candidate["content"]["parts"] + audio_response = ( + VertexGeminiConfig()._extract_audio_response_from_parts( + parts=candidate["content"]["parts"] + ) + ) + image_response = ( + VertexGeminiConfig()._extract_image_response_from_parts( + parts=candidate["content"]["parts"] + ) + ) + + thinking_blocks = ( + VertexGeminiConfig()._extract_thinking_blocks_from_parts( + parts=candidate["content"]["parts"] + ) ) # Extract thoughtSignatures from parts (can exist without thought: true) - thought_signatures = VertexGeminiConfig()._extract_thought_signatures_from_parts( - parts=candidate["content"]["parts"] + thought_signatures = ( + VertexGeminiConfig()._extract_thought_signatures_from_parts( + parts=candidate["content"]["parts"] + ) ) # Extract server-side tool invocations (context circulation) - server_side_tool_invocations = VertexGeminiConfig._extract_server_side_tool_invocations( - parts=candidate["content"]["parts"] + server_side_tool_invocations = ( + VertexGeminiConfig._extract_server_side_tool_invocations( + parts=candidate["content"]["parts"] + ) ) if audio_response is not None: - cast(dict[str, Any], chat_completion_message)["audio"] = audio_response + cast(dict[str, Any], chat_completion_message)[ + "audio" + ] = audio_response chat_completion_message["content"] = None # OpenAI spec if image_response is not None: # Handle image response - combine with text content into structured format - cast(dict[str, Any], chat_completion_message)["images"] = image_response + cast(dict[str, Any], chat_completion_message)[ + "images" + ] = image_response if content is not None: chat_completion_message["content"] = content @@ -2265,9 +2481,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): chat_completion_message["reasoning_content"] = reasoning_content if candidate_grounding_metadata: - annotations = VertexGeminiConfig._convert_grounding_metadata_to_annotations( - grounding_metadata=candidate_grounding_metadata, - content_text=content, + annotations = ( + VertexGeminiConfig._convert_grounding_metadata_to_annotations( + grounding_metadata=candidate_grounding_metadata, + content_text=content, + ) ) if annotations: chat_completion_message["annotations"] = annotations # type: ignore @@ -2297,7 +2515,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Convert thinking_blocks to reasoning_content for streaming # This ensures reasoning_content is available in streaming responses - if isinstance(model_response, ModelResponseStream) and reasoning_content is None: + if ( + isinstance(model_response, ModelResponseStream) + and reasoning_content is None + ): reasoning_content_parts = [] for block in thinking_blocks: thinking_text = block.get("thinking") @@ -2318,9 +2539,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if server_side_tool_invocations is not None: if "provider_specific_fields" not in chat_completion_message: chat_completion_message["provider_specific_fields"] = {} - chat_completion_message["provider_specific_fields"]["server_side_tool_invocations"] = ( - server_side_tool_invocations # type: ignore - ) + chat_completion_message["provider_specific_fields"][ + "server_side_tool_invocations" + ] = server_side_tool_invocations # type: ignore if isinstance(model_response, ModelResponseStream): choice = VertexGeminiConfig._create_streaming_choice( @@ -2413,7 +2634,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): model_response.model = model ## CHECK IF RESPONSE FLAGGED - if "promptFeedback" in completion_response and "blockReason" in completion_response["promptFeedback"]: + if ( + "promptFeedback" in completion_response + and "blockReason" in completion_response["promptFeedback"] + ): return self._handle_blocked_response( model_response=model_response, completion_response=completion_response, @@ -2421,8 +2645,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _candidates = completion_response.get("candidates") if _candidates and len(_candidates) > 0: - content_policy_violations = VertexGeminiConfig().get_flagged_finish_reasons() - if "finishReason" in _candidates[0] and _candidates[0]["finishReason"] in content_policy_violations.keys(): + content_policy_violations = ( + VertexGeminiConfig().get_flagged_finish_reasons() + ) + if ( + "finishReason" in _candidates[0] + and _candidates[0]["finishReason"] in content_policy_violations.keys() + ): return self._handle_content_policy_violation( model_response=model_response, completion_response=completion_response, @@ -2444,24 +2673,38 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): safety_ratings, citation_metadata, _, # cumulative_tool_call_index not needed in non-streaming - ) = VertexGeminiConfig._process_candidates(_candidates, model_response, logging_obj.optional_params) + ) = VertexGeminiConfig._process_candidates( + _candidates, model_response, logging_obj.optional_params + ) - usage = VertexGeminiConfig._calculate_usage(completion_response=completion_response) + usage = VertexGeminiConfig._calculate_usage( + completion_response=completion_response + ) - web_search_requests = VertexGeminiConfig._calculate_web_search_requests(grounding_metadata) + web_search_requests = VertexGeminiConfig._calculate_web_search_requests( + grounding_metadata + ) if web_search_requests is not None: - cast(PromptTokensDetailsWrapper, usage.prompt_tokens_details).web_search_requests = web_search_requests + cast( + PromptTokensDetailsWrapper, usage.prompt_tokens_details + ).web_search_requests = web_search_requests setattr(model_response, "usage", usage) ## ADD METADATA TO RESPONSE ## setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) - model_response._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata + model_response._hidden_params["vertex_ai_grounding_metadata"] = ( + grounding_metadata + ) - setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) + setattr( + model_response, "vertex_ai_url_context_metadata", url_context_metadata + ) - model_response._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata + model_response._hidden_params["vertex_ai_url_context_metadata"] = ( + url_context_metadata + ) setattr(model_response, "vertex_ai_safety_results", safety_ratings) model_response._hidden_params["vertex_ai_safety_results"] = ( @@ -2475,9 +2718,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) ## ADD TRAFFIC TYPE ## - traffic_type = completion_response.get("usageMetadata", {}).get("trafficType") + traffic_type = completion_response.get("usageMetadata", {}).get( + "trafficType" + ) if traffic_type: - model_response._hidden_params.setdefault("provider_specific_fields", {})["traffic_type"] = traffic_type + model_response._hidden_params.setdefault( + "provider_specific_fields", {} + )["traffic_type"] = traffic_type ## ADD SERVICE TIER ## if getattr(raw_response, "headers", None): @@ -2514,7 +2761,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def get_error_class( self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] ) -> BaseLLMException: - return VertexAIError(message=error_message, status_code=status_code, headers=headers) + return VertexAIError( + message=error_message, status_code=status_code, headers=headers + ) def transform_request( self, @@ -2524,7 +2773,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): litellm_params: dict, headers: dict, ) -> dict: - raise NotImplementedError("Vertex AI has a custom implementation of transform_request. Needs sync + async.") + raise NotImplementedError( + "Vertex AI has a custom implementation of transform_request. Needs sync + async." + ) def validate_environment( self, @@ -2567,7 +2818,9 @@ async def make_call( ) try: - response = await client.post(api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj) + response = await client.post( + api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj + ) response.raise_for_status() except httpx.HTTPStatusError as e: exception_string = str(await e.response.aread()) @@ -2616,7 +2869,9 @@ def make_sync_call( if client is None: client = HTTPHandler() # Create a new client if none provided - response = client.post(api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj) + response = client.post( + api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj + ) if response.status_code != 200 and response.status_code != 201: raise VertexAIError( @@ -2673,7 +2928,9 @@ class VertexLLM(VertexBase): gemini_api_key: str | None = None, extra_headers: dict | None = None, ) -> CustomStreamWrapper: - should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) + should_use_v1beta1_features = self.is_using_v1beta1_features( + optional_params=optional_params + ) _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, @@ -2730,7 +2987,11 @@ class VertexLLM(VertexBase): completion_stream=None, make_call=partial( make_call, - gemini_client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None), + gemini_client=( + client + if client is not None and isinstance(client, AsyncHTTPHandler) + else None + ), api_base=api_base, headers=headers, data=request_body_str, @@ -2769,7 +3030,9 @@ class VertexLLM(VertexBase): gemini_api_key: str | None = None, extra_headers: dict | None = None, ) -> Union[ModelResponse, CustomStreamWrapper]: - should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) + should_use_v1beta1_features = self.is_using_v1beta1_features( + optional_params=optional_params + ) _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, @@ -2814,7 +3077,9 @@ class VertexLLM(VertexBase): if timeout: _async_client_params["timeout"] = timeout if client is None or not isinstance(client, AsyncHTTPHandler): - client = get_async_httpx_client(params=_async_client_params, llm_provider=litellm.LlmProviders.VERTEX_AI) + client = get_async_httpx_client( + params=_async_client_params, llm_provider=litellm.LlmProviders.VERTEX_AI + ) else: client = client # type: ignore ## LOGGING @@ -2953,7 +3218,9 @@ class VertexLLM(VertexBase): extra_headers=extra_headers, ) - should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) + should_use_v1beta1_features = self.is_using_v1beta1_features( + optional_params=optional_params + ) _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, @@ -3012,7 +3279,11 @@ class VertexLLM(VertexBase): completion_stream=None, make_call=partial( make_sync_call, - gemini_client=(client if client is not None and isinstance(client, HTTPHandler) else None), + gemini_client=( + client + if client is not None and isinstance(client, HTTPHandler) + else None + ), api_base=url, data=request_data_str, model=model, @@ -3142,7 +3413,11 @@ class ModelResponseIterator: # to correctly set finish_reason="tool_calls" per the OpenAI spec. if not self.has_seen_tool_calls: for choice in model_response.choices: - if hasattr(choice, "delta") and choice.delta and choice.delta.tool_calls: + if ( + hasattr(choice, "delta") + and choice.delta + and choice.delta.tool_calls + ): self.has_seen_tool_calls = True break @@ -3162,7 +3437,9 @@ class ModelResponseIterator: if self.has_seen_tool_calls: mapped_finish_reason = "tool_calls" else: - mapped_finish_reason = VertexGeminiConfig._check_finish_reason(None, finish_reason_str) + mapped_finish_reason = VertexGeminiConfig._check_finish_reason( + None, finish_reason_str + ) choice = StreamingChoices( finish_reason=mapped_finish_reason, index=candidate.get("index", 0), @@ -3210,13 +3487,19 @@ class ModelResponseIterator: completion_response=processed_chunk, ) - web_search_requests = VertexGeminiConfig._calculate_web_search_requests(grounding_metadata) + web_search_requests = VertexGeminiConfig._calculate_web_search_requests( + grounding_metadata + ) if web_search_requests is not None: - cast(PromptTokensDetailsWrapper, usage.prompt_tokens_details).web_search_requests = web_search_requests + cast( + PromptTokensDetailsWrapper, usage.prompt_tokens_details + ).web_search_requests = web_search_requests traffic_type = processed_chunk.get("usageMetadata", {}).get("trafficType") if traffic_type: - model_response._hidden_params.setdefault("provider_specific_fields", {})["traffic_type"] = traffic_type + model_response._hidden_params.setdefault("provider_specific_fields", {})[ + "traffic_type" + ] = traffic_type service_tier = self.response_headers.get("x-gemini-service-tier") if service_tier: @@ -3263,7 +3546,9 @@ class ModelResponseIterator: citation_metadata, ) = self._apply_stream_candidates(_candidates, model_response) - usage = self._apply_stream_usage_metadata(processed_chunk, model_response, grounding_metadata) + usage = self._apply_stream_usage_metadata( + processed_chunk, model_response, grounding_metadata + ) setattr(model_response, "usage", usage) # type: ignore @@ -3295,7 +3580,9 @@ class ModelResponseIterator: return self.chunk_parser(chunk=json_chunk) - def handle_accumulated_json_chunk(self, chunk: str) -> Optional["ModelResponseStream"]: + def handle_accumulated_json_chunk( + self, chunk: str + ) -> Optional["ModelResponseStream"]: chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "" message = chunk.replace("\n\n", "") @@ -3311,7 +3598,9 @@ class ModelResponseIterator: # If it's not valid JSON yet, continue to the next event return None - def _common_chunk_parsing_logic(self, chunk: str) -> Optional["ModelResponseStream"]: + def _common_chunk_parsing_logic( + self, chunk: str + ) -> Optional["ModelResponseStream"]: try: chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "" if len(chunk) > 0: @@ -3378,12 +3667,16 @@ class ModelResponseIterator: try: await iterator.aclose() except Exception as e: # noqa: BLE001 - verbose_logger.debug("ModelResponseIterator.aclose: error closing iterator: %s", e) + verbose_logger.debug( + "ModelResponseIterator.aclose: error closing iterator: %s", e + ) if self.response is not None: try: await self.response.aclose() except Exception as e: # noqa: BLE001 - verbose_logger.debug("ModelResponseIterator.aclose: error closing response: %s", e) + verbose_logger.debug( + "ModelResponseIterator.aclose: error closing response: %s", e + ) def close(self) -> None: iterator = getattr(self, "response_iterator", self.streaming_response) @@ -3391,9 +3684,13 @@ class ModelResponseIterator: try: iterator.close() except Exception as e: # noqa: BLE001 - verbose_logger.debug("ModelResponseIterator.close: error closing iterator: %s", e) + verbose_logger.debug( + "ModelResponseIterator.close: error closing iterator: %s", e + ) if self.response is not None: try: self.response.close() except Exception as e: # noqa: BLE001 - verbose_logger.debug("ModelResponseIterator.close: error closing response: %s", e) + verbose_logger.debug( + "ModelResponseIterator.close: error closing response: %s", e + )