Support defer_loading field at function level for Chat Completions API (#22702)
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
92f28e9ba8
commit
28cfd3d272
@@ -33,7 +33,13 @@ class AnthropicContentBlock(BaseModel):
|
|||||||
"""Content block in message"""
|
"""Content block in message"""
|
||||||
|
|
||||||
type: Literal[
|
type: Literal[
|
||||||
"text", "image", "tool_use", "tool_result", "thinking", "redacted_thinking"
|
"text",
|
||||||
|
"image",
|
||||||
|
"tool_use",
|
||||||
|
"tool_result",
|
||||||
|
"tool_reference",
|
||||||
|
"thinking",
|
||||||
|
"redacted_thinking",
|
||||||
]
|
]
|
||||||
text: Optional[str] = None
|
text: Optional[str] = None
|
||||||
# For image content
|
# For image content
|
||||||
@@ -63,6 +69,7 @@ class AnthropicTool(BaseModel):
|
|||||||
name: str
|
name: str
|
||||||
description: Optional[str] = None
|
description: Optional[str] = None
|
||||||
input_schema: dict[str, Any]
|
input_schema: dict[str, Any]
|
||||||
|
defer_loading: Optional[bool] = None
|
||||||
|
|
||||||
@field_validator("input_schema")
|
@field_validator("input_schema")
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -146,6 +146,14 @@ class AnthropicServing:
|
|||||||
)
|
)
|
||||||
if image_part is not None:
|
if image_part is not None:
|
||||||
tool_content_parts.append(image_part)
|
tool_content_parts.append(image_part)
|
||||||
|
elif item_type == "tool_reference":
|
||||||
|
# Anthropic uses `tool_name`; the SGLang chat template
|
||||||
|
# matches on `name`. Translate at the boundary.
|
||||||
|
ref_name = item.get("tool_name") or item.get("name")
|
||||||
|
if ref_name:
|
||||||
|
tool_content_parts.append(
|
||||||
|
{"type": "tool_reference", "name": ref_name}
|
||||||
|
)
|
||||||
|
|
||||||
tool_text = "\n".join(tool_text_parts)
|
tool_text = "\n".join(tool_text_parts)
|
||||||
if (
|
if (
|
||||||
@@ -270,39 +278,41 @@ class AnthropicServing:
|
|||||||
|
|
||||||
chat_request = ChatCompletionRequest(**request_data)
|
chat_request = ChatCompletionRequest(**request_data)
|
||||||
|
|
||||||
# Convert tools
|
# Convert tools. Deferred tools stay in the list with defer_loading=True;
|
||||||
|
# the chat template hides them from the initial <tools> block and renders
|
||||||
|
# them on demand when a tool_reference block names them.
|
||||||
if anthropic_request.tools:
|
if anthropic_request.tools:
|
||||||
tools = []
|
chat_request.tools = [
|
||||||
for tool in anthropic_request.tools:
|
|
||||||
tools.append(
|
|
||||||
Tool(
|
Tool(
|
||||||
type="function",
|
type="function",
|
||||||
|
defer_loading=tool.defer_loading,
|
||||||
function={
|
function={
|
||||||
"name": tool.name,
|
"name": tool.name,
|
||||||
"description": tool.description or "",
|
"description": tool.description or "",
|
||||||
"parameters": tool.input_schema,
|
"parameters": tool.input_schema,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
)
|
for tool in anthropic_request.tools
|
||||||
chat_request.tools = tools
|
]
|
||||||
|
|
||||||
# Convert tool choice
|
# Convert tool choice
|
||||||
if anthropic_request.tool_choice is not None:
|
if anthropic_request.tool_choice is not None:
|
||||||
if anthropic_request.tool_choice.type == "none":
|
tc_type = anthropic_request.tool_choice.type
|
||||||
|
if tc_type == "none":
|
||||||
chat_request.tool_choice = "none"
|
chat_request.tool_choice = "none"
|
||||||
elif anthropic_request.tool_choice.type == "auto":
|
elif chat_request.tools:
|
||||||
|
if tc_type == "auto":
|
||||||
chat_request.tool_choice = "auto"
|
chat_request.tool_choice = "auto"
|
||||||
elif anthropic_request.tool_choice.type == "any":
|
elif tc_type == "any":
|
||||||
chat_request.tool_choice = "required"
|
chat_request.tool_choice = "required"
|
||||||
elif anthropic_request.tool_choice.type == "tool":
|
elif tc_type == "tool":
|
||||||
chat_request.tool_choice = ToolChoice(
|
chat_request.tool_choice = ToolChoice(
|
||||||
type="function",
|
type="function",
|
||||||
function=ToolChoiceFuncName(
|
function=ToolChoiceFuncName(
|
||||||
name=anthropic_request.tool_choice.name
|
name=anthropic_request.tool_choice.name
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
elif anthropic_request.tools:
|
elif chat_request.tools:
|
||||||
# Default to auto when tools are provided
|
|
||||||
chat_request.tool_choice = "auto"
|
chat_request.tool_choice = "auto"
|
||||||
|
|
||||||
return chat_request
|
return chat_request
|
||||||
|
|||||||
@@ -454,11 +454,23 @@ class ChatCompletionMessageContentAudioPart(BaseModel):
|
|||||||
audio_url: ChatCompletionMessageContentAudioURL
|
audio_url: ChatCompletionMessageContentAudioURL
|
||||||
|
|
||||||
|
|
||||||
|
class ChatCompletionMessageContentToolReferenceBlock(BaseModel):
|
||||||
|
# GLM-specific extension used alongside `defer_loading` tools. The chat
|
||||||
|
# template looks up `tools[*].function.name == tr.name` and renders the
|
||||||
|
# referenced tool schemas inline for the current turn. Not part of any
|
||||||
|
# OpenAI API; included here so Pydantic accepts the content through the
|
||||||
|
# Chat Completions path (the Anthropic endpoint translates its
|
||||||
|
# `tool_name` field to `name` before forwarding).
|
||||||
|
type: Literal["tool_reference"]
|
||||||
|
name: str
|
||||||
|
|
||||||
|
|
||||||
ChatCompletionMessageContentPart = Union[
|
ChatCompletionMessageContentPart = Union[
|
||||||
ChatCompletionMessageContentTextPart,
|
ChatCompletionMessageContentTextPart,
|
||||||
ChatCompletionMessageContentImagePart,
|
ChatCompletionMessageContentImagePart,
|
||||||
ChatCompletionMessageContentVideoPart,
|
ChatCompletionMessageContentVideoPart,
|
||||||
ChatCompletionMessageContentAudioPart,
|
ChatCompletionMessageContentAudioPart,
|
||||||
|
ChatCompletionMessageContentToolReferenceBlock,
|
||||||
]
|
]
|
||||||
|
|
||||||
# Rerank content types for multimodal reranking (e.g., Qwen3-VL-Reranker)
|
# Rerank content types for multimodal reranking (e.g., Qwen3-VL-Reranker)
|
||||||
@@ -528,6 +540,14 @@ class Function(BaseModel):
|
|||||||
name: str
|
name: str
|
||||||
parameters: Optional[object] = None
|
parameters: Optional[object] = None
|
||||||
strict: bool = False
|
strict: bool = False
|
||||||
|
defer_loading: Optional[bool] = None
|
||||||
|
|
||||||
|
@model_serializer(mode="wrap")
|
||||||
|
def _serialize(self, handler):
|
||||||
|
data = handler(self)
|
||||||
|
if self.defer_loading is None:
|
||||||
|
data.pop("defer_loading", None)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
class Tool(BaseModel):
|
class Tool(BaseModel):
|
||||||
@@ -535,6 +555,13 @@ class Tool(BaseModel):
|
|||||||
|
|
||||||
type: str = Field(default="function", examples=["function"])
|
type: str = Field(default="function", examples=["function"])
|
||||||
function: Function
|
function: Function
|
||||||
|
defer_loading: Optional[bool] = None
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def _propagate_defer_loading(self) -> "Tool":
|
||||||
|
if self.defer_loading is not None and self.function.defer_loading is None:
|
||||||
|
self.function.defer_loading = self.defer_loading
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
class ToolChoiceFuncName(BaseModel):
|
class ToolChoiceFuncName(BaseModel):
|
||||||
|
|||||||
@@ -201,6 +201,11 @@ def process_content_for_template_format(
|
|||||||
else:
|
else:
|
||||||
# Keep text content as-is for openai format
|
# Keep text content as-is for openai format
|
||||||
processed_content_parts.append(chunk)
|
processed_content_parts.append(chunk)
|
||||||
|
elif chunk_type == "tool_reference":
|
||||||
|
# GLM-specific extension: pass through so the chat template
|
||||||
|
# can match tool_reference.name against tools[*].function.name
|
||||||
|
# and render the referenced tool schemas inline.
|
||||||
|
processed_content_parts.append(chunk)
|
||||||
|
|
||||||
new_msg = {
|
new_msg = {
|
||||||
k: v for k, v in msg_dict.items() if v is not None and k != "content"
|
k: v for k, v in msg_dict.items() if v is not None and k != "content"
|
||||||
|
|||||||
@@ -24,8 +24,10 @@ from sglang.srt.entrypoints.openai.protocol import (
|
|||||||
ChatCompletionResponseChoice,
|
ChatCompletionResponseChoice,
|
||||||
ChatMessage,
|
ChatMessage,
|
||||||
CompletionRequest,
|
CompletionRequest,
|
||||||
|
Function,
|
||||||
ModelCard,
|
ModelCard,
|
||||||
ModelList,
|
ModelList,
|
||||||
|
Tool,
|
||||||
UsageInfo,
|
UsageInfo,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -336,6 +338,72 @@ class TestModelSerialization(unittest.TestCase):
|
|||||||
self.assertEqual(data["choices"][0]["hidden_states"], [0.1, 0.2, 0.3])
|
self.assertEqual(data["choices"][0]["hidden_states"], [0.1, 0.2, 0.3])
|
||||||
|
|
||||||
|
|
||||||
|
class TestFunctionDeferLoading(unittest.TestCase):
|
||||||
|
"""Test defer_loading field behavior on Function/Tool."""
|
||||||
|
|
||||||
|
def test_function_defaults_preserve_strict(self):
|
||||||
|
"""strict must default to False and be present in dumps so downstream
|
||||||
|
code (function_call_parser, chat templates) sees the expected shape."""
|
||||||
|
f = Function(name="foo")
|
||||||
|
data = f.model_dump()
|
||||||
|
self.assertEqual(data["name"], "foo")
|
||||||
|
self.assertEqual(data["strict"], False)
|
||||||
|
self.assertNotIn("defer_loading", data)
|
||||||
|
|
||||||
|
def test_function_defer_loading_true_serialized(self):
|
||||||
|
f = Function(name="foo", defer_loading=True)
|
||||||
|
data = f.model_dump()
|
||||||
|
self.assertTrue(data["defer_loading"])
|
||||||
|
self.assertEqual(data["strict"], False)
|
||||||
|
|
||||||
|
def test_function_defer_loading_false_serialized(self):
|
||||||
|
"""defer_loading=False is an explicit value and must be preserved."""
|
||||||
|
f = Function(name="foo", defer_loading=False)
|
||||||
|
data = f.model_dump()
|
||||||
|
self.assertIn("defer_loading", data)
|
||||||
|
self.assertFalse(data["defer_loading"])
|
||||||
|
|
||||||
|
def test_tool_level_defer_loading_propagates_to_function(self):
|
||||||
|
"""defer_loading at the Tool level should propagate to Function."""
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
defer_loading=True,
|
||||||
|
function={"name": "search_db"},
|
||||||
|
)
|
||||||
|
self.assertTrue(tool.function.defer_loading)
|
||||||
|
data = tool.model_dump()
|
||||||
|
self.assertTrue(data["function"]["defer_loading"])
|
||||||
|
|
||||||
|
def test_function_level_defer_loading_wins_over_tool_level(self):
|
||||||
|
"""Explicit function-level value is preserved when both set."""
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
defer_loading=True,
|
||||||
|
function={"name": "search_db", "defer_loading": False},
|
||||||
|
)
|
||||||
|
self.assertFalse(tool.function.defer_loading)
|
||||||
|
|
||||||
|
def test_tool_reference_content_part_accepted(self):
|
||||||
|
"""Chat completion should accept tool_reference content on tool-role
|
||||||
|
messages (GLM-specific extension consumed by the chat template)."""
|
||||||
|
messages = [
|
||||||
|
{
|
||||||
|
"role": "tool",
|
||||||
|
"tool_call_id": "call_1",
|
||||||
|
"content": [
|
||||||
|
{"type": "tool_reference", "name": "search_db"},
|
||||||
|
{"type": "text", "text": "ok"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
]
|
||||||
|
request = ChatCompletionRequest(model="test-model", messages=messages)
|
||||||
|
parts = request.messages[0].content
|
||||||
|
self.assertEqual(len(parts), 2)
|
||||||
|
self.assertEqual(parts[0].type, "tool_reference")
|
||||||
|
self.assertEqual(parts[0].name, "search_db")
|
||||||
|
self.assertEqual(parts[1].type, "text")
|
||||||
|
|
||||||
|
|
||||||
class TestValidationEdgeCases(unittest.TestCase):
|
class TestValidationEdgeCases(unittest.TestCase):
|
||||||
"""Test edge cases and validation scenarios"""
|
"""Test edge cases and validation scenarios"""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user