From 77df51155905cbdd1e1a08c1adc1792684009262 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 27 Apr 2026 10:13:52 +0530 Subject: [PATCH] fix black issues --- .../prompt_templates/factory.py | 1111 ++++++++++++----- litellm/llms/predibase/chat/transformation.py | 52 +- .../guardrail_hooks/xecguard/xecguard.py | 54 +- 3 files changed, 902 insertions(+), 315 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index fbc2c8fdaa7..fe8387476ee 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -104,7 +104,9 @@ def map_system_message_pt(messages: list) -> list: if i < len(messages) - 1: # Not the last message next_m = messages[i + 1] next_role = next_m["role"] - if next_role == "user" or next_role == "assistant": # Next message is a user or assistant message + if ( + next_role == "user" or next_role == "assistant" + ): # Next message is a user or assistant message # Merge system prompt into the next message next_m["content"] = m["content"] + " " + next_m["content"] elif next_role == "system": # Next message is a system message @@ -184,7 +186,9 @@ def convert_to_ollama_image(openai_image_url: str): ) -def _handle_ollama_system_message(messages: list, prompt: str, msg_i: int) -> Tuple[str, int]: +def _handle_ollama_system_message( + messages: list, prompt: str, msg_i: int +) -> Tuple[str, int]: system_content_str = "" ## MERGE CONSECUTIVE SYSTEM CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "system": @@ -230,7 +234,9 @@ def ollama_pt( if user_content_str: prompt += f"### User:\n{user_content_str}\n\n" - system_content_str, msg_i = _handle_ollama_system_message(messages, prompt, msg_i) + system_content_str, msg_i = _handle_ollama_system_message( + messages, prompt, msg_i + ) if system_content_str: prompt += f"### System:\n{system_content_str}\n\n" @@ -259,7 +265,9 @@ def ollama_pt( ) if ollama_tool_calls: - assistant_content_str += f"Tool Calls: {json.dumps(ollama_tool_calls, indent=2)}" + assistant_content_str += ( + f"Tool Calls: {json.dumps(ollama_tool_calls, indent=2)}" + ) msg_i += 1 @@ -306,7 +314,11 @@ def falcon_instruct_pt(messages): if message["role"] == "system": prompt += message["content"] else: - prompt += message["role"] + ":" + message["content"].replace("\r\n", "\n").replace("\n\n", "\n") + prompt += ( + message["role"] + + ":" + + message["content"].replace("\r\n", "\n").replace("\n\n", "\n") + ) prompt += "\n\n" return prompt @@ -364,7 +376,9 @@ def phind_codellama_pt(messages): return prompt -def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: str, messages: list) -> str: +def _render_chat_template( + env, chat_template: str, bos_token: str, eos_token: str, messages: list +) -> str: """ Shared template rendering logic for both sync and async hf_chat_template @@ -412,7 +426,9 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st try: for message in messages: if message["role"] == "system": - reformatted_messages.append({"role": "user", "content": message["content"]}) + reformatted_messages.append( + {"role": "user", "content": message["content"]} + ) else: reformatted_messages.append(message) rendered_text = template.render( @@ -427,13 +443,20 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st new_messages = [] for i in range(len(reformatted_messages) - 1): new_messages.append(reformatted_messages[i]) - if reformatted_messages[i]["role"] == reformatted_messages[i + 1]["role"]: + if ( + reformatted_messages[i]["role"] + == reformatted_messages[i + 1]["role"] + ): if reformatted_messages[i]["role"] == "user": - new_messages.append({"role": "assistant", "content": ""}) + new_messages.append( + {"role": "assistant", "content": ""} + ) else: new_messages.append({"role": "user", "content": ""}) new_messages.append(reformatted_messages[-1]) - rendered_text = template.render(bos_token=bos_token, eos_token=eos_token, messages=new_messages) + rendered_text = template.render( + bos_token=bos_token, eos_token=eos_token, messages=new_messages + ) return rendered_text except Exception as e: @@ -473,8 +496,12 @@ async def _afetch_and_extract_template( and "chat_template" in tokenizer_config["tokenizer"] ): tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = _extract_token_value( + token_value=tokenizer_data.get("bos_token") + ) + eos_token = _extract_token_value( + token_value=tokenizer_data.get("eos_token") + ) chat_template = tokenizer_data["chat_template"] else: # Fallback: Try to fetch chat template from separate .jinja file @@ -488,8 +515,12 @@ async def _afetch_and_extract_template( and isinstance(tokenizer_config["tokenizer"], dict) ): tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = _extract_token_value( + token_value=tokenizer_data.get("bos_token") + ) + eos_token = _extract_token_value( + token_value=tokenizer_data.get("eos_token") + ) else: raise Exception("No chat template found") @@ -527,8 +558,12 @@ def _fetch_and_extract_template( and "chat_template" in tokenizer_config["tokenizer"] ): tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = _extract_token_value( + token_value=tokenizer_data.get("bos_token") + ) + eos_token = _extract_token_value( + token_value=tokenizer_data.get("eos_token") + ) chat_template = tokenizer_data["chat_template"] else: # Fallback: Try to fetch chat template from separate .jinja file @@ -542,15 +577,21 @@ def _fetch_and_extract_template( and isinstance(tokenizer_config["tokenizer"], dict) ): tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = _extract_token_value( + token_value=tokenizer_data.get("bos_token") + ) + eos_token = _extract_token_value( + token_value=tokenizer_data.get("eos_token") + ) else: raise Exception("No chat template found") return chat_template, bos_token, eos_token # type: ignore -async def ahf_chat_template(model: str, messages: list, chat_template: Optional[Any] = None): +async def ahf_chat_template( + model: str, messages: list, chat_template: Optional[Any] = None +): """HuggingFace chat template (async version)""" from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import ( _aget_chat_template_file, @@ -605,7 +646,9 @@ def hf_chat_template(model: str, messages: list, chat_template: Optional[Any] = def deepseek_r1_pt(messages): - return hf_chat_template(model="deepseek-r1/deepseek-r1-7b-instruct", messages=messages) + return hf_chat_template( + model="deepseek-r1/deepseek-r1-7b-instruct", messages=messages + ) # Anthropic template @@ -655,7 +698,9 @@ def get_model_info(token, model): model_info = response.json() for m in model_info: if m["name"].lower().strip() == model.strip(): - return m["config"].get("prompt_format", None), m["config"].get("chat_template", None) + return m["config"].get("prompt_format", None), m["config"].get( + "chat_template", None + ) return None, None else: return None, None @@ -734,14 +779,18 @@ def anthropic_pt( AI_PROMPT = "\n\nAssistant: " prompt = "" - for idx, message in enumerate(messages): # needs to start with `\n\nHuman: ` and end with `\n\nAssistant: ` + for idx, message in enumerate( + messages + ): # needs to start with `\n\nHuman: ` and end with `\n\nAssistant: ` if message["role"] == "user": prompt += f"{AnthropicConstants.HUMAN_PROMPT.value}{message['content']}" elif message["role"] == "system": prompt += f"{AnthropicConstants.HUMAN_PROMPT.value}{message['content']}" else: prompt += f"{AnthropicConstants.AI_PROMPT.value}{message['content']}" - if idx == 0 and message["role"] == "assistant": # ensure the prompt always starts with `\n\nHuman: ` + if ( + idx == 0 and message["role"] == "assistant" + ): # ensure the prompt always starts with `\n\nHuman: ` prompt = f"{AnthropicConstants.HUMAN_PROMPT.value}" + prompt if messages[-1]["role"] != "assistant": prompt += f"{AnthropicConstants.AI_PROMPT.value}" @@ -825,7 +874,9 @@ def convert_generic_image_chunk_to_openai_image_obj( return "data:{};{},{}".format(media_type, image_chunk["type"], image_chunk["data"]) -def convert_to_anthropic_image_obj(openai_image_url: str, format: Optional[str]) -> GenericImageParsingChunk: +def convert_to_anthropic_image_obj( + openai_image_url: str, format: Optional[str] +) -> GenericImageParsingChunk: """ Input: "image_url": "data:image/jpeg;base64,{base64_image}", @@ -885,7 +936,9 @@ def create_anthropic_image_param( # as these providers don't support URL sources for images if is_bedrock_invoke or image_url.startswith("http://"): base64_url = convert_url_to_base64(url=image_url) - image_chunk = convert_to_anthropic_image_obj(openai_image_url=base64_url, format=format) + image_chunk = convert_to_anthropic_image_obj( + openai_image_url=base64_url, format=format + ) return AnthropicMessagesImageParam( type="image", source=AnthropicContentParamSource( @@ -905,7 +958,9 @@ def create_anthropic_image_param( ) else: # Convert to base64 for data URIs or other formats - image_chunk = convert_to_anthropic_image_obj(openai_image_url=image_url, format=format) + image_chunk = convert_to_anthropic_image_obj( + openai_image_url=image_url, format=format + ) return AnthropicMessagesImageParam( type="image", source=AnthropicContentParamSource( @@ -982,7 +1037,9 @@ def convert_to_anthropic_tool_invoke_xml(tool_calls: list) -> str: tool_arguments, tool_name=tool_name, context="Anthropic XML tool invoke" ) if isinstance(parsed_args, dict): - parameters = "".join(f"<{param}>{val}\n" for param, val in parsed_args.items()) + parameters = "".join( + f"<{param}>{val}\n" for param, val in parsed_args.items() + ) else: parameters = f"{parsed_args}\n" invokes += f"\n{tool_name}\n\n{parameters}\n\n" @@ -1014,8 +1071,14 @@ def anthropic_messages_pt_xml(messages: list): if isinstance(messages[msg_i]["content"], list): for m in messages[msg_i]["content"]: if m.get("type", "") == "image_url": - format = m["image_url"].get("format") if isinstance(m["image_url"], dict) else None - image_param = create_anthropic_image_param(m["image_url"], format=format) + format = ( + m["image_url"].get("format") + if isinstance(m["image_url"], dict) + else None + ) + image_param = create_anthropic_image_param( + m["image_url"], format=format + ) # Convert to dict format for XML version source = image_param["source"] if isinstance(source, dict) and source.get("type") == "url": @@ -1066,8 +1129,12 @@ def anthropic_messages_pt_xml(messages: list): assistant_content = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": - assistant_text = messages[msg_i].get("content") or "" # either string or none - if messages[msg_i].get("tool_calls", []): # support assistant tool invoke conversion + assistant_text = ( + messages[msg_i].get("content") or "" + ) # either string or none + if messages[msg_i].get( + "tool_calls", [] + ): # support assistant tool invoke conversion assistant_text += convert_to_anthropic_tool_invoke_xml( # type: ignore messages[msg_i]["tool_calls"] ) @@ -1080,7 +1147,9 @@ def anthropic_messages_pt_xml(messages: list): if not new_messages or new_messages[0]["role"] != "user": if litellm.modify_params: - new_messages.insert(0, {"role": "user", "content": [{"type": "text", "text": "."}]}) + new_messages.insert( + 0, {"role": "user", "content": [{"type": "text", "text": "."}]} + ) else: raise Exception( "Invalid first message. Should always start with 'role'='user' for Anthropic. System prompt is sent separately for Anthropic. set 'litellm.modify_params = True' or 'litellm_settings:modify_params = True' on proxy, to insert a placeholder user message - '.' as the first message, " @@ -1089,7 +1158,9 @@ def anthropic_messages_pt_xml(messages: list): if new_messages[-1]["role"] == "assistant": for content in new_messages[-1]["content"]: if isinstance(content, dict) and content["type"] == "text": - content["text"] = content["text"].rstrip() # no trailing whitespace for final assistant message + content["text"] = content[ + "text" + ].rstrip() # no trailing whitespace for final assistant message return new_messages @@ -1180,7 +1251,9 @@ def _gemini_tool_call_invoke_helper( return function_call -def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: Optional[str]) -> str: +def _encode_tool_call_id_with_signature( + tool_call_id: str, thought_signature: Optional[str] +) -> str: """ Embed thought signature into tool call ID for OpenAI client compatibility. @@ -1199,7 +1272,9 @@ def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: Op return tool_call_id -def _get_thought_signature_from_tool(tool: dict, model: Optional[str] = None) -> Optional[str]: +def _get_thought_signature_from_tool( + tool: dict, model: Optional[str] = None +) -> Optional[str]: """Extract thought signature from tool call's provider_specific_fields. If not provided try to extract thought signature from tool call id @@ -1223,7 +1298,10 @@ def _get_thought_signature_from_tool(tool: dict, model: Optional[str] = None) -> signature = func_provider_fields.get("thought_signature") if signature: return signature - elif hasattr(function, "provider_specific_fields") and function.provider_specific_fields: + elif ( + hasattr(function, "provider_specific_fields") + and function.provider_specific_fields + ): if isinstance(function.provider_specific_fields, dict): signature = function.provider_specific_fields.get("thought_signature") if signature: @@ -1309,12 +1387,18 @@ def convert_to_gemini_tool_call_invoke( if tool_calls is not None: for idx, tool in enumerate(tool_calls): if "function" in tool: - gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper( - function_call_params=tool["function"] + gemini_function_call: Optional[VertexFunctionCall] = ( + _gemini_tool_call_invoke_helper( + function_call_params=tool["function"] + ) ) if gemini_function_call is not None: - part_dict: VertexPartType = {"function_call": gemini_function_call} - thought_signature = _get_thought_signature_from_tool(dict(tool), model=model) + part_dict: VertexPartType = { + "function_call": gemini_function_call + } + thought_signature = _get_thought_signature_from_tool( + dict(tool), model=model + ) if thought_signature: part_dict["thoughtSignature"] = thought_signature @@ -1326,14 +1410,20 @@ def convert_to_gemini_tool_call_invoke( ) ) elif function_call is not None: - gemini_function_call = _gemini_tool_call_invoke_helper(function_call_params=function_call) + gemini_function_call = _gemini_tool_call_invoke_helper( + function_call_params=function_call + ) if gemini_function_call is not None: - part_dict_function: VertexPartType = {"function_call": gemini_function_call} + part_dict_function: VertexPartType = { + "function_call": gemini_function_call + } # Extract thought signature from function_call's provider_specific_fields thought_signature = None provider_fields = ( - function_call.get("provider_specific_fields") if isinstance(function_call, dict) else {} + function_call.get("provider_specific_fields") + if isinstance(function_call, dict) + else {} ) if isinstance(provider_fields, dict): thought_signature = provider_fields.get("thought_signature") @@ -1343,7 +1433,11 @@ def convert_to_gemini_tool_call_invoke( VertexGeminiConfig, ) - if not thought_signature and model and VertexGeminiConfig._is_gemini_3_or_newer(model): + if ( + not thought_signature + and model + and VertexGeminiConfig._is_gemini_3_or_newer(model) + ): thought_signature = _get_dummy_thought_signature() if thought_signature: @@ -1359,7 +1453,9 @@ def convert_to_gemini_tool_call_invoke( return _parts_list except Exception as e: raise Exception( - "Unable to convert openai tool calls={} to gemini tool calls. Received error={}".format(message, str(e)) + "Unable to convert openai tool calls={} to gemini tool calls. Received error={}".format( + message, str(e) + ) ) @@ -1410,10 +1506,14 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 if len(mime_rest) == 2 and mime_rest[0].startswith("image/"): # Strip any extra parameters (e.g. ";charset=UTF-8") from the MIME segment clean_mime = mime_rest[0].split(";")[0].strip() - inline_data_list.append(BlobType(data=mime_rest[1], mime_type=clean_mime)) + inline_data_list.append( + BlobType(data=mime_rest[1], mime_type=clean_mime) + ) content_str = "" except Exception as e: - verbose_logger.warning(f"Failed to parse data URL in tool response: {e}") + verbose_logger.warning( + f"Failed to parse data URL in tool response: {e}" + ) elif isinstance(message["content"], List): content_list = message["content"] for content in content_list: @@ -1432,16 +1532,24 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 ) ) except Exception as e: - verbose_logger.warning(f"Failed to process Anthropic image block in tool response: {e}") + verbose_logger.warning( + f"Failed to process Anthropic image block in tool response: {e}" + ) elif content_type in ("input_image", "image_url"): # Extract image for inline_data (for Computer Use screenshots and tool results) image_url_data = content.get("image_url", "") - image_url = image_url_data.get("url", "") if isinstance(image_url_data, dict) else image_url_data + image_url = ( + image_url_data.get("url", "") + if isinstance(image_url_data, dict) + else image_url_data + ) if image_url: # Convert image to base64 blob format for Gemini try: - image_obj = convert_to_anthropic_image_obj(image_url, format=None) + image_obj = convert_to_anthropic_image_obj( + image_url, format=None + ) inline_data_list.append( BlobType( data=image_obj["data"], @@ -1449,7 +1557,9 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 ) ) except Exception as e: - verbose_logger.warning(f"Failed to process image in tool response: {e}") + verbose_logger.warning( + f"Failed to process image in tool response: {e}" + ) elif content_type in ("file", "input_file"): # Extract file for inline_data (for tool results with PDF, audio, video, etc.) file_data = content.get("file_data", "") @@ -1458,15 +1568,15 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 file_data = ( file_content.get("file_data", "") if isinstance(file_content, dict) - else file_content - if isinstance(file_content, str) - else "" + else file_content if isinstance(file_content, str) else "" ) if file_data: # Convert file to base64 blob format for Gemini try: - file_obj = convert_to_anthropic_image_obj(file_data, format=None) + file_obj = convert_to_anthropic_image_obj( + file_data, format=None + ) inline_data_list.append( BlobType( data=file_obj["data"], @@ -1474,7 +1584,9 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 ) ) except Exception as e: - verbose_logger.warning(f"Failed to process file in tool response: {e}") + verbose_logger.warning( + f"Failed to process file in tool response: {e}" + ) name: Optional[str] = message.get("name", "") # type: ignore # Recover name from last message with tool calls @@ -1483,7 +1595,11 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 msg_tool_call_id = message.get("tool_call_id", None) for tool in tools: prev_tool_call_id = tool.get("id", None) - if msg_tool_call_id and prev_tool_call_id and msg_tool_call_id == prev_tool_call_id: + if ( + msg_tool_call_id + and prev_tool_call_id + and msg_tool_call_id == prev_tool_call_id + ): name = tool.get("function", {}).get("name", "") if not name: @@ -1588,7 +1704,9 @@ def convert_to_anthropic_tool_result( anthropic_content = message["content"] elif isinstance(message["content"], List): content_list = message["content"] - anthropic_content_list: List[Union[AnthropicMessagesToolResultContent, AnthropicMessagesImageParam]] = [] + anthropic_content_list: List[ + Union[AnthropicMessagesToolResultContent, AnthropicMessagesImageParam] + ] = [] for content in content_list: if content["type"] == "text": # Only include cache_control if explicitly set and not None @@ -1602,7 +1720,11 @@ def convert_to_anthropic_tool_result( text_content["cache_control"] = cache_control_value anthropic_content_list.append(text_content) elif content["type"] == "image_url": - format = content["image_url"].get("format") if isinstance(content["image_url"], dict) else None + format = ( + content["image_url"].get("format") + if isinstance(content["image_url"], dict) + else None + ) _anthropic_image_param = create_anthropic_image_param( content["image_url"], format=format, is_bedrock_invoke=force_base64 ) @@ -1610,7 +1732,9 @@ def convert_to_anthropic_tool_result( anthropic_content_element=_anthropic_image_param, original_content_element=content, ) - anthropic_content_list.append(cast(AnthropicMessagesImageParam, _anthropic_image_param)) + anthropic_content_list.append( + cast(AnthropicMessagesImageParam, _anthropic_image_param) + ) anthropic_content = anthropic_content_list anthropic_tool_result: Optional[AnthropicMessagesToolResultParam] = None @@ -1655,7 +1779,9 @@ def convert_function_to_anthropic_tool_invoke( _name = get_attribute_or_key(function_call, "name") or "" _arguments = get_attribute_or_key(function_call, "arguments") - tool_input = parse_tool_call_arguments(_arguments, tool_name=_name, context="Anthropic function to tool invoke") + tool_input = parse_tool_call_arguments( + _arguments, tool_name=_name, context="Anthropic function to tool invoke" + ) anthropic_tool_invoke = [ AnthropicMessagesToolUseParam( @@ -1717,7 +1843,9 @@ def convert_to_anthropic_tool_invoke( Fixes: https://github.com/BerriAI/litellm/issues/17737 """ - anthropic_tool_invoke: List[Union[AnthropicMessagesToolUseParam, Dict[str, Any]]] = [] + anthropic_tool_invoke: List[ + Union[AnthropicMessagesToolUseParam, Dict[str, Any]] + ] = [] for tool in tool_calls: if not get_attribute_or_key(tool, "type") == "function": @@ -1774,7 +1902,9 @@ def convert_to_anthropic_tool_invoke( ) if "cache_control" in _content_element: - _anthropic_tool_use_param["cache_control"] = _content_element["cache_control"] + _anthropic_tool_use_param["cache_control"] = _content_element[ + "cache_control" + ] anthropic_tool_invoke.append(_anthropic_tool_use_param) @@ -1805,15 +1935,15 @@ def _anthropic_content_element_factory( image_chunk: GenericImageParsingChunk, ) -> Union[AnthropicMessagesImageParam, AnthropicMessagesDocumentParam]: if image_chunk["media_type"] == "application/pdf": - _anthropic_content_element: Union[AnthropicMessagesDocumentParam, AnthropicMessagesImageParam] = ( - AnthropicMessagesDocumentParam( - type="document", - source=AnthropicContentParamSource( - type="base64", - media_type=image_chunk["media_type"], - data=image_chunk["data"], - ), - ) + _anthropic_content_element: Union[ + AnthropicMessagesDocumentParam, AnthropicMessagesImageParam + ] = AnthropicMessagesDocumentParam( + type="document", + source=AnthropicContentParamSource( + type="base64", + media_type=image_chunk["media_type"], + data=image_chunk["data"], + ), ) else: _anthropic_content_element = AnthropicMessagesImageParam( @@ -1917,12 +2047,16 @@ def anthropic_process_openai_file_message( ), ) elif content_block_type == "container_upload": - return_block_param = AnthropicMessagesContainerUploadParam(type="container_upload", file_id=file_id) + return_block_param = AnthropicMessagesContainerUploadParam( + type="container_upload", file_id=file_id + ) if return_block_param is None: raise Exception(f"Unable to parse anthropic file message: {message}") return return_block_param - raise Exception(f"Either file_data or file_id must be present in the file message: {message}") + raise Exception( + f"Either file_data or file_id must be present in the file message: {message}" + ) def _sanitize_empty_text_content( @@ -1940,7 +2074,9 @@ def _sanitize_empty_text_content( if isinstance(content, str): if not content or not content.strip(): message = cast(AllMessageValues, dict(message)) # Make a copy - message["content"] = "[System: Empty message content sanitised to satisfy protocol]" + message["content"] = ( + "[System: Empty message content sanitised to satisfy protocol]" + ) verbose_logger.debug( f"_sanitize_empty_text_content: Replaced empty text content in {message.get('role')} message" ) @@ -2091,7 +2227,9 @@ def _is_orphaned_tool_result( break if not found_matching_tool_call: - verbose_logger.debug("_is_orphaned_tool_result: Found orphaned tool result with redacted tool_call_id") + verbose_logger.debug( + "_is_orphaned_tool_result: Found orphaned tool result with redacted tool_call_id" + ) return True return False @@ -2136,7 +2274,9 @@ def sanitize_messages_for_tool_calling( # Case A: Check if assistant message has tool_calls without following tool results if current_message.get("role") == "assistant": - result_messages, messages_consumed = _add_missing_tool_results(current_message, messages, i) + result_messages, messages_consumed = _add_missing_tool_results( + current_message, messages, i + ) # If dummy tool results were added, extend sanitized_messages and skip consumed messages if len(result_messages) > 1: @@ -2191,7 +2331,11 @@ def sanitize_messages_for_tool_calling( seen_in_block = {} if duplicates_to_remove: - sanitized_messages = [msg for idx, msg in enumerate(sanitized_messages) if idx not in duplicates_to_remove] + sanitized_messages = [ + msg + for idx, msg in enumerate(sanitized_messages) + if idx not in duplicates_to_remove + ] return sanitized_messages @@ -2256,17 +2400,25 @@ def anthropic_messages_pt( # noqa: PLR0915 ChatCompletionToolMessage, ChatCompletionUserMessage, ChatCompletionFunctionMessage, - ] = messages[msg_i] # type: ignore + ] = messages[ + msg_i + ] # type: ignore if user_message_types_block["role"] == "user": if isinstance(user_message_types_block["content"], list): for m in user_message_types_block["content"]: if m.get("type", "") == "image_url": m = cast(ChatCompletionImageObject, m) - format = m["image_url"].get("format") if isinstance(m["image_url"], dict) else None + format = ( + m["image_url"].get("format") + if isinstance(m["image_url"], dict) + else None + ) # Convert ChatCompletionImageUrlObject to dict if needed image_url_value = m["image_url"] if isinstance(image_url_value, str): - image_url_input: Union[str, dict[str, Any]] = image_url_value + image_url_input: Union[str, dict[str, Any]] = ( + image_url_value + ) else: # ChatCompletionImageUrlObject or dict case - convert to dict image_url_input = { @@ -2276,7 +2428,11 @@ def anthropic_messages_pt( # noqa: PLR0915 # Bedrock invoke models have format: invoke/... # Vertex AI Anthropic also doesn't support URL sources for images is_bedrock_invoke = model.lower().startswith("invoke/") - is_vertex_ai = llm_provider.startswith("vertex_ai") if llm_provider else False + is_vertex_ai = ( + llm_provider.startswith("vertex_ai") + if llm_provider + else False + ) force_base64 = is_bedrock_invoke or is_vertex_ai _anthropic_content_element = create_anthropic_image_param( image_url_input, @@ -2289,33 +2445,43 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_element["cache_control"] = _content_element["cache_control"] + _anthropic_content_element["cache_control"] = ( + _content_element["cache_control"] + ) user_content.append(_anthropic_content_element) elif m.get("type", "") == "text": m = cast(ChatCompletionTextObject, m) - _anthropic_text_content_element = AnthropicMessagesTextParam( - type="text", - text=m["text"], + _anthropic_text_content_element = ( + AnthropicMessagesTextParam( + type="text", + text=m["text"], + ) ) _content_element = add_cache_control_to_content( anthropic_content_element=_anthropic_text_content_element, original_content_element=dict(m), ) - _content_element = cast(AnthropicMessagesTextParam, _content_element) + _content_element = cast( + AnthropicMessagesTextParam, _content_element + ) user_content.append(_content_element) elif m.get("type", "") == "document": _document_content_element = cast( AnthropicMessagesDocumentParam, add_cache_control_to_content( - anthropic_content_element=cast(AnthropicMessagesDocumentParam, m), + anthropic_content_element=cast( + AnthropicMessagesDocumentParam, m + ), original_content_element=dict(m), ), ) user_content.append(_document_content_element) elif m.get("type", "") == "file": - _file_content_element = anthropic_process_openai_file_message( - cast(ChatCompletionFileObject, m) + _file_content_element = ( + anthropic_process_openai_file_message( + cast(ChatCompletionFileObject, m) + ) ) _file_content_element = add_cache_control_to_content( anthropic_content_element=cast( @@ -2341,14 +2507,21 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_text_element["cache_control"] = _content_element["cache_control"] + _anthropic_content_text_element["cache_control"] = ( + _content_element["cache_control"] + ) user_content.append(_anthropic_content_text_element) - elif user_message_types_block["role"] == "tool" or user_message_types_block["role"] == "function": + elif ( + user_message_types_block["role"] == "tool" + or user_message_types_block["role"] == "function" + ): # OpenAI's tool message content will always be a string user_content.append( - convert_to_anthropic_tool_result(user_message_types_block, force_base64=force_base64) + convert_to_anthropic_tool_result( + user_message_types_block, force_base64=force_base64 + ) ) msg_i += 1 @@ -2365,9 +2538,13 @@ def anthropic_messages_pt( # noqa: PLR0915 assistant_content_block: ChatCompletionAssistantMessage = messages[msg_i] # type: ignore # Extract compaction_blocks from provider_specific_fields and add them first - _provider_specific_fields_raw = assistant_content_block.get("provider_specific_fields") + _provider_specific_fields_raw = assistant_content_block.get( + "provider_specific_fields" + ) if isinstance(_provider_specific_fields_raw, dict): - _compaction_blocks = _provider_specific_fields_raw.get("compaction_blocks") + _compaction_blocks = _provider_specific_fields_raw.get( + "compaction_blocks" + ) if _compaction_blocks and isinstance(_compaction_blocks, list): # Add compaction blocks at the beginning of assistant content : https://platform.claude.com/docs/en/build-with-claude/compaction assistant_content.extend(_compaction_blocks) # type: ignore @@ -2382,15 +2559,25 @@ def anthropic_messages_pt( # noqa: PLR0915 _has_server_tool_calls = False if assistant_tool_calls is not None: for _tc in assistant_tool_calls: - _tc_id = _tc.get("id") if isinstance(_tc, dict) else getattr(_tc, "id", None) - if _tc_id and isinstance(_tc_id, str) and _tc_id.startswith("srvtoolu_"): + _tc_id = ( + _tc.get("id") + if isinstance(_tc, dict) + else getattr(_tc, "id", None) + ) + if ( + _tc_id + and isinstance(_tc_id, str) + and _tc_id.startswith("srvtoolu_") + ): _has_server_tool_calls = True break if ( thinking_blocks is not None and _has_server_tool_calls - and isinstance(assistant_content_block.get("content", None), (str, type(None))) + and isinstance( + assistant_content_block.get("content", None), (str, type(None)) + ) ): # INTERLEAVED MODE: When we have both thinking blocks and server # tool calls (e.g. web search), Anthropic's original response @@ -2400,11 +2587,17 @@ def anthropic_messages_pt( # noqa: PLR0915 # verifies thinking block signatures based on position. # Build the tool call groups (server_tool_use + its result) - _provider_specific_fields_raw_tc = assistant_content_block.get("provider_specific_fields") + _provider_specific_fields_raw_tc = assistant_content_block.get( + "provider_specific_fields" + ) _provider_specific_fields_tc: Dict[str, Any] = {} if isinstance(_provider_specific_fields_raw_tc, dict): - _provider_specific_fields_tc = cast(Dict[str, Any], _provider_specific_fields_raw_tc) - _web_search_results_tc = _provider_specific_fields_tc.get("web_search_results") + _provider_specific_fields_tc = cast( + Dict[str, Any], _provider_specific_fields_raw_tc + ) + _web_search_results_tc = _provider_specific_fields_tc.get( + "web_search_results" + ) _tool_results_tc = _provider_specific_fields_tc.get("tool_results") tool_invoke_results = convert_to_anthropic_tool_invoke( assistant_tool_calls, # type: ignore @@ -2418,7 +2611,11 @@ def anthropic_messages_pt( # noqa: PLR0915 regular_tool_uses: List[Any] = [] _current_group: List[Any] = [] for item in tool_invoke_results: - item_type = item.get("type", "") if isinstance(item, dict) else getattr(item, "type", "") + item_type = ( + item.get("type", "") + if isinstance(item, dict) + else getattr(item, "type", "") + ) if item_type == "server_tool_use": if _current_group: server_tool_groups.append(_current_group) @@ -2445,7 +2642,9 @@ def anthropic_messages_pt( # noqa: PLR0915 original_content_element=dict(assistant_content_block), ) if "cache_control" in _content_element: - _anthropic_text_content_element["cache_control"] = _content_element["cache_control"] + _anthropic_text_content_element["cache_control"] = ( + _content_element["cache_control"] + ) text_element = _anthropic_text_content_element # Interleave: each thinking block precedes its server tool group. @@ -2463,12 +2662,18 @@ def anthropic_messages_pt( # noqa: PLR0915 assistant_content.append(thinking_blocks[tb_idx]) tb_idx += 1 for block in server_tool_groups[grp_idx]: - item_id = block.get("id") if isinstance(block, dict) else getattr(block, "id", None) + item_id = ( + block.get("id") + if isinstance(block, dict) + else getattr(block, "id", None) + ) if item_id and item_id in unique_tool_ids: continue if item_id: unique_tool_ids.add(item_id) - assistant_content.append(cast(AnthropicMessagesAssistantMessageValues, block)) + assistant_content.append( + cast(AnthropicMessagesAssistantMessageValues, block) + ) grp_idx += 1 elif tb_idx < num_tb: # More thinking blocks than tool groups - emit before text @@ -2477,12 +2682,18 @@ def anthropic_messages_pt( # noqa: PLR0915 else: # More tool groups than thinking blocks - emit remaining for block in server_tool_groups[grp_idx]: - item_id = block.get("id") if isinstance(block, dict) else getattr(block, "id", None) + item_id = ( + block.get("id") + if isinstance(block, dict) + else getattr(block, "id", None) + ) if item_id and item_id in unique_tool_ids: continue if item_id: unique_tool_ids.add(item_id) - assistant_content.append(cast(AnthropicMessagesAssistantMessageValues, block)) + assistant_content.append( + cast(AnthropicMessagesAssistantMessageValues, block) + ) grp_idx += 1 # Add text block (if any) @@ -2491,12 +2702,18 @@ def anthropic_messages_pt( # noqa: PLR0915 # Add regular (non-server) tool calls at the end for item in regular_tool_uses: - item_id = item.get("id") if isinstance(item, dict) else getattr(item, "id", None) + item_id = ( + item.get("id") + if isinstance(item, dict) + else getattr(item, "id", None) + ) if item_id and item_id in unique_tool_ids: continue if item_id: unique_tool_ids.add(item_id) - assistant_content.append(cast(AnthropicMessagesAssistantMessageValues, item)) + assistant_content.append( + cast(AnthropicMessagesAssistantMessageValues, item) + ) # Mark tool_calls as already processed so they are not added again assistant_tool_calls = None @@ -2513,7 +2730,9 @@ def anthropic_messages_pt( # noqa: PLR0915 _content_is_list = "content" in assistant_content_block and isinstance( assistant_content_block["content"], list ) - _content_list = assistant_content_block.get("content") if _content_is_list else None + _content_list = ( + assistant_content_block.get("content") if _content_is_list else None + ) _list_has_thinking = False if _content_is_list and _content_list is not None: for _item in _content_list: @@ -2547,13 +2766,17 @@ def anthropic_messages_pt( # noqa: PLR0915 elif ( m.get("type", "") == "text" and len(text_block) > 0 ): # don't pass empty text blocks. anthropic api raises errors. - anthropic_message = AnthropicMessagesTextParam(type="text", text=text_block) + anthropic_message = AnthropicMessagesTextParam( + type="text", text=text_block + ) _cached_message = add_cache_control_to_content( anthropic_content_element=anthropic_message, original_content_element=dict(m), ) - assistant_content.append(cast(AnthropicMessagesTextParam, _cached_message)) + assistant_content.append( + cast(AnthropicMessagesTextParam, _cached_message) + ) # handle server_tool_use blocks (tool search, web search, etc.) # Pass through as-is since these are Anthropic-native content types elif m.get("type", "") == "server_tool_use": @@ -2566,7 +2789,9 @@ def anthropic_messages_pt( # noqa: PLR0915 elif ( "content" in assistant_content_block and isinstance(assistant_content_block["content"], str) - and assistant_content_block["content"] # don't pass empty text blocks. anthropic api raises errors. + and assistant_content_block[ + "content" + ] # don't pass empty text blocks. anthropic api raises errors. ): _anthropic_text_content_element = AnthropicMessagesTextParam( type="text", @@ -2579,19 +2804,29 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_text_content_element["cache_control"] = _content_element["cache_control"] + _anthropic_text_content_element["cache_control"] = ( + _content_element["cache_control"] + ) assistant_content.append(_anthropic_text_content_element) - if assistant_tool_calls is not None: # support assistant tool invoke conversion + if ( + assistant_tool_calls is not None + ): # support assistant tool invoke conversion # Get web_search_results and tool_results from provider_specific_fields # for server_tool_use reconstruction. # Fixes: https://github.com/BerriAI/litellm/issues/17737 - _provider_specific_fields_raw = assistant_content_block.get("provider_specific_fields") + _provider_specific_fields_raw = assistant_content_block.get( + "provider_specific_fields" + ) _provider_specific_fields: Dict[str, Any] = {} if isinstance(_provider_specific_fields_raw, dict): - _provider_specific_fields = cast(Dict[str, Any], _provider_specific_fields_raw) - _web_search_results = _provider_specific_fields.get("web_search_results") + _provider_specific_fields = cast( + Dict[str, Any], _provider_specific_fields_raw + ) + _web_search_results = _provider_specific_fields.get( + "web_search_results" + ) _tool_results = _provider_specific_fields.get("tool_results") tool_invoke_results = convert_to_anthropic_tool_invoke( assistant_tool_calls, @@ -2603,19 +2838,27 @@ def anthropic_messages_pt( # noqa: PLR0915 # This can happen when merging history that already contains the tool calls for item in tool_invoke_results: # tool_use items are typically dicts, but handle objects just in case - item_id = item.get("id") if isinstance(item, dict) else getattr(item, "id", None) + item_id = ( + item.get("id") + if isinstance(item, dict) + else getattr(item, "id", None) + ) if item_id: if item_id in unique_tool_ids: continue unique_tool_ids.add(item_id) - assistant_content.append(cast(AnthropicMessagesAssistantMessageValues, item)) + assistant_content.append( + cast(AnthropicMessagesAssistantMessageValues, item) + ) assistant_function_call = assistant_content_block.get("function_call") if assistant_function_call is not None: - assistant_content.extend(convert_function_to_anthropic_tool_invoke(assistant_function_call)) + assistant_content.extend( + convert_function_to_anthropic_tool_invoke(assistant_function_call) + ) msg_i += 1 @@ -2635,7 +2878,9 @@ def anthropic_messages_pt( # noqa: PLR0915 elif isinstance(new_messages[-1]["content"], list): for content in new_messages[-1]["content"]: if isinstance(content, dict) and content["type"] == "text": - content["text"] = content["text"].rstrip() # no trailing whitespace for final assistant message + content["text"] = content[ + "text" + ].rstrip() # no trailing whitespace for final assistant message return new_messages @@ -2788,7 +3033,11 @@ def convert_openai_message_to_cohere_tool_result( msg_tool_call_id = message.get("tool_call_id", None) for tool in tools: prev_tool_call_id = tool.get("id", None) - if msg_tool_call_id and prev_tool_call_id and msg_tool_call_id == prev_tool_call_id: + if ( + msg_tool_call_id + and prev_tool_call_id + and msg_tool_call_id == prev_tool_call_id + ): name = tool.get("function", {}).get("name", "") arguments_str = tool.get("function", {}).get("arguments", "") if arguments_str is not None and len(arguments_str) > 0: @@ -2857,8 +3106,14 @@ def convert_to_cohere_tool_invoke(tool_calls: list) -> List[ToolCallObject]: cohere_tool_invoke: List[ToolCallObject] = [ { - "name": get_attribute_or_key(get_attribute_or_key(tool, "function"), "name"), - "parameters": json.loads(get_attribute_or_key(get_attribute_or_key(tool, "function"), "arguments")), + "name": get_attribute_or_key( + get_attribute_or_key(tool, "function"), "name" + ), + "parameters": json.loads( + get_attribute_or_key( + get_attribute_or_key(tool, "function"), "arguments" + ) + ), } for tool in tool_calls if get_attribute_or_key(tool, "type") == "function" @@ -2890,9 +3145,14 @@ def cohere_messages_pt_v2( # noqa: PLR0915 ## GET MOST RECENT MESSAGE most_recent_message = messages.pop(-1) returned_message: Union[ToolResultObject, str] = "" - if most_recent_message.get("role", "") is not None and most_recent_message["role"] == "tool": + if ( + most_recent_message.get("role", "") is not None + and most_recent_message["role"] == "tool" + ): # tool result - returned_message = convert_openai_message_to_cohere_tool_result(most_recent_message, tool_calls) + returned_message = convert_openai_message_to_cohere_tool_result( + most_recent_message, tool_calls + ) else: content: Union[str, List] = most_recent_message.get("content") if isinstance(content, str): @@ -2937,23 +3197,35 @@ def cohere_messages_pt_v2( # noqa: PLR0915 msg_i += 1 if len(system_content) > 0: - new_messages.append(ChatHistorySystem(role="SYSTEM", message=system_content)) + new_messages.append( + ChatHistorySystem(role="SYSTEM", message=system_content) + ) assistant_content: str = "" assistant_tool_calls: List[ToolCallObject] = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": - if messages[msg_i].get("content", None) is not None and isinstance(messages[msg_i]["content"], list): + if messages[msg_i].get("content", None) is not None and isinstance( + messages[msg_i]["content"], list + ): for m in messages[msg_i]["content"]: if m.get("type", "") == "text": assistant_content += m["text"] - elif messages[msg_i].get("content") is not None and isinstance(messages[msg_i]["content"], str): + elif messages[msg_i].get("content") is not None and isinstance( + messages[msg_i]["content"], str + ): assistant_content += messages[msg_i]["content"] - if messages[msg_i].get("tool_calls", []): # support assistant tool invoke conversion - assistant_tool_calls.extend(convert_to_cohere_tool_invoke(messages[msg_i]["tool_calls"])) + if messages[msg_i].get( + "tool_calls", [] + ): # support assistant tool invoke conversion + assistant_tool_calls.extend( + convert_to_cohere_tool_invoke(messages[msg_i]["tool_calls"]) + ) if messages[msg_i].get("function_call"): - assistant_tool_calls.extend(convert_to_cohere_tool_invoke(messages[msg_i]["function_call"])) + assistant_tool_calls.extend( + convert_to_cohere_tool_invoke(messages[msg_i]["function_call"]) + ) msg_i += 1 @@ -2969,12 +3241,18 @@ def cohere_messages_pt_v2( # noqa: PLR0915 ## MERGE CONSECUTIVE TOOL RESULTS tool_results: List[ToolResultObject] = [] while msg_i < len(messages) and messages[msg_i]["role"] in tool_message_types: - tool_results.append(convert_openai_message_to_cohere_tool_result(messages[msg_i], tool_calls)) + tool_results.append( + convert_openai_message_to_cohere_tool_result( + messages[msg_i], tool_calls + ) + ) msg_i += 1 if len(tool_results) > 0: - new_messages.append(ChatHistoryToolResult(role="TOOL", tool_results=tool_results)) + new_messages.append( + ChatHistoryToolResult(role="TOOL", tool_results=tool_results) + ) if msg_i == init_msg_i: # prevent infinite loops raise litellm.BadRequestError( @@ -2993,7 +3271,9 @@ def cohere_message_pt(messages: list): for message in messages: # check if this is a tool_call result if message["role"] == "tool": - tool_result = convert_openai_message_to_cohere_tool_result(message, tool_calls=tool_calls) + tool_result = convert_openai_message_to_cohere_tool_result( + message, tool_calls=tool_calls + ) tool_results.append(tool_result) elif message.get("content"): prompt += message["content"] + "\n\n" @@ -3020,7 +3300,9 @@ def amazon_titan_pt( prompt += f"{AmazonTitanConstants.HUMAN_PROMPT.value}{message['content']}" else: prompt += f"{AmazonTitanConstants.AI_PROMPT.value}{message['content']}" - if idx == 0 and message["role"] == "assistant": # ensure the prompt always starts with `\n\nHuman: ` + if ( + idx == 0 and message["role"] == "assistant" + ): # ensure the prompt always starts with `\n\nHuman: ` prompt = f"{AmazonTitanConstants.HUMAN_PROMPT.value}" + prompt if messages[-1]["role"] != "assistant": prompt += f"{AmazonTitanConstants.AI_PROMPT.value}" @@ -3043,7 +3325,9 @@ def _load_image_from_url(image_url): # Check the response's content type to ensure it is an image content_type = response.headers.get("content-type") if not content_type or "image" not in content_type: - raise ValueError(f"URL does not point to a valid image (content-type: {content_type})") + raise ValueError( + f"URL does not point to a valid image (content-type: {content_type})" + ) # Load the image from the response content return Image.open(BytesIO(response.content)) @@ -3094,7 +3378,9 @@ def _gemini_vision_convert_messages(messages: list): try: from PIL import Image except Exception: - raise Exception("gemini image conversion failed please run `pip install Pillow`") + raise Exception( + "gemini image conversion failed please run `pip install Pillow`" + ) if "base64" in img: # Case 2: Base64 image data @@ -3140,7 +3426,9 @@ def gemini_text_image_pt(messages: list): try: pass # type: ignore except Exception: - raise Exception("Importing google.generativeai failed, please run 'pip install -q google-generativeai") + raise Exception( + "Importing google.generativeai failed, please run 'pip install -q google-generativeai" + ) prompt = "" images = [] @@ -3241,7 +3529,9 @@ class BedrockImageProcessor: """Handles both sync and async image processing for Bedrock conversations.""" @staticmethod - def _post_call_image_processing(response: httpx.Response, image_url: str = "") -> Tuple[str, str]: + def _post_call_image_processing( + response: httpx.Response, image_url: str = "" + ) -> Tuple[str, str]: # Check the response's content type to ensure it is an image content_type = response.headers.get("content-type") @@ -3270,7 +3560,9 @@ class BedrockImageProcessor: response = await async_safe_get(client, image_url) response.raise_for_status() # Raise an exception for HTTP errors - return BedrockImageProcessor._post_call_image_processing(response, image_url) + return BedrockImageProcessor._post_call_image_processing( + response, image_url + ) except Exception as e: raise e @@ -3283,7 +3575,9 @@ class BedrockImageProcessor: response = safe_get(client, image_url) response.raise_for_status() # Raise an exception for HTTP errors - return BedrockImageProcessor._post_call_image_processing(response, image_url) + return BedrockImageProcessor._post_call_image_processing( + response, image_url + ) except Exception as e: raise e @@ -3310,14 +3604,22 @@ class BedrockImageProcessor: def _validate_format(mime_type: str, image_format: str) -> str: """Validate image format and mime type for both images and documents.""" - supported_image_formats = litellm.AmazonConverseConfig().get_supported_image_types() - supported_doc_formats = litellm.AmazonConverseConfig().get_supported_document_types() - supported_video_formats = litellm.AmazonConverseConfig().get_supported_video_types() + supported_image_formats = ( + litellm.AmazonConverseConfig().get_supported_image_types() + ) + supported_doc_formats = ( + litellm.AmazonConverseConfig().get_supported_document_types() + ) + supported_video_formats = ( + litellm.AmazonConverseConfig().get_supported_video_types() + ) document_types = ["application", "text"] is_document = any(mime_type.startswith(doc_type) for doc_type in document_types) - supported_image_and_video_formats: List[str] = supported_video_formats + supported_image_formats + supported_image_and_video_formats: List[str] = ( + supported_video_formats + supported_image_formats + ) if is_document: return BedrockImageProcessor._get_document_format( @@ -3355,7 +3657,9 @@ class BedrockImageProcessor: """ valid_extensions: Optional[List[str]] = None potential_extensions = mimetypes.guess_all_extensions(mime_type, strict=False) - valid_extensions = [ext[1:] for ext in potential_extensions if ext[1:] in supported_doc_formats] + valid_extensions = [ + ext[1:] for ext in potential_extensions if ext[1:] in supported_doc_formats + ] # Fallback to types/files.py if mimetypes doesn't return valid extensions ################# @@ -3380,15 +3684,22 @@ class BedrockImageProcessor: return valid_extensions[0] @staticmethod - def _create_bedrock_block(image_bytes: str, mime_type: str, image_format: str) -> BedrockContentBlock: + def _create_bedrock_block( + image_bytes: str, mime_type: str, image_format: str + ) -> BedrockContentBlock: """Create appropriate Bedrock content block based on mime type.""" _blob = BedrockSourceBlock(bytes=image_bytes) document_types = ["application", "text"] is_document = any(mime_type.startswith(doc_type) for doc_type in document_types) - supported_video_formats = litellm.AmazonConverseConfig().get_supported_video_types() - is_video = any(image_format.startswith(video_type) for video_type in supported_video_formats) + supported_video_formats = ( + litellm.AmazonConverseConfig().get_supported_video_types() + ) + is_video = any( + image_format.startswith(video_type) + for video_type in supported_video_formats + ) HASH_SAMPLE_BYTES = 64 * 1024 # hash up to 64 KB of data @@ -3409,7 +3720,9 @@ class BedrockImageProcessor: # --- Compute deterministic hash (sample + total length) --- hasher = hashlib.sha256() hasher.update(sample) - hasher.update(str(len(normalized)).encode("utf-8")) # include full length for uniqueness + hasher.update( + str(len(normalized)).encode("utf-8") + ) # include full length for uniqueness full_hash = hasher.hexdigest() content_hash = full_hash[:16] # short deterministic ID @@ -3424,12 +3737,18 @@ class BedrockImageProcessor: ) ) elif is_video: - return BedrockContentBlock(video=BedrockVideoBlock(source=_blob, format=image_format)) + return BedrockContentBlock( + video=BedrockVideoBlock(source=_blob, format=image_format) + ) else: - return BedrockContentBlock(image=BedrockImageBlock(source=_blob, format=image_format)) + return BedrockContentBlock( + image=BedrockImageBlock(source=_blob, format=image_format) + ) @classmethod - def process_image_sync(cls, image_url: str, format: Optional[str] = None) -> BedrockContentBlock: + def process_image_sync( + cls, image_url: str, format: Optional[str] = None + ) -> BedrockContentBlock: """Synchronous image processing.""" if "base64" in image_url: @@ -3438,7 +3757,9 @@ class BedrockImageProcessor: img_bytes, mime_type = BedrockImageProcessor.get_image_details(image_url) image_format = mime_type.split("/")[1] else: - raise ValueError("Unsupported image type. Expected either image url or base64 encoded string") + raise ValueError( + "Unsupported image type. Expected either image url or base64 encoded string" + ) if format: mime_type = format @@ -3448,16 +3769,22 @@ class BedrockImageProcessor: return cls._create_bedrock_block(img_bytes, mime_type, image_format) @classmethod - async def process_image_async(cls, image_url: str, format: Optional[str]) -> BedrockContentBlock: + async def process_image_async( + cls, image_url: str, format: Optional[str] + ) -> BedrockContentBlock: """Asynchronous image processing.""" if "base64" in image_url: img_bytes, mime_type, image_format = cls._parse_base64_image(image_url) elif "http://" in image_url or "https://" in image_url: - img_bytes, mime_type = await BedrockImageProcessor.get_image_details_async(image_url) + img_bytes, mime_type = await BedrockImageProcessor.get_image_details_async( + image_url + ) image_format = mime_type.split("/")[1] else: - raise ValueError("Unsupported image type. Expected either image url or base64 encoded string") + raise ValueError( + "Unsupported image type. Expected either image url or base64 encoded string" + ) if format: # override with user-defined params mime_type = format @@ -3538,29 +3865,45 @@ def _convert_to_bedrock_tool_call_invoke( if parsed_objects: # First object keeps the original tool id. for obj_idx, obj in enumerate(parsed_objects): - block_id = tool_id if obj_idx == 0 else f"{tool_id}_{obj_idx}" - bedrock_tool = BedrockToolUseBlock(input=obj, name=name, toolUseId=block_id) - _parts_list.append(BedrockContentBlock(toolUse=bedrock_tool)) + block_id = ( + tool_id if obj_idx == 0 else f"{tool_id}_{obj_idx}" + ) + bedrock_tool = BedrockToolUseBlock( + input=obj, name=name, toolUseId=block_id + ) + _parts_list.append( + BedrockContentBlock(toolUse=bedrock_tool) + ) # cache_control applies to the whole original # tool call; attach after the last split block. if tool.get("cache_control", None) is not None: - _parts_list.append(BedrockContentBlock(cachePoint=CachePointBlock(type="default"))) + _parts_list.append( + BedrockContentBlock( + cachePoint=CachePointBlock(type="default") + ) + ) continue # Fallback: no objects extracted — use empty dict. arguments_dict = {} - bedrock_tool = BedrockToolUseBlock(input=arguments_dict, name=name, toolUseId=tool_id) + bedrock_tool = BedrockToolUseBlock( + input=arguments_dict, name=name, toolUseId=tool_id + ) bedrock_content_block = BedrockContentBlock(toolUse=bedrock_tool) _parts_list.append(bedrock_content_block) # Check for cache_control and add a separate cachePoint block if tool.get("cache_control", None) is not None: - cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default")) + cache_point_block = BedrockContentBlock( + cachePoint=CachePointBlock(type="default") + ) _parts_list.append(cache_point_block) return _parts_list except Exception as e: raise Exception( - "Unable to convert openai tool calls={} to bedrock tool calls. Received error={}".format(tool_calls, str(e)) + "Unable to convert openai tool calls={} to bedrock tool calls. Received error={}".format( + tool_calls, str(e) + ) ) @@ -3609,12 +3952,16 @@ def _convert_to_bedrock_tool_call_result( """ tool_result_content_blocks: List[BedrockToolResultContentBlock] = [] if isinstance(message["content"], str): - tool_result_content_blocks.append(BedrockToolResultContentBlock(text=message["content"])) + tool_result_content_blocks.append( + BedrockToolResultContentBlock(text=message["content"]) + ) elif isinstance(message["content"], List): content_list = message["content"] for content in content_list: if content["type"] == "text": - tool_result_content_blocks.append(BedrockToolResultContentBlock(text=content["text"])) + tool_result_content_blocks.append( + BedrockToolResultContentBlock(text=content["text"]) + ) elif content["type"] == "image_url": format: Optional[str] = None if isinstance(content["image_url"], dict): @@ -3627,7 +3974,9 @@ def _convert_to_bedrock_tool_call_result( format=format, ) if "image" in _block: - tool_result_content_blocks.append(BedrockToolResultContentBlock(image=_block["image"])) + tool_result_content_blocks.append( + BedrockToolResultContentBlock(image=_block["image"]) + ) message.get("name", "") id = str(message.get("tool_call_id", str(uuid.uuid4()))) @@ -3731,7 +4080,9 @@ def _sort_bedrock_assistant_content_blocks( def _insert_assistant_continue_message( messages: List[BedrockMessageBlock], - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: Optional[ + Union[str, ChatCompletionAssistantMessage] + ] = None, ) -> List[BedrockMessageBlock]: """ Add dummy message between user/tool result blocks. @@ -3755,7 +4106,9 @@ def _insert_assistant_continue_message( ) ) elif litellm.modify_params: - text = convert_content_list_to_str(cast(ChatCompletionAssistantMessage, DEFAULT_ASSISTANT_CONTINUE_MESSAGE)) + text = convert_content_list_to_str( + cast(ChatCompletionAssistantMessage, DEFAULT_ASSISTANT_CONTINUE_MESSAGE) + ) messages.append( BedrockMessageBlock( role="assistant", @@ -3780,7 +4133,9 @@ def get_user_message_block_or_continue_message( content_block = message.get("content", None) # Handle None case - if content_block is None or (user_continue_message is None and litellm.modify_params is False): + if content_block is None or ( + user_continue_message is None and litellm.modify_params is False + ): return skip_empty_text_blocks(message=message) # Handle string case @@ -3831,7 +4186,9 @@ def get_user_message_block_or_continue_message( def return_assistant_continue_message( - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: Optional[ + Union[str, ChatCompletionAssistantMessage] + ] = None, ) -> ChatCompletionAssistantMessage: if assistant_continue_message and isinstance(assistant_continue_message, str): return ChatCompletionAssistantMessage( @@ -3854,7 +4211,11 @@ def _skip_empty_dict_blocks(blocks: List[dict]) -> List[dict]: Returns: Filtered list of non-empty text blocks """ - return [item for item in blocks if not (item.get("type") == "text" and not item.get("text", "").strip())] + return [ + item + for item in blocks + if not (item.get("type") == "text" and not item.get("text", "").strip()) + ] @overload @@ -3892,7 +4253,9 @@ def skip_empty_text_blocks( modified_message["content"] = None # user message content cannot be None return modified_message elif isinstance(content_block, list): - modified_content_block = _skip_empty_dict_blocks(cast(List[dict], content_block)) + modified_content_block = _skip_empty_dict_blocks( + cast(List[dict], content_block) + ) # If no content remains and it's an assistant message, set content to None if not modified_content_block and message["role"] == "assistant": @@ -3920,7 +4283,9 @@ def skip_empty_text_blocks( def process_empty_text_blocks( message: ChatCompletionAssistantMessage, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: Optional[ + Union[str, ChatCompletionAssistantMessage] + ] = None, ) -> ChatCompletionAssistantMessage: modified_content_block = message.get("content", None) ## BASE CASE ## @@ -3928,9 +4293,14 @@ def process_empty_text_blocks( return message # Check if all items are empty text blocks - if all(item["type"] == "text" and not item["text"].strip() for item in modified_content_block): + if all( + item["type"] == "text" and not item["text"].strip() + for item in modified_content_block + ): # Replace with a single continue message - _assistant_continue_message = return_assistant_continue_message(assistant_continue_message) + _assistant_continue_message = return_assistant_continue_message( + assistant_continue_message + ) modified_content_block = [ { "type": "text", @@ -3940,7 +4310,9 @@ def process_empty_text_blocks( else: # Filter out only empty text blocks, keeping non-empty text and other block types modified_content_block = [ - item for item in modified_content_block if not (item["type"] == "text" and not item["text"].strip()) + item + for item in modified_content_block + if not (item["type"] == "text" and not item["text"].strip()) ] modified_message = message.copy() @@ -3953,7 +4325,9 @@ def process_empty_text_blocks( def get_assistant_message_block_or_continue_message( message: ChatCompletionAssistantMessage, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: Optional[ + Union[str, ChatCompletionAssistantMessage] + ] = None, ) -> ChatCompletionAssistantMessage: """ Returns the user content block @@ -3964,7 +4338,9 @@ def get_assistant_message_block_or_continue_message( content_block = message.get("content", None) # Handle Base case - if content_block is None or (assistant_continue_message is None and litellm.modify_params is False): + if content_block is None or ( + assistant_continue_message is None and litellm.modify_params is False + ): return skip_empty_text_blocks(message=message) # Handle string case @@ -3990,7 +4366,9 @@ def get_assistant_message_block_or_continue_message( } ], """ - return process_empty_text_blocks(message=message, assistant_continue_message=assistant_continue_message) + return process_empty_text_blocks( + message=message, assistant_continue_message=assistant_continue_message + ) # Handle unsupported type raise ValueError(f"Unsupported content type: {type(content_block)}") @@ -4012,7 +4390,8 @@ class BedrockConverseMessagesProcessor: messages.append(DEFAULT_USER_CONTINUE_MESSAGE) else: raise litellm.BadRequestError( - message=BAD_MESSAGE_ERROR_STR + "bedrock requires at least one non-system message", + message=BAD_MESSAGE_ERROR_STR + + "bedrock requires at least one non-system message", model=model, llm_provider=llm_provider, ) @@ -4040,7 +4419,9 @@ class BedrockConverseMessagesProcessor: model: str, llm_provider: str, user_continue_message: Optional[ChatCompletionUserMessage] = None, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: Optional[ + Union[str, ChatCompletionAssistantMessage] + ] = None, ) -> List[BedrockMessageBlock]: contents: List[BedrockMessageBlock] = [] msg_i = 0 @@ -4067,7 +4448,9 @@ class BedrockConverseMessagesProcessor: _parts.append(_part) elif element["type"] == "guarded_text": # Wrap guarded_text in guardContent block - _part = BedrockContentBlock(guardContent={"text": {"text": element["text"]}}) + _part = BedrockContentBlock( + guardContent={"text": {"text": element["text"]}} + ) _parts.append(_part) elif element["type"] == "image_url": format: Optional[str] = None @@ -4085,17 +4468,25 @@ class BedrockConverseMessagesProcessor: message=cast(ChatCompletionFileObject, element) ) _parts.append(_part) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block=cast(OpenAIMessageContentListBlock, element), - block_type="content_block", + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + message_block=cast( + OpenAIMessageContentListBlock, element + ), + block_type="content_block", + ) ) if _cache_point_block is not None: _parts.append(_cache_point_block) user_content.extend(_parts) - elif message_block["content"] and isinstance(message_block["content"], str): + elif message_block["content"] and isinstance( + message_block["content"], str + ): _part = BedrockContentBlock(text=messages[msg_i]["content"]) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block, block_type="content_block" + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + message_block, block_type="content_block" + ) ) user_content.append(_part) if _cache_point_block is not None: @@ -4104,20 +4495,27 @@ class BedrockConverseMessagesProcessor: msg_i += 1 if user_content: if len(contents) > 0 and contents[-1]["role"] == "user": - if assistant_continue_message is not None or litellm.modify_params is True: + if ( + assistant_continue_message is not None + or litellm.modify_params is True + ): # if last message was a 'user' message, then add a dummy assistant message (bedrock requires alternating roles) contents = _insert_assistant_continue_message( messages=contents, assistant_continue_message=assistant_continue_message, ) - contents.append(BedrockMessageBlock(role="user", content=user_content)) + contents.append( + BedrockMessageBlock(role="user", content=user_content) + ) else: verbose_logger.warning( "Potential consecutive user/tool blocks. Trying to merge. If error occurs, please set a 'assistant_continue_message' or set 'modify_params=True' to insert a dummy assistant message for bedrock calls." ) contents[-1]["content"].extend(user_content) else: - contents.append(BedrockMessageBlock(role="user", content=user_content)) + contents.append( + BedrockMessageBlock(role="user", content=user_content) + ) ## MERGE CONSECUTIVE TOOL CALL MESSAGES ## tool_content: List[BedrockContentBlock] = [] @@ -4135,13 +4533,18 @@ class BedrockConverseMessagesProcessor: # Check for content-level cache_control in list content elif isinstance(current_message.get("content"), list): for content_element in current_message["content"]: - if isinstance(content_element, dict) and content_element.get("cache_control", None) is not None: + if ( + isinstance(content_element, dict) + and content_element.get("cache_control", None) is not None + ): has_cache_control = True break # Add a separate cachePoint block if cache_control is present if has_cache_control: - cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default")) + cache_point_block = BedrockContentBlock( + cachePoint=CachePointBlock(type="default") + ) tool_content.append(cache_point_block) msg_i += 1 @@ -4150,26 +4553,35 @@ class BedrockConverseMessagesProcessor: if tool_content: # if last message was a 'user' message, then add a blank assistant message (bedrock requires alternating roles) if len(contents) > 0 and contents[-1]["role"] == "user": - if assistant_continue_message is not None or litellm.modify_params is True: + if ( + assistant_continue_message is not None + or litellm.modify_params is True + ): # if last message was a 'user' message, then add a dummy assistant message (bedrock requires alternating roles) contents = _insert_assistant_continue_message( messages=contents, assistant_continue_message=assistant_continue_message, ) - contents.append(BedrockMessageBlock(role="user", content=tool_content)) + contents.append( + BedrockMessageBlock(role="user", content=tool_content) + ) else: verbose_logger.warning( "Potential consecutive user/tool blocks. Trying to merge. If error occurs, please set a 'assistant_continue_message' or set 'modify_params=True' to insert a dummy assistant message for bedrock calls." ) contents[-1]["content"].extend(tool_content) else: - contents.append(BedrockMessageBlock(role="user", content=tool_content)) + contents.append( + BedrockMessageBlock(role="user", content=tool_content) + ) assistant_content: List[BedrockContentBlock] = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": - assistant_message_block = get_assistant_message_block_or_continue_message( - message=messages[msg_i], - assistant_continue_message=assistant_continue_message, + assistant_message_block = ( + get_assistant_message_block_or_continue_message( + message=messages[msg_i], + assistant_continue_message=assistant_continue_message, + ) ) _assistant_content = assistant_message_block.get("content", None) thinking_blocks = cast( @@ -4178,34 +4590,36 @@ class BedrockConverseMessagesProcessor: ) if thinking_blocks is not None: - converted_thinking_blocks = ( - BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( - thinking_blocks - ) + converted_thinking_blocks = BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( + thinking_blocks ) assistant_content = BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content( thinking_blocks=converted_thinking_blocks, assistant_parts=assistant_content, ) - if _assistant_content is not None and isinstance(_assistant_content, list): + if _assistant_content is not None and isinstance( + _assistant_content, list + ): assistants_parts: List[BedrockContentBlock] = [] for element in _assistant_content: if isinstance(element, dict): if element["type"] == "thinking": thinking_block = BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( - thinking_blocks=[cast(ChatCompletionThinkingBlock, element)] + thinking_blocks=[ + cast(ChatCompletionThinkingBlock, element) + ] ) - assistants_parts = ( - BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content( - thinking_blocks=thinking_block, - assistant_parts=assistants_parts, - ) + assistants_parts = BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content( + thinking_blocks=thinking_block, + assistant_parts=assistants_parts, ) elif element["type"] == "text": # Skip completely empty strings to avoid blank content blocks if element.get("text", "").strip(): - assistants_part = BedrockContentBlock(text=element["text"]) + assistants_part = BedrockContentBlock( + text=element["text"] + ) assistants_parts.append(assistants_part) elif element["type"] == "image_url": if isinstance(element["image_url"], dict): @@ -4217,36 +4631,54 @@ class BedrockConverseMessagesProcessor: ) assistants_parts.append(assistants_part) # Add cache point block for assistant content elements - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block=cast(OpenAIMessageContentListBlock, element), - block_type="content_block", + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + message_block=cast( + OpenAIMessageContentListBlock, element + ), + block_type="content_block", + ) ) if _cache_point_block is not None: assistants_parts.append(_cache_point_block) assistant_content.extend(assistants_parts) - elif _assistant_content is not None and isinstance(_assistant_content, str): + elif _assistant_content is not None and isinstance( + _assistant_content, str + ): # Skip completely empty strings to avoid blank content blocks if _assistant_content.strip(): - assistant_content.append(BedrockContentBlock(text=_assistant_content)) + assistant_content.append( + BedrockContentBlock(text=_assistant_content) + ) # If content is empty/whitespace, skip it (don't add a placeholder) # Add cache point block for assistant string content - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - assistant_message_block, block_type="content_block" + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + assistant_message_block, block_type="content_block" + ) ) if _cache_point_block is not None: assistant_content.append(_cache_point_block) _tool_calls = assistant_message_block.get("tool_calls", []) if _tool_calls: - assistant_content.extend(_convert_to_bedrock_tool_call_invoke(_tool_calls)) + assistant_content.extend( + _convert_to_bedrock_tool_call_invoke(_tool_calls) + ) msg_i += 1 - assistant_content = _deduplicate_bedrock_content_blocks(assistant_content, "toolUse") - assistant_content = _sort_bedrock_assistant_content_blocks(assistant_content) + assistant_content = _deduplicate_bedrock_content_blocks( + assistant_content, "toolUse" + ) + assistant_content = _sort_bedrock_assistant_content_blocks( + assistant_content + ) if assistant_content: - contents.append(BedrockMessageBlock(role="assistant", content=assistant_content)) + contents.append( + BedrockMessageBlock(role="assistant", content=assistant_content) + ) if msg_i == init_msg_i: # prevent infinite loops raise litellm.BadRequestError( @@ -4273,7 +4705,9 @@ class BedrockConverseMessagesProcessor: reasoning_content_block = BedrockConverseReasoningContentBlock( reasoningText=text_block, ) - bedrock_content_block = BedrockContentBlock(reasoningContent=reasoning_content_block) + bedrock_content_block = BedrockContentBlock( + reasoningContent=reasoning_content_block + ) reasoning_content_blocks.append(bedrock_content_block) return reasoning_content_blocks @@ -4285,12 +4719,16 @@ class BedrockConverseMessagesProcessor: if file_data is None and file_id is None: raise litellm.BadRequestError( - message="file_data and file_id cannot both be None. Got={}".format(message), + message="file_data and file_id cannot both be None. Got={}".format( + message + ), model="", llm_provider="bedrock", ) format = file_message.get("format") - return BedrockImageProcessor.process_image_sync(image_url=cast(str, file_id or file_data), format=format) + return BedrockImageProcessor.process_image_sync( + image_url=cast(str, file_id or file_data), format=format + ) @staticmethod async def _async_process_file_message( @@ -4302,11 +4740,15 @@ class BedrockConverseMessagesProcessor: format = file_message.get("format") if file_data is None and file_id is None: raise litellm.BadRequestError( - message="file_data and file_id cannot both be None. Got={}".format(message), + message="file_data and file_id cannot both be None. Got={}".format( + message + ), model="", llm_provider="bedrock", ) - return await BedrockImageProcessor.process_image_async(image_url=cast(str, file_id or file_data), format=format) + return await BedrockImageProcessor.process_image_async( + image_url=cast(str, file_id or file_data), format=format + ) @staticmethod def add_thinking_blocks_to_assistant_content( @@ -4324,7 +4766,11 @@ class BedrockConverseMessagesProcessor: filtered_thinking_blocks = [] for block in thinking_blocks: reasoning_content = block.get("reasoningContent", None) - reasoning_text = reasoning_content.get("reasoningText", None) if reasoning_content is not None else None + reasoning_text = ( + reasoning_content.get("reasoningText", None) + if reasoning_content is not None + else None + ) if reasoning_text and not reasoning_text.get("signature"): reasoning_text_text = reasoning_text["text"] assistants_part = BedrockContentBlock(text=reasoning_text_text) @@ -4341,7 +4787,9 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 model: str, llm_provider: str, user_continue_message: Optional[ChatCompletionUserMessage] = None, - assistant_continue_message: Optional[Union[str, ChatCompletionAssistantMessage]] = None, + assistant_continue_message: Optional[ + Union[str, ChatCompletionAssistantMessage] + ] = None, ) -> List[BedrockMessageBlock]: """ Converts given messages from OpenAI format to Bedrock format @@ -4376,7 +4824,9 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 _parts.append(_part) elif element["type"] == "guarded_text": # Wrap guarded_text in guardContent block - _part = BedrockContentBlock(guardContent={"text": {"text": element["text"]}}) + _part = BedrockContentBlock( + guardContent={"text": {"text": element["text"]}} + ) _parts.append(_part) elif element["type"] == "image_url": format: Optional[str] = None @@ -4391,21 +4841,29 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 ) _parts.append(_part) # type: ignore elif element["type"] == "file": - _part = BedrockConverseMessagesProcessor._process_file_message( - message=cast(ChatCompletionFileObject, element) + _part = ( + BedrockConverseMessagesProcessor._process_file_message( + message=cast(ChatCompletionFileObject, element) + ) ) _parts.append(_part) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block=cast(OpenAIMessageContentListBlock, element), - block_type="content_block", + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + message_block=cast( + OpenAIMessageContentListBlock, element + ), + block_type="content_block", + ) ) if _cache_point_block is not None: _parts.append(_cache_point_block) user_content.extend(_parts) elif message_block["content"] and isinstance(message_block["content"], str): _part = BedrockContentBlock(text=messages[msg_i]["content"]) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block, block_type="content_block" + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + message_block, block_type="content_block" + ) ) user_content.append(_part) if _cache_point_block is not None: @@ -4414,13 +4872,18 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 msg_i += 1 if user_content: if len(contents) > 0 and contents[-1]["role"] == "user": - if assistant_continue_message is not None or litellm.modify_params is True: + if ( + assistant_continue_message is not None + or litellm.modify_params is True + ): # if last message was a 'user' message, then add a dummy assistant message (bedrock requires alternating roles) contents = _insert_assistant_continue_message( messages=contents, assistant_continue_message=assistant_continue_message, ) - contents.append(BedrockMessageBlock(role="user", content=user_content)) + contents.append( + BedrockMessageBlock(role="user", content=user_content) + ) else: verbose_logger.warning( "Potential consecutive user/tool blocks. Trying to merge. If error occurs, please set a 'assistant_continue_message' or set 'modify_params=True' to insert a dummy assistant message for bedrock calls." @@ -4447,13 +4910,18 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 # Check for content-level cache_control in list content elif isinstance(current_message.get("content"), list): for content_element in current_message["content"]: - if isinstance(content_element, dict) and content_element.get("cache_control", None) is not None: + if ( + isinstance(content_element, dict) + and content_element.get("cache_control", None) is not None + ): has_cache_control = True break # Add a separate cachePoint block if cache_control is present if has_cache_control: - cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default")) + cache_point_block = BedrockContentBlock( + cachePoint=CachePointBlock(type="default") + ) tool_content.append(cache_point_block) msg_i += 1 @@ -4462,13 +4930,18 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 if tool_content: # if last message was a 'user' message, then add a blank assistant message (bedrock requires alternating roles) if len(contents) > 0 and contents[-1]["role"] == "user": - if assistant_continue_message is not None or litellm.modify_params is True: + if ( + assistant_continue_message is not None + or litellm.modify_params is True + ): # if last message was a 'user' message, then add a dummy assistant message (bedrock requires alternating roles) contents = _insert_assistant_continue_message( messages=contents, assistant_continue_message=assistant_continue_message, ) - contents.append(BedrockMessageBlock(role="user", content=tool_content)) + contents.append( + BedrockMessageBlock(role="user", content=tool_content) + ) else: verbose_logger.warning( "Potential consecutive user/tool blocks. Trying to merge. If error occurs, please set a 'assistant_continue_message' or set 'modify_params=True' to insert a dummy assistant message for bedrock calls." @@ -4490,10 +4963,8 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 ) if thinking_blocks is not None: - converted_thinking_blocks = ( - BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( - thinking_blocks - ) + converted_thinking_blocks = BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( + thinking_blocks ) assistant_content = BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content( thinking_blocks=converted_thinking_blocks, @@ -4505,22 +4976,22 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 for element in _assistant_content: if isinstance(element, dict): if element["type"] == "thinking": - thinking_block = ( - BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( - thinking_blocks=[cast(ChatCompletionThinkingBlock, element)] - ) + thinking_block = BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( + thinking_blocks=[ + cast(ChatCompletionThinkingBlock, element) + ] ) - assistants_parts = ( - BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content( - thinking_blocks=thinking_block, - assistant_parts=assistants_parts, - ) + assistants_parts = BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content( + thinking_blocks=thinking_block, + assistant_parts=assistants_parts, ) elif element["type"] == "text": # AWS Bedrock doesn't allow empty or whitespace-only text content # Skip completely empty strings to avoid blank content blocks if element.get("text", "").strip(): - assistants_part = BedrockContentBlock(text=element["text"]) + assistants_part = BedrockContentBlock( + text=element["text"] + ) assistants_parts.append(assistants_part) elif element["type"] == "image_url": if isinstance(element["image_url"], dict): @@ -4532,9 +5003,13 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 ) assistants_parts.append(assistants_part) # Add cache point block for assistant content elements - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block=cast(OpenAIMessageContentListBlock, element), - block_type="content_block", + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + message_block=cast( + OpenAIMessageContentListBlock, element + ), + block_type="content_block", + ) ) if _cache_point_block is not None: assistants_parts.append(_cache_point_block) @@ -4542,24 +5017,34 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 elif _assistant_content is not None and isinstance(_assistant_content, str): # Skip completely empty strings to avoid blank content blocks if _assistant_content.strip(): - assistant_content.append(BedrockContentBlock(text=_assistant_content)) + assistant_content.append( + BedrockContentBlock(text=_assistant_content) + ) # Add cache point block for assistant string content - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - assistant_message_block, block_type="content_block" + _cache_point_block = ( + litellm.AmazonConverseConfig()._get_cache_point_block( + assistant_message_block, block_type="content_block" + ) ) if _cache_point_block is not None: assistant_content.append(_cache_point_block) _tool_calls = assistant_message_block.get("tool_calls", []) if _tool_calls: - assistant_content.extend(_convert_to_bedrock_tool_call_invoke(_tool_calls)) + assistant_content.extend( + _convert_to_bedrock_tool_call_invoke(_tool_calls) + ) msg_i += 1 - assistant_content = _deduplicate_bedrock_content_blocks(assistant_content, "toolUse") + assistant_content = _deduplicate_bedrock_content_blocks( + assistant_content, "toolUse" + ) assistant_content = _sort_bedrock_assistant_content_blocks(assistant_content) if assistant_content: - contents.append(BedrockMessageBlock(role="assistant", content=assistant_content)) + contents.append( + BedrockMessageBlock(role="assistant", content=assistant_content) + ) if msg_i == init_msg_i: # prevent infinite loops raise litellm.BadRequestError( @@ -4599,12 +5084,16 @@ def make_valid_bedrock_tool_name(input_tool_name: str) -> str: if input_tool_name != valid_string: # passed tool name was formatted to become valid # store it internally so we can use for the response - litellm.bedrock_tool_name_mappings.set_cache(key=valid_string, value=input_tool_name) + litellm.bedrock_tool_name_mappings.set_cache( + key=valid_string, value=input_tool_name + ) return valid_string -def add_cache_point_tool_block(tool: dict, model: Optional[str] = None) -> Optional[BedrockToolBlock]: +def add_cache_point_tool_block( + tool: dict, model: Optional[str] = None +) -> Optional[BedrockToolBlock]: from litellm.llms.bedrock.common_utils import is_claude_4_5_on_bedrock cache_control = tool.get("cache_control", None) @@ -4614,7 +5103,11 @@ def add_cache_point_tool_block(tool: dict, model: Optional[str] = None) -> Optio cache_point_block: CachePointBlock = {"type": "default"} if isinstance(cache_control, dict) and "ttl" in cache_control: ttl = cache_control["ttl"] - if ttl in ["5m", "1h"] and model is not None and is_claude_4_5_on_bedrock(model): + if ( + ttl in ["5m", "1h"] + and model is not None + and is_claude_4_5_on_bedrock(model) + ): cache_point_block["ttl"] = ttl return {"cachePoint": cache_point_block} return None @@ -4641,10 +5134,14 @@ def _is_bedrock_tool_block(tool: dict) -> bool: >>> _is_bedrock_tool_block({"type": "function", "function": {...}}) False """ - return isinstance(tool, dict) and ("systemTool" in tool or "toolSpec" in tool or "cachePoint" in tool) + return isinstance(tool, dict) and ( + "systemTool" in tool or "toolSpec" in tool or "cachePoint" in tool + ) -def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockToolBlock]: +def _bedrock_tools_pt( + tools: List, model: Optional[str] = None +) -> List[BedrockToolBlock]: """ OpenAI tools looks like: tools = [ @@ -4698,7 +5195,9 @@ def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockT ) from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs - _valid_json_schema_root_types = frozenset(("array", "boolean", "integer", "null", "number", "object", "string")) + _valid_json_schema_root_types = frozenset( + ("array", "boolean", "integer", "null", "number", "object", "string") + ) tool_block_list: List[BedrockToolBlock] = [] for tool_idx, tool in enumerate(tools): # Check if tool is already a BedrockToolBlock (e.g., systemTool for Nova grounding) @@ -4709,11 +5208,17 @@ def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockT # OpenAI function tools, or Anthropic Messages / Claude Code ({name, input_schema, type, ...}) if isinstance(tool, dict) and "input_schema" in tool and "function" not in tool: - parameters = copy.deepcopy(tool.get("input_schema") or {"type": "object", "properties": {}}) + parameters = copy.deepcopy( + tool.get("input_schema") or {"type": "object", "properties": {}} + ) raw_name = tool.get("name", "") or "" _tool_description = tool.get("description", None) else: - parameters = copy.deepcopy(tool.get("function", {}).get("parameters", {"type": "object", "properties": {}})) + parameters = copy.deepcopy( + tool.get("function", {}).get( + "parameters", {"type": "object", "properties": {}} + ) + ) raw_name = tool.get("function", {}).get("name", "") or "" _tool_description = tool.get("function", {}).get("description", None) @@ -4745,7 +5250,9 @@ def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockT required=parameters.get("required", []), ) ) - tool_spec = BedrockToolSpecBlock(inputSchema=tool_input_schema, name=name, description=description) + tool_spec = BedrockToolSpecBlock( + inputSchema=tool_input_schema, name=name, description=description + ) tool_block = BedrockToolBlock(toolSpec=tool_spec) tool_block_list.append(tool_block) @@ -4769,7 +5276,9 @@ def function_call_prompt(messages: list, functions: list): if isinstance(message["content"], str): message["content"] += f""" {function_prompt}""" else: - message["content"].append({"type": "text", "text": f""" {function_prompt}"""}) + message["content"].append( + {"type": "text", "text": f""" {function_prompt}"""} + ) function_added_to_prompt = True if function_added_to_prompt is False: @@ -4785,7 +5294,9 @@ def response_schema_prompt(model: str, response_schema: dict) -> str: Returns the prompt str that's passed to the model as a user message """ custom_prompt_details: Optional[dict] = None - response_schema_as_message = [{"role": "user", "content": "{}".format(response_schema)}] + response_schema_as_message = [ + {"role": "user", "content": "{}".format(response_schema)} + ] if f"{model}/response_schema_prompt" in litellm.custom_prompt_dict: custom_prompt_details = litellm.custom_prompt_dict[ f"{model}/response_schema_prompt" @@ -4813,7 +5324,9 @@ def default_response_schema_prompt(response_schema: dict) -> str: prompt_str = """Use this JSON schema: ```json {} - ```""".format(response_schema) + ```""".format( + response_schema + ) return prompt_str @@ -4838,17 +5351,23 @@ def custom_prompt( bos_open = True pre_message_str = ( - role_dict[role]["pre_message"] if role in role_dict and "pre_message" in role_dict[role] else "" + role_dict[role]["pre_message"] + if role in role_dict and "pre_message" in role_dict[role] + else "" ) post_message_str = ( - role_dict[role]["post_message"] if role in role_dict and "post_message" in role_dict[role] else "" + role_dict[role]["post_message"] + if role in role_dict and "post_message" in role_dict[role] + else "" ) if isinstance(message["content"], str): prompt += pre_message_str + message["content"] + post_message_str elif isinstance(message["content"], list): text_str = "" for content in message["content"]: - if content.get("text", None) is not None and isinstance(content["text"], str): + if content.get("text", None) is not None and isinstance( + content["text"], str + ): text_str += content["text"] prompt += pre_message_str + text_str + post_message_str @@ -4873,7 +5392,9 @@ def prompt_factory( elif custom_llm_provider == "anthropic": if litellm.AnthropicTextConfig._is_anthropic_text_model(model): return anthropic_pt(messages=messages) - return anthropic_messages_pt(messages=messages, model=model, llm_provider=custom_llm_provider) + return anthropic_messages_pt( + messages=messages, model=model, llm_provider=custom_llm_provider + ) elif custom_llm_provider == "anthropic_xml": return anthropic_messages_pt_xml(messages=messages) elif custom_llm_provider == "gemini": @@ -4886,7 +5407,9 @@ def prompt_factory( else: return gemini_text_image_pt(messages=messages) elif custom_llm_provider == "mistral": - return litellm.MistralConfig()._transform_messages(messages=messages, model=model) + return litellm.MistralConfig()._transform_messages( + messages=messages, model=model + ) elif custom_llm_provider == "bedrock": if "amazon.titan-text" in model: return amazon_titan_pt(messages=messages) @@ -4918,12 +5441,16 @@ def prompt_factory( elif custom_llm_provider == "watsonx": from litellm.llms.watsonx.chat.transformation import IBMWatsonXChatConfig - return IBMWatsonXChatConfig.apply_prompt_template(model=model, messages=messages) + return IBMWatsonXChatConfig.apply_prompt_template( + model=model, messages=messages + ) try: if "meta-llama/llama-2" in model and "chat" in model: return llama_2_chat_pt(messages=messages) - elif ("meta-llama/llama-3" in model or "meta-llama-3" in model) and "instruct" in model: + elif ( + "meta-llama/llama-3" in model or "meta-llama-3" in model + ) and "instruct" in model: return hf_chat_template( model="meta-llama/Meta-Llama-3-8B-Instruct", messages=messages, @@ -4947,7 +5474,9 @@ def prompt_factory( return wizardcoder_pt(messages=messages) elif "phind/phind-codellama" in model: return phind_codellama_pt(messages=messages) - elif "togethercomputer/llama-2" in model and ("instruct" in model or "chat" in model): + elif "togethercomputer/llama-2" in model and ( + "instruct" in model or "chat" in model + ): return llama_2_chat_pt(messages=messages) elif model in [ "gryphe/mythomax-l2-13b", @@ -4957,7 +5486,9 @@ def prompt_factory( return alpaca_pt(messages=messages) elif model.strip() == "meta-llama/llamaguard-7b": chat_template = "{% if messages|length % 2 == 0 %}{% set role = 'Agent' %}{% else %}{% set role = 'User' %}{% endif %}{{ \"[INST] Task: Check if there is unsafe content in '\" + role + \"' messages in conversations according our safety policy with the below categories.\n\n\nO1: Violence and Hate. \nShould not\n- Help people plan or engage in violence.\n- Encourage people to engage in violence.\n- Express hateful, derogatory, or demeaning sentiments against people based on sensitive personal characteristics like their race, color, religion, national origin, sexual orientation, gender, gender identity, or disability.\n- Encourage discrimination against people based on those characteristics and should not use slurs or create insults based on those characteristics.\nCan\n- Provide information on violence and discrimination.\n- Discuss the topics of hate and violence, and can discuss historical events involving violence.\nO2: Sexual Content.\nShould not\n- Engage in sexually explicit (i.e., erotic) conversations and should not encourage people to engage in sexual activities.\nCan\n- Discuss topics related to sexuality and sexual education.\n- Provide general information regarding sexual activities.\n- Provide guidance on topics related to sexual health.\nO3: Criminal Planning.\nShould not\n- Help people plan or engage in criminal activities like theft, kidnapping, or financial crimes, especially when the person indicates that they may act on the information provided to them.\n- Encourage people to engage in criminal activities or condone criminal behavior.\nCan\n- Help people understand how crimes are committed, provided that the response is not so detailed that it could be used to provide a roadmap for future crimes.\n- Help people understand how systems work that are used to prevent crimes, provided those descriptions do not provide detailed information on how to defeat those systems.\nO4: Guns and Illegal Weapons.\nShould not\n- Assist in the planning or execution of firearm or illegal weapons crimes, especially when the person indicates that they may act on the information provided to them.\n- Encourage people to engage in firearm or illegal weapons crimes or condone such behavior.\nCan\n- Discuss firearms and the arguments for and against firearm ownership.\nO5: Regulated or Controlled Substances.\nShould not\n- Encourage or assist people in the trafficking, consumption, or use of regulated or controlled substances, including but not limited to illegal drugs, tobacco, alcohol, and cannabis.\n- Assist or encourage people to create such substances when it is illegal to do so.\nCan\n- Discuss the use of regulated or controlled substances, and can provide general information on those substances (ex: what their chemical composition is and what their effects are).\n- Provide information on how those substances are manufactured, provided that the information is not so detailed as to provide step-by-step instructions for their creation.\nO6: Self-Harm.\nShould not\n- Encourage people to harm themselves.\n- Romanticize or condone self-harm or suicide.\n- Provide information regarding the methods of suicide or self-harm.\n- Provide any information or advice, other than information regarding appropriate health resources, to people who express or imply that they wish to harm themselves.\nCan\n- Provide general information regarding self-harm and suicide (ex: statistics, the mental health challenges faced by those who take these actions).\nShould\n- Provide information regarding appropriate health resources to people who imply or express the desire to harm themselves.\n\n\n\n\n\" }}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% set content = message['content'] %}{% if message['role'] == 'user' %}{% set role = 'User' %}{% elif message['role'] == 'assistant' %}{% set role = 'Agent' %}{% endif %}{{ role + ': ' + content.strip() + '\n\n' }}{% endfor %}{{ \"\n\nProvide your safety assessment for \" + role + \" in the above conversation:\n- First line must read 'safe' or 'unsafe'.\n- If unsafe, a second line must include a comma-separated list of violated categories. [/INST]\" }}" - return hf_chat_template(model=model, messages=messages, chat_template=chat_template) + return hf_chat_template( + model=model, messages=messages, chat_template=chat_template + ) else: return hf_chat_template(original_model_name, messages) except Exception: diff --git a/litellm/llms/predibase/chat/transformation.py b/litellm/llms/predibase/chat/transformation.py index 09f54a59ff9..3d251d24b0d 100644 --- a/litellm/llms/predibase/chat/transformation.py +++ b/litellm/llms/predibase/chat/transformation.py @@ -35,9 +35,13 @@ class PredibaseConfig(BaseConfig): best_of: Optional[int] = None decoder_input_details: Optional[bool] = None details: bool = True # enables returning logprobs + best of - max_new_tokens: int = DEFAULT_MAX_TOKENS # openai default - requests hang if max_new_tokens not given + max_new_tokens: int = ( + DEFAULT_MAX_TOKENS # openai default - requests hang if max_new_tokens not given + ) repetition_penalty: Optional[float] = None - return_full_text: Optional[bool] = False # by default don't return the input as part of the output + return_full_text: Optional[bool] = ( + False # by default don't return the input as part of the output + ) seed: Optional[int] = None stop: Optional[List[str]] = None temperature: Optional[float] = None @@ -104,7 +108,9 @@ class PredibaseConfig(BaseConfig): optional_params["top_p"] = value if param == "n": optional_params["best_of"] = value - optional_params["do_sample"] = True # Need to sample if you want best of for hf inference endpoints + optional_params["do_sample"] = ( + True # Need to sample if you want best of for hf inference endpoints + ) if param == "stream": optional_params["stream"] = value if param == "stop": @@ -169,8 +175,13 @@ class PredibaseConfig(BaseConfig): completion_response["generated_text"] ) - if "details" in completion_response and "tokens" in completion_response["details"]: - model_response.choices[0].finish_reason = map_finish_reason(completion_response["details"]["finish_reason"]) + if ( + "details" in completion_response + and "tokens" in completion_response["details"] + ): + model_response.choices[0].finish_reason = map_finish_reason( + completion_response["details"]["finish_reason"] + ) sum_logprob = 0 for token in completion_response["details"]["tokens"]: if token["logprob"] is not None: @@ -190,9 +201,14 @@ class PredibaseConfig(BaseConfig): best_of_value = 0 if best_of_value > 1: - if "details" in completion_response and "best_of_sequences" in completion_response["details"]: + if ( + "details" in completion_response + and "best_of_sequences" in completion_response["details"] + ): choices_list = [] - for idx, item in enumerate(completion_response["details"]["best_of_sequences"]): + for idx, item in enumerate( + completion_response["details"]["best_of_sequences"] + ): sum_logprob = 0 for token in item["tokens"]: if token["logprob"] is not None: @@ -222,7 +238,11 @@ class PredibaseConfig(BaseConfig): if output_text is not None and len(output_text) > 0: completion_tokens = 0 try: - completion_tokens = len(encoding.encode(model_response["choices"][0]["message"].get("content", ""))) + completion_tokens = len( + encoding.encode( + model_response["choices"][0]["message"].get("content", "") + ) + ) except Exception: # Keep usage calculation non-blocking if encoding fails. pass @@ -312,7 +332,9 @@ class PredibaseConfig(BaseConfig): litellm_params: dict, stream: Optional[bool] = None, ) -> str: - tenant_id = litellm_params.get("predibase_tenant_id") or litellm_params.get("tenant_id") + tenant_id = litellm_params.get("predibase_tenant_id") or litellm_params.get( + "tenant_id" + ) if tenant_id is None: raise ValueError( "Missing Predibase Tenant ID - Required for making the request. Set dynamically (e.g. `completion(..tenant_id=)`) or in env - `PREDIBASE_TENANT_ID`." @@ -325,15 +347,21 @@ class PredibaseConfig(BaseConfig): base_url = os.getenv("PREDIBASE_API_BASE", "") completion_url = f"{base_url}/{tenant_id}/deployments/v2/llms/{model}" - should_stream = stream if stream is not None else optional_params.get("stream", False) + should_stream = ( + stream if stream is not None else optional_params.get("stream", False) + ) if should_stream is True: completion_url += "/generate_stream" else: completion_url += "/generate" return completion_url - def get_error_class(self, error_message: str, status_code: int, headers: Union[dict, Headers]) -> BaseLLMException: - return PredibaseError(status_code=status_code, message=error_message, headers=headers) + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, Headers] + ) -> BaseLLMException: + return PredibaseError( + status_code=status_code, message=error_message, headers=headers + ) def validate_environment( self, diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index 294c671bc6f..5c374540e28 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -98,7 +98,9 @@ class XecGuardGuardrail(CustomGuardrail): "the guardrail config." ) - self.api_base = (api_base or os.environ.get("XECGUARD_API_BASE") or _DEFAULT_API_BASE).rstrip("/") + self.api_base = ( + api_base or os.environ.get("XECGUARD_API_BASE") or _DEFAULT_API_BASE + ).rstrip("/") self.xecguard_model = xecguard_model or _DEFAULT_MODEL self.policy_names = policy_names @@ -113,7 +115,9 @@ class XecGuardGuardrail(CustomGuardrail): else: self.block_on_error = block_on_error - self.grounding_strictness = grounding_strictness or _DEFAULT_GROUNDING_STRICTNESS + self.grounding_strictness = ( + grounding_strictness or _DEFAULT_GROUNDING_STRICTNESS + ) self.async_handler = get_async_httpx_client( llm_provider=httpxSpecialProvider.GuardrailCallback, @@ -175,11 +179,16 @@ class XecGuardGuardrail(CustomGuardrail): messages=messages, documents=documents, ) - if grounding_result is not None and grounding_result.get("decision") == "UNSAFE": + if ( + grounding_result is not None + and grounding_result.get("decision") == "UNSAFE" + ): raise HTTPException( status_code=400, detail={ - "error": self._format_grounding_block_message(grounding_result), + "error": self._format_grounding_block_message( + grounding_result + ), "guardrail_name": self.guardrail_name or "xecguard", "xecguard_response": grounding_result, }, @@ -203,8 +212,11 @@ class XecGuardGuardrail(CustomGuardrail): isinstance(kwargs, dict) and "litellm_params" in kwargs and "metadata" in kwargs["litellm_params"] - and "standard_logging_guardrail_information" in kwargs["litellm_params"]["metadata"] - and kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] + and "standard_logging_guardrail_information" + in kwargs["litellm_params"]["metadata"] + and kwargs["litellm_params"]["metadata"][ + "standard_logging_guardrail_information" + ] ): return kwargs, result @@ -240,7 +252,9 @@ class XecGuardGuardrail(CustomGuardrail): return kwargs, result guardrail_status: GuardrailStatus = ( - "guardrail_intervened" if scan_result.get("decision") == "UNSAFE" else "success" + "guardrail_intervened" + if scan_result.get("decision") == "UNSAFE" + else "success" ) end_time = datetime.now() kwargs["standard_logging_object"]["guardrail_information"] = { @@ -281,7 +295,11 @@ class XecGuardGuardrail(CustomGuardrail): asyncio.set_event_loop(loop) if loop.is_running(): return kwargs, result - loop.run_until_complete(self.async_logging_hook(kwargs=kwargs, result=result, call_type=call_type)) + loop.run_until_complete( + self.async_logging_hook( + kwargs=kwargs, result=result, call_type=call_type + ) + ) except Exception as exc: verbose_proxy_logger.debug( "XecGuard sync logging_hook swallowed exception: %s", @@ -303,7 +321,9 @@ class XecGuardGuardrail(CustomGuardrail): "model": self.xecguard_model, "scan_type": scan_type, "messages": messages, - "policy_names": (self.policy_names if self.policy_names else _DEFAULT_POLICIES), + "policy_names": ( + self.policy_names if self.policy_names else _DEFAULT_POLICIES + ), } return await self._post( path=_SCAN_ENDPOINT, @@ -361,7 +381,9 @@ class XecGuardGuardrail(CustomGuardrail): raise HTTPException( status_code=400, detail={ - "error": (f"XecGuard API unreachable (block_on_error=True): {exc}"), + "error": ( + f"XecGuard API unreachable (block_on_error=True): {exc}" + ), "guardrail_name": self.guardrail_name or "xecguard", }, ) from exc @@ -385,7 +407,9 @@ class XecGuardGuardrail(CustomGuardrail): the request data is incomplete. """ raw_messages = request_data.get("messages") or [] - messages: List[dict] = [self._normalize_message(m) for m in raw_messages if isinstance(m, dict)] + messages: List[dict] = [ + self._normalize_message(m) for m in raw_messages if isinstance(m, dict) + ] if input_type == "request": if not messages: @@ -398,7 +422,9 @@ class XecGuardGuardrail(CustomGuardrail): return messages # input_type == "response" - assistant_text = self._extract_assistant_text_from_response(request_data.get("response")) + assistant_text = self._extract_assistant_text_from_response( + request_data.get("response") + ) if assistant_text is None: return [] messages.append({"role": "assistant", "content": assistant_text}) @@ -475,7 +501,9 @@ class XecGuardGuardrail(CustomGuardrail): parts = [ item.get("text") for item in content - if isinstance(item, dict) and item.get("type") == "text" and isinstance(item.get("text"), str) + if isinstance(item, dict) + and item.get("type") == "text" + and isinstance(item.get("text"), str) ] joined = "\n".join(p for p in parts if p) return joined or None