diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index efc890d9ee2..14d4835d332 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -3,10 +3,12 @@ Translating between OpenAI's `/chat/completion` format and Amazon's `/converse` """ import copy +import hashlib import json +import re import time import types -from typing import List, Literal, Optional, Tuple, Union, cast, overload +from typing import Any, List, Literal, Optional, Tuple, Union, cast, overload import httpx @@ -1777,6 +1779,136 @@ class AmazonConverseConfig(BaseConfig): tool_set.add(_name) return list(tool_set) + def _create_text_tool_call( + self, tool_name: str, arguments: dict[str, Any] + ) -> ChatCompletionMessageToolCall: + arguments_str = json.dumps(arguments, ensure_ascii=False) + tool_call_id = hashlib.sha256( + f"{tool_name}:{arguments_str}".encode() + ).hexdigest()[:16] + return ChatCompletionMessageToolCall( + id="call_" + tool_call_id, + type="function", + function=Function( + name=tool_name, + arguments=arguments_str, + ), + ) + + def _resolve_text_tool_call_name( + self, tool_name: Optional[str], tool_call_names: List[str] + ) -> Optional[str]: + if not tool_name: + return None + + if tool_name in tool_call_names: + return tool_name + + short_name = tool_name.split(".")[-1] + for allowed_tool_name in tool_call_names: + if allowed_tool_name == short_name: + return allowed_tool_name + if allowed_tool_name.endswith("." + short_name): + return allowed_tool_name + if allowed_tool_name.endswith("_" + short_name): + return allowed_tool_name + + return None + + def _parse_function_text_tool_call( + self, content: str + ) -> Tuple[Optional[str], dict[str, Any]]: + function_match = re.search( + r"(.*?)", content, re.DOTALL | re.IGNORECASE + ) + if function_match is None: + return None, {} + + params = { + match.group(1): match.group(2).strip() + for match in re.finditer( + r"""(.*?)""", + function_match.group(1), + re.DOTALL | re.IGNORECASE, + ) + } + tool_name = ( + params.pop("command", None) + or params.pop("name", None) + or params.pop("tool", None) + ) + return tool_name, params + + def _parse_tool_use_text_tool_call( + self, content: str + ) -> Tuple[Optional[str], dict[str, Any]]: + tool_use_match = re.search( + r"(.*?)", content, re.DOTALL | re.IGNORECASE + ) + if tool_use_match is None: + return None, {} + + body = tool_use_match.group(1).strip() + tool_name_match = re.search( + r"(.*?)", body, re.DOTALL | re.IGNORECASE + ) + input_match = re.search( + r"(.*?)", body, re.DOTALL | re.IGNORECASE + ) + if tool_name_match is not None: + return tool_name_match.group(1).strip(), self._parse_tool_call_json_arguments( + input_match.group(1).strip() if input_match is not None else "" + ) + + return self._parse_bare_text_tool_call(body) + + def _parse_bare_text_tool_call( + self, content: str + ) -> Tuple[Optional[str], dict[str, Any]]: + lines = [line.strip() for line in content.strip().splitlines() if line.strip()] + if len(lines) < 2: + return None, {} + + tool_name = lines[0] + arguments = self._parse_tool_call_json_arguments("\n".join(lines[1:])) + if arguments == {}: + return None, {} + + return tool_name, arguments + + def _parse_tool_call_json_arguments(self, json_text: str) -> dict[str, Any]: + if not json_text: + return {} + try: + parsed_arguments = json.loads(json_text) + except Exception: + return {} + if not isinstance(parsed_arguments, dict): + return {} + return parsed_arguments + + def _text_content_tool_call_transformation( + self, content: str, tools: List[ToolBlock] + ) -> Optional[ChatCompletionMessageToolCall]: + tool_call_names = self.get_tool_call_names(tools) + if not tool_call_names: + return None + + parsers = ( + self._parse_function_text_tool_call, + self._parse_tool_use_text_tool_call, + self._parse_bare_text_tool_call, + ) + for parser in parsers: + tool_name, arguments = parser(content) + resolved_tool_name = self._resolve_text_tool_call_name( + tool_name, tool_call_names + ) + if resolved_tool_name is not None: + return self._create_text_tool_call(resolved_tool_name, arguments) + + return None + def apply_tool_call_transformation_if_needed( self, message: Message, @@ -1810,7 +1942,13 @@ class AmazonConverseConfig(BaseConfig): message.content = None returned_finish_reason = "tool_calls" except Exception: - pass + tool_call = self._text_content_tool_call_transformation( + message.content, tools + ) + if tool_call is not None: + message.tool_calls = [tool_call] + message.content = None + returned_finish_reason = "tool_calls" return message, returned_finish_reason diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index 5f2ed3dc00f..932572c36b8 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -206,6 +206,128 @@ def test_apply_tool_call_transformation_if_needed(): ) +def _read_file_tool(): + return [ + { + "type": "function", + "function": { + "name": "read_file", + "description": "Read a file", + "parameters": { + "type": "object", + "properties": { + "path": {"type": "string"}, + "offset": {"type": "integer"}, + "length": {"type": "integer"}, + }, + "required": ["path"], + }, + }, + } + ] + + +def _assert_read_file_tool_call(transformed_message, finish_reason): + assert finish_reason == "tool_calls" + assert transformed_message.content is None + assert transformed_message.tool_calls is not None + assert len(transformed_message.tool_calls) == 1 + tool_call = transformed_message.tool_calls[0] + assert tool_call.type == "function" + assert tool_call.function.name == "read_file" + arguments = json.loads(tool_call.function.arguments) + assert arguments["path"] == "C:\\Projects\\redaigo\\scripts\\run_etf_v13.py" + assert arguments["offset"] in (0, "0") + assert arguments["length"] in (3000, "3000") + + +def test_apply_tool_call_transformation_parses_function_parameter_text(): + from litellm.types.utils import Message + + config = AmazonConverseConfig() + message = Message( + role="assistant", + content=( + "\n\n" + 'read_file\n' + 'C:\\Projects\\redaigo\\scripts\\run_etf_v13.py\n' + '0\n' + '3000\n' + "" + ), + ) + + transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( + message, _read_file_tool(), initial_finish_reason="stop" + ) + + _assert_read_file_tool_call(transformed_message, finish_reason) + + +def test_apply_tool_call_transformation_parses_tool_use_xml_text(): + from litellm.types.utils import Message + + config = AmazonConverseConfig() + message = Message( + role="assistant", + content=( + "\n" + "desktop-commander\n" + "read_file\n" + '{"path": "C:\\\\Projects\\\\redaigo\\\\scripts\\\\run_etf_v13.py", "offset": 0, "length": 3000}\n' + "" + ), + ) + + transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( + message, _read_file_tool(), initial_finish_reason="stop" + ) + + _assert_read_file_tool_call(transformed_message, finish_reason) + + +def test_apply_tool_call_transformation_parses_bare_tool_name_json_text(): + from litellm.types.utils import Message + + config = AmazonConverseConfig() + message = Message( + role="assistant", + content=( + "\nread_file\n" + '{"path": "C:\\\\Projects\\\\redaigo\\\\scripts\\\\run_etf_v13.py", "offset": 0, "length": 3000}' + ), + ) + + transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( + message, _read_file_tool(), initial_finish_reason="stop" + ) + + _assert_read_file_tool_call(transformed_message, finish_reason) + + +def test_apply_tool_call_transformation_ignores_text_for_unknown_tool_name(): + from litellm.types.utils import Message + + config = AmazonConverseConfig() + original_content = "\nread_file\n{}" + message = Message(role="assistant", content=original_content) + + transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( + message, + [ + { + "type": "function", + "function": {"name": "write_file", "parameters": {}}, + } + ], + initial_finish_reason="stop", + ) + + assert finish_reason == "stop" + assert transformed_message.content == original_content + assert transformed_message.tool_calls is None + + def test_transform_tool_call_with_cache_control(): from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig