diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 410a92abb4a..76da11dd337 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -527,9 +527,9 @@ class AmazonConverseConfig(BaseConfig): ## Filter out 'cross-region' from model name base_model = BedrockModelInfo.get_base_model(model) - is_anthropic_model = base_model.startswith("anthropic") or base_model.startswith( - "claude-" - ) + is_anthropic_model = base_model.startswith( + "anthropic" + ) or base_model.startswith("claude-") if ( is_anthropic_model @@ -1832,7 +1832,11 @@ class AmazonConverseConfig(BaseConfig): if tool_call_match is None: return None, {} - tool_call = self._parse_tool_call_json_arguments(tool_call_match.group(1).strip()) + tool_call = self._parse_tool_call_json_arguments( + tool_call_match.group(1).strip() + ) + if tool_call is None: + return None, {} tool_name = tool_call.get("name") arguments = tool_call.get("arguments", {}) if not isinstance(tool_name, str): @@ -1883,9 +1887,12 @@ class AmazonConverseConfig(BaseConfig): 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( + arguments = self._parse_tool_call_json_arguments( input_match.group(1).strip() if input_match is not None else "" ) + if arguments is None: + return None, {} + return tool_name_match.group(1).strip(), arguments return self._parse_bare_text_tool_call(body) @@ -1898,20 +1905,22 @@ class AmazonConverseConfig(BaseConfig): tool_name = lines[0] arguments = self._parse_tool_call_json_arguments("\n".join(lines[1:])) - if arguments == {}: + if arguments is None: return None, {} return tool_name, arguments - def _parse_tool_call_json_arguments(self, json_text: str) -> dict[str, Any]: + def _parse_tool_call_json_arguments( + self, json_text: str + ) -> Optional[dict[str, Any]]: if not json_text: return {} try: parsed_arguments = json.loads(json_text) except Exception: - return {} + return None if not isinstance(parsed_arguments, dict): - return {} + return None return parsed_arguments def _text_content_tool_call_transformation( 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 e355b171765..609e49bcbd8 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -257,8 +257,10 @@ def test_apply_tool_call_transformation_parses_function_parameter_text(): ), ) - transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( - message, _read_file_tool(), initial_finish_reason="stop" + 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) @@ -279,8 +281,10 @@ def test_apply_tool_call_transformation_parses_tool_use_xml_text(): ), ) - transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( - message, _read_file_tool(), initial_finish_reason="stop" + 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) @@ -298,13 +302,34 @@ def test_apply_tool_call_transformation_parses_bare_tool_name_json_text(): ), ) - transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( - message, _read_file_tool(), initial_finish_reason="stop" + 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_zero_arg_tool_call(): + from litellm.types.utils import Message + + config = AmazonConverseConfig() + message = Message(role="assistant", content="\nread_file\n{}") + + transformed_message, finish_reason = ( + config.apply_tool_call_transformation_if_needed( + message, _read_file_tool(), initial_finish_reason="stop" + ) + ) + + assert finish_reason == "tool_calls" + assert transformed_message.content is None + assert transformed_message.tool_calls is not None + assert transformed_message.tool_calls[0].function.name == "read_file" + assert json.loads(transformed_message.tool_calls[0].function.arguments) == {} + + def test_apply_tool_call_transformation_parses_tool_call_json_tag_text(): from litellm.types.utils import Message @@ -318,8 +343,10 @@ def test_apply_tool_call_transformation_parses_tool_call_json_tag_text(): ), ) - transformed_message, finish_reason = config.apply_tool_call_transformation_if_needed( - message, _read_file_tool(), initial_finish_reason="stop" + 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) @@ -332,15 +359,17 @@ def test_apply_tool_call_transformation_ignores_text_for_unknown_tool_name(): 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", + 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"