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}{param}>\n" for param, val in parsed_args.items())
+ parameters = "".join(
+ f"<{param}>{val}{param}>\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