mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(bedrock): parse text tool calls in converse responses
This commit is contained in:
parent
d04373f4ce
commit
fcfa5f636e
2 changed files with 262 additions and 2 deletions
|
|
@ -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"<function>(.*?)</function>", 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"""<parameter\s+name=[\"']([^\"']+)[\"']>(.*?)</parameter>""",
|
||||
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"<tool_use>(.*?)</tool_use>", 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"<tool_name>(.*?)</tool_name>", body, re.DOTALL | re.IGNORECASE
|
||||
)
|
||||
input_match = re.search(
|
||||
r"<input>(.*?)</input>", 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
|
||||
|
||||
|
|
|
|||
|
|
@ -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<function>\n"
|
||||
'<parameter name="command">read_file</parameter>\n'
|
||||
'<parameter name="path">C:\\Projects\\redaigo\\scripts\\run_etf_v13.py</parameter>\n'
|
||||
'<parameter name="offset">0</parameter>\n'
|
||||
'<parameter name="length">3000</parameter>\n'
|
||||
"</function>"
|
||||
),
|
||||
)
|
||||
|
||||
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=(
|
||||
"<tool_use>\n"
|
||||
"<server_name>desktop-commander</server_name>\n"
|
||||
"<tool_name>read_file</tool_name>\n"
|
||||
'<input>{"path": "C:\\\\Projects\\\\redaigo\\\\scripts\\\\run_etf_v13.py", "offset": 0, "length": 3000}</input>\n'
|
||||
"</tool_use>"
|
||||
),
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue