Fix Responses API request handling (#25881)
Co-authored-by: Kai-Hsun Chen <kaihsun@apache.org> Co-authored-by: Kristin Cowalcijk <kristincowalcijk@gmail.com> Co-authored-by: aerosta <63026763+aerosta@users.noreply.github.com> Co-authored-by: glaziermag <glaziermag@users.noreply.github.com> Co-authored-by: Blake Ledden <blake.ledden@gmail.com> Co-authored-by: PanJason <pyyjason@gmail.com> Co-authored-by: Leoyzen <leoyzen@gmail.com> Co-authored-by: kennyu <966806+kennyu@users.noreply.github.com>
This commit is contained in:
co-authored by
Kai-Hsun Chen
Kristin Cowalcijk
aerosta
glaziermag
Blake Ledden
PanJason
Leoyzen
kennyu
parent
b3270264e4
commit
85712fa5b0
@@ -3,6 +3,7 @@
|
|||||||
# Adapted from vLLM: https://github.com/vllm-project/vllm/blob/1b9902806915040ac9b3029f2ab7522ec505afc3/vllm/entrypoints/harmony_utils.py
|
# Adapted from vLLM: https://github.com/vllm-project/vllm/blob/1b9902806915040ac9b3029f2ab7522ec505afc3/vllm/entrypoints/harmony_utils.py
|
||||||
# Slight differences in processing chat messages
|
# Slight differences in processing chat messages
|
||||||
import datetime
|
import datetime
|
||||||
|
import logging
|
||||||
from collections.abc import Iterable
|
from collections.abc import Iterable
|
||||||
from typing import Literal, Optional, Union
|
from typing import Literal, Optional, Union
|
||||||
|
|
||||||
@@ -42,6 +43,8 @@ from openai_harmony import (
|
|||||||
from sglang.srt.entrypoints.openai.protocol import ResponseInputOutputItem
|
from sglang.srt.entrypoints.openai.protocol import ResponseInputOutputItem
|
||||||
from sglang.srt.utils import random_uuid
|
from sglang.srt.utils import random_uuid
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
REASONING_EFFORT = {
|
REASONING_EFFORT = {
|
||||||
"high": ReasoningEffort.HIGH,
|
"high": ReasoningEffort.HIGH,
|
||||||
"medium": ReasoningEffort.MEDIUM,
|
"medium": ReasoningEffort.MEDIUM,
|
||||||
@@ -92,13 +95,22 @@ def get_developer_message(
|
|||||||
if tools is not None:
|
if tools is not None:
|
||||||
function_tools = []
|
function_tools = []
|
||||||
for tool in tools:
|
for tool in tools:
|
||||||
if tool.type in ("web_search_preview", "code_interpreter"):
|
if tool.type in (
|
||||||
|
"web_search",
|
||||||
|
"web_search_preview",
|
||||||
|
"code_interpreter",
|
||||||
|
):
|
||||||
# These are built-in tools that are added to the system message.
|
# These are built-in tools that are added to the system message.
|
||||||
pass
|
pass
|
||||||
elif tool.type == "function":
|
elif tool.type == "function":
|
||||||
function_tools.append(tool)
|
function_tools.append(tool)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"tool type {tool.type} not supported")
|
# No harmony prompt template for the remaining built-ins;
|
||||||
|
# drop them so the request still runs.
|
||||||
|
logger.debug(
|
||||||
|
"harmony: ignoring unsupported response tool type %r",
|
||||||
|
tool.type,
|
||||||
|
)
|
||||||
if function_tools:
|
if function_tools:
|
||||||
function_tool_descriptions = [
|
function_tool_descriptions = [
|
||||||
ToolDescription.new(
|
ToolDescription.new(
|
||||||
@@ -139,7 +151,16 @@ def parse_response_input(
|
|||||||
if isinstance(content, str):
|
if isinstance(content, str):
|
||||||
msg = Message.from_role_and_content(role, text_prefix + content)
|
msg = Message.from_role_and_content(role, text_prefix + content)
|
||||||
else:
|
else:
|
||||||
contents = [TextContent(text=text_prefix + c["text"]) for c in content]
|
# Filter to text parts first, then enumerate, so the surviving first
|
||||||
|
# text chunk always carries the system→developer text_prefix even if
|
||||||
|
# earlier parts were non-text (image/audio) and got dropped.
|
||||||
|
text_chunks = [
|
||||||
|
c for c in content if c.get("type") in ("text", "input_text")
|
||||||
|
]
|
||||||
|
contents = [
|
||||||
|
TextContent(text=(text_prefix if i == 0 else "") + c.get("text", ""))
|
||||||
|
for i, c in enumerate(text_chunks)
|
||||||
|
]
|
||||||
msg = Message.from_role_and_contents(role, contents)
|
msg = Message.from_role_and_contents(role, contents)
|
||||||
elif response_msg["type"] == "function_call_output":
|
elif response_msg["type"] == "function_call_output":
|
||||||
call_id = response_msg["call_id"]
|
call_id = response_msg["call_id"]
|
||||||
|
|||||||
@@ -1799,12 +1799,11 @@ async def v1_score_request(request: ScoringRequest, raw_request: Request):
|
|||||||
|
|
||||||
|
|
||||||
@app.post("/v1/responses", dependencies=[Depends(validate_json_request)])
|
@app.post("/v1/responses", dependencies=[Depends(validate_json_request)])
|
||||||
async def v1_responses_request(request: dict, raw_request: Request):
|
async def v1_responses_request(request: ResponsesRequest, raw_request: Request):
|
||||||
"""Endpoint for the responses API with reasoning support."""
|
"""Endpoint for the responses API with reasoning support."""
|
||||||
|
|
||||||
request_obj = ResponsesRequest(**request)
|
|
||||||
result = await raw_request.app.state.openai_serving_responses.create_responses(
|
result = await raw_request.app.state.openai_serving_responses.create_responses(
|
||||||
request_obj, raw_request
|
request, raw_request
|
||||||
)
|
)
|
||||||
|
|
||||||
# Handle streaming responses
|
# Handle streaming responses
|
||||||
|
|||||||
@@ -1284,14 +1284,46 @@ class ResponseReasoningParam(BaseModel):
|
|||||||
default="medium",
|
default="medium",
|
||||||
description="Constrains effort on reasoning for reasoning models.",
|
description="Constrains effort on reasoning for reasoning models.",
|
||||||
)
|
)
|
||||||
|
summary: Optional[Literal["auto", "concise", "detailed"]] = Field(
|
||||||
|
default=None,
|
||||||
|
description="Include a summary of the model's reasoning trace on the response.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Only ``function`` / ``web_search*`` / ``code_interpreter`` are wired to
|
||||||
|
# execution paths; the rest pass validation so clients aren't rejected.
|
||||||
|
RESPONSE_TOOL_TYPES = Literal[
|
||||||
|
"function",
|
||||||
|
"web_search",
|
||||||
|
"web_search_preview",
|
||||||
|
"code_interpreter",
|
||||||
|
"file_search",
|
||||||
|
"image_generation",
|
||||||
|
"computer_use_preview",
|
||||||
|
"local_shell",
|
||||||
|
"mcp",
|
||||||
|
"custom",
|
||||||
|
"namespace",
|
||||||
|
"tool_search",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
class ResponseTool(BaseModel):
|
class ResponseTool(BaseModel):
|
||||||
"""Tool definition for responses."""
|
"""Tool definition for responses."""
|
||||||
|
|
||||||
type: Literal["web_search_preview", "code_interpreter"] = Field(
|
type: RESPONSE_TOOL_TYPES = Field(description="Type of tool to enable")
|
||||||
description="Type of tool to enable"
|
name: Optional[str] = None
|
||||||
)
|
description: Optional[str] = None
|
||||||
|
parameters: Optional[Dict[str, Any]] = None
|
||||||
|
strict: bool = False
|
||||||
|
# Inner schemas for ``namespace`` tools.
|
||||||
|
tools: Optional[List[Dict[str, Any]]] = None
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_function_tool(self) -> "ResponseTool":
|
||||||
|
if self.type == "function" and not self.name:
|
||||||
|
raise ValueError("Function tools must include a name.")
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
ResponseInputOutputItem: TypeAlias = Union[
|
ResponseInputOutputItem: TypeAlias = Union[
|
||||||
@@ -1318,7 +1350,9 @@ class ResponsesRequest(BaseModel):
|
|||||||
]
|
]
|
||||||
]
|
]
|
||||||
] = None
|
] = None
|
||||||
input: Union[str, List[ResponseInputOutputItem]]
|
# Accept dict-shaped items as the loose arm; downstream normalization
|
||||||
|
# handles replayed shapes that don't satisfy every openai TypedDict.
|
||||||
|
input: Union[str, List[ResponseInputOutputItem], List[Dict[str, Any]]]
|
||||||
instructions: Optional[str] = None
|
instructions: Optional[str] = None
|
||||||
max_output_tokens: Optional[int] = None
|
max_output_tokens: Optional[int] = None
|
||||||
max_tool_calls: Optional[int] = None
|
max_tool_calls: Optional[int] = None
|
||||||
@@ -1352,13 +1386,13 @@ class ResponsesRequest(BaseModel):
|
|||||||
default=None, description="Cache salt for request caching"
|
default=None, description="Cache salt for request caching"
|
||||||
)
|
)
|
||||||
|
|
||||||
# SGLang-specific sampling parameters
|
# SGLang sampling extras. ``None`` defers to ``--preferred-sampling-params``.
|
||||||
frequency_penalty: float = 0.0
|
frequency_penalty: float = 0.0
|
||||||
presence_penalty: float = 0.0
|
presence_penalty: float = 0.0
|
||||||
stop: Optional[Union[str, List[str]]] = None
|
stop: Optional[Union[str, List[str]]] = None
|
||||||
top_k: int = -1
|
top_k: Optional[int] = None
|
||||||
min_p: float = 0.0
|
min_p: Optional[float] = None
|
||||||
repetition_penalty: float = 1.0
|
repetition_penalty: Optional[float] = None
|
||||||
|
|
||||||
# Default sampling parameters
|
# Default sampling parameters
|
||||||
_DEFAULT_SAMPLING_PARAMS = {
|
_DEFAULT_SAMPLING_PARAMS = {
|
||||||
@@ -1369,8 +1403,57 @@ class ResponsesRequest(BaseModel):
|
|||||||
"repetition_penalty": 1.0,
|
"repetition_penalty": 1.0,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@model_validator(mode="before")
|
||||||
|
@classmethod
|
||||||
|
def normalize_responses_input(cls, values):
|
||||||
|
if not isinstance(values, dict):
|
||||||
|
return values
|
||||||
|
|
||||||
|
input_value = values.get("input")
|
||||||
|
if not isinstance(input_value, list):
|
||||||
|
return values
|
||||||
|
|
||||||
|
values = values.copy()
|
||||||
|
values["input"] = [
|
||||||
|
cls._normalize_input_item_for_validation(item) for item in input_value
|
||||||
|
]
|
||||||
|
return values
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_input_item_for_validation(item):
|
||||||
|
if not isinstance(item, dict):
|
||||||
|
return item
|
||||||
|
|
||||||
|
content = item.get("content")
|
||||||
|
if not isinstance(content, list):
|
||||||
|
return item
|
||||||
|
|
||||||
|
item = item.copy()
|
||||||
|
item["content"] = [
|
||||||
|
ResponsesRequest._normalize_content_part_for_validation(part)
|
||||||
|
for part in content
|
||||||
|
]
|
||||||
|
return item
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_content_part_for_validation(part):
|
||||||
|
if not isinstance(part, dict):
|
||||||
|
return part
|
||||||
|
|
||||||
|
part_type = part.get("type")
|
||||||
|
if part_type != "input_image" or part.get("detail") is not None:
|
||||||
|
return part
|
||||||
|
|
||||||
|
part = part.copy()
|
||||||
|
part["detail"] = "auto"
|
||||||
|
return part
|
||||||
|
|
||||||
def to_sampling_params(
|
def to_sampling_params(
|
||||||
self, default_max_tokens: int, default_params: Optional[Dict] = None
|
self,
|
||||||
|
default_max_tokens: int,
|
||||||
|
default_params: Optional[Dict] = None,
|
||||||
|
stop: Optional[Union[str, List[str]]] = None,
|
||||||
|
tool_call_constraint: Optional[ToolCallConstraint] = None,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""Convert to sampling parameters for generation."""
|
"""Convert to sampling parameters for generation."""
|
||||||
if default_params is None:
|
if default_params is None:
|
||||||
@@ -1382,10 +1465,9 @@ class ResponsesRequest(BaseModel):
|
|||||||
else:
|
else:
|
||||||
max_tokens = default_max_tokens
|
max_tokens = default_max_tokens
|
||||||
|
|
||||||
# Avoid exceed the context length by minus 2 token
|
# Headroom for BOS/EOS the engine appends on top of prompt+budget.
|
||||||
max_tokens -= 2
|
max_tokens -= 2
|
||||||
|
|
||||||
# Get parameters with defaults
|
|
||||||
temperature = self.temperature
|
temperature = self.temperature
|
||||||
if temperature is None:
|
if temperature is None:
|
||||||
temperature = default_params.get(
|
temperature = default_params.get(
|
||||||
@@ -1396,23 +1478,51 @@ class ResponsesRequest(BaseModel):
|
|||||||
if top_p is None:
|
if top_p is None:
|
||||||
top_p = default_params.get("top_p", self._DEFAULT_SAMPLING_PARAMS["top_p"])
|
top_p = default_params.get("top_p", self._DEFAULT_SAMPLING_PARAMS["top_p"])
|
||||||
|
|
||||||
params = {
|
# Omit None entries so they fall through to ``--preferred-sampling-params``
|
||||||
|
# rather than overriding it with a literal default.
|
||||||
|
params: dict[str, Any] = {
|
||||||
"max_new_tokens": max_tokens,
|
"max_new_tokens": max_tokens,
|
||||||
"temperature": temperature,
|
"temperature": temperature,
|
||||||
"top_p": top_p,
|
"top_p": top_p,
|
||||||
"frequency_penalty": self.frequency_penalty,
|
"frequency_penalty": self.frequency_penalty,
|
||||||
"presence_penalty": self.presence_penalty,
|
"presence_penalty": self.presence_penalty,
|
||||||
"stop": self.stop,
|
"stop": self.stop if stop is None else stop,
|
||||||
"top_k": self.top_k,
|
|
||||||
"min_p": self.min_p,
|
|
||||||
"repetition_penalty": self.repetition_penalty,
|
|
||||||
}
|
}
|
||||||
|
if self.top_k is not None:
|
||||||
|
params["top_k"] = self.top_k
|
||||||
|
if self.min_p is not None:
|
||||||
|
params["min_p"] = self.min_p
|
||||||
|
if self.repetition_penalty is not None:
|
||||||
|
params["repetition_penalty"] = self.repetition_penalty
|
||||||
|
|
||||||
# Apply any additional default parameters
|
# Apply any additional default parameters
|
||||||
for key, value in default_params.items():
|
for key, value in default_params.items():
|
||||||
if key not in params or params[key] is None:
|
if key not in params or params[key] is None:
|
||||||
params[key] = value
|
params[key] = value
|
||||||
|
|
||||||
|
has_existing_constraints = (
|
||||||
|
params.get("regex")
|
||||||
|
or params.get("ebnf")
|
||||||
|
or params.get("structural_tag")
|
||||||
|
or params.get("json_schema")
|
||||||
|
)
|
||||||
|
if tool_call_constraint and has_existing_constraints:
|
||||||
|
# Refuse rather than silently drop the tool-call grammar.
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot combine tool calls with constrained decoding "
|
||||||
|
"(regex / ebnf / structural_tag / json_schema). Remove one."
|
||||||
|
)
|
||||||
|
if tool_call_constraint:
|
||||||
|
constraint_type, constraint_value = tool_call_constraint
|
||||||
|
if constraint_type in ("structural_tag", "json_schema"):
|
||||||
|
params[constraint_type] = convert_json_schema_to_str(
|
||||||
|
constraint_value.model_dump(by_alias=True)
|
||||||
|
if hasattr(constraint_value, "model_dump")
|
||||||
|
else constraint_value
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
params[constraint_type] = constraint_value
|
||||||
|
|
||||||
return params
|
return params
|
||||||
|
|
||||||
|
|
||||||
@@ -1515,7 +1625,11 @@ class ResponsesResponse(BaseModel):
|
|||||||
output=output,
|
output=output,
|
||||||
status=status,
|
status=status,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
parallel_tool_calls=request.parallel_tool_calls or True,
|
parallel_tool_calls=(
|
||||||
|
request.parallel_tool_calls
|
||||||
|
if request.parallel_tool_calls is not None
|
||||||
|
else True
|
||||||
|
),
|
||||||
tool_choice=request.tool_choice,
|
tool_choice=request.tool_choice,
|
||||||
tools=request.tools,
|
tools=request.tools,
|
||||||
# fields for parity with v1/responses
|
# fields for parity with v1/responses
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -157,14 +157,16 @@ def process_content_for_template_format(
|
|||||||
if isinstance(chunk, dict):
|
if isinstance(chunk, dict):
|
||||||
chunk_type = chunk.get("type")
|
chunk_type = chunk.get("type")
|
||||||
|
|
||||||
if chunk_type == "image_url":
|
if chunk_type in ("image_url", "input_image"):
|
||||||
image_obj = chunk.get("image_url") or {}
|
image_obj = chunk.get("image_url") or {}
|
||||||
|
if isinstance(image_obj, str):
|
||||||
|
image_obj = {"url": image_obj, "detail": chunk.get("detail")}
|
||||||
mdp = image_obj.get("max_dynamic_patch", None)
|
mdp = image_obj.get("max_dynamic_patch", None)
|
||||||
# Also allow flat style: chunk["max_dynamic_patch"]
|
# Also allow flat style: chunk["max_dynamic_patch"]
|
||||||
image_data.append(
|
image_data.append(
|
||||||
ImageData(
|
ImageData(
|
||||||
url=image_obj["url"],
|
url=image_obj["url"],
|
||||||
detail=image_obj.get("detail", "auto"),
|
detail=image_obj.get("detail") or "auto",
|
||||||
max_dynamic_patch=mdp,
|
max_dynamic_patch=mdp,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -194,13 +196,15 @@ def process_content_for_template_format(
|
|||||||
audio_data.append(chunk["audio_url"]["url"])
|
audio_data.append(chunk["audio_url"]["url"])
|
||||||
# Normalize to simple 'audio' type
|
# Normalize to simple 'audio' type
|
||||||
processed_content_parts.append({"type": "audio"})
|
processed_content_parts.append({"type": "audio"})
|
||||||
elif chunk_type == "text":
|
elif chunk_type in ("text", "input_text"):
|
||||||
# For v32 encoding, collect text parts separately
|
# For v32 encoding, collect text parts separately
|
||||||
if use_dpsk_v32_encoding:
|
if use_dpsk_v32_encoding:
|
||||||
text_parts.append(chunk["text"])
|
text_parts.append(chunk["text"])
|
||||||
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(
|
||||||
|
{"type": "text", "text": chunk["text"]}
|
||||||
|
)
|
||||||
elif chunk_type == "tool_reference":
|
elif chunk_type == "tool_reference":
|
||||||
# GLM-specific extension: pass through so the chat template
|
# GLM-specific extension: pass through so the chat template
|
||||||
# can match tool_reference.name against tools[*].function.name
|
# can match tool_reference.name against tools[*].function.name
|
||||||
@@ -220,7 +224,7 @@ def process_content_for_template_format(
|
|||||||
# String format: flatten to text only (for templates like DeepSeek)
|
# String format: flatten to text only (for templates like DeepSeek)
|
||||||
text_parts = []
|
text_parts = []
|
||||||
for chunk in msg_dict["content"]:
|
for chunk in msg_dict["content"]:
|
||||||
if isinstance(chunk, dict) and chunk.get("type") == "text":
|
if isinstance(chunk, dict) and chunk.get("type") in ("text", "input_text"):
|
||||||
text_parts.append(chunk["text"])
|
text_parts.append(chunk["text"])
|
||||||
# Note: For string format, we ignore images/audio since the template
|
# Note: For string format, we ignore images/audio since the template
|
||||||
# doesn't expect structured content - multimodal placeholders would
|
# doesn't expect structured content - multimodal placeholders would
|
||||||
|
|||||||
@@ -257,6 +257,9 @@ class Qwen3Detector(BaseReasoningFormatDetector):
|
|||||||
think_excluded_tokens=think_excluded_tokens,
|
think_excluded_tokens=think_excluded_tokens,
|
||||||
force_reasoning=force_reasoning,
|
force_reasoning=force_reasoning,
|
||||||
stream_reasoning=stream_reasoning,
|
stream_reasoning=stream_reasoning,
|
||||||
|
# Qwen3.5 sometimes opens ``<tool_call>`` without closing
|
||||||
|
# ``</think>``; treat it as an implicit reasoning close.
|
||||||
|
tool_start_token="<tool_call>",
|
||||||
continue_final_message=continue_final_message,
|
continue_final_message=continue_final_message,
|
||||||
previous_content=previous_content,
|
previous_content=previous_content,
|
||||||
thinks_internally=True,
|
thinks_internally=True,
|
||||||
|
|||||||
@@ -0,0 +1,146 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from utils import make_serving # noqa: F401 — bootstrap import
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import ResponsesRequest, UsageInfo
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=4, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class ResponsesRequestTestCase(unittest.TestCase):
|
||||||
|
def test_function_tool_accepted(self):
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="call the tool",
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "lookup",
|
||||||
|
"description": "Look up a value.",
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"key": {"type": "string"}},
|
||||||
|
"required": ["key"],
|
||||||
|
},
|
||||||
|
"strict": True,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(request.tools[0].type, "function")
|
||||||
|
self.assertEqual(request.tools[0].name, "lookup")
|
||||||
|
self.assertTrue(request.tools[0].strict)
|
||||||
|
|
||||||
|
def test_function_tool_requires_name(self):
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
ResponsesRequest(
|
||||||
|
model="x", input="hi", tools=[{"type": "function"}], store=False
|
||||||
|
)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
tools=[{"type": "function", "name": ""}],
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_extended_tool_types_accepted(self):
|
||||||
|
for tool_type in (
|
||||||
|
"web_search",
|
||||||
|
"web_search_preview",
|
||||||
|
"code_interpreter",
|
||||||
|
"file_search",
|
||||||
|
"image_generation",
|
||||||
|
"computer_use_preview",
|
||||||
|
"local_shell",
|
||||||
|
"mcp",
|
||||||
|
"custom",
|
||||||
|
"namespace",
|
||||||
|
):
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
tools=[{"type": tool_type}],
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
self.assertEqual(request.tools[0].type, tool_type)
|
||||||
|
|
||||||
|
def test_namespace_tool_carries_inner_tools_list(self):
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "namespace",
|
||||||
|
"name": "codex",
|
||||||
|
"tools": [
|
||||||
|
{"type": "function", "name": "apply_patch"},
|
||||||
|
{"type": "function", "name": "shell"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
self.assertEqual(request.tools[0].type, "namespace")
|
||||||
|
self.assertEqual(len(request.tools[0].tools), 2)
|
||||||
|
self.assertEqual(request.tools[0].tools[0]["name"], "apply_patch")
|
||||||
|
|
||||||
|
|
||||||
|
class ResponsesSamplingParamsTestCase(unittest.TestCase):
|
||||||
|
def test_processed_stop_and_tool_constraint_propagate(self):
|
||||||
|
request = ResponsesRequest(model="x", input="call the tool", store=False)
|
||||||
|
params = request.to_sampling_params(
|
||||||
|
default_max_tokens=128,
|
||||||
|
default_params={},
|
||||||
|
stop=["</s>"],
|
||||||
|
tool_call_constraint=("json_schema", {"type": "object"}),
|
||||||
|
)
|
||||||
|
self.assertEqual(params["stop"], ["</s>"])
|
||||||
|
self.assertEqual(params["json_schema"], '{"type": "object"}')
|
||||||
|
|
||||||
|
def test_constraint_conflict_raises(self):
|
||||||
|
request = ResponsesRequest(model="x", input="hi", store=False)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
request.to_sampling_params(
|
||||||
|
default_max_tokens=128,
|
||||||
|
default_params={"json_schema": '{"type": "object"}'},
|
||||||
|
tool_call_constraint=("json_schema", {"type": "object"}),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_structural_tag_with_model_dump(self):
|
||||||
|
class _FakeStructuralTag:
|
||||||
|
def model_dump(self, by_alias=False):
|
||||||
|
return {"type": "structural_tag"}
|
||||||
|
|
||||||
|
request = ResponsesRequest(model="x", input="hi", store=False)
|
||||||
|
params = request.to_sampling_params(
|
||||||
|
default_max_tokens=128,
|
||||||
|
default_params={},
|
||||||
|
tool_call_constraint=("structural_tag", _FakeStructuralTag()),
|
||||||
|
)
|
||||||
|
self.assertEqual(params["structural_tag"], '{"type": "structural_tag"}')
|
||||||
|
|
||||||
|
|
||||||
|
class ResponsesResponseFromRequestTestCase(unittest.TestCase):
|
||||||
|
def test_parallel_tool_calls_false_preserved(self):
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import ResponsesResponse
|
||||||
|
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x", input="hi", parallel_tool_calls=False, store=False
|
||||||
|
)
|
||||||
|
response = ResponsesResponse.from_request(
|
||||||
|
request,
|
||||||
|
sampling_params={},
|
||||||
|
model_name="x",
|
||||||
|
created_time=0,
|
||||||
|
output=[],
|
||||||
|
status="completed",
|
||||||
|
usage=UsageInfo(prompt_tokens=1, completion_tokens=1, total_tokens=2),
|
||||||
|
)
|
||||||
|
self.assertFalse(response.parallel_tool_calls)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,529 @@
|
|||||||
|
import asyncio
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
|
from openai.types.responses import (
|
||||||
|
ResponseOutputMessage,
|
||||||
|
ResponseOutputText,
|
||||||
|
ResponseReasoningItem,
|
||||||
|
)
|
||||||
|
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
|
||||||
|
from utils import make_serving
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.context import SimpleContext
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
|
MessageProcessingResult,
|
||||||
|
RequestResponseMetadata,
|
||||||
|
ResponsesRequest,
|
||||||
|
)
|
||||||
|
from sglang.srt.entrypoints.openai.serving_responses import OpenAIServingResponses
|
||||||
|
from sglang.srt.function_call.core_types import ToolCallItem
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class InputMessageConstructionTestCase(unittest.TestCase):
|
||||||
|
def test_previous_response_replays_assistant_text_not_instructions(self):
|
||||||
|
serving = make_serving()
|
||||||
|
prev_response = Mock(id="resp_prev")
|
||||||
|
prev_response.output = [
|
||||||
|
ResponseReasoningItem(
|
||||||
|
id="rs_prev", summary=[], type="reasoning", content=None, status=None
|
||||||
|
),
|
||||||
|
ResponseOutputMessage(
|
||||||
|
id="msg_prev",
|
||||||
|
content=[
|
||||||
|
ResponseOutputText(
|
||||||
|
text="first answer part",
|
||||||
|
annotations=[],
|
||||||
|
type="output_text",
|
||||||
|
logprobs=None,
|
||||||
|
),
|
||||||
|
ResponseOutputText(
|
||||||
|
text="second answer part",
|
||||||
|
annotations=[],
|
||||||
|
type="output_text",
|
||||||
|
logprobs=None,
|
||||||
|
),
|
||||||
|
],
|
||||||
|
role="assistant",
|
||||||
|
status="completed",
|
||||||
|
type="message",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
serving.msg_store["resp_prev"] = [{"role": "user", "content": "old input"}]
|
||||||
|
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
instructions="Be brief",
|
||||||
|
previous_response_id="resp_prev",
|
||||||
|
input="new input",
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
messages = serving._construct_input_messages(request, prev_response)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
messages,
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "Be brief"},
|
||||||
|
{"role": "user", "content": "old input"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": "first answer part\nsecond answer part",
|
||||||
|
},
|
||||||
|
{"role": "user", "content": "new input"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_input_parts_normalized_for_chat_templates(self):
|
||||||
|
serving = make_serving()
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input=[
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "input_text", "text": "what is this?"},
|
||||||
|
{
|
||||||
|
"type": "input_image",
|
||||||
|
"image_url": "http://example.com/cat.png",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
messages = serving._construct_input_messages(request)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
messages,
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "text", "text": "what is this?"},
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {
|
||||||
|
"url": "http://example.com/cat.png",
|
||||||
|
"detail": "auto",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_previous_response_id_input_list_does_not_call_copy_module(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.use_harmony = True
|
||||||
|
prev = Mock(id="resp_prev")
|
||||||
|
prev.output = [
|
||||||
|
ResponseFunctionToolCall(
|
||||||
|
arguments="{}",
|
||||||
|
call_id="call_x",
|
||||||
|
name="t",
|
||||||
|
type="function_call",
|
||||||
|
id="fc_x",
|
||||||
|
status="completed",
|
||||||
|
)
|
||||||
|
]
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input=[{"role": "user", "content": "hi"}],
|
||||||
|
previous_response_id="resp_prev",
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
serving._construct_input_messages_with_harmony(request, prev)
|
||||||
|
except TypeError as exc:
|
||||||
|
self.fail(f"copy() module-call regression: {exc}")
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class ChatToolForwardingTestCase(unittest.TestCase):
|
||||||
|
def test_make_request_passes_function_tools_to_chat_processing(self):
|
||||||
|
serving = make_serving()
|
||||||
|
seen = {}
|
||||||
|
|
||||||
|
def fake_process(chat_request, is_multimodal):
|
||||||
|
seen["tools"] = chat_request.tools
|
||||||
|
seen["tool_choice"] = chat_request.tool_choice
|
||||||
|
seen["parallel_tool_calls"] = chat_request.parallel_tool_calls
|
||||||
|
return MessageProcessingResult(
|
||||||
|
prompt="prompt",
|
||||||
|
prompt_ids=[1, 2, 3],
|
||||||
|
image_data=None,
|
||||||
|
audio_data=None,
|
||||||
|
video_data=None,
|
||||||
|
modalities=[],
|
||||||
|
stop=["</s>"],
|
||||||
|
tool_call_constraint=("json_schema", {"type": "object"}),
|
||||||
|
)
|
||||||
|
|
||||||
|
serving._process_messages = Mock(side_effect=fake_process)
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="call the tool",
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "lookup",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
tool_choice="required",
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
messages, request_prompts, engine_prompts, processed = asyncio.run(
|
||||||
|
serving._make_request(request, None, serving.tokenizer_manager.tokenizer)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(messages, [{"role": "user", "content": "call the tool"}])
|
||||||
|
self.assertEqual(request_prompts, [[1, 2, 3]])
|
||||||
|
self.assertEqual(engine_prompts, [[1, 2, 3]])
|
||||||
|
self.assertEqual(seen["tools"][0].function.name, "lookup")
|
||||||
|
self.assertEqual(seen["tool_choice"], "required")
|
||||||
|
self.assertFalse(seen["parallel_tool_calls"])
|
||||||
|
self.assertEqual(processed.tool_call_constraint[0], "json_schema")
|
||||||
|
|
||||||
|
def test_required_tool_choice_without_function_tool_returns_400(self):
|
||||||
|
serving = make_serving()
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
tool_choice="required",
|
||||||
|
tools=[{"type": "web_search"}, {"type": "mcp"}],
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
result = asyncio.run(serving.create_responses(request, raw_request=None))
|
||||||
|
self.assertEqual(getattr(result, "status_code", None), 400)
|
||||||
|
|
||||||
|
|
||||||
|
class InputItemNormalizationTestCase(unittest.TestCase):
|
||||||
|
def test_function_call_becomes_assistant_tool_call(self):
|
||||||
|
normalized = OpenAIServingResponses._normalize_response_message_for_chat(
|
||||||
|
{
|
||||||
|
"type": "function_call",
|
||||||
|
"id": "fc_1",
|
||||||
|
"call_id": "call_abc",
|
||||||
|
"name": "lookup",
|
||||||
|
"arguments": '{"key": "val"}',
|
||||||
|
"status": "completed",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
normalized,
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": "call_abc",
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "lookup",
|
||||||
|
"arguments": '{"key": "val"}',
|
||||||
|
},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_developer_role_becomes_system(self):
|
||||||
|
normalized = OpenAIServingResponses._normalize_response_message_for_chat(
|
||||||
|
{"role": "developer", "content": "Be terse."}
|
||||||
|
)
|
||||||
|
self.assertEqual(normalized, {"role": "system", "content": "Be terse."})
|
||||||
|
|
||||||
|
def test_function_call_output_becomes_tool_message(self):
|
||||||
|
normalized = OpenAIServingResponses._normalize_response_message_for_chat(
|
||||||
|
{
|
||||||
|
"type": "function_call_output",
|
||||||
|
"call_id": "call_abc",
|
||||||
|
"output": "42",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
normalized,
|
||||||
|
{"role": "tool", "tool_call_id": "call_abc", "content": "42"},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unknown_input_item_type_raises(self):
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
OpenAIServingResponses._normalize_response_message_for_chat(
|
||||||
|
{"type": "web_search_call", "id": "ws_1"}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class FullResponseUsageTestCase(unittest.TestCase):
|
||||||
|
def test_full_response_uses_dict_meta_info_for_usage(self):
|
||||||
|
serving = make_serving()
|
||||||
|
context = SimpleContext()
|
||||||
|
context.last_output = {
|
||||||
|
"text": "done",
|
||||||
|
"meta_info": {
|
||||||
|
"prompt_tokens": 11,
|
||||||
|
"completion_tokens": 7,
|
||||||
|
"cached_tokens": 3,
|
||||||
|
"reasoning_tokens": 2,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x", input="hello", request_id="resp_usage", store=False
|
||||||
|
)
|
||||||
|
metadata = RequestResponseMetadata(request_id=request.request_id)
|
||||||
|
|
||||||
|
async def empty_generator():
|
||||||
|
if False:
|
||||||
|
yield None
|
||||||
|
|
||||||
|
response = asyncio.run(
|
||||||
|
serving.responses_full_generator(
|
||||||
|
request,
|
||||||
|
sampling_params={},
|
||||||
|
result_generator=empty_generator(),
|
||||||
|
context=context,
|
||||||
|
model_name="x",
|
||||||
|
tokenizer=serving.tokenizer_manager.tokenizer,
|
||||||
|
request_metadata=metadata,
|
||||||
|
created_time=123,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.usage.prompt_tokens, 11)
|
||||||
|
self.assertEqual(response.usage.completion_tokens, 7)
|
||||||
|
self.assertEqual(response.usage.reasoning_tokens, 2)
|
||||||
|
self.assertEqual(metadata.final_usage_info, response.usage)
|
||||||
|
|
||||||
|
|
||||||
|
class MultimodalRequestTestCase(unittest.TestCase):
|
||||||
|
def test_multimodal_create_responses_sends_text_and_media_to_engine(self):
|
||||||
|
serving = make_serving(is_multimodal=True)
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
serving._process_messages = Mock(
|
||||||
|
return_value=MessageProcessingResult(
|
||||||
|
prompt="rendered multimodal prompt",
|
||||||
|
prompt_ids=[9, 9, 9],
|
||||||
|
image_data=["http://example.com/cat.png"],
|
||||||
|
audio_data=None,
|
||||||
|
video_data=None,
|
||||||
|
modalities=["image"],
|
||||||
|
stop=[],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
async def fake_generate(
|
||||||
|
request_id,
|
||||||
|
request_prompt,
|
||||||
|
adapted_request,
|
||||||
|
sampling_params,
|
||||||
|
context,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
captured["request_prompt"] = request_prompt
|
||||||
|
captured["adapted_request"] = adapted_request
|
||||||
|
context.append_output(
|
||||||
|
{
|
||||||
|
"text": "looks like a cat",
|
||||||
|
"meta_info": {
|
||||||
|
"prompt_tokens": 5,
|
||||||
|
"completion_tokens": 4,
|
||||||
|
"cached_tokens": 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
yield context
|
||||||
|
|
||||||
|
serving._generate_with_builtin_tools = fake_generate
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input=[
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "input_text", "text": "describe it"},
|
||||||
|
{
|
||||||
|
"type": "input_image",
|
||||||
|
"image_url": "http://example.com/cat.png",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
request_id="resp_mm",
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = asyncio.run(serving.create_responses(request))
|
||||||
|
|
||||||
|
self.assertEqual(response.status, "completed")
|
||||||
|
self.assertEqual(captured["request_prompt"], "rendered multimodal prompt")
|
||||||
|
self.assertEqual(captured["adapted_request"].text, "rendered multimodal prompt")
|
||||||
|
self.assertIsNone(captured["adapted_request"].input_ids)
|
||||||
|
self.assertEqual(
|
||||||
|
captured["adapted_request"].image_data, ["http://example.com/cat.png"]
|
||||||
|
)
|
||||||
|
self.assertEqual(captured["adapted_request"].modalities, ["image"])
|
||||||
|
|
||||||
|
|
||||||
|
class OutputItemsTestCase(unittest.TestCase):
|
||||||
|
def _function_tool_request(self):
|
||||||
|
return ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="weather?",
|
||||||
|
store=False,
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "get_weather",
|
||||||
|
"description": "Get weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_function_tool_call_extracted_via_parser(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.tool_call_parser = "qwen3_coder"
|
||||||
|
fake_call = ToolCallItem(
|
||||||
|
tool_index=0, name="get_weather", parameters='{"city": "Beijing"}'
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.entrypoints.openai.serving_responses.FunctionCallParser"
|
||||||
|
) as parser_cls:
|
||||||
|
instance = parser_cls.return_value
|
||||||
|
instance.has_tool_call.return_value = True
|
||||||
|
instance.parse_non_stream.return_value = ("trailing text", [fake_call])
|
||||||
|
output_items = serving._make_response_output_items(
|
||||||
|
self._function_tool_request(),
|
||||||
|
"raw model output with <tool_call>",
|
||||||
|
tokenizer=Mock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
tool_calls = [
|
||||||
|
item for item in output_items if isinstance(item, ResponseFunctionToolCall)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(tool_calls), 1)
|
||||||
|
self.assertEqual(tool_calls[0].name, "get_weather")
|
||||||
|
self.assertEqual(tool_calls[0].arguments, '{"city": "Beijing"}')
|
||||||
|
|
||||||
|
message_items = [
|
||||||
|
item for item in output_items if isinstance(item, ResponseOutputMessage)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(message_items), 1)
|
||||||
|
self.assertEqual(message_items[0].content[0].text, "trailing text")
|
||||||
|
|
||||||
|
def test_prose_emitted_before_tool_call_item(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.tool_call_parser = "qwen3_coder"
|
||||||
|
fake_call = ToolCallItem(
|
||||||
|
tool_index=0, name="get_weather", parameters='{"city": "Beijing"}'
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.entrypoints.openai.serving_responses.FunctionCallParser"
|
||||||
|
) as parser_cls:
|
||||||
|
instance = parser_cls.return_value
|
||||||
|
instance.has_tool_call.return_value = True
|
||||||
|
instance.parse_non_stream.return_value = (
|
||||||
|
"I'll check the weather.",
|
||||||
|
[fake_call],
|
||||||
|
)
|
||||||
|
output_items = serving._make_response_output_items(
|
||||||
|
self._function_tool_request(), "raw model output", tokenizer=Mock()
|
||||||
|
)
|
||||||
|
|
||||||
|
types = [type(item).__name__ for item in output_items]
|
||||||
|
self.assertEqual(types, ["ResponseOutputMessage", "ResponseFunctionToolCall"])
|
||||||
|
|
||||||
|
def test_required_tool_choice_parses_json_array_without_native_parser(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.tool_call_parser = None
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
tool_choice="required",
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "get_weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
raw = '[{"name": "get_weather", "parameters": {"city": "Beijing"}}]'
|
||||||
|
|
||||||
|
output_items = serving._make_response_output_items(
|
||||||
|
request, raw, tokenizer=Mock()
|
||||||
|
)
|
||||||
|
|
||||||
|
tool_calls = [
|
||||||
|
item for item in output_items if isinstance(item, ResponseFunctionToolCall)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(tool_calls), 1)
|
||||||
|
self.assertEqual(tool_calls[0].name, "get_weather")
|
||||||
|
self.assertEqual(tool_calls[0].arguments, '{"city": "Beijing"}')
|
||||||
|
self.assertEqual(
|
||||||
|
[item for item in output_items if isinstance(item, ResponseOutputMessage)],
|
||||||
|
[],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_no_tool_call_extraction_when_tool_choice_none(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.tool_call_parser = "qwen3_coder"
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
store=False,
|
||||||
|
tool_choice="none",
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "get_weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.entrypoints.openai.serving_responses.FunctionCallParser"
|
||||||
|
) as parser_cls:
|
||||||
|
output_items = serving._make_response_output_items(
|
||||||
|
request, "just a plain answer", tokenizer=Mock()
|
||||||
|
)
|
||||||
|
parser_cls.assert_not_called()
|
||||||
|
|
||||||
|
self.assertEqual(len(output_items), 1)
|
||||||
|
self.assertIsInstance(output_items[0], ResponseOutputMessage)
|
||||||
|
|
||||||
|
|
||||||
|
class HarmonyResponsesTestCase(unittest.TestCase):
|
||||||
|
def test_developer_message_skips_unsupported_tool_types(self):
|
||||||
|
from sglang.srt.entrypoints.harmony_utils import get_developer_message
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import ResponseTool
|
||||||
|
|
||||||
|
tools = [
|
||||||
|
ResponseTool(
|
||||||
|
type="function",
|
||||||
|
name="get_weather",
|
||||||
|
description="Look up weather.",
|
||||||
|
parameters={"type": "object"},
|
||||||
|
),
|
||||||
|
ResponseTool(type="web_search"),
|
||||||
|
ResponseTool(type="namespace", name="codex"),
|
||||||
|
ResponseTool(type="mcp"),
|
||||||
|
]
|
||||||
|
msg = get_developer_message(instructions="be helpful", tools=tools)
|
||||||
|
self.assertIsNotNone(msg)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,211 @@
|
|||||||
|
import asyncio
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
|
from utils import (
|
||||||
|
collect_stream_events,
|
||||||
|
event_payloads,
|
||||||
|
event_types,
|
||||||
|
find_completed_event,
|
||||||
|
make_serving,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
|
RequestResponseMetadata,
|
||||||
|
ResponsesRequest,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=4, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _StreamFixture:
|
||||||
|
def __init__(self, serving, request):
|
||||||
|
self.serving = serving
|
||||||
|
self.request = request
|
||||||
|
self.request_metadata = RequestResponseMetadata(request_id=request.request_id)
|
||||||
|
|
||||||
|
def run(self, chunks):
|
||||||
|
async def gen():
|
||||||
|
for ch in chunks:
|
||||||
|
yield ch
|
||||||
|
|
||||||
|
async def collect():
|
||||||
|
return await collect_stream_events(
|
||||||
|
self.serving.responses_stream_generator_non_harmony(
|
||||||
|
self.request,
|
||||||
|
sampling_params={},
|
||||||
|
result_generator=gen(),
|
||||||
|
model_name="x",
|
||||||
|
tokenizer=Mock(),
|
||||||
|
request_metadata=self.request_metadata,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return asyncio.run(collect())
|
||||||
|
|
||||||
|
|
||||||
|
def _engine_chunk(text, completion_tokens, *, finish=False):
|
||||||
|
return {
|
||||||
|
"text": text,
|
||||||
|
"meta_info": {
|
||||||
|
"id": "rid",
|
||||||
|
"prompt_tokens": 5,
|
||||||
|
"completion_tokens": completion_tokens,
|
||||||
|
"cached_tokens": 0,
|
||||||
|
"reasoning_tokens": 0,
|
||||||
|
"finish_reason": {"type": "stop"} if finish else None,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class NonHarmonyStreamTestCase(unittest.TestCase):
|
||||||
|
def test_emits_typed_sse_events_in_order(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.reasoning_parser = None
|
||||||
|
serving.tool_call_parser = None
|
||||||
|
|
||||||
|
request = ResponsesRequest(model="x", input="hi", stream=True, store=False)
|
||||||
|
fixture = _StreamFixture(serving, request)
|
||||||
|
events = fixture.run(
|
||||||
|
[
|
||||||
|
_engine_chunk("Hel", 1),
|
||||||
|
_engine_chunk("Hello", 2),
|
||||||
|
_engine_chunk("Hello world", 4, finish=True),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
types = event_types(events)
|
||||||
|
self.assertEqual(types[0], "response.created")
|
||||||
|
self.assertEqual(types[1], "response.in_progress")
|
||||||
|
for ev in (
|
||||||
|
"response.output_item.added",
|
||||||
|
"response.content_part.added",
|
||||||
|
"response.output_text.delta",
|
||||||
|
"response.output_text.done",
|
||||||
|
"response.content_part.done",
|
||||||
|
"response.output_item.done",
|
||||||
|
):
|
||||||
|
self.assertIn(ev, types)
|
||||||
|
self.assertEqual(types[-1], "response.completed")
|
||||||
|
|
||||||
|
seqs = [p["sequence_number"] for p in event_payloads(events)]
|
||||||
|
self.assertEqual(seqs, list(range(len(seqs))))
|
||||||
|
|
||||||
|
def test_required_tool_choice_emits_function_call_events(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.reasoning_parser = None
|
||||||
|
serving.tool_call_parser = None
|
||||||
|
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
stream=True,
|
||||||
|
store=False,
|
||||||
|
tool_choice="required",
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "get_weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
payload = '[{"name": "get_weather", "parameters": {"city": "Beijing"}}]'
|
||||||
|
|
||||||
|
chunks = []
|
||||||
|
sent = 0
|
||||||
|
while sent < len(payload):
|
||||||
|
sent += min(8, len(payload) - sent)
|
||||||
|
chunks.append(
|
||||||
|
_engine_chunk(payload[:sent], sent, finish=sent == len(payload))
|
||||||
|
)
|
||||||
|
|
||||||
|
fixture = _StreamFixture(serving, request)
|
||||||
|
events = fixture.run(chunks)
|
||||||
|
types = event_types(events)
|
||||||
|
|
||||||
|
self.assertIn("response.function_call_arguments.delta", types)
|
||||||
|
self.assertIn("response.function_call_arguments.done", types)
|
||||||
|
self.assertIn("response.output_item.added", types)
|
||||||
|
self.assertIn("response.output_item.done", types)
|
||||||
|
self.assertNotIn("response.output_text.delta", types)
|
||||||
|
|
||||||
|
added_kinds = [
|
||||||
|
payload["item"]["type"]
|
||||||
|
for payload in event_payloads(events)
|
||||||
|
if payload.get("type") == "response.output_item.added"
|
||||||
|
]
|
||||||
|
self.assertIn("function_call", added_kinds)
|
||||||
|
|
||||||
|
def test_final_output_preserves_text_tool_text_order(self):
|
||||||
|
from sglang.srt.function_call.core_types import (
|
||||||
|
StreamingParseResult,
|
||||||
|
ToolCallItem,
|
||||||
|
)
|
||||||
|
|
||||||
|
serving = make_serving()
|
||||||
|
serving.reasoning_parser = None
|
||||||
|
serving.tool_call_parser = "qwen3_coder"
|
||||||
|
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="hi",
|
||||||
|
stream=True,
|
||||||
|
store=False,
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "get_weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
scripted = [
|
||||||
|
StreamingParseResult(normal_text="I'll check.", calls=[]),
|
||||||
|
StreamingParseResult(
|
||||||
|
normal_text="",
|
||||||
|
calls=[
|
||||||
|
ToolCallItem(
|
||||||
|
tool_index=0,
|
||||||
|
name="get_weather",
|
||||||
|
parameters='{"city": "Beijing"}',
|
||||||
|
)
|
||||||
|
],
|
||||||
|
),
|
||||||
|
StreamingParseResult(normal_text="It's sunny.", calls=[]),
|
||||||
|
]
|
||||||
|
chunks = [
|
||||||
|
_engine_chunk(" " * 3, 3),
|
||||||
|
_engine_chunk(" " * 10, 10),
|
||||||
|
_engine_chunk(" " * 14, 14, finish=True),
|
||||||
|
]
|
||||||
|
|
||||||
|
script_iter = iter(scripted)
|
||||||
|
|
||||||
|
def fake_parse_stream_chunk(delta):
|
||||||
|
sp = next(script_iter)
|
||||||
|
return sp.normal_text, sp.calls
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.entrypoints.openai.serving_responses.FunctionCallParser"
|
||||||
|
) as parser_cls:
|
||||||
|
parser_cls.return_value.detector.supports_structural_tag.return_value = True
|
||||||
|
parser_cls.return_value.parse_stream_chunk.side_effect = (
|
||||||
|
fake_parse_stream_chunk
|
||||||
|
)
|
||||||
|
fixture = _StreamFixture(serving, request)
|
||||||
|
events = fixture.run(chunks)
|
||||||
|
|
||||||
|
completed = find_completed_event(events)
|
||||||
|
output = completed["response"]["output"]
|
||||||
|
kinds = [item["type"] for item in output]
|
||||||
|
self.assertEqual(kinds, ["message", "function_call", "message"])
|
||||||
|
self.assertEqual(output[0]["content"][0]["text"], "I'll check.")
|
||||||
|
self.assertEqual(output[1]["name"], "get_weather")
|
||||||
|
self.assertEqual(output[2]["content"][0]["text"], "It's sunny.")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
"""Stub CUDA-only deps before importing sglang.srt serving modules. Must
|
||||||
|
be imported first by every /v1/responses test that runs on CPU."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
_ORIGINAL_TORCH_COMPILE = torch.compile
|
||||||
|
|
||||||
|
def _identity_compile(fn=None, **kwargs):
|
||||||
|
if fn is None:
|
||||||
|
return lambda inner_fn: inner_fn
|
||||||
|
return fn
|
||||||
|
|
||||||
|
torch.compile = _identity_compile
|
||||||
|
except ImportError:
|
||||||
|
torch = None
|
||||||
|
_ORIGINAL_TORCH_COMPILE = None
|
||||||
|
|
||||||
|
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import AsyncIterator
|
||||||
|
from unittest.mock import Mock
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.serving_responses import OpenAIServingResponses
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(
|
||||||
|
est_time=0,
|
||||||
|
suite="base-a-test-cpu",
|
||||||
|
disabled="helper module — exported fixtures, not a test",
|
||||||
|
)
|
||||||
|
|
||||||
|
if torch is not None:
|
||||||
|
torch.compile = _ORIGINAL_TORCH_COMPILE
|
||||||
|
|
||||||
|
|
||||||
|
class MockTokenizerManager:
|
||||||
|
def __init__(self, *, is_multimodal: bool = False):
|
||||||
|
self.model_config = Mock(is_multimodal=is_multimodal, context_len=4096)
|
||||||
|
self.model_config.get_default_sampling_params.return_value = {}
|
||||||
|
self.model_config.hf_config = Mock(
|
||||||
|
model_type="llama", architectures=["LlamaForCausalLM"]
|
||||||
|
)
|
||||||
|
self.server_args = Mock(
|
||||||
|
enable_cache_report=False,
|
||||||
|
reasoning_parser=None,
|
||||||
|
stream_response_default_include_usage=False,
|
||||||
|
tokenizer_metrics_allowed_custom_labels=None,
|
||||||
|
tool_call_parser=None,
|
||||||
|
incremental_streaming_output=False,
|
||||||
|
)
|
||||||
|
self.tokenizer = Mock()
|
||||||
|
self.tokenizer.encode.return_value = [1, 2, 3]
|
||||||
|
self.tokenizer.chat_template = None
|
||||||
|
self.tokenizer.bos_token_id = 1
|
||||||
|
self.num_reserved_tokens = 0
|
||||||
|
self.generate_request = Mock()
|
||||||
|
self.create_abort_task = Mock()
|
||||||
|
|
||||||
|
|
||||||
|
class MockTemplateManager:
|
||||||
|
def __init__(self):
|
||||||
|
self.chat_template_name = "llama-3"
|
||||||
|
self.jinja_template_content_format = None
|
||||||
|
self.completion_template_name = None
|
||||||
|
self.reasoning_config = None
|
||||||
|
self.force_reasoning = False
|
||||||
|
|
||||||
|
|
||||||
|
def make_serving(*, is_multimodal: bool = False) -> OpenAIServingResponses:
|
||||||
|
return OpenAIServingResponses(
|
||||||
|
MockTokenizerManager(is_multimodal=is_multimodal), MockTemplateManager()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def collect_stream_events(stream: AsyncIterator[str]) -> list[str]:
|
||||||
|
events = []
|
||||||
|
async for chunk in stream:
|
||||||
|
events.append(chunk)
|
||||||
|
return events
|
||||||
|
|
||||||
|
|
||||||
|
def event_types(events: list[str]) -> list[str]:
|
||||||
|
return [
|
||||||
|
line[len("event: ") :].strip()
|
||||||
|
for chunk in events
|
||||||
|
for line in chunk.splitlines()
|
||||||
|
if line.startswith("event: ")
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def event_payloads(events: list[str]) -> list[dict]:
|
||||||
|
return [
|
||||||
|
json.loads(line[len("data: ") :])
|
||||||
|
for chunk in events
|
||||||
|
for line in chunk.splitlines()
|
||||||
|
if line.startswith("data: ")
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def find_completed_event(events: list[str]) -> dict:
|
||||||
|
for chunk in events:
|
||||||
|
lines = chunk.splitlines()
|
||||||
|
if lines and lines[0] == "event: response.completed":
|
||||||
|
return json.loads(lines[1][len("data: ") :])
|
||||||
|
raise AssertionError("response.completed event missing from stream")
|
||||||
@@ -383,6 +383,37 @@ class TestTemplateContentFormatDetection(CustomTestCase):
|
|||||||
# Image data is still extracted
|
# Image data is still extracted
|
||||||
self.assertEqual(len(image_data), 1)
|
self.assertEqual(len(image_data), 1)
|
||||||
|
|
||||||
|
def test_process_content_v32_encoding_accepts_responses_input_text(self):
|
||||||
|
msg_dict = {
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "input_text", "text": "Hello"},
|
||||||
|
{
|
||||||
|
"type": "input_image",
|
||||||
|
"image_url": "http://example.com/img.jpg",
|
||||||
|
"detail": "auto",
|
||||||
|
},
|
||||||
|
{"type": "input_text", "text": "World"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
image_data = []
|
||||||
|
video_data = []
|
||||||
|
audio_data = []
|
||||||
|
modalities = []
|
||||||
|
result = process_content_for_template_format(
|
||||||
|
msg_dict,
|
||||||
|
"openai",
|
||||||
|
image_data,
|
||||||
|
video_data,
|
||||||
|
audio_data,
|
||||||
|
modalities,
|
||||||
|
use_dpsk_v32_encoding=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(result["content"], "Hello World")
|
||||||
|
self.assertEqual(len(image_data), 1)
|
||||||
|
self.assertEqual(image_data[0].url, "http://example.com/img.jpg")
|
||||||
|
|
||||||
def test_process_content_invalid_format_raises(self):
|
def test_process_content_invalid_format_raises(self):
|
||||||
"""Test that invalid content_format raises ValueError."""
|
"""Test that invalid content_format raises ValueError."""
|
||||||
msg_dict = {
|
msg_dict = {
|
||||||
|
|||||||
Reference in New Issue
Block a user