[Fix]: render tool_reference schema regardless of tool_result part order (#32522)
Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
e23ccb15f0
commit
09193bf36f
@@ -297,7 +297,7 @@ class AnthropicServing:
|
|||||||
|
|
||||||
def _convert_tool_result_content(
|
def _convert_tool_result_content(
|
||||||
content: Any,
|
content: Any,
|
||||||
) -> tuple[Union[str, list[dict]], str]:
|
) -> tuple[list[Union[str, list[dict]]], str]:
|
||||||
if isinstance(content, list):
|
if isinstance(content, list):
|
||||||
tool_content_parts = []
|
tool_content_parts = []
|
||||||
tool_text_parts = []
|
tool_text_parts = []
|
||||||
@@ -342,17 +342,29 @@ class AnthropicServing:
|
|||||||
)
|
)
|
||||||
|
|
||||||
tool_text = "\n".join(tool_text_parts)
|
tool_text = "\n".join(tool_text_parts)
|
||||||
|
# GLM templates expand references only at the start of a tool
|
||||||
|
# message, so isolate reference runs without changing part order.
|
||||||
|
tool_content_groups: list[list[dict]] = []
|
||||||
|
for part in tool_content_parts:
|
||||||
|
is_reference = part["type"] == "tool_reference"
|
||||||
if (
|
if (
|
||||||
len(tool_content_parts) == 1
|
not tool_content_groups
|
||||||
and tool_content_parts[0]["type"] == "text"
|
or (tool_content_groups[-1][0]["type"] == "tool_reference")
|
||||||
|
!= is_reference
|
||||||
):
|
):
|
||||||
return tool_content_parts[0]["text"], tool_text
|
tool_content_groups.append([])
|
||||||
if tool_content_parts:
|
tool_content_groups[-1].append(part)
|
||||||
return tool_content_parts, tool_text
|
|
||||||
return "", tool_text
|
tool_contents: list[Union[str, list[dict]]] = []
|
||||||
|
for group in tool_content_groups:
|
||||||
|
if len(group) == 1 and group[0]["type"] == "text":
|
||||||
|
tool_contents.append(group[0]["text"])
|
||||||
|
else:
|
||||||
|
tool_contents.append(group)
|
||||||
|
return tool_contents or [""], tool_text
|
||||||
|
|
||||||
tool_text = str(content) if content else ""
|
tool_text = str(content) if content else ""
|
||||||
return tool_text, tool_text
|
return [tool_text], tool_text
|
||||||
|
|
||||||
def _convert_assistant_thinking_blocks(
|
def _convert_assistant_thinking_blocks(
|
||||||
blocks: list[AnthropicContentBlock],
|
blocks: list[AnthropicContentBlock],
|
||||||
@@ -484,7 +496,7 @@ class AnthropicServing:
|
|||||||
tool_calls.append(tool_call)
|
tool_calls.append(tool_call)
|
||||||
|
|
||||||
elif block.type == "tool_result":
|
elif block.type == "tool_result":
|
||||||
tool_content, tool_text = _convert_tool_result_content(
|
tool_contents, tool_text = _convert_tool_result_content(
|
||||||
block.content
|
block.content
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -497,6 +509,7 @@ class AnthropicServing:
|
|||||||
# block must come AFTER that text in OpenAI form too).
|
# block must come AFTER that text in OpenAI form too).
|
||||||
if msg.role == "user":
|
if msg.role == "user":
|
||||||
_emit_user_message(content_parts)
|
_emit_user_message(content_parts)
|
||||||
|
for tool_content in tool_contents:
|
||||||
openai_messages.append(
|
openai_messages.append(
|
||||||
{
|
{
|
||||||
"role": "tool",
|
"role": "tool",
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from sglang.test.test_utils import maybe_stub_sgl_kernel
|
|||||||
maybe_stub_sgl_kernel() # must precede imports that may pull in sgl_kernel
|
maybe_stub_sgl_kernel() # must precede imports that may pull in sgl_kernel
|
||||||
|
|
||||||
from fastapi.responses import JSONResponse # noqa: E402
|
from fastapi.responses import JSONResponse # noqa: E402
|
||||||
|
from jinja2 import Environment # noqa: E402
|
||||||
|
|
||||||
from sglang.srt.entrypoints.anthropic.protocol import ( # noqa: E402
|
from sglang.srt.entrypoints.anthropic.protocol import ( # noqa: E402
|
||||||
AnthropicMessage,
|
AnthropicMessage,
|
||||||
@@ -140,6 +141,22 @@ class TestAnthropicServing(unittest.TestCase):
|
|||||||
"{{- message.role }}: {{ message.content }}\n"
|
"{{- message.role }}: {{ message.content }}\n"
|
||||||
"{%- endfor %}"
|
"{%- endfor %}"
|
||||||
)
|
)
|
||||||
|
GLM_TOOL_RESULT_TEMPLATE = """
|
||||||
|
{%- for message in messages if message.role == "tool" -%}
|
||||||
|
{%- if loop.first -%}<|observation|>{%- endif -%}
|
||||||
|
{%- if message.content is string -%}
|
||||||
|
<tool_response>{{ message.content }}</tool_response>
|
||||||
|
{%- elif message.content.0.type == "tool_reference" -%}
|
||||||
|
<tool_response><tools>
|
||||||
|
{%- for reference in message.content -%}
|
||||||
|
{%- for tool in tools if tool.function.name == reference.name -%}
|
||||||
|
{{ tool.function.name }}
|
||||||
|
{%- endfor -%}
|
||||||
|
{%- endfor -%}
|
||||||
|
</tools></tool_response>
|
||||||
|
{%- endif -%}
|
||||||
|
{%- endfor -%}
|
||||||
|
"""
|
||||||
|
|
||||||
def _serving(self, stream_lines=None, chat_template=None):
|
def _serving(self, stream_lines=None, chat_template=None):
|
||||||
return AnthropicServing(_FakeOpenAIServingChat(stream_lines, chat_template))
|
return AnthropicServing(_FakeOpenAIServingChat(stream_lines, chat_template))
|
||||||
@@ -154,6 +171,26 @@ class TestAnthropicServing(unittest.TestCase):
|
|||||||
data.update(overrides)
|
data.update(overrides)
|
||||||
return AnthropicMessagesRequest.model_validate(data)
|
return AnthropicMessagesRequest.model_validate(data)
|
||||||
|
|
||||||
|
def _tool_result_request(self, content, tools=None):
|
||||||
|
overrides = {
|
||||||
|
"stream": False,
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{
|
||||||
|
"type": "tool_result",
|
||||||
|
"tool_use_id": "call_1",
|
||||||
|
"content": content,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
if tools is not None:
|
||||||
|
overrides["tools"] = tools
|
||||||
|
return self._anthropic_request(**overrides)
|
||||||
|
|
||||||
def test_stream_closes_tool_block_before_text_delta(self):
|
def test_stream_closes_tool_block_before_text_delta(self):
|
||||||
serving = self._serving(
|
serving = self._serving(
|
||||||
[
|
[
|
||||||
@@ -373,6 +410,74 @@ class TestAnthropicServing(unittest.TestCase):
|
|||||||
self.assertIn("https://docs.sglang.ai", tool_message["content"])
|
self.assertIn("https://docs.sglang.ai", tool_message["content"])
|
||||||
self.assertIn("Anthropic API notes", tool_message["content"])
|
self.assertIn("Anthropic API notes", tool_message["content"])
|
||||||
|
|
||||||
|
def test_mixed_tool_reference_content_preserves_part_order(self):
|
||||||
|
request = self._tool_result_request(
|
||||||
|
[
|
||||||
|
{"type": "text", "text": "Tool loaded: Bash"},
|
||||||
|
{"type": "tool_reference", "tool_name": "Bash"},
|
||||||
|
{"type": "text", "text": "Ready"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
chat_request = self._serving()._convert_to_chat_completion_request(request)
|
||||||
|
messages = chat_request.model_dump(exclude_none=True)["messages"]
|
||||||
|
|
||||||
|
self.assertEqual([message["role"] for message in messages], ["tool"] * 3)
|
||||||
|
self.assertEqual(
|
||||||
|
[message["tool_call_id"] for message in messages], ["call_1"] * 3
|
||||||
|
)
|
||||||
|
self.assertEqual(messages[0]["content"], "Tool loaded: Bash")
|
||||||
|
self.assertEqual(
|
||||||
|
messages[1]["content"],
|
||||||
|
[{"type": "tool_reference", "name": "Bash"}],
|
||||||
|
)
|
||||||
|
self.assertEqual(messages[2]["content"], "Ready")
|
||||||
|
|
||||||
|
def test_mixed_tool_reference_content_renders_text_and_schema(self):
|
||||||
|
template = Environment().from_string(self.GLM_TOOL_RESULT_TEMPLATE)
|
||||||
|
request = self._tool_result_request(
|
||||||
|
[
|
||||||
|
{"type": "text", "text": "Tool loaded: Bash"},
|
||||||
|
{"type": "tool_reference", "tool_name": "Bash"},
|
||||||
|
],
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"name": "Bash",
|
||||||
|
"description": "Run a shell command",
|
||||||
|
"input_schema": {"type": "object", "properties": {}},
|
||||||
|
"defer_loading": True,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
chat_request = self._serving()._convert_to_chat_completion_request(request)
|
||||||
|
payload = chat_request.model_dump(exclude_none=True)
|
||||||
|
prompt = template.render(messages=payload["messages"], tools=payload["tools"])
|
||||||
|
|
||||||
|
self.assertIn("<tool_response>Tool loaded: Bash</tool_response>", prompt)
|
||||||
|
self.assertIn("<tools>Bash</tools>", prompt)
|
||||||
|
self.assertEqual(prompt.count("<|observation|>"), 1)
|
||||||
|
|
||||||
|
def test_reference_only_tool_result_remains_one_message(self):
|
||||||
|
request = self._tool_result_request(
|
||||||
|
[
|
||||||
|
{"type": "tool_reference", "tool_name": "Bash"},
|
||||||
|
{"type": "tool_reference", "tool_name": "Read"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
chat_request = self._serving()._convert_to_chat_completion_request(request)
|
||||||
|
messages = chat_request.model_dump(exclude_none=True)["messages"]
|
||||||
|
|
||||||
|
self.assertEqual(len(messages), 1)
|
||||||
|
self.assertEqual(
|
||||||
|
messages[0]["content"],
|
||||||
|
[
|
||||||
|
{"type": "tool_reference", "name": "Bash"},
|
||||||
|
{"type": "tool_reference", "name": "Read"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
def test_builtin_web_search_tool_without_schema_is_skipped(self):
|
def test_builtin_web_search_tool_without_schema_is_skipped(self):
|
||||||
request = AnthropicMessagesRequest.model_validate(
|
request = AnthropicMessagesRequest.model_validate(
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user