4596 lines
180 KiB
Python
4596 lines
180 KiB
Python
"""
|
||
Unit-tests for OpenAIServingChat -- rewritten to use only the std-lib 'unittest'.
|
||
Run with either:
|
||
python tests/test_serving_chat_unit.py -v
|
||
or
|
||
python -m unittest discover -s tests -p "test_*unit.py" -v
|
||
"""
|
||
|
||
from sglang.test.test_utils import CustomTestCase, enter_override, maybe_stub_sgl_kernel
|
||
|
||
maybe_stub_sgl_kernel() # must precede any import that pulls in sgl_kernel
|
||
|
||
import json
|
||
import re
|
||
import tempfile
|
||
import unittest
|
||
import uuid
|
||
from http import HTTPStatus
|
||
from pathlib import Path
|
||
from typing import Optional
|
||
from unittest.mock import Mock, patch
|
||
|
||
from fastapi import Request
|
||
|
||
from sglang.srt.entrypoints.openai import chat_encoding
|
||
from sglang.srt.entrypoints.openai.chat_encoding import (
|
||
resolve_dsv4_reasoning_effort_profile,
|
||
)
|
||
from sglang.srt.entrypoints.openai.protocol import (
|
||
ChatCompletionRequest,
|
||
MessageProcessingResult,
|
||
ToolChoice,
|
||
ToolChoiceFuncName,
|
||
)
|
||
from sglang.srt.entrypoints.openai.serving_chat import (
|
||
OpenAIServingChat,
|
||
normalize_tool_content,
|
||
)
|
||
from sglang.srt.environ import envs
|
||
from sglang.srt.function_call.kimik3_format import TOOLS_CLOSE, TOOLS_OPEN
|
||
from sglang.srt.managers.io_struct import GenerateReqInput
|
||
from sglang.srt.parser.jinja_template_utils import (
|
||
jinja_template_may_reorder_tool_results,
|
||
)
|
||
from sglang.srt.parser.template_detection import ReasoningToggleConfig
|
||
from sglang.srt.runtime_context import get_context, publish, reset_context
|
||
from sglang.srt.sampling.sampling_params import (
|
||
REQUEST_REASONING_END_TOKEN_IDS_KEY,
|
||
)
|
||
from sglang.srt.server_args import ServerArgs
|
||
from sglang.srt.utils import get_or_create_event_loop
|
||
from sglang.test.ci.ci_register import register_cpu_ci
|
||
|
||
register_cpu_ci(est_time=13, suite="base-a-test-cpu")
|
||
|
||
# Every spec resolve_chat_encoding_spec can return; pinned by the guard below.
|
||
_ALL_CHAT_ENCODING_SPECS = ("dsv41", "dsv4", "dsv32", "inkling", "kimi_k3")
|
||
|
||
|
||
def _spec_result(index):
|
||
return {
|
||
"text": f"choice-{index}",
|
||
"meta_info": {
|
||
"id": "chatcmpl-spec-test",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop"},
|
||
"weight_version": "default",
|
||
"spec_accept_rate": 0.5,
|
||
"spec_accept_length": 2.0,
|
||
"spec_cap_length": index + 1.0,
|
||
"spec_block_accept_length": index + 0.5,
|
||
"spec_num_correct_drafts": 1,
|
||
"spec_num_proposed_drafts": 2,
|
||
"spec_verify_ct": 1,
|
||
"spec_correct_drafts_histogram": [0, 1],
|
||
"spec_cap_lens_histogram": [index, 1],
|
||
},
|
||
"index": index,
|
||
}
|
||
|
||
|
||
_DSV4_PREVIEW_ENCODER = 'REASONING_EFFORT_MAX = "preview"\n'
|
||
_DSV4_OFFICIAL_ENCODER = (
|
||
"REASONING_EFFORT_PROMPTS: Dict[str, str] = "
|
||
'{"low": "", "high": "h", "max": "m"}\n'
|
||
'DEFAULT_REASONING_EFFORT = "low"\n'
|
||
)
|
||
|
||
_TOOL_RESULT_REORDER_TEMPLATE = """
|
||
{%- for assistant in messages
|
||
if assistant.role == 'assistant' and assistant.tool_calls -%}
|
||
{%- for tool_call in assistant.tool_calls -%}
|
||
{%- for result in messages
|
||
if result.role == 'tool' and result.tool_call_id == tool_call.id -%}
|
||
{%- for part in result.content -%}
|
||
{%- if part.type == 'text' -%}
|
||
{{- part.text -}}
|
||
{%- else -%}
|
||
{{- '<' + part.type + '>' -}}
|
||
{%- endif -%}
|
||
{%- endfor -%}
|
||
{%- endfor -%}
|
||
{%- endfor -%}
|
||
{%- endfor -%}
|
||
"""
|
||
|
||
|
||
def _create_dsv4_checkpoint(test_case: unittest.TestCase, source: str) -> str:
|
||
model_dir = tempfile.TemporaryDirectory()
|
||
test_case.addCleanup(model_dir.cleanup)
|
||
encoder_path = Path(model_dir.name) / "encoding" / "encoding_dsv4.py"
|
||
encoder_path.parent.mkdir(parents=True)
|
||
encoder_path.write_text(source, encoding="utf-8")
|
||
return model_dir.name
|
||
|
||
|
||
class _MockTokenizerManager:
|
||
"""Minimal mock that satisfies OpenAIServingChat."""
|
||
|
||
def __init__(self):
|
||
self.model_config = Mock(is_multimodal=False)
|
||
self.server_args = Mock(
|
||
model_path="deepseek-ai/DeepSeek-V4-Flash",
|
||
revision=None,
|
||
enable_cache_report=False,
|
||
tool_call_parser="hermes",
|
||
reasoning_parser=None,
|
||
stream_response_default_include_usage=False,
|
||
default_chat_template_kwargs=None,
|
||
return_input_ids=False,
|
||
return_output_ids=False,
|
||
incremental_streaming_output=False,
|
||
)
|
||
self.model_path = self.server_args.model_path
|
||
# The manager tracks the served name itself; a weight update rewrites it.
|
||
self.served_model_name = "test-model"
|
||
# Stands in for the context's resolved leaves: an override replaces the
|
||
# field's one live value, the seed stays on server_args.
|
||
self._config_overrides = {}
|
||
|
||
# Mock hf_config for _resolve_chat_encoding_spec check
|
||
mock_hf_config = Mock()
|
||
mock_hf_config.architectures = ["LlamaForCausalLM"]
|
||
mock_hf_config.to_dict.return_value = {}
|
||
self.model_config.hf_config = mock_hf_config
|
||
|
||
self.chat_template_name: Optional[str] = "llama-3"
|
||
|
||
# tokenizer stub
|
||
self.tokenizer = Mock()
|
||
self.tokenizer.encode.return_value = [1, 2, 3, 4, 5]
|
||
self.tokenizer.decode.return_value = "Test response"
|
||
self.tokenizer.chat_template = None
|
||
self.tokenizer.bos_token_id = 1
|
||
|
||
# async generator stub for generate_request
|
||
async def _mock_generate():
|
||
yield {
|
||
"text": "Test response",
|
||
"meta_info": {
|
||
"id": f"chatcmpl-{uuid.uuid4()}",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 5,
|
||
"cached_tokens": 0,
|
||
"weight_version": "test-version",
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"output_token_logprobs": [(0.1, 1, "Test"), (0.2, 2, "response")],
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.generate_request = Mock(return_value=_mock_generate())
|
||
self.create_abort_task = Mock()
|
||
self.request_logger = Mock(log_requests=False, log_requests_level=0)
|
||
|
||
def config_value(self, name: str):
|
||
"""The value in effect for one config field."""
|
||
if name in self._config_overrides:
|
||
return self._config_overrides[name]
|
||
return getattr(self.server_args, name)
|
||
|
||
|
||
class _MockTemplateManager:
|
||
"""Minimal mock for TemplateManager."""
|
||
|
||
def __init__(self):
|
||
self.chat_template_name: Optional[str] = "llama-3"
|
||
self.jinja_template_content_format: Optional[str] = None
|
||
self.completion_template_name: Optional[str] = None
|
||
self.reasoning_config = None
|
||
self.force_reasoning = False
|
||
self.jinja_template_may_reorder_tool_results = False
|
||
|
||
|
||
class TestChatTemplateCache(CustomTestCase):
|
||
def setUp(self):
|
||
super().setUp()
|
||
reset_context()
|
||
self.addCleanup(reset_context)
|
||
publish(
|
||
ServerArgs(model_path="dummy", default_chat_template_kwargs=None),
|
||
role="tokenizer",
|
||
)
|
||
self.tokenizer_manager = _MockTokenizerManager()
|
||
self.chat = OpenAIServingChat(
|
||
self.tokenizer_manager,
|
||
_MockTemplateManager(),
|
||
)
|
||
self.tokenizer_manager.tokenizer.apply_chat_template.return_value = "rendered"
|
||
self.tokenizer_manager.tokenizer.encode.return_value = [11, 12]
|
||
self.tokenizer_manager.tokenizer.decode.return_value = "decoded"
|
||
self.tokenizer_manager.tokenizer.reset_mock()
|
||
|
||
def _render(self, **overrides):
|
||
kwargs = {
|
||
"messages": [{"role": "user", "content": "same text prefix"}],
|
||
"tools": None,
|
||
"template_kwargs": {"enable_thinking": False},
|
||
"encode_kwargs": {"add_special_tokens": False},
|
||
"use_cache": True,
|
||
}
|
||
kwargs.update(overrides)
|
||
return self.chat._render_and_encode_chat_template(**kwargs)
|
||
|
||
def test_cache_hit_reuses_render_encode_and_returns_an_owned_id_list(self):
|
||
first = self._render()
|
||
first[1].append(99)
|
||
second = self._render()
|
||
|
||
self.assertEqual(second, ("rendered", [11, 12], "decoded"))
|
||
self.tokenizer_manager.tokenizer.apply_chat_template.assert_called_once()
|
||
self.tokenizer_manager.tokenizer.encode.assert_called_once()
|
||
self.tokenizer_manager.tokenizer.decode.assert_called_once()
|
||
|
||
def test_cache_key_includes_template_and_encode_options(self):
|
||
self._render()
|
||
self._render(template_kwargs={"enable_thinking": True})
|
||
self._render(encode_kwargs={"add_special_tokens": True})
|
||
|
||
self.assertEqual(
|
||
self.tokenizer_manager.tokenizer.apply_chat_template.call_count,
|
||
3,
|
||
)
|
||
self.assertEqual(self.tokenizer_manager.tokenizer.encode.call_count, 3)
|
||
|
||
def test_cache_key_tracks_tokenizer_chat_template_updates(self):
|
||
self.tokenizer_manager.tokenizer.chat_template = "template-v1"
|
||
self._render()
|
||
self.tokenizer_manager.tokenizer.chat_template = "template-v2"
|
||
self._render()
|
||
|
||
self.assertEqual(
|
||
self.tokenizer_manager.tokenizer.apply_chat_template.call_count,
|
||
2,
|
||
)
|
||
|
||
def test_non_serializable_input_bypasses_cache(self):
|
||
messages = [{"role": "user", "content": object()}]
|
||
self._render(messages=messages)
|
||
self._render(messages=messages)
|
||
|
||
self.assertEqual(
|
||
self.tokenizer_manager.tokenizer.apply_chat_template.call_count,
|
||
2,
|
||
)
|
||
self.assertEqual(self.tokenizer_manager.tokenizer.encode.call_count, 2)
|
||
self.tokenizer_manager.tokenizer.decode.assert_not_called()
|
||
|
||
|
||
class ServingChatTestCase(unittest.TestCase):
|
||
# ------------- common fixtures -------------
|
||
def setUp(self):
|
||
# The serving layer reads its config from the bags, so the fixture has
|
||
# to publish one rather than hang the values off a mock manager.
|
||
reset_context()
|
||
self.addCleanup(reset_context)
|
||
publish(
|
||
ServerArgs(
|
||
model_path="dummy",
|
||
revision=None,
|
||
enable_cache_report=False,
|
||
tool_call_parser="hermes",
|
||
reasoning_parser=None,
|
||
stream_response_default_include_usage=False,
|
||
default_chat_template_kwargs=None,
|
||
),
|
||
role="tokenizer",
|
||
)
|
||
self.tm = _MockTokenizerManager()
|
||
self.template_manager = _MockTemplateManager()
|
||
self.chat = OpenAIServingChat(self.tm, self.template_manager)
|
||
|
||
# frequently reused requests
|
||
self.basic_req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
session_id="session-1",
|
||
temperature=0.7,
|
||
max_tokens=100,
|
||
stream=False,
|
||
)
|
||
self.stream_req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
temperature=0.7,
|
||
max_tokens=100,
|
||
stream=True,
|
||
)
|
||
|
||
self.fastapi_request = Mock(spec=Request)
|
||
self.fastapi_request.headers = {}
|
||
|
||
@staticmethod
|
||
def _render_tool_results_in_call_order(messages, **kwargs):
|
||
"""Block-level tool_call_id association, like the GLM chat templates."""
|
||
del kwargs
|
||
rendered = []
|
||
index = 0
|
||
while index < len(messages):
|
||
message = messages[index]
|
||
index += 1
|
||
tool_calls = message.get("tool_calls") or []
|
||
if message.get("role") != "assistant" or not tool_calls:
|
||
continue
|
||
run = []
|
||
while index < len(messages) and messages[index].get("role") == "tool":
|
||
run.append(messages[index])
|
||
index += 1
|
||
by_id = {result.get("tool_call_id"): result for result in run}
|
||
for tool_call in tool_calls:
|
||
result = by_id.get(tool_call.get("id"))
|
||
if result is None:
|
||
continue
|
||
for part in result.get("content") or []:
|
||
if part.get("type") == "text":
|
||
rendered.append(part.get("text", ""))
|
||
else:
|
||
rendered.append(f"<{part.get('type')}>")
|
||
return "".join(rendered)
|
||
|
||
@staticmethod
|
||
def _tool_round(call_ids, result_ids, part_type="image_url"):
|
||
assistant = {
|
||
"role": "assistant",
|
||
"content": "",
|
||
"tool_calls": [
|
||
{"id": call_id, "function": {"name": call_id, "arguments": {}}}
|
||
for call_id in call_ids
|
||
],
|
||
}
|
||
part_key = {"image_url": "image_url", "video_url": "video_url"}[part_type]
|
||
results = [
|
||
{
|
||
"role": "tool",
|
||
"tool_call_id": result_id,
|
||
"content": [{"type": part_type, part_key: {"url": result_id}}],
|
||
}
|
||
for result_id in result_ids
|
||
]
|
||
return [assistant] + results
|
||
|
||
def test_canonicalize_tool_message_order_sorts_media_runs(self):
|
||
messages = self._tool_round(
|
||
["call-a", "call-b"], ["call-b", "call-a"]
|
||
) + self._tool_round(["call-c", "call-d"], ["call-d", "call-c"], "video_url")
|
||
|
||
canonical = self.chat._canonicalize_tool_message_order(messages)
|
||
|
||
self.assertEqual(
|
||
[
|
||
message["tool_call_id"]
|
||
for message in canonical
|
||
if message.get("role") == "tool"
|
||
],
|
||
["call-a", "call-b", "call-c", "call-d"],
|
||
)
|
||
|
||
def test_canonicalize_tool_message_order_keeps_unassociable_runs(self):
|
||
cases = {
|
||
"unknown_id": (["call-a", "call-b"], ["call-b", "call-z"]),
|
||
"duplicate_result_id": (["call-a", "call-b"], ["call-b", "call-b"]),
|
||
"missing_call_id": ([None, "call-b"], ["call-b", None]),
|
||
}
|
||
for name, (call_ids, result_ids) in cases.items():
|
||
with self.subTest(name=name):
|
||
messages = self._tool_round(call_ids, result_ids)
|
||
canonical = self.chat._canonicalize_tool_message_order(messages)
|
||
self.assertEqual(canonical, messages)
|
||
|
||
def test_canonicalize_tool_message_order_keeps_text_only_runs(self):
|
||
messages = self._tool_round(["call-a", "call-b"], ["call-b", "call-a"])
|
||
for message in messages[1:]:
|
||
message["content"] = [{"type": "text", "text": "done"}]
|
||
|
||
canonical = self.chat._canonicalize_tool_message_order(messages)
|
||
|
||
self.assertEqual(canonical, messages)
|
||
|
||
def test_jinja_path_recovers_tool_result_images_only_when_template_needs_it(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "openai"
|
||
self.template_manager.jinja_template_may_reorder_tool_results = (
|
||
jinja_template_may_reorder_tool_results(_TOOL_RESULT_REORDER_TEMPLATE)
|
||
)
|
||
self.assertTrue(self.template_manager.jinja_template_may_reorder_tool_results)
|
||
self.tm.tokenizer.apply_chat_template.side_effect = (
|
||
self._render_tool_results_in_call_order
|
||
)
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "inspect"},
|
||
{
|
||
"role": "assistant",
|
||
"content": "",
|
||
"tool_calls": [
|
||
{
|
||
"id": "call-a",
|
||
"type": "function",
|
||
"function": {"name": "a", "arguments": {}},
|
||
},
|
||
{
|
||
"id": "call-b",
|
||
"type": "function",
|
||
"function": {"name": "b", "arguments": {}},
|
||
},
|
||
],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"tool_call_id": "call-b",
|
||
"content": [{"type": "image_url", "image_url": {"url": "image-b"}}],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"tool_call_id": "call-a",
|
||
"content": [{"type": "image_url", "image_url": {"url": "image-a"}}],
|
||
},
|
||
],
|
||
)
|
||
|
||
result = self.chat._apply_jinja_template(request, None, is_multimodal=True)
|
||
|
||
self.assertEqual(
|
||
[item.url for item in result.image_data], ["image-a", "image-b"]
|
||
)
|
||
self.assertEqual(self.tm.tokenizer.apply_chat_template.call_count, 1)
|
||
rendered_messages = self.tm.tokenizer.apply_chat_template.call_args[0][0]
|
||
self.assertEqual(
|
||
[m.get("tool_call_id") for m in rendered_messages if "tool_call_id" in m],
|
||
["call-a", "call-b"],
|
||
)
|
||
|
||
# An already-ordered request must produce the exact same render input.
|
||
ordered_request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=request.messages[:2] + request.messages[2:][::-1],
|
||
)
|
||
self.tm.tokenizer.apply_chat_template.reset_mock()
|
||
self.chat._apply_jinja_template(ordered_request, None, is_multimodal=True)
|
||
self.tm.tokenizer.apply_chat_template.assert_not_called()
|
||
|
||
self.template_manager.jinja_template_may_reorder_tool_results = False
|
||
self.tm.tokenizer.apply_chat_template.reset_mock()
|
||
result = self.chat._apply_jinja_template(request, None, is_multimodal=True)
|
||
self.assertEqual(
|
||
[item.url for item in result.image_data], ["image-b", "image-a"]
|
||
)
|
||
self.assertEqual(self.tm.tokenizer.apply_chat_template.call_count, 1)
|
||
|
||
def test_parsers_follow_the_control_plane_overlay(self):
|
||
"""Template detection records the parsers through `override`, so they
|
||
answer from the bags; `ServerArgs` keeps the launcher's seed."""
|
||
self.tm.server_args.tool_call_parser = "auto"
|
||
self.tm.server_args.reasoning_parser = "auto"
|
||
self.tm._config_overrides.update(
|
||
{"tool_call_parser": "qwen25", "reasoning_parser": None}
|
||
)
|
||
|
||
chat = OpenAIServingChat(self.tm, self.template_manager)
|
||
|
||
self.assertEqual(chat.tool_call_parser, "qwen25")
|
||
self.assertIsNone(chat.reasoning_parser)
|
||
self.assertEqual(self.tm.server_args.tool_call_parser, "auto")
|
||
|
||
def test_the_xgrammar_gate_follows_the_overlay(self):
|
||
"""A detected `reasoning_parser` must gate xgrammar, not the seed's "auto"."""
|
||
self.tm.server_args.reasoning_parser = "auto"
|
||
self.tm._config_overrides["reasoning_parser"] = "qwen3"
|
||
chat = OpenAIServingChat(self.tm, self.template_manager)
|
||
self.assertEqual(chat.reasoning_parser, "qwen3")
|
||
# the gate reads the same value the parser was built from
|
||
self.assertIsNotNone(chat.reasoning_parser)
|
||
|
||
def test_text_only_model_rejects_media_before_generation(self):
|
||
media_parts = {
|
||
"image_url": {
|
||
"type": "image_url",
|
||
"image_url": {"url": "https://example.com/image.png"},
|
||
},
|
||
"video_url": {
|
||
"type": "video_url",
|
||
"video_url": {"url": "https://example.com/video.mp4"},
|
||
},
|
||
"audio_url": {
|
||
"type": "audio_url",
|
||
"audio_url": {"url": "https://example.com/audio.wav"},
|
||
},
|
||
}
|
||
|
||
for media_type, media_part in media_parts.items():
|
||
with self.subTest(media_type=media_type):
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{
|
||
"role": "user",
|
||
"content": [
|
||
{"type": "text", "text": "describe"},
|
||
media_part,
|
||
],
|
||
}
|
||
],
|
||
)
|
||
response = get_or_create_event_loop().run_until_complete(
|
||
self.chat.handle_request(request, self.fastapi_request)
|
||
)
|
||
error = json.loads(response.body)
|
||
self.assertEqual(response.status_code, HTTPStatus.BAD_REQUEST)
|
||
self.assertEqual(error["type"], "BadRequestError")
|
||
self.assertIn(media_type, error["message"])
|
||
self.tm.generate_request.assert_not_called()
|
||
|
||
def test_media_validation_does_not_reject_supported_content(self):
|
||
text_request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{
|
||
"role": "user",
|
||
"content": [{"type": "tool_reference", "name": "get_weather"}],
|
||
}
|
||
],
|
||
)
|
||
self.assertIsNone(self.chat._validate_request(text_request))
|
||
|
||
self.tm.model_config.is_multimodal = True
|
||
multimodal_request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{
|
||
"role": "user",
|
||
"content": [
|
||
{
|
||
"type": "image_url",
|
||
"image_url": {"url": "https://example.com/image.png"},
|
||
}
|
||
],
|
||
}
|
||
],
|
||
)
|
||
self.assertIsNone(self.chat._validate_request(multimodal_request))
|
||
|
||
# ------------- conversion tests -------------
|
||
def test_convert_to_internal_request_single(self):
|
||
with (
|
||
patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock,
|
||
patch.object(self.chat, "_process_messages") as proc_mock,
|
||
):
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_ins.image_data = conv_ins.audio_data = None
|
||
conv_ins.modalities = []
|
||
conv_ins.stop_str = ["</s>"]
|
||
conv_mock.return_value = conv_ins
|
||
|
||
proc_mock.return_value = MessageProcessingResult(
|
||
"Test prompt",
|
||
[1, 2, 3],
|
||
None,
|
||
None,
|
||
[],
|
||
["</s>"],
|
||
None,
|
||
)
|
||
|
||
self.basic_req.return_sampling_mask = True
|
||
self.basic_req.return_meta_info = True
|
||
adapted, processed = self.chat._convert_to_internal_request(self.basic_req)
|
||
self.assertIsInstance(adapted, GenerateReqInput)
|
||
self.assertFalse(adapted.stream)
|
||
self.assertTrue(adapted.return_sampling_mask)
|
||
self.assertEqual(adapted.session_id, "session-1")
|
||
self.assertEqual(processed, self.basic_req)
|
||
|
||
def test_chat_applies_pd_header_overrides(self):
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
rid="body-rid",
|
||
routed_dp_rank=3,
|
||
disagg_prefill_dp_rank=4,
|
||
priority=5,
|
||
)
|
||
self.fastapi_request.headers = {
|
||
"x-override-rid": "header-rid",
|
||
"x-override-bootstrap-host": "header-host",
|
||
"x-override-bootstrap-port": "8998",
|
||
"x-override-bootstrap-room": "456",
|
||
"x-override-conversation-id": "conversation-1",
|
||
"x-override-routed-dp-rank": "6",
|
||
"x-override-disagg-prefill-dp-rank": "7",
|
||
"x-override-priority": "8",
|
||
}
|
||
body = request.model_dump()
|
||
|
||
processed_messages = MessageProcessingResult(
|
||
"Test prompt", [1, 2, 3], None, None, [], [], None
|
||
)
|
||
with (
|
||
envs.SGLANG_ENABLE_REQUEST_HEADER_OVERRIDES.override(True),
|
||
patch.object(
|
||
self.chat, "_process_messages", return_value=processed_messages
|
||
),
|
||
):
|
||
response = get_or_create_event_loop().run_until_complete(
|
||
self.chat.handle_request(request, self.fastapi_request)
|
||
)
|
||
|
||
self.assertEqual(response.choices[0].message.content, "Test response")
|
||
adapted_request = self.tm.generate_request.call_args.args[0]
|
||
self.assertEqual(adapted_request.bootstrap_room, 456)
|
||
self.assertEqual(adapted_request.bootstrap_host, "header-host")
|
||
self.assertEqual(adapted_request.bootstrap_port, 8998)
|
||
self.assertEqual(adapted_request.rid, "header-rid")
|
||
self.assertEqual(adapted_request.conversation_id, "conversation-1")
|
||
self.assertEqual(adapted_request.routed_dp_rank, 6)
|
||
self.assertEqual(adapted_request.disagg_prefill_dp_rank, 7)
|
||
self.assertEqual(adapted_request.priority, 8)
|
||
self.assertEqual(request.model_dump(), body)
|
||
self.assertFalse(hasattr(request, "conversation_id"))
|
||
|
||
def test_convert_to_internal_request_rejects_stream_token_ids(self):
|
||
for field in ("return_prompt_token_ids", "return_token_ids"):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
stream=True,
|
||
**{field: True},
|
||
)
|
||
with self.subTest(field=field), self.assertRaisesRegex(ValueError, field):
|
||
self.chat._convert_to_internal_request(req, self.fastapi_request)
|
||
|
||
def test_validate_request_rejects_sampling_mask_without_meta_info(self):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
return_sampling_mask=True,
|
||
)
|
||
|
||
self.assertEqual(
|
||
self.chat._validate_request(req),
|
||
"return_sampling_mask requires return_meta_info=true.",
|
||
)
|
||
|
||
def test_convert_to_internal_request_rejects_stream_return_meta_info(self):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
stream=True,
|
||
return_meta_info=True,
|
||
)
|
||
|
||
with self.assertRaisesRegex(
|
||
ValueError, "return_meta_info is not supported with streaming"
|
||
):
|
||
self.chat._convert_to_internal_request(req, self.fastapi_request)
|
||
|
||
def test_convert_to_internal_request_input_ids_bypasses_template(self):
|
||
self.tm.tokenizer = None
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
input_ids=[101, 102, 103],
|
||
stop=["STOP"],
|
||
return_prompt_token_ids=True,
|
||
cache_salt="tenant-a",
|
||
extra_key="classification",
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
adapted, processed = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
|
||
self.assertEqual(processed, req)
|
||
self.assertEqual(adapted.input_ids, [101, 102, 103])
|
||
self.assertTrue(adapted.return_prompt_token_ids)
|
||
self.assertEqual(adapted.sampling_params["stop"], ["STOP"])
|
||
self.assertEqual(adapted.cache_salt, "tenant-a")
|
||
self.assertEqual(adapted.extra_key, "classification")
|
||
conv_mock.assert_not_called()
|
||
|
||
def test_kimi_k3_usage_excludes_assistant_generation_stub(self):
|
||
self.chat.chat_encoding_spec = "kimi_k3"
|
||
ret = [
|
||
{
|
||
"text": "Answer",
|
||
"meta_info": {
|
||
"id": "chatcmpl-kimi-k3-usage",
|
||
"prompt_tokens": 2075,
|
||
"completion_tokens": 1,
|
||
"cached_tokens": 0,
|
||
"image_tokens": 2035,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"weight_version": "default",
|
||
},
|
||
}
|
||
]
|
||
|
||
response = self.chat._build_chat_response(self.basic_req, ret, created=123)
|
||
|
||
self.assertEqual(response.usage.prompt_tokens, 2072)
|
||
self.assertEqual(response.usage.total_tokens, 2073)
|
||
self.assertEqual(response.usage.prompt_tokens_details.image_tokens, 2035)
|
||
|
||
def test_kimi_tool_call_keeps_default_reasoning(self):
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=True
|
||
)
|
||
self.tm.server_args.reasoning_parser = "kimi_k2"
|
||
self.tm.server_args.tool_call_parser = "kimi_k2"
|
||
self.chat.reasoning_parser = "kimi_k2"
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"a": {"type": "integer"}},
|
||
},
|
||
},
|
||
}
|
||
],
|
||
tool_choice="required",
|
||
)
|
||
|
||
with patch.object(self.chat, "_process_messages") as proc_mock:
|
||
proc_mock.return_value = MessageProcessingResult(
|
||
"",
|
||
[1, 2, 3],
|
||
None,
|
||
None,
|
||
[],
|
||
[],
|
||
None,
|
||
require_reasoning=True,
|
||
)
|
||
|
||
adapted, _ = self.chat._convert_to_internal_request(req)
|
||
|
||
self.assertTrue(adapted.require_reasoning)
|
||
|
||
def test_process_messages_records_template_reasoning_state(self):
|
||
self.chat.default_chat_template_kwargs = {"thinking": True}
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=False
|
||
)
|
||
self.chat.reasoning_parser = "deepseek-v3"
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
)
|
||
rendered = MessageProcessingResult(
|
||
prompt="prompt",
|
||
prompt_ids=[1, 2, 3],
|
||
image_data=None,
|
||
audio_data=None,
|
||
video_data=None,
|
||
modalities=[],
|
||
stop=[],
|
||
)
|
||
|
||
with patch.object(
|
||
self.chat, "_apply_conversation_template", return_value=rendered
|
||
):
|
||
processed = self.chat._process_messages(request, is_multimodal=False)
|
||
|
||
self.assertTrue(processed.require_reasoning)
|
||
|
||
def test_kimi_tool_call_respects_explicit_reasoning_disable(self):
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=True
|
||
)
|
||
self.tm.server_args.reasoning_parser = "kimi_k2"
|
||
self.tm.server_args.tool_call_parser = "kimi_k2"
|
||
self.chat.reasoning_parser = "kimi_k2"
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"a": {"type": "integer"}},
|
||
},
|
||
},
|
||
}
|
||
],
|
||
tool_choice="required",
|
||
chat_template_kwargs={"thinking": False},
|
||
)
|
||
|
||
with patch.object(self.chat, "_process_messages") as proc_mock:
|
||
proc_mock.return_value = MessageProcessingResult(
|
||
"",
|
||
[1, 2, 3],
|
||
None,
|
||
None,
|
||
[],
|
||
[],
|
||
None,
|
||
)
|
||
|
||
adapted, _ = self.chat._convert_to_internal_request(req)
|
||
|
||
self.assertFalse(adapted.require_reasoning)
|
||
|
||
def test_default_chat_template_kwargs_applied_when_request_unset(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
self.chat.default_chat_template_kwargs = {"enable_thinking": False}
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertIs(kwargs["enable_thinking"], False)
|
||
|
||
def test_default_chat_template_kwargs_overridden_per_request(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
self.chat.default_chat_template_kwargs = {"enable_thinking": False}
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
chat_template_kwargs={"enable_thinking": True},
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertIs(kwargs["enable_thinking"], True)
|
||
|
||
def test_default_chat_template_kwargs_mirrors_reasoning_effort(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
self.chat.default_chat_template_kwargs = {"reasoning_effort": "high"}
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
self.assertEqual(req.reasoning_effort, "high")
|
||
|
||
def test_hunyuan_default_reasoning_effort_is_normalized_for_template(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
special_case="hunyuan_effort"
|
||
)
|
||
self.chat.reasoning_parser = "hunyuan"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
|
||
cases = [
|
||
("none", None, "no_think"),
|
||
("none", "xhigh", "high"),
|
||
]
|
||
for default_effort, request_effort, normalized_effort in cases:
|
||
with self.subTest(
|
||
default_effort=default_effort, request_effort=request_effort
|
||
):
|
||
self.chat.default_chat_template_kwargs = {
|
||
"reasoning_effort": default_effort
|
||
}
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
reasoning_effort=request_effort,
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertEqual(req.reasoning_effort, normalized_effort)
|
||
self.assertEqual(kwargs["reasoning_effort"], normalized_effort)
|
||
self.assertNotIn("reasoning_effort", req.chat_template_kwargs)
|
||
|
||
def test_hunyuan_reasoning_effort_precedence_survives_conversion(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
special_case="hunyuan_effort"
|
||
)
|
||
self.chat.reasoning_parser = "hunyuan"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
|
||
cases = [
|
||
("xhigh", "none", "high"),
|
||
(None, "no_think", "no_think"),
|
||
]
|
||
for request_effort, template_effort, normalized_effort in cases:
|
||
with self.subTest(
|
||
request_effort=request_effort, template_effort=template_effort
|
||
):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
reasoning_effort=request_effort,
|
||
chat_template_kwargs={"reasoning_effort": template_effort},
|
||
)
|
||
|
||
self.chat._convert_to_internal_request(req)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertEqual(req.reasoning_effort, normalized_effort)
|
||
self.assertEqual(kwargs["reasoning_effort"], normalized_effort)
|
||
self.assertNotIn("reasoning_effort", req.chat_template_kwargs)
|
||
|
||
def test_non_hunyuan_default_reasoning_effort_is_unchanged(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
self.chat.default_chat_template_kwargs = {"reasoning_effort": "medium"}
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertEqual(req.reasoning_effort, "medium")
|
||
self.assertEqual(kwargs["reasoning_effort"], "medium")
|
||
self.assertEqual(req.chat_template_kwargs["reasoning_effort"], "medium")
|
||
|
||
def test_k2_selected_terminator_reaches_sampling_params(self):
|
||
self.tm._config_overrides["reasoning_parser"] = "k2_horizon"
|
||
self.chat = OpenAIServingChat(self.tm, self.template_manager)
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
self.tm.tokenizer.encode.side_effect = lambda text, **_: (
|
||
[8, 9] if text == "</ifm|think_fast>" else [1, 2, 3]
|
||
)
|
||
req = ChatCompletionRequest(
|
||
model="IFM/K2-Horizon-7B",
|
||
messages=[{"role": "user", "content": "hi"}],
|
||
chat_template_kwargs={"reasoning_effort": "medium"},
|
||
)
|
||
|
||
processed = self.chat._process_messages(req, is_multimodal=False)
|
||
self.assertEqual(processed.reasoning_end_token_ids, [8, 9])
|
||
|
||
with patch.object(self.chat, "_process_messages", return_value=processed):
|
||
adapted, _ = self.chat._convert_to_internal_request(req)
|
||
self.assertEqual(
|
||
adapted.sampling_params["custom_params"][
|
||
REQUEST_REASONING_END_TOKEN_IDS_KEY
|
||
],
|
||
[8, 9],
|
||
)
|
||
|
||
def test_kimi_tool_call_keeps_template_default_thinking(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=True
|
||
)
|
||
self.tm.server_args.reasoning_parser = "kimi_k2"
|
||
self.tm.server_args.tool_call_parser = "kimi_k2"
|
||
self.chat.reasoning_parser = "kimi_k2"
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"a": {"type": "integer"}},
|
||
},
|
||
},
|
||
}
|
||
],
|
||
tool_choice="required",
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertNotIn("thinking", kwargs)
|
||
|
||
def test_kimi_tool_call_keeps_explicit_template_thinking(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=True
|
||
)
|
||
self.tm.server_args.reasoning_parser = "kimi_k2"
|
||
self.tm.server_args.tool_call_parser = "kimi_k2"
|
||
self.chat.reasoning_parser = "kimi_k2"
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"a": {"type": "integer"}},
|
||
},
|
||
},
|
||
}
|
||
],
|
||
tool_choice="required",
|
||
chat_template_kwargs={"thinking": True},
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertTrue(kwargs["thinking"])
|
||
|
||
def test_kimi_tool_call_keeps_explicit_template_thinking_false(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=True
|
||
)
|
||
self.tm.server_args.reasoning_parser = "kimi_k2"
|
||
self.tm.server_args.tool_call_parser = "kimi_k2"
|
||
self.chat.reasoning_parser = "kimi_k2"
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"a": {"type": "integer"}},
|
||
},
|
||
},
|
||
}
|
||
],
|
||
tool_choice="required",
|
||
chat_template_kwargs={"thinking": False},
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertFalse(kwargs["thinking"])
|
||
|
||
def test_jinja_uses_openai_tool_schema_first(self):
|
||
"""Ensure Jinja chat templates receive OpenAI-shaped tools by default."""
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"description": "Add two numbers.",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {
|
||
"a": {"type": "integer"},
|
||
"b": {"type": "integer"},
|
||
},
|
||
"required": ["a", "b"],
|
||
},
|
||
},
|
||
}
|
||
],
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
expected_tools = [tool.model_dump() for tool in req.tools]
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertEqual(kwargs["tools"], expected_tools)
|
||
|
||
def test_jinja_tool_schema_fallback_to_flat_function(self):
|
||
"""Fallback to function-only schema when template rejects OpenAI wrapper."""
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"description": "Add two numbers.",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {
|
||
"a": {"type": "integer"},
|
||
"b": {"type": "integer"},
|
||
},
|
||
"required": ["a", "b"],
|
||
},
|
||
},
|
||
}
|
||
],
|
||
)
|
||
|
||
self.tm.tokenizer.apply_chat_template.side_effect = [
|
||
RuntimeError("template expects flat tools format"),
|
||
[1, 2, 3],
|
||
]
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
first_tools = self.tm.tokenizer.apply_chat_template.call_args_list[0].kwargs[
|
||
"tools"
|
||
]
|
||
second_tools = self.tm.tokenizer.apply_chat_template.call_args_list[1].kwargs[
|
||
"tools"
|
||
]
|
||
self.assertEqual(first_tools, [tool.model_dump() for tool in req.tools])
|
||
self.assertEqual(
|
||
second_tools, [tool.function.model_dump() for tool in req.tools]
|
||
)
|
||
|
||
def test_xgrammar_tag_omits_reasoning_when_parser_owns_it(self):
|
||
"""ReasonerGrammarBackend owns the thinking prefix when a parser is set."""
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=True
|
||
)
|
||
self.tm.server_args.reasoning_parser = "kimi_k2"
|
||
self.tm.server_args.tool_call_parser = "kimi_k2"
|
||
self.chat.reasoning_parser = "kimi_k2"
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "add",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {
|
||
"a": {"type": "integer"},
|
||
"b": {"type": "integer"},
|
||
},
|
||
"required": ["a", "b"],
|
||
},
|
||
"strict": True,
|
||
},
|
||
}
|
||
],
|
||
tool_choice="required",
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||
) as parser_cls:
|
||
parser = parser_cls.return_value
|
||
parser.get_structure_constraint.return_value = ("structural_tag", "tag")
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
parser.get_structure_constraint.assert_called_once()
|
||
self.assertFalse(
|
||
parser.get_structure_constraint.call_args.kwargs["thinking_mode"]
|
||
)
|
||
|
||
def test_kimi_k3_constraint_failure_keeps_native_stop_format(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.chat.chat_encoding_spec = "kimi_k3"
|
||
self.chat.tool_call_parser = "kimi_k3"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
tool = {
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"location": {"type": "string"}},
|
||
"required": ["location"],
|
||
},
|
||
"strict": True,
|
||
},
|
||
}
|
||
|
||
cases = (
|
||
(None, [TOOLS_CLOSE]),
|
||
("USER_STOP", ["USER_STOP", TOOLS_CLOSE]),
|
||
(["USER_STOP"], ["USER_STOP", TOOLS_CLOSE]),
|
||
([TOOLS_CLOSE], [TOOLS_CLOSE]),
|
||
)
|
||
for request_stop, expected in cases:
|
||
with (
|
||
self.subTest(request_stop=request_stop),
|
||
patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||
) as parser_cls,
|
||
):
|
||
parser = parser_cls.return_value
|
||
parser.detector.eot_token = TOOLS_CLOSE
|
||
parser.detector.parses_required_natively.return_value = False
|
||
parser.get_structure_constraint.return_value = None
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Weather in Paris?"}],
|
||
tools=[tool],
|
||
tool_choice="required",
|
||
stop=request_stop,
|
||
)
|
||
original_stop = (
|
||
list(request.stop)
|
||
if isinstance(request.stop, list)
|
||
else request.stop
|
||
)
|
||
|
||
result = self.chat._process_messages(request, is_multimodal=False)
|
||
|
||
self.assertEqual(result.stop, expected)
|
||
self.assertEqual(request.stop, original_stop)
|
||
self.assertIsNone(result.tool_call_constraint)
|
||
|
||
def test_kimi_k3_tool_call_stop_is_scoped_to_active_tools(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.chat.chat_encoding_spec = "kimi_k3"
|
||
self.chat.tool_call_parser = "kimi_k3"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Weather in Paris?"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"parameters": {"type": "object"},
|
||
},
|
||
}
|
||
],
|
||
tool_choice="none",
|
||
)
|
||
|
||
result = self.chat._process_messages(request, is_multimodal=False)
|
||
|
||
self.assertIsNone(result.stop)
|
||
|
||
def test_kimi_k3_encoder_receives_wire_request_fields(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.chat.chat_encoding_spec = "kimi_k3"
|
||
self.tm.model_config.is_multimodal = True
|
||
self.tm.tokenizer.apply_chat_template.return_value = [7, 8, 9]
|
||
tool = {
|
||
"type": "function",
|
||
"function": {
|
||
"name": "weather",
|
||
"parameters": {"type": "object"},
|
||
},
|
||
}
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{
|
||
"role": "developer",
|
||
"content": "<|kimi_image_placeholder|>",
|
||
"tools": [tool],
|
||
},
|
||
{
|
||
"role": "user",
|
||
"content": [
|
||
{
|
||
"type": "text",
|
||
"text": "Explain <|kimi_image_placeholder|>",
|
||
},
|
||
{"type": "image_url", "image_url": {"url": "image-1"}},
|
||
],
|
||
},
|
||
{
|
||
"role": "assistant",
|
||
"content": None,
|
||
"reasoning_content": "Inspect <|kimi_image_placeholder|>",
|
||
"tool_calls": [
|
||
{
|
||
"id": "call-1",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "inspect",
|
||
"arguments": {
|
||
"source": "<|kimi_image_placeholder|>",
|
||
"nested": ["<|kimi_image_placeholder|>"],
|
||
},
|
||
},
|
||
}
|
||
],
|
||
},
|
||
],
|
||
tools=[tool],
|
||
tool_choice="required",
|
||
response_format={
|
||
"type": "json_schema",
|
||
"json_schema": {
|
||
"name": "answer",
|
||
"schema": {"type": "object"},
|
||
"strict": False,
|
||
},
|
||
},
|
||
)
|
||
|
||
result = self.chat._process_messages(request, is_multimodal=True)
|
||
|
||
call = self.tm.tokenizer.apply_chat_template.call_args
|
||
rendered_messages = call.args[0]
|
||
self.assertEqual(rendered_messages[0]["role"], "system")
|
||
self.assertEqual(
|
||
rendered_messages[0]["content"], "<| kimi_image_placeholder |>"
|
||
)
|
||
self.assertNotIn("strict", rendered_messages[0]["tools"][0]["function"])
|
||
self.assertEqual(
|
||
rendered_messages[1]["content"][0]["text"],
|
||
"Explain <| kimi_image_placeholder |>",
|
||
)
|
||
self.assertEqual(
|
||
rendered_messages[2]["reasoning_content"],
|
||
"Inspect <| kimi_image_placeholder |>",
|
||
)
|
||
self.assertEqual(
|
||
rendered_messages[2]["tool_calls"][0]["function"]["arguments"],
|
||
{
|
||
"source": "<| kimi_image_placeholder |>",
|
||
"nested": ["<| kimi_image_placeholder |>"],
|
||
},
|
||
)
|
||
self.assertEqual(call.kwargs["image_prompts"], ["<|media_pad|>"])
|
||
self.assertEqual(call.kwargs["tool_choice"], "required")
|
||
self.assertNotIn("strict", call.kwargs["tools"][0]["function"])
|
||
self.assertEqual(
|
||
call.kwargs["response_format"]["json_schema"]["schema"],
|
||
{"type": "object"},
|
||
)
|
||
self.assertNotIn("schema_", call.kwargs["response_format"]["json_schema"])
|
||
self.assertEqual(result.prompt_ids, [7, 8, 9])
|
||
self.assertEqual(result.image_data[0].url, "image-1")
|
||
|
||
def test_kimi_k3_neutralizes_text_only_assistant_history(self):
|
||
self.template_manager.chat_template_name = None
|
||
self.chat.chat_encoding_spec = "kimi_k3"
|
||
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "Run it"},
|
||
{
|
||
"role": "assistant",
|
||
"content": None,
|
||
"reasoning_content": "Read <|kimi_image_placeholder|>",
|
||
"tool_calls": [
|
||
{
|
||
"id": "call-1",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "shell",
|
||
"arguments": "not-json <|kimi_image_placeholder|>",
|
||
},
|
||
}
|
||
],
|
||
},
|
||
],
|
||
)
|
||
|
||
self.chat._process_messages(request, is_multimodal=False)
|
||
|
||
messages = self.tm.tokenizer.apply_chat_template.call_args.args[0]
|
||
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||
self.assertEqual(messages[-1]["role"], "assistant")
|
||
self.assertEqual(
|
||
messages[-1]["reasoning_content"],
|
||
"Read <| kimi_image_placeholder |>",
|
||
)
|
||
self.assertEqual(
|
||
messages[-1]["tool_calls"][0]["function"]["arguments"],
|
||
"not-json <| kimi_image_placeholder |>",
|
||
)
|
||
self.assertNotIn("image_prompts", kwargs)
|
||
|
||
def test_message_tools_participate_in_validation_across_encodings(self):
|
||
tool = {
|
||
"type": "function",
|
||
"function": {
|
||
"name": "weather",
|
||
"parameters": {"type": "object"},
|
||
},
|
||
}
|
||
messages = [
|
||
{"role": "system", "content": "", "tools": [tool]},
|
||
{"role": "user", "content": "Weather?"},
|
||
]
|
||
for chat_encoding_spec in (None, "dsv4", "dsv32", "kimi_k3"):
|
||
with self.subTest(chat_encoding_spec=chat_encoding_spec):
|
||
self.chat.chat_encoding_spec = chat_encoding_spec
|
||
automatic = ChatCompletionRequest(
|
||
model="x", messages=messages, tool_choice=None
|
||
)
|
||
self.assertEqual(automatic.tool_choice, "auto")
|
||
self.assertIsNone(self.chat._validate_request(automatic))
|
||
|
||
required = ChatCompletionRequest(
|
||
model="x", messages=messages, tool_choice="required"
|
||
)
|
||
self.assertIsNone(self.chat._validate_request(required))
|
||
|
||
duplicate = ChatCompletionRequest(
|
||
model="x",
|
||
messages=messages,
|
||
tools=[tool],
|
||
tool_choice="required",
|
||
)
|
||
self.assertEqual(
|
||
self.chat._validate_request(duplicate),
|
||
"Tool names must be unique across request and message tools.",
|
||
)
|
||
|
||
def test_jinja_rejects_non_object_tool_call_arguments(self):
|
||
"""History tool call arguments must parse to a JSON object."""
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
|
||
for arguments in ['"Beijing"', '["Beijing"]']:
|
||
with self.subTest(arguments=arguments):
|
||
self.tm.tokenizer.apply_chat_template.reset_mock()
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "Where is it raining?"},
|
||
{
|
||
"role": "assistant",
|
||
"content": None,
|
||
"tool_calls": [
|
||
{
|
||
"id": "call_1",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"arguments": arguments,
|
||
},
|
||
}
|
||
],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"tool_call_id": "call_1",
|
||
"content": "sunny",
|
||
},
|
||
],
|
||
)
|
||
|
||
with self.assertRaisesRegex(ValueError, "must be a JSON object"):
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
self.tm.tokenizer.apply_chat_template.assert_not_called()
|
||
|
||
def test_jinja_accepts_object_tool_call_arguments_string(self):
|
||
"""OpenAI JSON string arguments are converted to dicts for templates."""
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "Where is it raining?"},
|
||
{
|
||
"role": "assistant",
|
||
"content": None,
|
||
"tool_calls": [
|
||
{
|
||
"id": "call_1",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"arguments": '{"city": "Beijing"}',
|
||
},
|
||
}
|
||
],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"tool_call_id": "call_1",
|
||
"content": "sunny",
|
||
},
|
||
],
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
messages = self.tm.tokenizer.apply_chat_template.call_args.args[0]
|
||
self.assertEqual(
|
||
messages[1]["tool_calls"][0]["function"]["arguments"],
|
||
{"city": "Beijing"},
|
||
)
|
||
|
||
def test_dsv_encoders_reject_non_object_tool_call_arguments(self):
|
||
"""DeepSeek encoders should reject history tool call scalars as BadRequest."""
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.chat._dsv4_reasoning_effort_profile = "preview"
|
||
|
||
for chat_encoding_spec in ("dsv4", "dsv32"):
|
||
with self.subTest(chat_encoding_spec=chat_encoding_spec):
|
||
self.chat.chat_encoding_spec = chat_encoding_spec
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "Where is it raining?"},
|
||
{
|
||
"role": "assistant",
|
||
"content": None,
|
||
"tool_calls": [
|
||
{
|
||
"id": "call_1",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"arguments": '"Beijing"',
|
||
},
|
||
}
|
||
],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"tool_call_id": "call_1",
|
||
"content": "sunny",
|
||
},
|
||
],
|
||
)
|
||
|
||
with self.assertRaisesRegex(ValueError, "must be a JSON object"):
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
def test_dsv_encoders_accept_object_tool_call_arguments_string(self):
|
||
"""DeepSeek encoders accept object-shaped OpenAI JSON string arguments."""
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.jinja_template_content_format = "string"
|
||
self.chat._dsv4_reasoning_effort_profile = "preview"
|
||
|
||
for chat_encoding_spec in ("dsv4", "dsv32"):
|
||
with self.subTest(chat_encoding_spec=chat_encoding_spec):
|
||
self.chat.chat_encoding_spec = chat_encoding_spec
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "Where is it raining?"},
|
||
{
|
||
"role": "assistant",
|
||
"content": None,
|
||
"tool_calls": [
|
||
{
|
||
"id": "call_1",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"arguments": '{"city": "Beijing"}',
|
||
},
|
||
}
|
||
],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"tool_call_id": "call_1",
|
||
"content": "sunny",
|
||
},
|
||
],
|
||
)
|
||
|
||
self.chat._process_messages(req, is_multimodal=False)
|
||
|
||
def test_stop_str_isolation_between_requests(self):
|
||
"""Test that stop strings from one request don't affect subsequent requests.
|
||
|
||
This tests the fix for the bug where conv.stop_str was being mutated globally,
|
||
causing stop strings from one request to persist in subsequent requests.
|
||
"""
|
||
# Mock conversation template with initial stop_str
|
||
initial_stop_str = ["\n"]
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
# Create a mock conversation object that will be returned by generate_chat_conv
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_ins.image_data = None
|
||
conv_ins.audio_data = None
|
||
conv_ins.modalities = []
|
||
conv_ins.stop_str = (
|
||
initial_stop_str.copy()
|
||
) # Template's default stop strings
|
||
conv_mock.return_value = conv_ins
|
||
|
||
# First request with additional stop string
|
||
req1 = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "First request"}],
|
||
stop=["CUSTOM_STOP"],
|
||
)
|
||
|
||
# Call the actual _apply_conversation_template method (not mocked)
|
||
result1 = self.chat._apply_conversation_template(req1, is_multimodal=False)
|
||
|
||
# Verify first request has both stop strings
|
||
expected_stop1 = initial_stop_str + ["CUSTOM_STOP"]
|
||
self.assertEqual(result1.stop, expected_stop1)
|
||
|
||
# Verify the original template's stop_str wasn't mutated after first request
|
||
self.assertEqual(conv_ins.stop_str, initial_stop_str)
|
||
|
||
# Second request without additional stop string
|
||
req2 = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Second request"}],
|
||
# No custom stop strings
|
||
)
|
||
result2 = self.chat._apply_conversation_template(req2, is_multimodal=False)
|
||
|
||
# Verify second request only has original stop strings (no CUSTOM_STOP from req1)
|
||
self.assertEqual(result2.stop, initial_stop_str)
|
||
self.assertNotIn("CUSTOM_STOP", result2.stop)
|
||
self.assertEqual(conv_ins.stop_str, initial_stop_str)
|
||
|
||
def test_unstreamed_tool_args_completion(self):
|
||
"""Test that remaining tool call arguments are sent when generation finishes."""
|
||
|
||
# Mock FunctionCallParser with detector that has partial tool call data
|
||
mock_parser = Mock()
|
||
mock_detector = Mock()
|
||
|
||
# Simulate a tool call that was partially streamed
|
||
mock_detector.prev_tool_call_arr = [
|
||
{
|
||
"name": "get_weather",
|
||
"arguments": {"location": "San Francisco", "unit": "celsius"},
|
||
}
|
||
]
|
||
mock_detector.streamed_args_for_tool = [
|
||
'{"location": "San Francisco"' # Partial arguments streamed so far
|
||
]
|
||
mock_parser.detector = mock_detector
|
||
|
||
content = {
|
||
"meta_info": {
|
||
"id": "chatcmpl-test123",
|
||
}
|
||
}
|
||
|
||
request = ChatCompletionRequest(
|
||
model="test",
|
||
messages=[{"role": "user", "content": "What's the weather?"}],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
)
|
||
|
||
# Test the completion method
|
||
result = self.chat._check_for_unstreamed_tool_args(
|
||
parser=mock_parser,
|
||
content=content,
|
||
request=request,
|
||
index=0,
|
||
)
|
||
|
||
# Should return a chunk with remaining arguments
|
||
self.assertIsNotNone(result, "Should return chunk with remaining arguments")
|
||
|
||
# Parse the result to verify content
|
||
self.assertTrue(result.startswith("data: "))
|
||
chunk = json.loads(result[6:])
|
||
tool_calls = chunk["choices"][0]["delta"]["tool_calls"]
|
||
self.assertEqual(len(tool_calls), 1)
|
||
arguments = tool_calls[0]["function"]["arguments"]
|
||
self.assertIn(', "unit": "celsius"}', arguments)
|
||
|
||
self.assertIn(
|
||
'"finish_reason":null',
|
||
result,
|
||
"Should not include finish_reason in completion chunk",
|
||
)
|
||
|
||
def test_unstreamed_tool_args_no_completion_needed(self):
|
||
"""Test that no completion chunk is sent when all arguments were already streamed."""
|
||
|
||
# Mock FunctionCallParser with detector that has complete tool call data
|
||
mock_parser = Mock()
|
||
mock_detector = Mock()
|
||
|
||
# Simulate a tool call that was completely streamed
|
||
mock_detector.prev_tool_call_arr = [
|
||
{"name": "get_weather", "arguments": {"location": "San Francisco"}}
|
||
]
|
||
mock_detector.streamed_args_for_tool = [
|
||
'{"location": "San Francisco"}' # All arguments already streamed
|
||
]
|
||
mock_parser.detector = mock_detector
|
||
|
||
content = {
|
||
"meta_info": {
|
||
"id": "chatcmpl-test123",
|
||
}
|
||
}
|
||
|
||
request = ChatCompletionRequest(
|
||
model="test",
|
||
messages=[{"role": "user", "content": "What's the weather?"}],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
)
|
||
|
||
# Test the completion method
|
||
result = self.chat._check_for_unstreamed_tool_args(
|
||
parser=mock_parser,
|
||
content=content,
|
||
request=request,
|
||
index=0,
|
||
)
|
||
|
||
# Should return None since no completion is needed
|
||
self.assertIsNone(result, "Should return None when no completion is needed")
|
||
|
||
def test_unstreamed_tool_args_raw_string_no_completion_needed(self):
|
||
"""Test raw JSON string arguments are not JSON-encoded again at finish."""
|
||
|
||
mock_parser = Mock()
|
||
mock_detector = Mock()
|
||
mock_detector.prev_tool_call_arr = [
|
||
{"name": "report_template_recognition_commit", "arguments": "{}"}
|
||
]
|
||
mock_detector.streamed_args_for_tool = ["{}"]
|
||
mock_parser.detector = mock_detector
|
||
|
||
content = {
|
||
"meta_info": {
|
||
"id": "chatcmpl-test123",
|
||
}
|
||
}
|
||
|
||
request = ChatCompletionRequest(
|
||
model="test",
|
||
messages=[{"role": "user", "content": "commit"}],
|
||
tools=[
|
||
{
|
||
"type": "function",
|
||
"function": {"name": "report_template_recognition_commit"},
|
||
}
|
||
],
|
||
)
|
||
|
||
result = self.chat._check_for_unstreamed_tool_args(
|
||
parser=mock_parser,
|
||
content=content,
|
||
request=request,
|
||
index=0,
|
||
)
|
||
|
||
self.assertIsNone(result, "Should not append encoded quotes")
|
||
|
||
def test_unstreamed_tool_args_raw_string_completion(self):
|
||
"""Test remaining raw JSON string arguments are sent at finish."""
|
||
|
||
mock_parser = Mock()
|
||
mock_detector = Mock()
|
||
mock_detector.prev_tool_call_arr = [
|
||
{
|
||
"name": "get_weather",
|
||
"arguments": '{"location": "San Francisco", "unit": "celsius"}',
|
||
}
|
||
]
|
||
mock_detector.streamed_args_for_tool = ['{"location": "San Francisco"']
|
||
mock_parser.detector = mock_detector
|
||
|
||
content = {
|
||
"meta_info": {
|
||
"id": "chatcmpl-test123",
|
||
}
|
||
}
|
||
|
||
request = ChatCompletionRequest(
|
||
model="test",
|
||
messages=[{"role": "user", "content": "What's the weather?"}],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
)
|
||
|
||
result = self.chat._check_for_unstreamed_tool_args(
|
||
parser=mock_parser,
|
||
content=content,
|
||
request=request,
|
||
index=0,
|
||
)
|
||
|
||
self.assertIsNotNone(result, "Should return chunk with remaining arguments")
|
||
chunk = json.loads(result[6:])
|
||
tool_calls = chunk["choices"][0]["delta"]["tool_calls"]
|
||
self.assertEqual(tool_calls[0]["function"]["arguments"], ', "unit": "celsius"}')
|
||
|
||
def test_unstreamed_tool_args_no_parser_data(self):
|
||
"""Test that no completion chunk is sent when parser has no tool call data."""
|
||
|
||
# Mock FunctionCallParser with empty detector
|
||
mock_parser = Mock()
|
||
mock_detector = Mock()
|
||
mock_detector.prev_tool_call_arr = []
|
||
mock_detector.streamed_args_for_tool = []
|
||
mock_parser.detector = mock_detector
|
||
|
||
content = {
|
||
"meta_info": {
|
||
"id": "chatcmpl-test123",
|
||
}
|
||
}
|
||
|
||
request = ChatCompletionRequest(
|
||
model="test",
|
||
messages=[{"role": "user", "content": "What's the weather?"}],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
)
|
||
|
||
# Test the completion method
|
||
result = self.chat._check_for_unstreamed_tool_args(
|
||
parser=mock_parser,
|
||
content=content,
|
||
request=request,
|
||
index=0,
|
||
)
|
||
|
||
# Should return None since there's no parser data
|
||
self.assertIsNone(
|
||
result, "Should return None when parser has no tool call data"
|
||
)
|
||
|
||
# ------------- kimi_k2 tool_call_id formatting -------------
|
||
def test_kimi_k2_non_streaming_tool_call_id_with_history(self):
|
||
"""Ensure non-streaming tool_call.id increase with tool calls history for kimi_k2 parser."""
|
||
|
||
# Force kimi_k2 parser
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
|
||
# Prepare request with tool calls history
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "What's the weather today in paris?"},
|
||
{
|
||
"role": "assistant",
|
||
"content": "Let me do some search first.",
|
||
"tool_calls": [
|
||
{
|
||
"id": "functions.get_weather:0",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"arguments": '{"city": "Paris"}',
|
||
},
|
||
}
|
||
],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"content": "It's rainy in paris now.",
|
||
"tool_call_id": "functions.get_weather:0",
|
||
},
|
||
{
|
||
"role": "assistant",
|
||
"content": "It's rainy now.",
|
||
},
|
||
{
|
||
"role": "user",
|
||
"content": "What about LA and Tokyo?",
|
||
},
|
||
],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
stream=False,
|
||
)
|
||
|
||
# Mock FunctionCallParser.parse_non_stream to return one tool call
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||
) as ParserMock:
|
||
parser_instance = ParserMock.return_value
|
||
|
||
# Build a mock ToolCallItem-like object
|
||
call_info = Mock()
|
||
call_info.name = "get_weather"
|
||
call_info.parameters = '{"city":"Loa Angeles"}'
|
||
# Kimi-K2 series models might generate fixed number tool_indx,
|
||
# ignoring the tool calls history and mess up all the following tool calls
|
||
call_info.tool_index = 0
|
||
|
||
call_info2 = Mock()
|
||
call_info2.name = "get_weather"
|
||
call_info2.parameters = '{"city":"Tokyo"}'
|
||
call_info2.tool_index = 1
|
||
|
||
parser_instance.has_tool_call.return_value = True
|
||
parser_instance.parse_non_stream.return_value = (
|
||
"",
|
||
[call_info, call_info2],
|
||
)
|
||
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
tools = [
|
||
{"type": "function", "function": {"name": "get_weather"}},
|
||
]
|
||
|
||
history_tool_calls_cnt = self.chat._get_history_tool_calls_cnt(req)
|
||
tool_calls, remaining_text, _ = self.chat._process_tool_calls(
|
||
text="<|tool_calls_section_begin|>...",
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
history_tool_calls_cnt=history_tool_calls_cnt,
|
||
)
|
||
|
||
self.assertEqual(history_tool_calls_cnt, 1)
|
||
self.assertIsNotNone(tool_calls)
|
||
self.assertEqual(len(tool_calls), 2)
|
||
self.assertEqual(tool_calls[0].id, "functions.get_weather:1")
|
||
self.assertEqual(tool_calls[0].function.name, "get_weather")
|
||
self.assertEqual(tool_calls[1].id, "functions.get_weather:2")
|
||
self.assertEqual(tool_calls[1].function.name, "get_weather")
|
||
|
||
def test_non_streaming_tool_call_index_is_the_call_ordinal(self):
|
||
"""Two calls to one tool are numbered 0 and 1, as in the streaming deltas,
|
||
not by the detector's tool_index (0 for both)."""
|
||
self.chat.tool_call_parser = "deepseekv4"
|
||
tools = [{"type": "function", "function": {"name": "get_weather"}}]
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||
) as ParserMock:
|
||
parser_instance = ParserMock.return_value
|
||
calls = []
|
||
for city in ("San Francisco", "London"):
|
||
call_info = Mock()
|
||
call_info.name = "get_weather"
|
||
call_info.parameters = json.dumps({"location": city})
|
||
call_info.tool_index = 0
|
||
calls.append(call_info)
|
||
parser_instance.has_tool_call.return_value = True
|
||
parser_instance.parse_non_stream.return_value = ("", calls)
|
||
|
||
tool_calls, _, finish_reason = self.chat._process_tool_calls(
|
||
text="<|DSML|tool_calls>...",
|
||
tools=tools,
|
||
finish_reason={"type": "stop", "matched": None},
|
||
history_tool_calls_cnt=0,
|
||
)
|
||
|
||
self.assertEqual([tc.index for tc in tool_calls], [0, 1])
|
||
self.assertEqual(
|
||
[tc.function.arguments for tc in tool_calls],
|
||
[
|
||
json.dumps({"location": "San Francisco"}),
|
||
json.dumps({"location": "London"}),
|
||
],
|
||
)
|
||
self.assertEqual(finish_reason["type"], "tool_calls")
|
||
|
||
def test_required_tool_choice_skips_json_fallback_for_native_parser(self):
|
||
"""A structural-tag parser owns the output format, so a missing tool
|
||
call must not be pushed through the json_schema array fallback."""
|
||
self.chat.tool_call_parser = "kimi_k3"
|
||
tools = [
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"city": {"type": "string"}},
|
||
},
|
||
},
|
||
}
|
||
]
|
||
texts = {
|
||
"prose": "<|open|>response<|sep|>I'll check that.<|close|>response<|sep|>",
|
||
"empty": "",
|
||
"json_object": '{"name": "get_weather", "parameters": {"city": "Paris"}}',
|
||
}
|
||
for choice in (
|
||
"required",
|
||
ToolChoice(function=ToolChoiceFuncName(name="get_weather")),
|
||
):
|
||
for label, text in texts.items():
|
||
with self.subTest(tool_choice=choice, payload=label):
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
with self.assertLogs(
|
||
"sglang.srt.entrypoints.openai.serving_chat", level="WARNING"
|
||
) as logs:
|
||
tool_calls, remaining, finish_reason = (
|
||
self.chat._process_tool_calls(
|
||
text=text,
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
tool_choice=choice,
|
||
)
|
||
)
|
||
self.assertIsNone(tool_calls)
|
||
self.assertEqual(remaining, text)
|
||
self.assertEqual(finish_reason["type"], "stop")
|
||
self.assertNotIn("Tool call parsing error", "\n".join(logs.output))
|
||
|
||
def test_truncated_native_tool_call_logs_and_drops(self):
|
||
"""A tools section cut off before its closing tag parses to zero calls
|
||
without raising; the sync path used to drop it with no log at all."""
|
||
self.chat.tool_call_parser = "kimi_k3"
|
||
tools = [
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"city": {"type": "string"}},
|
||
},
|
||
},
|
||
}
|
||
]
|
||
truncated = TOOLS_OPEN + '<|open|>call tool="get_weather" index="1"<|sep|>'
|
||
for choice in ("auto", "required"):
|
||
with self.subTest(tool_choice=choice):
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
with self.assertLogs(
|
||
"sglang.srt.entrypoints.openai.serving_chat", level="WARNING"
|
||
) as logs:
|
||
tool_calls, remaining, finish_reason = (
|
||
self.chat._process_tool_calls(
|
||
text=truncated,
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
tool_choice=choice,
|
||
)
|
||
)
|
||
self.assertIsNone(tool_calls)
|
||
self.assertEqual(remaining, "")
|
||
self.assertEqual(finish_reason["type"], "stop")
|
||
self.assertIn("no complete call", "\n".join(logs.output))
|
||
|
||
def test_required_tool_choice_json_fallback_tolerates_odd_shapes(self):
|
||
"""Parsers without a structural tag keep the JSON array fallback, but a
|
||
non-array payload degrades instead of raising an opaque TypeError."""
|
||
self.chat.tool_call_parser = "glm45"
|
||
tools = [
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {"city": {"type": "string"}},
|
||
},
|
||
},
|
||
}
|
||
]
|
||
cases = [
|
||
(
|
||
'[{"name": "get_weather", "parameters": {"city": "Paris"}}]',
|
||
'{"city": "Paris"}',
|
||
),
|
||
(
|
||
'{"name": "get_weather", "parameters": {"city": "Paris"}}',
|
||
'{"city": "Paris"}',
|
||
),
|
||
('{"name": "get_weather"}', "{}"),
|
||
]
|
||
for text, expected_args in cases:
|
||
with self.subTest(text=text):
|
||
tool_calls, _, finish_reason = self.chat._process_tool_calls(
|
||
text=text,
|
||
tools=tools,
|
||
finish_reason={"type": "stop", "matched": None},
|
||
tool_choice="required",
|
||
)
|
||
self.assertEqual(len(tool_calls), 1)
|
||
self.assertEqual(tool_calls[0].function.name, "get_weather")
|
||
self.assertEqual(tool_calls[0].function.arguments, expected_args)
|
||
self.assertEqual(finish_reason["type"], "tool_calls")
|
||
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
with self.assertLogs(
|
||
"sglang.srt.entrypoints.openai.serving_chat", level="ERROR"
|
||
):
|
||
tool_calls, remaining, finish_reason = self.chat._process_tool_calls(
|
||
text='["get_weather"]',
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
tool_choice="required",
|
||
)
|
||
self.assertIsNone(tool_calls)
|
||
self.assertEqual(remaining, '["get_weather"]')
|
||
self.assertEqual(finish_reason["type"], "stop")
|
||
|
||
def test_required_tool_choice_rejects_conflicting_output_constraint(self):
|
||
"""response_format and a forced tool call cannot both be honored: the
|
||
tool-call constraint was dropped with only a warning, so the model was
|
||
constrained to a shape that can never contain a tool call."""
|
||
tools = [{"type": "function", "function": {"name": "get_weather"}}]
|
||
constraint = ("structural_tag", None)
|
||
conflicting = [
|
||
{"type": "json_object"},
|
||
{
|
||
"type": "json_schema",
|
||
"json_schema": {"name": "a", "schema": {"type": "object"}},
|
||
},
|
||
]
|
||
for tool_choice in (
|
||
"required",
|
||
ToolChoice(function=ToolChoiceFuncName(name="get_weather")),
|
||
):
|
||
for response_format in conflicting:
|
||
with self.subTest(tool_choice=tool_choice, rf=response_format["type"]):
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "hi"}],
|
||
tools=tools,
|
||
tool_choice=tool_choice,
|
||
response_format=response_format,
|
||
)
|
||
with self.assertRaises(ValueError) as ctx:
|
||
request.to_sampling_params(
|
||
stop=[],
|
||
model_generation_config={},
|
||
tool_call_constraint=constraint,
|
||
)
|
||
self.assertIn("cannot be combined", str(ctx.exception))
|
||
|
||
def test_auto_tool_choice_keeps_response_format_without_raising(self):
|
||
""" "auto" means the model need not call a tool, so dropping the
|
||
tool-call constraint still leaves a satisfiable request."""
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "hi"}],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
tool_choice="auto",
|
||
response_format={"type": "json_object"},
|
||
)
|
||
sampling_params = request.to_sampling_params(
|
||
stop=[],
|
||
model_generation_config={},
|
||
tool_call_constraint=("structural_tag", None),
|
||
)
|
||
self.assertEqual(sampling_params["json_schema"], '{"type": "object"}')
|
||
|
||
def test_kimi_k2_streaming_tool_call_id_with_history(self):
|
||
"""Ensure streaming first chunk tool_call.id increase with tool calls history for kimi_k2 parser."""
|
||
|
||
# Force kimi_k2 parser
|
||
self.chat.tool_call_parser = "kimi_k2"
|
||
|
||
# Prepare request with tool calls history
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "user", "content": "What's the weather today in paris?"},
|
||
{
|
||
"role": "assistant",
|
||
"content": "Let me do some search first.",
|
||
"tool_calls": [
|
||
{
|
||
"id": "functions.get_weather:0",
|
||
"type": "function",
|
||
"function": {
|
||
"name": "get_weather",
|
||
"arguments": '{"city": "Paris"}',
|
||
},
|
||
}
|
||
],
|
||
},
|
||
{
|
||
"role": "tool",
|
||
"content": "It's rainy in paris now.",
|
||
"tool_call_id": "functions.get_weather:0",
|
||
},
|
||
{
|
||
"role": "assistant",
|
||
"content": "It's rainy now.",
|
||
},
|
||
{
|
||
"role": "user",
|
||
"content": "What about LA?",
|
||
},
|
||
],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
stream=True,
|
||
)
|
||
|
||
# Patch FunctionCallParser used inside _process_tool_call_stream
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||
) as ParserMock:
|
||
parser_instance = ParserMock.return_value
|
||
|
||
# First call returns one ToolCallItem-like chunk (with name)
|
||
first_chunk_call = Mock()
|
||
# Kimi-K2 series models might generate fixed number tool_indx,
|
||
# ignoring the tool calls history and mess up all the following tool calls
|
||
first_chunk_call.tool_index = 0
|
||
first_chunk_call.name = "get_weather"
|
||
first_chunk_call.parameters = ""
|
||
parser_instance.parse_stream_chunk.side_effect = [
|
||
("", [first_chunk_call]),
|
||
("", []),
|
||
]
|
||
|
||
async def collect_first_tool_chunk():
|
||
gen = self.chat._process_tool_call_stream(
|
||
index=0,
|
||
delta="irrelevant",
|
||
parser_dict={},
|
||
content={"meta_info": {"id": "chatcmpl-test"}},
|
||
request=req,
|
||
has_tool_calls={},
|
||
)
|
||
# Get first yielded SSE line
|
||
line = None
|
||
async for emitted in gen:
|
||
line = emitted
|
||
break
|
||
return line
|
||
|
||
loop = get_or_create_event_loop()
|
||
line = loop.run_until_complete(collect_first_tool_chunk())
|
||
self.assertIsNotNone(line)
|
||
self.assertTrue(line.startswith("data: "))
|
||
|
||
payload = json.loads(line[len("data: ") :])
|
||
tool_calls = payload["choices"][0]["delta"]["tool_calls"]
|
||
self.assertEqual(tool_calls[0]["id"], "functions.get_weather:1")
|
||
|
||
def test_dpsk_v32_encoding_path(self):
|
||
"""Test DeepSeek V3.2 encoding path detection and application."""
|
||
from sglang.srt.parser.template_manager import TemplateManager
|
||
|
||
# Only mock the fields that _use_dpsk_v32_encoding() actually reads:
|
||
# tokenizer.chat_template and hf_config.architectures
|
||
tm = _MockTokenizerManager()
|
||
|
||
mock_hf_config = Mock()
|
||
mock_hf_config.architectures = ["DeepseekV32ForCausalLM"]
|
||
mock_hf_config.to_dict.return_value = {}
|
||
tm.model_config.hf_config = mock_hf_config
|
||
|
||
# Case 1: No chat template + DeepSeek V3.2 arch -> should use dsv32 encoding
|
||
tm.tokenizer.chat_template = None
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
self.assertEqual(serving_chat.chat_encoding_spec, "dsv32")
|
||
|
||
# Case 2: Chat template exists -> should NOT use dsv32 encoding
|
||
tm.tokenizer.chat_template = "some template"
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
self.assertIsNone(serving_chat.chat_encoding_spec)
|
||
|
||
# Case 3: Not DeepSeek V3.2 architecture -> should NOT use dsv32 encoding
|
||
tm.tokenizer.chat_template = None
|
||
mock_hf_config.architectures = ["LlamaForCausalLM"]
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
self.assertIsNone(serving_chat.chat_encoding_spec)
|
||
|
||
# Case 4: DeepseekV4 arch -> always dsv4, even with chat_template
|
||
# (release ships a stale V3 jinja we deliberately override).
|
||
mock_hf_config.architectures = ["DeepseekV4ForCausalLM"]
|
||
mock_hf_config.to_dict.return_value = {
|
||
"dsv4_reasoning_effort_profile": "preview"
|
||
}
|
||
tm.model_path = "deepseek-ai/DeepSeek-V4-Flash"
|
||
tm.tokenizer.chat_template = "stale v3 jinja"
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
self.assertEqual(serving_chat.chat_encoding_spec, "dsv4")
|
||
|
||
tm.tokenizer.chat_template = None
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
self.assertEqual(serving_chat.chat_encoding_spec, "dsv4")
|
||
|
||
def test_kimi_k3_encoding_detection(self):
|
||
from sglang.srt.parser.template_manager import TemplateManager
|
||
|
||
tm = _MockTokenizerManager()
|
||
tm.model_config.hf_config.architectures = ["KimiK3ForConditionalGeneration"]
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
self.assertEqual(serving_chat.chat_encoding_spec, "kimi_k3")
|
||
|
||
tm.model_config.hf_config.architectures = ["LlamaForCausalLM"]
|
||
tm.server_args.tool_call_parser = "kimi_k3"
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
self.assertEqual(serving_chat.chat_encoding_spec, "kimi_k3")
|
||
|
||
def test_custom_encoders_own_reasoning_history(self):
|
||
"""Every custom encoding spec takes reasoning history natively.
|
||
|
||
The alternative splices a detector's markers into content, which an
|
||
encoder that frames its own channels turns into visible raw markers.
|
||
A new spec must not silently default to that path.
|
||
"""
|
||
for spec in _ALL_CHAT_ENCODING_SPECS:
|
||
with self.subTest(chat_encoding_spec=spec):
|
||
self.chat.chat_encoding_spec = spec
|
||
self.assertTrue(self.chat.supports_native_reasoning_history())
|
||
|
||
# The HF chat-template path keeps the wrap-into-content behaviour.
|
||
self.chat.chat_encoding_spec = None
|
||
self.assertFalse(self.chat.supports_native_reasoning_history())
|
||
|
||
def test_all_chat_encoding_specs_are_enumerated(self):
|
||
"""Guard the spec list this file asserts capabilities over."""
|
||
source = Path(chat_encoding.__file__).read_text()
|
||
returned = set(
|
||
re.findall(
|
||
r'^\s+return "(\w+)"$',
|
||
source[source.index("def resolve_chat_encoding_spec") :],
|
||
re.MULTILINE,
|
||
)
|
||
)
|
||
self.assertEqual(returned, set(_ALL_CHAT_ENCODING_SPECS))
|
||
|
||
# ------------- dsv4 task + latest_reminder -------------
|
||
def test_dsv4_task_field_schema(self):
|
||
"""Top-level `task` accepts the 6 DS task tokens and rejects others."""
|
||
for valid in ("action", "query", "authority", "domain", "title", "read_url"):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "hi"}],
|
||
task=valid,
|
||
)
|
||
self.assertEqual(req.task, valid)
|
||
|
||
# None / unset is fine
|
||
self.assertIsNone(self.basic_req.task)
|
||
|
||
# Bogus value rejected at validation time
|
||
from pydantic import ValidationError
|
||
|
||
with self.assertRaises(ValidationError):
|
||
ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "hi"}],
|
||
task="bogus",
|
||
)
|
||
|
||
def test_attach_task_to_last_user_message(self):
|
||
"""Helper attaches task to the nearest user/developer message."""
|
||
from sglang.srt.entrypoints.openai import encoding_dsv4
|
||
|
||
messages = [{"role": "user", "content": "Hi"}]
|
||
encoding_dsv4.attach_task_to_last_user_message(messages, "domain")
|
||
self.assertEqual(messages[0]["task"], "domain")
|
||
|
||
# Prefers the LAST user message across a multi-turn conversation.
|
||
messages = [
|
||
{"role": "user", "content": "first"},
|
||
{"role": "assistant", "content": "ok"},
|
||
{"role": "user", "content": "second"},
|
||
]
|
||
encoding_dsv4.attach_task_to_last_user_message(messages, "query")
|
||
self.assertNotIn("task", messages[0])
|
||
self.assertEqual(messages[2]["task"], "query")
|
||
|
||
# `developer` role is treated like `user` (matches encoder semantics).
|
||
messages = [{"role": "developer", "content": "dev"}]
|
||
encoding_dsv4.attach_task_to_last_user_message(messages, "authority")
|
||
self.assertEqual(messages[0]["task"], "authority")
|
||
|
||
# No user/developer present -> raises.
|
||
with self.assertRaises(ValueError):
|
||
encoding_dsv4.attach_task_to_last_user_message(
|
||
[{"role": "system", "content": "s"}], "domain"
|
||
)
|
||
|
||
def test_dsv4_content_parts_list_normalized(self):
|
||
"""OpenAI list-of-parts content flattens to text before reaching the encoder."""
|
||
from sglang.srt.entrypoints.openai import encoding_dsv4
|
||
from sglang.srt.parser.jinja_template_utils import (
|
||
process_content_for_template_format,
|
||
)
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{
|
||
"role": "user",
|
||
"content": [{"type": "text", "text": "say hi"}],
|
||
}
|
||
],
|
||
)
|
||
messages = [m.model_dump() for m in req.messages]
|
||
# Mirror the boundary normalization _process_messages does for any
|
||
# non-None chat_encoding_spec.
|
||
for i, msg in enumerate(messages):
|
||
if isinstance(msg.get("content"), list):
|
||
messages[i] = process_content_for_template_format(
|
||
msg, "string", [], [], [], []
|
||
)
|
||
out = encoding_dsv4.encode_messages(messages, thinking_mode="chat")
|
||
self.assertIn("<|User|>say hi", out)
|
||
|
||
# Multiple text parts concat with single space; non-text parts dropped.
|
||
messages = [
|
||
{
|
||
"role": "user",
|
||
"content": [
|
||
{"type": "text", "text": "describe"},
|
||
{"type": "image_url", "image_url": {"url": "x"}},
|
||
],
|
||
}
|
||
]
|
||
for i, msg in enumerate(messages):
|
||
if isinstance(msg.get("content"), list):
|
||
messages[i] = process_content_for_template_format(
|
||
msg, "string", [], [], [], []
|
||
)
|
||
out = encoding_dsv4.encode_messages(messages, thinking_mode="chat")
|
||
self.assertIn("<|User|>describe", out)
|
||
self.assertNotIn("image_url", out)
|
||
|
||
def test_dsv4_task_and_reminder_encode_end_to_end(self):
|
||
"""Task + latest_reminder plumb through to the dsv4 encoder correctly."""
|
||
from sglang.srt.entrypoints.openai import encoding_dsv4
|
||
|
||
# 1) task='domain' in chat mode -> `<|domain|>` appended, no Assistant
|
||
# prefix (this is a single-shot classification, not a chat turn).
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "What is SGLang?"}],
|
||
task="domain",
|
||
)
|
||
messages = [m.model_dump() for m in req.messages]
|
||
encoding_dsv4.attach_task_to_last_user_message(messages, req.task)
|
||
out = encoding_dsv4.encode_messages(messages, thinking_mode="chat")
|
||
self.assertIn("<|domain|>", out)
|
||
self.assertTrue(out.rstrip().endswith("<|domain|>"))
|
||
self.assertNotIn("<|Assistant|>", out)
|
||
|
||
# 2) task='action' in thinking mode -> Assistant + <think> + <|action|>
|
||
# (action is the one task that still runs a reasoning pass).
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi"}],
|
||
task="action",
|
||
)
|
||
messages = [m.model_dump() for m in req.messages]
|
||
encoding_dsv4.attach_task_to_last_user_message(messages, req.task)
|
||
out = encoding_dsv4.encode_messages(messages, thinking_mode="thinking")
|
||
self.assertIn("<|Assistant|>", out)
|
||
self.assertIn("<think>", out)
|
||
self.assertTrue(out.rstrip().endswith("<|action|>"))
|
||
|
||
# 3) latest_reminder preceding user -> reminder renders before user,
|
||
# Assistant prefix still comes after user.
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[
|
||
{"role": "latest_reminder", "content": "Be terse."},
|
||
{"role": "user", "content": "Hello"},
|
||
],
|
||
)
|
||
messages = [m.model_dump() for m in req.messages]
|
||
out = encoding_dsv4.encode_messages(messages, thinking_mode="chat")
|
||
self.assertIn("<|latest_reminder|>Be terse.", out)
|
||
self.assertIn("<|User|>Hello", out)
|
||
self.assertLess(
|
||
out.index("<|latest_reminder|>"),
|
||
out.index("<|User|>"),
|
||
)
|
||
self.assertIn("<|Assistant|>", out)
|
||
|
||
def test_dsv4_reasoning_effort_profiles(self):
|
||
from sglang.srt.entrypoints.openai import encoding_dsv4
|
||
|
||
messages = [
|
||
{"role": "system", "content": ""},
|
||
{"role": "user", "content": "Solve this."},
|
||
]
|
||
absolute_maximum = (
|
||
"Reasoning Effort: Absolute maximum with no shortcuts permitted.\n"
|
||
"You MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\n"
|
||
"Explicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.\n\n"
|
||
)
|
||
beyond_maximum = (
|
||
"Reasoning Effort: Beyond maximum — exhaustive, relentless, and uncompromising.\n"
|
||
"You MUST reason with the utmost depth and rigor, leaving absolutely nothing to chance: exhaustively decompose the problem into its most fundamental components, trace every causal chain to its root, and resolve the underlying cause rather than any surface symptom.\n"
|
||
"Do not stop reasoning until you have independently verified the solution from multiple angles and are certain that no assumption remains unchecked and no error remains undiscovered.\n\n"
|
||
)
|
||
|
||
def encode(profile, effort):
|
||
return encoding_dsv4.encode_messages(
|
||
messages,
|
||
thinking_mode="thinking",
|
||
reasoning_effort=effort,
|
||
reasoning_effort_profile=profile,
|
||
)
|
||
|
||
preview_high = encode("preview", "high")
|
||
preview_max = encode("preview", "max")
|
||
official_low = encode("official", "low")
|
||
official_high = encode("official", "high")
|
||
official_max = encode("official", "max")
|
||
|
||
self.assertNotIn("Reasoning Effort:", preview_high)
|
||
self.assertTrue(
|
||
preview_max.startswith(encoding_dsv4.bos_token + absolute_maximum)
|
||
)
|
||
self.assertNotIn("Reasoning Effort:", official_low)
|
||
self.assertEqual(
|
||
official_high,
|
||
encoding_dsv4.bos_token
|
||
+ absolute_maximum
|
||
+ official_low.removeprefix(encoding_dsv4.bos_token),
|
||
)
|
||
self.assertEqual(
|
||
official_max,
|
||
encoding_dsv4.bos_token
|
||
+ beyond_maximum
|
||
+ official_low.removeprefix(encoding_dsv4.bos_token),
|
||
)
|
||
self.assertEqual(encode("preview", None), preview_high)
|
||
self.assertEqual(encode("official", None), official_low)
|
||
self.assertEqual(len({official_low, official_high, official_max}), 3)
|
||
|
||
with self.assertRaises(ValueError):
|
||
encode("preview", "low")
|
||
|
||
def test_dsv4_reasoning_effort_profile_resolution(self):
|
||
resolve = resolve_dsv4_reasoning_effort_profile
|
||
preview_model_path = _create_dsv4_checkpoint(self, _DSV4_PREVIEW_ENCODER)
|
||
official_model_path = _create_dsv4_checkpoint(self, _DSV4_OFFICIAL_ENCODER)
|
||
inconclusive_model_path = _create_dsv4_checkpoint(
|
||
self, 'UNRELATED_METADATA = "value"\n'
|
||
)
|
||
self.assertEqual(resolve(model_path=preview_model_path), "preview")
|
||
self.assertEqual(resolve(model_path=official_model_path), "official")
|
||
self.assertEqual(resolve(model_path=inconclusive_model_path), "preview")
|
||
|
||
self.assertEqual(
|
||
resolve(model_path="renamed/model", override="official"), "official"
|
||
)
|
||
self.assertEqual(
|
||
resolve(model_path="renamed/model", override="preview"), "preview"
|
||
)
|
||
with self.assertRaisesRegex(ValueError, "dsv4_reasoning_effort_profile"):
|
||
resolve(model_path="renamed/model", override="auto")
|
||
|
||
def test_dsv4_reasoning_effort_profile_from_checkpoint(self):
|
||
from sglang.srt.parser.template_manager import TemplateManager
|
||
|
||
official_model_path = _create_dsv4_checkpoint(self, _DSV4_OFFICIAL_ENCODER)
|
||
tm = _MockTokenizerManager()
|
||
tm.model_config.hf_config.architectures = ["DeepseekV4ForCausalLM"]
|
||
tm.model_config.hf_config.to_dict.return_value = {}
|
||
tm.model_config.hf_config.dspark_block_size = 5
|
||
tm.model_config.hf_config.dspark_markov_rank = 256
|
||
tm.model_path = official_model_path
|
||
tm.server_args.model_path = tm.model_path
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hello"}],
|
||
reasoning_effort="max",
|
||
)
|
||
serving_chat._process_messages(request, is_multimodal=False)
|
||
prompt = tm.tokenizer.encode.call_args.args[0]
|
||
self.assertIn("Reasoning Effort: Beyond maximum", prompt)
|
||
|
||
def test_dsv4_reasoning_effort_profile_override_from_model_config(self):
|
||
from sglang.srt.parser.template_manager import TemplateManager
|
||
|
||
tm = _MockTokenizerManager()
|
||
tm.model_config.hf_config.architectures = ["DeepseekV4ForCausalLM"]
|
||
tm.model_config.hf_config.to_dict.return_value = {
|
||
"dsv4_reasoning_effort_profile": "official"
|
||
}
|
||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||
request = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hello"}],
|
||
reasoning_effort="max",
|
||
)
|
||
|
||
serving_chat._process_messages(request, is_multimodal=False)
|
||
|
||
prompt = tm.tokenizer.encode.call_args.args[0]
|
||
self.assertIn("Reasoning Effort: Beyond maximum", prompt)
|
||
|
||
def test_dsv4_invalid_profile_override_fails_at_construction(self):
|
||
from sglang.srt.parser.template_manager import TemplateManager
|
||
|
||
tm = _MockTokenizerManager()
|
||
tm.model_config.hf_config.architectures = ["DeepseekV4ForCausalLM"]
|
||
tm.model_config.hf_config.to_dict.return_value = {
|
||
"dsv4_reasoning_effort_profile": "invalid"
|
||
}
|
||
with self.assertRaisesRegex(ValueError, "dsv4_reasoning_effort_profile"):
|
||
OpenAIServingChat(tm, TemplateManager())
|
||
|
||
def test_streaming_abort_yields_error(self):
|
||
"""Test that an abort finish reason during streaming correctly yields an error and stops."""
|
||
err_msg = "Aborted by scheduler"
|
||
err_code = HTTPStatus.INTERNAL_SERVER_ERROR
|
||
|
||
async def _mock_generate_abort():
|
||
yield {
|
||
"text": "Partial ",
|
||
"meta_info": {
|
||
"id": "chatcmpl-test",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {
|
||
"type": "abort",
|
||
"status_code": err_code,
|
||
"message": err_msg,
|
||
},
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = _mock_generate_abort()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
temperature=0.7,
|
||
max_tokens=100,
|
||
stream=True,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
# Create a mock conversation object
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
|
||
async def run_stream():
|
||
chunks = []
|
||
try:
|
||
async for chunk in self.chat._generate_chat_stream(
|
||
adapted_request, req, self.fastapi_request
|
||
):
|
||
chunks.append(chunk)
|
||
except Exception as e:
|
||
print(f"Error during stream iteration: {e}")
|
||
return chunks
|
||
|
||
loop = get_or_create_event_loop()
|
||
chunks = loop.run_until_complete(run_stream())
|
||
|
||
error_chunk_data = None
|
||
for c in chunks:
|
||
if "error" in c:
|
||
error_chunk_data = json.loads(c[len("data: ") :])
|
||
break
|
||
self.assertIsNotNone(error_chunk_data, "Error chunk not found in stream")
|
||
self.assertEqual(error_chunk_data["error"]["message"], err_msg)
|
||
self.assertEqual(error_chunk_data["error"]["code"], err_code.value)
|
||
|
||
# Ensure the stream stops after the abort error
|
||
# The last chunk should be "data: [DONE]\n\n"
|
||
self.assertEqual(chunks[-1], "data: [DONE]\n\n")
|
||
|
||
# Check that there is an error chunk and a DONE chunk
|
||
self.assertEqual(len(chunks), 2)
|
||
self.assertIn("error", chunks[0])
|
||
|
||
def test_streaming_abort_with_ids_enabled(self):
|
||
"""Test that a terminal abort with input/output ids enabled yields only an error and [DONE]."""
|
||
err_msg = "Aborted by scheduler"
|
||
err_code = HTTPStatus.INTERNAL_SERVER_ERROR
|
||
|
||
async def _mock_generate_abort():
|
||
yield {
|
||
"text": "Partial ",
|
||
"prompt_token_ids": [4, 5, 6],
|
||
"output_ids": [1, 2, 3],
|
||
"meta_info": {
|
||
"id": "chatcmpl-test",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {
|
||
"type": "abort",
|
||
"status_code": err_code,
|
||
"message": err_msg,
|
||
},
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = _mock_generate_abort()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
temperature=0.7,
|
||
max_tokens=100,
|
||
stream=True,
|
||
return_input_ids_in_sglext=True,
|
||
return_output_ids_in_sglext=True,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
|
||
async def run_stream():
|
||
chunks = []
|
||
try:
|
||
async for chunk in self.chat._generate_chat_stream(
|
||
adapted_request, req, self.fastapi_request
|
||
):
|
||
chunks.append(chunk)
|
||
except Exception as e:
|
||
print(f"Error during stream iteration: {e}")
|
||
return chunks
|
||
|
||
loop = get_or_create_event_loop()
|
||
chunks = loop.run_until_complete(run_stream())
|
||
|
||
# Exactly one error chunk followed by [DONE]; no sglext ids leak.
|
||
self.assertIn("error", chunks[0])
|
||
self.assertEqual(chunks[1], "data: [DONE]\n\n")
|
||
self.assertFalse(
|
||
any("input_ids" in c or "output_ids" in c for c in chunks),
|
||
"sglext ids event leaked after abort error",
|
||
)
|
||
|
||
error_chunk_data = json.loads(chunks[0][len("data: ") :])
|
||
self.assertEqual(error_chunk_data["error"]["message"], err_msg)
|
||
self.assertEqual(error_chunk_data["error"]["code"], err_code.value)
|
||
|
||
def test_streaming_error_abort_still_finalizes_other_choices(self):
|
||
"""An error abort ends generation but still runs finalization, so a
|
||
sibling choice's finish chunk and the usage chunk follow the error."""
|
||
err_code = HTTPStatus.SERVICE_UNAVAILABLE
|
||
|
||
async def _mock_generate():
|
||
yield {
|
||
"text": "hello",
|
||
"prompt_token_ids": [4, 5, 6],
|
||
"output_ids": [1, 2],
|
||
"meta_info": {
|
||
"id": "chatcmpl-multi-abort",
|
||
"prompt_tokens": 3,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop"},
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
yield {
|
||
"text": "partial",
|
||
"prompt_token_ids": [4, 5, 6],
|
||
"output_ids": [7],
|
||
"meta_info": {
|
||
"id": "chatcmpl-multi-abort",
|
||
"prompt_tokens": 3,
|
||
"completion_tokens": 1,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {
|
||
"type": "abort",
|
||
"status_code": err_code,
|
||
"message": "Aborted by scheduler",
|
||
},
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 1,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = _mock_generate()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
n=2,
|
||
stream=True,
|
||
stream_options={"include_usage": True},
|
||
return_input_ids_in_sglext=True,
|
||
return_output_ids_in_sglext=True,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
chunks = self._run_chat_stream(adapted_request, req)
|
||
|
||
error_idx = next(i for i, c in enumerate(chunks) if "error" in c)
|
||
after_error = self._parse_chunks(chunks[error_idx + 1 :])
|
||
|
||
finish_reasons = [
|
||
choice["finish_reason"]
|
||
for c in after_error
|
||
for choice in c.get("choices", [])
|
||
if choice.get("finish_reason") is not None
|
||
]
|
||
self.assertEqual(finish_reasons, ["stop"])
|
||
self.assertTrue(
|
||
any(c.get("usage") is not None for c in after_error),
|
||
"usage chunk dropped after error abort",
|
||
)
|
||
|
||
def _run_chat_stream(self, adapted_request, req):
|
||
async def run_stream():
|
||
chunks = []
|
||
async for chunk in self.chat._generate_chat_stream(
|
||
adapted_request, req, self.fastapi_request
|
||
):
|
||
chunks.append(chunk)
|
||
return chunks
|
||
|
||
return get_or_create_event_loop().run_until_complete(run_stream())
|
||
|
||
def _parse_chunks(self, chunks):
|
||
parsed = []
|
||
for c in chunks:
|
||
if c.startswith("data: ") and c != "data: [DONE]\n\n":
|
||
parsed.append(json.loads(c[len("data: ") :]))
|
||
return parsed
|
||
|
||
async def _collect_stream_content(self, content, choice_logprobs, req):
|
||
chunks = []
|
||
async for chunk in self.chat._generate_stream_content(
|
||
content=content,
|
||
index=0,
|
||
request=req,
|
||
stream_offsets={},
|
||
reasoning_parser_dict={},
|
||
parser_dict={},
|
||
has_tool_calls={},
|
||
choice_logprobs=choice_logprobs,
|
||
finish_reason_type="stop",
|
||
continuous_usage_stats=False,
|
||
prompt_tokens={0: 5},
|
||
reasoning_tokens={0: 0},
|
||
completion_tokens={0: 1},
|
||
):
|
||
chunks.append(chunk)
|
||
return chunks
|
||
|
||
def test_streaming_logprobs_attached_with_reasoning_parser(self):
|
||
"""Logprobs must ride on the reasoning chunk when a reasoning parser is active."""
|
||
self.chat.reasoning_parser = "qwen3"
|
||
|
||
content = {
|
||
"text": "thinking...",
|
||
"meta_info": {
|
||
"id": "chatcmpl-reasoning",
|
||
"prompt_tokens": 5,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"output_token_logprobs": [(0.1, 1, "think"), (0.2, 2, "ing")],
|
||
"output_top_logprobs": [],
|
||
"output_token_logprobs_length": 2,
|
||
},
|
||
"index": 0,
|
||
}
|
||
choice_logprobs = self.chat._process_streaming_logprobs(
|
||
content, 0, 2
|
||
).model_dump()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
stream=True,
|
||
logprobs=True,
|
||
separate_reasoning=True,
|
||
)
|
||
|
||
with patch.object(self.chat, "_process_reasoning_stream") as proc_mock:
|
||
proc_mock.return_value = ("Let me think", "")
|
||
chunks = get_or_create_event_loop().run_until_complete(
|
||
self._collect_stream_content(content, choice_logprobs, req)
|
||
)
|
||
|
||
parsed = self._parse_chunks(chunks)
|
||
reasoning_chunks = [
|
||
c
|
||
for c in parsed
|
||
if c["choices"][0]["delta"].get("reasoning_content") is not None
|
||
]
|
||
self.assertTrue(
|
||
reasoning_chunks, "reasoning_parser did not emit a reasoning_content chunk"
|
||
)
|
||
logprob_chunks = [
|
||
c for c in parsed if c["choices"][0].get("logprobs") is not None
|
||
]
|
||
self.assertTrue(
|
||
logprob_chunks,
|
||
"logprobs dropped: no chunk carried logprobs with reasoning_parser active",
|
||
)
|
||
|
||
def test_streaming_logprobs_flushed_when_tool_parser_buffers_delta(self):
|
||
"""Logprobs must be flushed on a standalone chunk when the tool parser emits no content delta."""
|
||
self.chat.tool_call_parser = "hermes"
|
||
|
||
content = {
|
||
"text": "(<",
|
||
"meta_info": {
|
||
"id": "chatcmpl-tool",
|
||
"prompt_tokens": 5,
|
||
"completion_tokens": 1,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"output_token_logprobs": [(0.3, 9, "(<")],
|
||
"output_top_logprobs": [],
|
||
"output_token_logprobs_length": 1,
|
||
},
|
||
"index": 0,
|
||
}
|
||
choice_logprobs = self.chat._process_streaming_logprobs(
|
||
content, 0, 1
|
||
).model_dump()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
tools=[{"type": "function", "function": {"name": "get_weather"}}],
|
||
stream=True,
|
||
logprobs=True,
|
||
)
|
||
|
||
async def _empty_tool_stream(*args, **kwargs):
|
||
return
|
||
yield # make it an async generator
|
||
|
||
with patch.object(self.chat, "_process_tool_call_stream", _empty_tool_stream):
|
||
chunks = get_or_create_event_loop().run_until_complete(
|
||
self._collect_stream_content(content, choice_logprobs, req)
|
||
)
|
||
|
||
parsed = self._parse_chunks(chunks)
|
||
logprob_chunks = [
|
||
c for c in parsed if c["choices"][0].get("logprobs") is not None
|
||
]
|
||
self.assertTrue(
|
||
logprob_chunks,
|
||
"logprobs dropped: no flush chunk carried logprobs when tool parser buffered the delta",
|
||
)
|
||
|
||
def test_streaming_logprobs_not_flushed_on_empty_delta_step_without_parser(self):
|
||
"""With no parser active, an empty-delta step must not emit a standalone
|
||
empty-delta logprobs chunk — clients expect each chunk to carry real
|
||
content/reasoning/tool_calls or a finish_reason."""
|
||
self.chat.reasoning_parser = None
|
||
self.chat.tool_call_parser = None
|
||
|
||
async def _mock_generate():
|
||
yield {
|
||
"text": "",
|
||
"meta_info": {
|
||
"id": "chatcmpl-empty",
|
||
"prompt_tokens": 5,
|
||
"completion_tokens": 1,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"output_token_logprobs": [(0.5, 7, "")],
|
||
"output_top_logprobs": [],
|
||
"output_token_logprobs_length": 1,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = _mock_generate()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
stream=True,
|
||
logprobs=True,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
chunks = self._run_chat_stream(adapted_request, req)
|
||
|
||
parsed = self._parse_chunks(chunks)
|
||
empty_logprob_chunks = [
|
||
c
|
||
for c in parsed
|
||
if c["choices"][0].get("logprobs") is not None
|
||
and not c["choices"][0]["delta"].get("content")
|
||
and not c["choices"][0]["delta"].get("reasoning_content")
|
||
and not c["choices"][0]["delta"].get("tool_calls")
|
||
and not c["choices"][0].get("finish_reason")
|
||
]
|
||
self.assertFalse(
|
||
empty_logprob_chunks,
|
||
"empty-delta logprobs chunk emitted without a parser; would break client chunk-shape assumptions",
|
||
)
|
||
|
||
def test_non_streaming_extension_fields_emit_sglext_without_meta_info(self):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
return_meta_info=False,
|
||
return_routed_experts=True,
|
||
return_cached_tokens_details=True,
|
||
)
|
||
ret = [
|
||
{
|
||
"text": "Cached response",
|
||
"meta_info": {
|
||
"id": "chatcmpl-cache-test",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": 6,
|
||
"cached_tokens_details": {
|
||
"device": 4,
|
||
"host": 1,
|
||
"storage": 1,
|
||
"storage_backend": "file",
|
||
},
|
||
"routed_experts": "cm91dGUtYQ==",
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"weight_version": "default",
|
||
},
|
||
}
|
||
]
|
||
|
||
response = self.chat._build_chat_response(req, ret, 1234567890)
|
||
|
||
self.assertIsNotNone(response.sglext)
|
||
self.assertEqual(response.sglext.routed_experts, "cm91dGUtYQ==")
|
||
self.assertEqual(
|
||
response.sglext.cached_tokens_details.model_dump(exclude_none=True),
|
||
{
|
||
"device": 4,
|
||
"host": 1,
|
||
"storage": 1,
|
||
"storage_backend": "file",
|
||
},
|
||
)
|
||
self.assertIsNone(response.choices[0].meta_info)
|
||
dumped_response = response.model_dump()
|
||
self.assertIn("sglext", dumped_response)
|
||
self.assertNotIn("meta_info", dumped_response["choices"][0])
|
||
|
||
def test_non_streaming_meta_info_omits_response_level_routed_experts(self):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
n=2,
|
||
return_meta_info=True,
|
||
return_routed_experts=True,
|
||
)
|
||
ret = [
|
||
{
|
||
"text": f"Response {index}",
|
||
"meta_info": {
|
||
"id": "chatcmpl-meta-test",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": index,
|
||
"routed_experts": routed_experts,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"weight_version": "default",
|
||
},
|
||
}
|
||
for index, routed_experts in enumerate(["cm91dGUtYQ==", "cm91dGUtYg=="])
|
||
]
|
||
|
||
response = self.chat._build_chat_response(req, ret, 1234567890)
|
||
|
||
self.assertIsNone(
|
||
response.sglext,
|
||
"sglext is absent only when routed_experts is the sole extension",
|
||
)
|
||
self.assertEqual(
|
||
[choice.meta_info for choice in response.choices],
|
||
[ret_item["meta_info"] for ret_item in ret],
|
||
)
|
||
dumped_response = response.model_dump()
|
||
self.assertNotIn("sglext", dumped_response)
|
||
self.assertEqual(
|
||
[choice["meta_info"] for choice in dumped_response["choices"]],
|
||
[ret_item["meta_info"] for ret_item in ret],
|
||
)
|
||
serialized_response = json.dumps(dumped_response)
|
||
self.assertEqual(serialized_response.count('"routed_experts"'), 2)
|
||
|
||
def test_non_streaming_meta_info_preserves_cache_and_spec_in_sglext(self):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
n=2,
|
||
return_meta_info=True,
|
||
return_routed_experts=True,
|
||
return_cached_tokens_details=True,
|
||
return_spec_tokens_details=True,
|
||
)
|
||
routed_experts = ["cm91dGUtYQ==", "cm91dGUtYg=="]
|
||
ret = [_spec_result(index) for index in range(2)]
|
||
for index, ret_item in enumerate(ret):
|
||
ret_item["meta_info"].update(
|
||
{
|
||
"cached_tokens_details": {
|
||
"device": 4 - index,
|
||
"host": index,
|
||
},
|
||
"routed_experts": routed_experts[index],
|
||
}
|
||
)
|
||
|
||
response = self.chat._build_chat_response(req, ret, 1234567890)
|
||
|
||
self.assertIsNotNone(response.sglext)
|
||
self.assertIsNone(response.sglext.routed_experts)
|
||
self.assertEqual(
|
||
response.sglext.cached_tokens_details.model_dump(exclude_none=True),
|
||
{"device": 4, "host": 0},
|
||
)
|
||
self.assertEqual(
|
||
[item.spec_cap_length for item in response.sglext.spec_tokens_details],
|
||
[1.0, 2.0],
|
||
)
|
||
self.assertEqual(
|
||
[choice.meta_info for choice in response.choices],
|
||
[ret_item["meta_info"] for ret_item in ret],
|
||
)
|
||
dumped_response = response.model_dump()
|
||
self.assertIn("sglext", dumped_response)
|
||
self.assertIn("spec_tokens_details", dumped_response["sglext"])
|
||
self.assertNotIn("routed_experts", dumped_response["sglext"])
|
||
self.assertEqual(
|
||
dumped_response["sglext"]["cached_tokens_details"],
|
||
{"device": 4, "host": 0},
|
||
)
|
||
serialized_response = json.dumps(dumped_response)
|
||
self.assertEqual(serialized_response.count('"routed_experts"'), 2)
|
||
|
||
def test_parallel_sampling_returns_spec_details_per_choice(self):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
n=2,
|
||
return_spec_tokens_details=True,
|
||
)
|
||
ret = [_spec_result(index) for index in range(2)]
|
||
|
||
response = self.chat._build_chat_response(req, ret, 1234567890)
|
||
|
||
details = response.sglext.spec_tokens_details
|
||
self.assertEqual([item.spec_cap_length for item in details], [1.0, 2.0])
|
||
self.assertEqual(
|
||
[item.spec_cap_lens_histogram for item in details],
|
||
[[0, 1], [1, 1]],
|
||
)
|
||
|
||
single_req = req.model_copy(update={"n": 1})
|
||
single_response = self.chat._build_chat_response(
|
||
single_req, ret[:1], 1234567890
|
||
)
|
||
self.assertEqual(
|
||
single_response.sglext.spec_tokens_details.spec_cap_length,
|
||
1.0,
|
||
)
|
||
|
||
def test_non_streaming_chat_response_returns_requested_token_ids_and_meta_info(
|
||
self,
|
||
):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
return_prompt_token_ids=True,
|
||
return_token_ids=True,
|
||
return_meta_info=True,
|
||
)
|
||
ret = [
|
||
{
|
||
"text": "Answer",
|
||
"output_ids": [21, 22],
|
||
"prompt_token_ids": [11, 12, 13],
|
||
"meta_info": {
|
||
"id": "chatcmpl-token-ids",
|
||
"prompt_tokens": 3,
|
||
"completion_tokens": 1,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"weight_version": "default",
|
||
},
|
||
}
|
||
]
|
||
|
||
response = self.chat._build_chat_response(req, ret, created=123)
|
||
choice = response.choices[0]
|
||
|
||
self.assertEqual(choice.prompt_token_ids, [11, 12, 13])
|
||
self.assertEqual(choice.response_token_ids, [21, 22])
|
||
self.assertEqual(choice.meta_info, ret[0]["meta_info"])
|
||
dumped_choice = response.model_dump()["choices"][0]
|
||
self.assertEqual(dumped_choice["prompt_token_ids"], [11, 12, 13])
|
||
self.assertEqual(dumped_choice["response_token_ids"], [21, 22])
|
||
self.assertEqual(dumped_choice["meta_info"], ret[0]["meta_info"])
|
||
|
||
def test_streaming_cached_tokens_details_emits_sglext(self):
|
||
"""Test that streaming chat responses emit cached token details in sglext."""
|
||
|
||
async def _mock_generate_with_cached_tokens_details():
|
||
yield {
|
||
"text": "Cached response",
|
||
"meta_info": {
|
||
"id": "chatcmpl-cache-test",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": 6,
|
||
"cached_tokens_details": {
|
||
"device": 4,
|
||
"host": 1,
|
||
"storage": 1,
|
||
"storage_backend": "file",
|
||
},
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = (
|
||
_mock_generate_with_cached_tokens_details()
|
||
)
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
stream=True,
|
||
return_cached_tokens_details=True,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
|
||
async def run_stream():
|
||
chunks = []
|
||
async for chunk in self.chat._generate_chat_stream(
|
||
adapted_request, req, self.fastapi_request
|
||
):
|
||
chunks.append(chunk)
|
||
return chunks
|
||
|
||
loop = get_or_create_event_loop()
|
||
chunks = loop.run_until_complete(run_stream())
|
||
|
||
sglext_chunks = []
|
||
for chunk in chunks:
|
||
if not chunk.startswith("data: ") or chunk.strip() == "data: [DONE]":
|
||
continue
|
||
data = json.loads(chunk[len("data: ") :])
|
||
if "sglext" in data:
|
||
sglext_chunks.append(data)
|
||
|
||
self.assertEqual(len(sglext_chunks), 1)
|
||
self.assertEqual(sglext_chunks[0]["choices"], [])
|
||
self.assertEqual(
|
||
sglext_chunks[0]["sglext"]["cached_tokens_details"],
|
||
{
|
||
"device": 4,
|
||
"host": 1,
|
||
"storage": 1,
|
||
"storage_backend": "file",
|
||
},
|
||
)
|
||
|
||
def _output_ids_ret(self, *output_ids_by_choice):
|
||
return [
|
||
{
|
||
"text": "Answer",
|
||
"prompt_token_ids": [1, 2, 3],
|
||
"output_ids": list(output_ids),
|
||
"meta_info": {
|
||
"id": "chatcmpl-output-ids",
|
||
"prompt_tokens": 3,
|
||
"completion_tokens": len(output_ids),
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"weight_version": "default",
|
||
},
|
||
}
|
||
for output_ids in output_ids_by_choice
|
||
]
|
||
|
||
def test_non_streaming_ids_emit_sglext(self):
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
n=2,
|
||
return_input_ids_in_sglext=True,
|
||
return_output_ids_in_sglext=True,
|
||
)
|
||
|
||
response = self.chat._build_chat_response(
|
||
req, self._output_ids_ret([5, 6, 7], [8, 9]), 123
|
||
)
|
||
|
||
self.assertIsNotNone(response.sglext)
|
||
self.assertEqual(response.sglext.input_ids, [1, 2, 3])
|
||
self.assertEqual(response.sglext.output_ids, [[5, 6, 7], [8, 9]])
|
||
dumped = json.loads(response.model_dump_json())
|
||
self.assertEqual(dumped["sglext"]["input_ids"], [1, 2, 3])
|
||
self.assertEqual(dumped["sglext"]["output_ids"], [[5, 6, 7], [8, 9]])
|
||
# None sglext fields stay stripped from the serialized response.
|
||
self.assertNotIn("routed_experts", dumped["sglext"])
|
||
|
||
def test_non_streaming_ids_not_returned_by_default(self):
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
|
||
response = self.chat._build_chat_response(
|
||
req, self._output_ids_ret([5, 6, 7]), 123
|
||
)
|
||
|
||
self.assertIsNone(response.sglext)
|
||
|
||
def test_non_streaming_ids_server_default_enables_flag(self):
|
||
enter_override(
|
||
self,
|
||
get_context().override_server_args(
|
||
return_input_ids=True, return_output_ids=True
|
||
),
|
||
)
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
|
||
response = self.chat._build_chat_response(
|
||
req, self._output_ids_ret([5, 6, 7]), 123
|
||
)
|
||
|
||
self.assertIsNotNone(response.sglext)
|
||
self.assertEqual(response.sglext.input_ids, [1, 2, 3])
|
||
self.assertEqual(response.sglext.output_ids, [[5, 6, 7]])
|
||
|
||
def test_ids_headers_enable_flags(self):
|
||
self.fastapi_request.headers = {
|
||
"x-sglext-return-input-ids": "1",
|
||
"x-sglext-return-output-ids": "1",
|
||
}
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
)
|
||
processed_messages = MessageProcessingResult(
|
||
"Test prompt",
|
||
[1, 2, 3],
|
||
None,
|
||
None,
|
||
[],
|
||
["</s>"],
|
||
None,
|
||
)
|
||
|
||
with patch.object(
|
||
self.chat, "_process_messages", return_value=processed_messages
|
||
):
|
||
_, processed_request = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
|
||
self.assertTrue(processed_request.return_input_ids_in_sglext)
|
||
self.assertTrue(processed_request.return_output_ids_in_sglext)
|
||
|
||
def _run_output_ids_stream(
|
||
self,
|
||
chunk_output_ids,
|
||
incremental,
|
||
framed=False,
|
||
return_raw=False,
|
||
cached_tokens_details=None,
|
||
):
|
||
"""Stream chunks with incremental or cumulative output_ids;
|
||
return parsed sglext chunks, or raw SSE strings when return_raw."""
|
||
enter_override(
|
||
self,
|
||
get_context().override_server_args(
|
||
incremental_streaming_output=incremental
|
||
),
|
||
)
|
||
if framed:
|
||
self.fastapi_request.headers["x-sglext-ids-framed"] = "1"
|
||
|
||
async def _mock_generate():
|
||
for i, ids in enumerate(chunk_output_ids):
|
||
finished = i == len(chunk_output_ids) - 1
|
||
yield {
|
||
"text": "chunk",
|
||
"prompt_token_ids": [1, 2, 3],
|
||
"output_ids": list(ids),
|
||
"meta_info": {
|
||
"id": "chatcmpl-output-ids-stream",
|
||
"prompt_tokens": 3,
|
||
"completion_tokens": 1 + i,
|
||
"cached_tokens": 0,
|
||
"cached_tokens_details": cached_tokens_details,
|
||
"finish_reason": (
|
||
{"type": "stop", "matched": None} if finished else None
|
||
),
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = _mock_generate()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
stream=True,
|
||
return_input_ids_in_sglext=True,
|
||
return_output_ids_in_sglext=True,
|
||
return_cached_tokens_details=cached_tokens_details is not None,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
chunks = self._run_chat_stream(adapted_request, req)
|
||
|
||
if return_raw:
|
||
return chunks
|
||
return [c for c in self._parse_chunks(chunks) if "sglext" in c]
|
||
|
||
def test_streaming_output_ids_incremental_accumulates_deltas(self):
|
||
sglext_chunks = self._run_output_ids_stream([[5, 6], [7]], incremental=True)
|
||
|
||
self.assertEqual(len(sglext_chunks), 1)
|
||
self.assertEqual(sglext_chunks[0]["choices"], [])
|
||
# input_ids is captured once as a flat shared prompt, not accumulated.
|
||
self.assertEqual(sglext_chunks[0]["sglext"]["input_ids"], [1, 2, 3])
|
||
self.assertEqual(sglext_chunks[0]["sglext"]["output_ids"], [[5, 6, 7]])
|
||
|
||
def test_streaming_output_ids_non_incremental_keeps_latest_full_list(self):
|
||
sglext_chunks = self._run_output_ids_stream(
|
||
[[5, 6], [5, 6, 7]], incremental=False
|
||
)
|
||
|
||
self.assertEqual(len(sglext_chunks), 1)
|
||
self.assertEqual(sglext_chunks[0]["sglext"]["input_ids"], [1, 2, 3])
|
||
self.assertEqual(sglext_chunks[0]["sglext"]["output_ids"], [[5, 6, 7]])
|
||
|
||
_SGLEXT_IDS_EVENT_PREFIX = "event: sglext_ids\ndata: "
|
||
|
||
def test_streaming_framed_ids_emits_named_event(self):
|
||
raw = self._run_output_ids_stream(
|
||
[[5, 6], [5, 6, 7]], incremental=False, framed=True, return_raw=True
|
||
)
|
||
|
||
named = [c for c in raw if c.startswith(self._SGLEXT_IDS_EVENT_PREFIX)]
|
||
self.assertEqual(len(named), 1)
|
||
payload = json.loads(named[0][len(self._SGLEXT_IDS_EVENT_PREFIX) :])
|
||
self.assertEqual(payload["choices"], [])
|
||
# The named event carries ONLY the id fields.
|
||
self.assertEqual(set(payload["sglext"]), {"input_ids", "output_ids"})
|
||
self.assertEqual(payload["sglext"]["input_ids"], [1, 2, 3])
|
||
self.assertEqual(payload["sglext"]["output_ids"], [[5, 6, 7]])
|
||
# All requested sglext fields were ids, so no plain sglext chunk.
|
||
self.assertEqual([c for c in self._parse_chunks(raw) if "sglext" in c], [])
|
||
|
||
def test_streaming_framed_splits_non_id_fields_into_plain_chunk(self):
|
||
raw = self._run_output_ids_stream(
|
||
[[5, 6], [5, 6, 7]],
|
||
incremental=False,
|
||
framed=True,
|
||
return_raw=True,
|
||
cached_tokens_details={"device": 2, "host": 1},
|
||
)
|
||
|
||
named = [c for c in raw if c.startswith(self._SGLEXT_IDS_EVENT_PREFIX)]
|
||
self.assertEqual(len(named), 1)
|
||
named_payload = json.loads(named[0][len(self._SGLEXT_IDS_EVENT_PREFIX) :])
|
||
self.assertEqual(set(named_payload["sglext"]), {"input_ids", "output_ids"})
|
||
|
||
plain_sglext = [c for c in self._parse_chunks(raw) if "sglext" in c]
|
||
self.assertEqual(len(plain_sglext), 1)
|
||
self.assertEqual(
|
||
plain_sglext[0]["sglext"]["cached_tokens_details"],
|
||
{"device": 2, "host": 1},
|
||
)
|
||
self.assertNotIn("input_ids", plain_sglext[0]["sglext"])
|
||
self.assertNotIn("output_ids", plain_sglext[0]["sglext"])
|
||
# The plain non-id chunk precedes the named ids event.
|
||
plain_raw_idx = next(
|
||
i
|
||
for i, c in enumerate(raw)
|
||
if c.startswith("data: ") and "cached_tokens_details" in c
|
||
)
|
||
self.assertLess(plain_raw_idx, raw.index(named[0]))
|
||
|
||
def _run_output_ids_stream_with_graceful_abort(
|
||
self, normal_chunks, abort_output_ids, incremental, abort_completion_tokens
|
||
):
|
||
"""Stream normal chunks then a graceful-abort chunk; return parsed sglext chunks."""
|
||
enter_override(
|
||
self,
|
||
get_context().override_server_args(
|
||
incremental_streaming_output=incremental
|
||
),
|
||
)
|
||
|
||
async def _mock_generate():
|
||
generated = 0
|
||
for ids in normal_chunks:
|
||
generated = generated + len(ids) if incremental else len(ids)
|
||
yield {
|
||
"text": "chunk",
|
||
"output_ids": list(ids),
|
||
"meta_info": {
|
||
"id": "chatcmpl-abort-ids-stream",
|
||
"prompt_tokens": 3,
|
||
"completion_tokens": generated,
|
||
"cached_tokens": 0,
|
||
"finish_reason": None,
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
# Graceful abort terminal chunk (no status_code): falls through to
|
||
# the normal finalization path.
|
||
yield {
|
||
"text": "chunk",
|
||
"output_ids": list(abort_output_ids),
|
||
"meta_info": {
|
||
"id": "chatcmpl-abort-ids-stream",
|
||
"prompt_tokens": 3,
|
||
"completion_tokens": abort_completion_tokens,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "abort", "message": "Aborted."},
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = _mock_generate()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
stream=True,
|
||
return_output_ids_in_sglext=True,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
chunks = self._run_chat_stream(adapted_request, req)
|
||
|
||
return [c for c in self._parse_chunks(chunks) if "sglext" in c]
|
||
|
||
def test_streaming_output_ids_incremental_ignores_abort_chunk(self):
|
||
# Abort re-sends the last token; it must not be appended again.
|
||
sglext_chunks = self._run_output_ids_stream_with_graceful_abort(
|
||
normal_chunks=[[5, 6], [7]],
|
||
abort_output_ids=[7],
|
||
incremental=True,
|
||
abort_completion_tokens=3,
|
||
)
|
||
|
||
self.assertEqual(len(sglext_chunks), 1)
|
||
self.assertEqual(sglext_chunks[0]["sglext"]["output_ids"], [[5, 6, 7]])
|
||
|
||
def test_streaming_output_ids_incremental_keeps_coalesced_abort_deltas(self):
|
||
# Coalesced abort chunk: keep real deltas, drop the repeated last token.
|
||
sglext_chunks = self._run_output_ids_stream_with_graceful_abort(
|
||
normal_chunks=[[5, 6]],
|
||
abort_output_ids=[7, 8, 8],
|
||
incremental=True,
|
||
abort_completion_tokens=4,
|
||
)
|
||
|
||
self.assertEqual(len(sglext_chunks), 1)
|
||
self.assertEqual(sglext_chunks[0]["sglext"]["output_ids"], [[5, 6, 7, 8]])
|
||
|
||
def test_streaming_output_ids_non_incremental_abort_supersedes_earlier(self):
|
||
sglext_chunks = self._run_output_ids_stream_with_graceful_abort(
|
||
normal_chunks=[[5, 6]],
|
||
abort_output_ids=[5, 6, 7],
|
||
incremental=False,
|
||
abort_completion_tokens=3,
|
||
)
|
||
|
||
self.assertEqual(len(sglext_chunks), 1)
|
||
self.assertEqual(sglext_chunks[0]["sglext"]["output_ids"], [[5, 6, 7]])
|
||
|
||
def test_streaming_parallel_sampling_orders_spec_details_by_choice(self):
|
||
async def mock_generate():
|
||
for index in (1, 0):
|
||
yield _spec_result(index)
|
||
|
||
self.tm.generate_request.return_value = mock_generate()
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
max_tokens=100,
|
||
n=2,
|
||
stream=True,
|
||
return_spec_tokens_details=True,
|
||
)
|
||
|
||
parsed = self._parse_chunks(self._run_chat_stream(Mock(), req))
|
||
details = next(chunk["sglext"] for chunk in parsed if "sglext" in chunk)[
|
||
"spec_tokens_details"
|
||
]
|
||
self.assertEqual([item["spec_cap_length"] for item in details], [1.0, 2.0])
|
||
self.assertEqual(
|
||
[item["spec_cap_lens_histogram"] for item in details],
|
||
[[0, 1], [1, 1]],
|
||
)
|
||
|
||
def _collect_continuous_usage(self, cached_tokens):
|
||
content = {
|
||
"text": "Hello",
|
||
"meta_info": {
|
||
"id": "chatcmpl-cont-usage",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 2,
|
||
"cached_tokens": cached_tokens,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
stream=True,
|
||
)
|
||
|
||
async def _collect():
|
||
chunks = []
|
||
async for chunk in self.chat._generate_stream_content(
|
||
content=content,
|
||
index=0,
|
||
request=req,
|
||
stream_offsets={},
|
||
reasoning_parser_dict={},
|
||
parser_dict={},
|
||
has_tool_calls={},
|
||
choice_logprobs=None,
|
||
finish_reason_type="stop",
|
||
continuous_usage_stats=True,
|
||
prompt_tokens={0: 10},
|
||
reasoning_tokens={0: 0},
|
||
completion_tokens={0: 2},
|
||
):
|
||
chunks.append(chunk)
|
||
return chunks
|
||
|
||
chunks = get_or_create_event_loop().run_until_complete(_collect())
|
||
return [c["usage"] for c in self._parse_chunks(chunks) if c.get("usage")]
|
||
|
||
def test_continuous_usage_reports_cached_tokens(self):
|
||
"""continuous_usage_stats chunks include cached tokens when cache reporting is on."""
|
||
enter_override(
|
||
self, get_context().override_server_args(enable_cache_report=True)
|
||
)
|
||
usages = self._collect_continuous_usage(cached_tokens=6)
|
||
self.assertTrue(usages, "continuous_usage_stats attached no usage")
|
||
self.assertEqual(usages[0]["prompt_tokens_details"]["cached_tokens"], 6)
|
||
|
||
def test_continuous_usage_omits_cached_tokens_when_report_disabled(self):
|
||
"""With cache reporting off, continuous_usage_stats must not leak cached tokens."""
|
||
enter_override(
|
||
self, get_context().override_server_args(enable_cache_report=False)
|
||
)
|
||
usages = self._collect_continuous_usage(cached_tokens=6)
|
||
self.assertTrue(usages, "continuous_usage_stats attached no usage")
|
||
self.assertIsNone(usages[0].get("prompt_tokens_details"))
|
||
|
||
# ------------- incremental streaming output tests -------------
|
||
def test_incremental_streaming_output_delta(self):
|
||
"""Test that streaming with incremental_streaming_output produces correct deltas.
|
||
|
||
When incremental_streaming_output is enabled, content["text"] is already the
|
||
incremental delta (not the full accumulated text). The delta computation must
|
||
use content["text"] directly instead of slicing by the accumulated buffer length.
|
||
|
||
Regression test for https://github.com/sgl-project/sglang/issues/22510.
|
||
"""
|
||
# Enable incremental_streaming_output on the mock
|
||
enter_override(
|
||
self, get_context().override_server_args(incremental_streaming_output=True)
|
||
)
|
||
|
||
# Simulate incremental streaming: each yield has ONLY the new text (delta),
|
||
# NOT the full accumulated text.
|
||
incremental_chunks = [
|
||
("I am", None),
|
||
(" a large", None),
|
||
(" language model", None),
|
||
(".", {"type": "stop", "matched": None}),
|
||
]
|
||
|
||
async def _mock_generate_incremental():
|
||
for text, finish_reason in incremental_chunks:
|
||
yield {
|
||
"text": text,
|
||
"meta_info": {
|
||
"id": "chatcmpl-incr-test",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 5,
|
||
"cached_tokens": 0,
|
||
"finish_reason": finish_reason,
|
||
"output_token_logprobs": None,
|
||
"output_top_logprobs": None,
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
self.tm.generate_request.return_value = _mock_generate_incremental()
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
temperature=0.7,
|
||
max_tokens=100,
|
||
stream=True,
|
||
)
|
||
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "Test prompt"
|
||
conv_mock.return_value = conv_ins
|
||
|
||
adapted_request, _ = self.chat._convert_to_internal_request(
|
||
req, self.fastapi_request
|
||
)
|
||
|
||
async def run_stream():
|
||
chunks = []
|
||
async for chunk in self.chat._generate_chat_stream(
|
||
adapted_request, req, self.fastapi_request
|
||
):
|
||
chunks.append(chunk)
|
||
return chunks
|
||
|
||
loop = get_or_create_event_loop()
|
||
chunks = loop.run_until_complete(run_stream())
|
||
|
||
# Extract content deltas from SSE chunks
|
||
deltas = []
|
||
for c in chunks:
|
||
if not c.startswith("data: ") or c.strip() == "data: [DONE]":
|
||
continue
|
||
data = json.loads(c[len("data: ") :])
|
||
if "choices" in data and data["choices"]:
|
||
content = data["choices"][0]["delta"].get("content")
|
||
if content:
|
||
deltas.append(content)
|
||
|
||
joined = "".join(deltas)
|
||
self.assertEqual(
|
||
joined,
|
||
"I am a large language model.",
|
||
f"Streaming deltas produced broken text: {deltas!r}",
|
||
)
|
||
|
||
# ------------- X-Data-Parallel-Rank header tests -------------
|
||
def test_extract_routed_dp_rank_from_header_no_header(self):
|
||
"""Test that None is returned when no header is present."""
|
||
self.fastapi_request.headers = {}
|
||
result = self.chat.extract_routed_dp_rank_from_header(
|
||
self.fastapi_request, body_routed_dp_rank=None
|
||
)
|
||
self.assertIsNone(result)
|
||
|
||
def test_extract_routed_dp_rank_header_overrides_body(self):
|
||
"""Test that header value has higher priority than body."""
|
||
self.fastapi_request.headers = {"x-data-parallel-rank": "3"}
|
||
result = self.chat.extract_routed_dp_rank_from_header(
|
||
self.fastapi_request, body_routed_dp_rank=1
|
||
)
|
||
self.assertEqual(result, 3) # header wins
|
||
|
||
def test_extract_routed_dp_rank_from_header_invalid(self):
|
||
"""Test that invalid header value raises HTTPException."""
|
||
from fastapi import HTTPException
|
||
|
||
self.fastapi_request.headers = {"x-data-parallel-rank": "abc"}
|
||
with self.assertRaises(HTTPException) as context:
|
||
self.chat.extract_routed_dp_rank_from_header(
|
||
self.fastapi_request, body_routed_dp_rank=None
|
||
)
|
||
self.assertEqual(context.exception.status_code, 400)
|
||
self.assertIn("must be an integer", context.exception.detail)
|
||
|
||
def test_hunyuan_reasoning_effort_dispatch(self):
|
||
tm = _MockTokenizerManager()
|
||
tm.server_args.reasoning_parser = "hunyuan"
|
||
chat = OpenAIServingChat(tm, _MockTemplateManager())
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "hi"}]
|
||
)
|
||
cases = [
|
||
("no_think", False),
|
||
("none", False),
|
||
(None, False),
|
||
("high", True),
|
||
("low", True),
|
||
]
|
||
for effort, expected in cases:
|
||
with self.subTest(effort=effort):
|
||
req.reasoning_effort = effort
|
||
self.assertEqual(chat._get_reasoning_from_request(req), expected)
|
||
|
||
def _setup_nemotron_super(self):
|
||
"""Drive _apply_jinja_template (chat_template_name=None) with a
|
||
Nemotron-3 Super reasoning_config carrying effort_kwarg."""
|
||
self.tm.server_args.reasoning_parser = "nemotron_3"
|
||
self.chat.reasoning_parser = "nemotron_3"
|
||
self.tm.server_args.served_model_name = (
|
||
"nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16"
|
||
)
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="enable_thinking",
|
||
default_enabled=True,
|
||
effort_kwarg="low_effort",
|
||
)
|
||
self.chat.chat_encoding_spec = None
|
||
|
||
def _run_jinja_with_effort(self, effort):
|
||
req = ChatCompletionRequest(
|
||
model="nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16",
|
||
messages=[{"role": "user", "content": "hi"}],
|
||
reasoning_effort=effort,
|
||
)
|
||
with (
|
||
patch.object(self.chat, "_encode_messages", return_value=None),
|
||
patch.object(
|
||
self.chat,
|
||
"_handle_last_assistant_message",
|
||
return_value=([{"role": "user", "content": "hi"}], None),
|
||
),
|
||
):
|
||
self.chat._process_messages(req, False)
|
||
_, kwargs = self.tm.tokenizer.apply_chat_template.call_args
|
||
return kwargs
|
||
|
||
def test_nemotron_super_low_effort_mapped_to_kwarg(self):
|
||
self._setup_nemotron_super()
|
||
kwargs = self._run_jinja_with_effort("low")
|
||
self.assertTrue(kwargs.get("low_effort"))
|
||
|
||
def test_nemotron_super_high_effort_warns_without_kwarg(self):
|
||
self._setup_nemotron_super()
|
||
with self.assertLogs(
|
||
"sglang.srt.entrypoints.openai.serving_chat", level="WARNING"
|
||
) as logs:
|
||
kwargs = self._run_jinja_with_effort("high")
|
||
self.assertNotIn("low_effort", kwargs)
|
||
self.assertTrue(any("only 'low' reasoning effort" in m for m in logs.output))
|
||
|
||
def test_nemotron_nano_no_effort_kwarg(self):
|
||
# Nano template has no low_effort, so effort_kwarg stays None and no
|
||
# warning is emitted even for non-low effort.
|
||
self.tm.server_args.reasoning_parser = "nemotron_3"
|
||
self.chat.reasoning_parser = "nemotron_3"
|
||
self.template_manager.chat_template_name = None
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="enable_thinking", default_enabled=True
|
||
)
|
||
self.chat.chat_encoding_spec = None
|
||
kwargs = self._run_jinja_with_effort("high")
|
||
self.assertNotIn("low_effort", kwargs)
|
||
|
||
def test_non_stream_reasoning_response_preserves_payload_whitespace(self):
|
||
self.chat.reasoning_parser = "qwen3"
|
||
self.template_manager.force_reasoning = False
|
||
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
stream=False,
|
||
separate_reasoning=True,
|
||
)
|
||
ret = [
|
||
{
|
||
"text": "<think>\nLet me think\n</think>\n\nThe answer is 42.\n",
|
||
"meta_info": {
|
||
"id": "chatcmpl-test",
|
||
"prompt_tokens": 5,
|
||
"completion_tokens": 8,
|
||
"cached_tokens": 0,
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
"weight_version": "test",
|
||
},
|
||
"index": 0,
|
||
}
|
||
]
|
||
|
||
response = self.chat._build_chat_response(req, ret, created=123)
|
||
|
||
message = response.choices[0].message
|
||
self.assertEqual(message.reasoning_content, "\nLet me think\n")
|
||
self.assertEqual(message.content, "\n\nThe answer is 42.\n")
|
||
|
||
# ------------- reasoning config tests -------------
|
||
def test_get_reasoning_from_request_default_true_toggle(self):
|
||
self.tm.server_args.reasoning_parser = "qwen3"
|
||
self.chat.reasoning_parser = "qwen3"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="enable_thinking", default_enabled=True
|
||
)
|
||
|
||
enabled_by_default = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
disabled_explicitly = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
chat_template_kwargs={"enable_thinking": False},
|
||
)
|
||
|
||
self.assertTrue(self.chat._get_reasoning_from_request(enabled_by_default))
|
||
self.assertFalse(self.chat._get_reasoning_from_request(disabled_explicitly))
|
||
|
||
def test_get_reasoning_from_request_default_false_toggle(self):
|
||
self.tm.server_args.reasoning_parser = "deepseek-v3"
|
||
self.chat.reasoning_parser = "deepseek-v3"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="thinking", default_enabled=False
|
||
)
|
||
|
||
disabled_by_default = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
enabled_explicitly = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
chat_template_kwargs={"thinking": True},
|
||
)
|
||
|
||
self.assertFalse(self.chat._get_reasoning_from_request(disabled_by_default))
|
||
self.assertTrue(self.chat._get_reasoning_from_request(enabled_explicitly))
|
||
|
||
def test_get_reasoning_from_request_special_cases(self):
|
||
self.tm.server_args.reasoning_parser = "mistral"
|
||
self.chat.reasoning_parser = "mistral"
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
special_case="always"
|
||
)
|
||
self.assertTrue(self.chat._get_reasoning_from_request(req))
|
||
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
special_case="mistral"
|
||
)
|
||
self.assertFalse(self.chat._get_reasoning_from_request(req))
|
||
req.reasoning_effort = "medium"
|
||
self.assertTrue(self.chat._get_reasoning_from_request(req))
|
||
|
||
# --- fallback path tests (config=None, uses reasoning_default) ---
|
||
|
||
def _setup_fallback(self, parser_name):
|
||
"""Set up reasoning with config=None to exercise the fallback path."""
|
||
self.tm.server_args.reasoning_parser = parser_name
|
||
self.chat = OpenAIServingChat(self.tm, self.template_manager)
|
||
self.chat.reasoning_parser = parser_name
|
||
self.template_manager.reasoning_config = None
|
||
|
||
def test_fallback_always_mode(self):
|
||
self._setup_fallback("deepseek-r1")
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
self.assertTrue(self.chat._get_reasoning_from_request(req))
|
||
|
||
def test_fallback_mistral_mode(self):
|
||
self._setup_fallback("mistral")
|
||
req_no_effort = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
self.assertFalse(self.chat._get_reasoning_from_request(req_no_effort))
|
||
|
||
req_with_effort = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
reasoning_effort="high",
|
||
)
|
||
self.assertTrue(self.chat._get_reasoning_from_request(req_with_effort))
|
||
|
||
def test_fallback_enable_thinking_mode_default_on(self):
|
||
self._setup_fallback("qwen3")
|
||
req_default = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
self.assertTrue(self.chat._get_reasoning_from_request(req_default))
|
||
|
||
req_disabled = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
chat_template_kwargs={"enable_thinking": False},
|
||
)
|
||
self.assertFalse(self.chat._get_reasoning_from_request(req_disabled))
|
||
|
||
def test_fallback_explicit_thinking_mode_default_off(self):
|
||
self._setup_fallback("deepseek-v3")
|
||
req_default = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
self.assertFalse(self.chat._get_reasoning_from_request(req_default))
|
||
|
||
req_enabled = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
chat_template_kwargs={"thinking": True},
|
||
)
|
||
self.assertTrue(self.chat._get_reasoning_from_request(req_enabled))
|
||
|
||
def test_fallback_explicit_enable_thinking_mode_default_off(self):
|
||
self._setup_fallback("mimo")
|
||
req_default = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
self.assertFalse(self.chat._get_reasoning_from_request(req_default))
|
||
|
||
req_enabled = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
chat_template_kwargs={"enable_thinking": True},
|
||
)
|
||
self.assertTrue(self.chat._get_reasoning_from_request(req_enabled))
|
||
|
||
def test_fallback_ling3_default_on(self):
|
||
"""Ling3 public checkpoints default `thinking_option='on'` in the chat
|
||
template when `enable_thinking` is omitted, and the template detector
|
||
cannot infer that indirect assignment. The parser fallback must mirror
|
||
the template default: omitted kwargs enable reasoning, only an explicit
|
||
`enable_thinking=False` disables it. Regression: the detector shipped
|
||
with `explicit_enable_thinking`, which left `reasoning_content` null on
|
||
default requests while the model was in fact thinking."""
|
||
self._setup_fallback("ling3")
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "hi"}]
|
||
)
|
||
cases = [
|
||
(None, True), # no chat_template_kwargs → thinking (template default)
|
||
({}, True), # empty kwargs → thinking
|
||
({"enable_thinking": True}, True), # explicit on
|
||
({"enable_thinking": False}, False), # explicit off
|
||
]
|
||
for kwargs, expected in cases:
|
||
with self.subTest(kwargs=kwargs):
|
||
req.chat_template_kwargs = kwargs
|
||
self.assertEqual(self.chat._get_reasoning_from_request(req), expected)
|
||
|
||
def test_fallback_no_detector_returns_false(self):
|
||
self.chat.reasoning_parser = "qwen3"
|
||
self.chat._reasoning_detector = None
|
||
self.template_manager.reasoning_config = None
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "Hi?"}]
|
||
)
|
||
self.assertFalse(self.chat._get_reasoning_from_request(req))
|
||
|
||
def test_build_chat_response_qwen3_thinking_forces_reasoning(self):
|
||
self.tm.server_args.reasoning_parser = "qwen3-thinking"
|
||
self.chat.reasoning_parser = "qwen3-thinking"
|
||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||
toggle_param="enable_thinking", default_enabled=True
|
||
)
|
||
|
||
req = ChatCompletionRequest(
|
||
model="Qwen/Qwen3-0.6B",
|
||
messages=[{"role": "user", "content": "Hi?"}],
|
||
separate_reasoning=True,
|
||
chat_template_kwargs={"enable_thinking": False},
|
||
)
|
||
ret_item = {
|
||
"text": "42",
|
||
"meta_info": {
|
||
"id": f"chatcmpl-{uuid.uuid4()}",
|
||
"prompt_tokens": 10,
|
||
"completion_tokens": 1,
|
||
"weight_version": "default",
|
||
"finish_reason": {"type": "stop", "matched": None},
|
||
},
|
||
"index": 0,
|
||
}
|
||
|
||
response = self.chat._build_chat_response(req, [ret_item], created=0)
|
||
msg = response.choices[0].message
|
||
self.assertEqual(msg.content, "")
|
||
self.assertEqual(msg.reasoning_content, "42")
|
||
|
||
# --- poolside_v1 (Laguna-XS.2) regression tests ---
|
||
|
||
def test_poolside_v1_enable_thinking_dispatch(self):
|
||
"""Laguna chat template defaults `enable_thinking=false`. Parser must
|
||
follow that default — must NOT return True via the generic fallback.
|
||
After the reasoning-config refactor, this is driven by
|
||
`_PoolsideV1Detector.reasoning_default = "explicit_enable_thinking"`."""
|
||
self._setup_fallback("poolside_v1")
|
||
req = ChatCompletionRequest(
|
||
model="x", messages=[{"role": "user", "content": "hi"}]
|
||
)
|
||
cases = [
|
||
(None, False), # no chat_template_kwargs → non-thinking (default)
|
||
({}, False), # empty kwargs → non-thinking
|
||
({"enable_thinking": False}, False), # explicit off
|
||
({"enable_thinking": True}, True), # explicit on
|
||
]
|
||
for kwargs, expected in cases:
|
||
with self.subTest(kwargs=kwargs):
|
||
req.chat_template_kwargs = kwargs
|
||
self.assertEqual(self.chat._get_reasoning_from_request(req), expected)
|
||
|
||
def test_poolside_v1_does_not_double_prepend_think(self):
|
||
"""When `enable_thinking=True` for poolside_v1, the HF chat template
|
||
already emits `<think>` via add_generation_prompt — server must NOT
|
||
append a second `<think>`. After the refactor this is guarded by
|
||
`_PoolsideV1Detector.thinks_internally = True` (inherited from Qwen3Detector).
|
||
"""
|
||
self._setup_fallback("poolside_v1")
|
||
req = ChatCompletionRequest(
|
||
model="x",
|
||
messages=[{"role": "user", "content": "hi"}],
|
||
chat_template_kwargs={"enable_thinking": True},
|
||
)
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv"
|
||
) as conv_mock:
|
||
conv_ins = Mock()
|
||
conv_ins.get_prompt.return_value = "BASE_PROMPT"
|
||
conv_ins.image_data = conv_ins.audio_data = conv_ins.video_data = None
|
||
conv_ins.modalities = []
|
||
conv_ins.stop_str = []
|
||
conv_mock.return_value = conv_ins
|
||
result = self.chat._apply_conversation_template(req, is_multimodal=False)
|
||
self.assertEqual(result.prompt, "BASE_PROMPT")
|
||
|
||
# ------------- hook method tests -------------
|
||
def test_encode_messages_returns_none_by_default(self):
|
||
"""Default _encode_messages returns None (use standard encoding)."""
|
||
result = self.chat._encode_messages([], Mock(), False)
|
||
self.assertIsNone(result)
|
||
|
||
def test_decode_response_returns_text(self):
|
||
"""Default _decode_response returns ret_item['text']."""
|
||
ret_item = {"text": "Hello world", "output_ids": [1, 2, 3]}
|
||
result = self.chat._decode_response(ret_item)
|
||
self.assertEqual(result, "Hello world")
|
||
|
||
def test_get_parsed_response_fields_passthrough(self):
|
||
"""Default _get_parsed_response_fields passes through values."""
|
||
reasoning = "thinking..."
|
||
tool_calls = [{"name": "foo"}]
|
||
r, t = self.chat._get_parsed_response_fields(reasoning, tool_calls)
|
||
self.assertEqual(r, reasoning)
|
||
self.assertEqual(t, tool_calls)
|
||
|
||
|
||
class TestProcessToolCallsWithRequiredToolChoice(unittest.TestCase):
|
||
"""Test _process_tool_calls with tool_choice='required' uses model-specific parser."""
|
||
|
||
def setUp(self):
|
||
reset_context()
|
||
self.addCleanup(reset_context)
|
||
publish(ServerArgs(model_path="dummy"), role="tokenizer")
|
||
tm = _MockTokenizerManager()
|
||
tm.server_args.tool_call_parser = "kimi_k2"
|
||
self.chat = OpenAIServingChat(tm, _MockTemplateManager())
|
||
|
||
def test_required_with_parser_uses_function_call_parser(self):
|
||
"""tool_choice='required' should use FunctionCallParser when tool_call_parser is set."""
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||
) as ParserMock:
|
||
call_info = Mock()
|
||
call_info.name = "get_weather"
|
||
call_info.parameters = '{"location":"Tokyo"}'
|
||
call_info.tool_index = 0
|
||
|
||
parser_instance = ParserMock.return_value
|
||
parser_instance.has_tool_call.return_value = True
|
||
parser_instance.parse_non_stream.return_value = ("", [call_info])
|
||
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
tools = [{"type": "function", "function": {"name": "get_weather"}}]
|
||
|
||
tool_calls, text, fr = self.chat._process_tool_calls(
|
||
text="<|tool_calls_section_begin|>...<|tool_calls_section_end|>",
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
tool_choice="required",
|
||
)
|
||
|
||
self.assertIsNotNone(tool_calls)
|
||
self.assertEqual(len(tool_calls), 1)
|
||
self.assertEqual(tool_calls[0].function.name, "get_weather")
|
||
self.assertEqual(fr["type"], "tool_calls")
|
||
|
||
def test_empty_parser_result_is_not_reported_as_tool_call(self):
|
||
with patch(
|
||
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||
) as ParserMock:
|
||
parser_instance = ParserMock.return_value
|
||
parser_instance.has_tool_call.return_value = True
|
||
parser_instance.detector.supports_structural_tag.return_value = True
|
||
parser_instance.parse_non_stream.return_value = ("Visible prefix.", [])
|
||
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
tools = [{"type": "function", "function": {"name": "get_weather"}}]
|
||
|
||
tool_calls, text, fr = self.chat._process_tool_calls(
|
||
text="<|malformed_tool_call|>",
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
tool_choice="required",
|
||
)
|
||
|
||
self.assertIsNone(tool_calls)
|
||
self.assertEqual(text, "Visible prefix.")
|
||
self.assertEqual(fr, {"type": "stop", "matched": None})
|
||
|
||
def test_required_without_parser_falls_back_to_json(self):
|
||
"""tool_choice='required' without parser should parse as JSON array."""
|
||
self.chat.tool_call_parser = None
|
||
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
tools = [{"type": "function", "function": {"name": "get_weather"}}]
|
||
|
||
tool_calls, text, fr = self.chat._process_tool_calls(
|
||
text='[{"name":"get_weather","parameters":{"location":"Tokyo"}}]',
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
tool_choice="required",
|
||
)
|
||
|
||
self.assertIsNotNone(tool_calls)
|
||
self.assertEqual(len(tool_calls), 1)
|
||
self.assertEqual(tool_calls[0].function.name, "get_weather")
|
||
|
||
def test_required_without_parser_invalid_json_returns_none(self):
|
||
"""tool_choice='required' without parser and invalid JSON returns tool_calls=None."""
|
||
self.chat.tool_call_parser = None
|
||
|
||
finish_reason = {"type": "stop", "matched": None}
|
||
tools = [{"type": "function", "function": {"name": "get_weather"}}]
|
||
|
||
tool_calls, text, fr = self.chat._process_tool_calls(
|
||
text="<|tool_calls_section_begin|>not json",
|
||
tools=tools,
|
||
finish_reason=finish_reason,
|
||
tool_choice="required",
|
||
)
|
||
|
||
self.assertIsNone(tool_calls)
|
||
|
||
|
||
class TestNormalizeToolContent(unittest.TestCase):
|
||
"""Unit tests for normalize_tool_content()."""
|
||
|
||
def test_multiple_text_parts_joined(self):
|
||
result = normalize_tool_content(
|
||
"tool",
|
||
[{"type": "text", "text": "hello"}, {"type": "text", "text": "world"}],
|
||
)
|
||
self.assertEqual(result, "hello world")
|
||
|
||
def test_non_text_part_list_preserved(self):
|
||
content = [{"name": "func", "output": "result"}]
|
||
result = normalize_tool_content("tool", content)
|
||
self.assertIs(result, content)
|
||
|
||
def test_string_content_unchanged(self):
|
||
self.assertEqual(normalize_tool_content("tool", "hello"), "hello")
|
||
|
||
def test_empty_list_returns_empty_string(self):
|
||
self.assertEqual(normalize_tool_content("tool", []), "")
|
||
|
||
def test_non_tool_role_unchanged(self):
|
||
content = [{"type": "text", "text": "hi"}]
|
||
result = normalize_tool_content("user", content)
|
||
self.assertIs(result, content)
|
||
|
||
def test_mixed_str_and_dict_parts(self):
|
||
result = normalize_tool_content(
|
||
"tool", ["plain", {"type": "text", "text": "rich"}]
|
||
)
|
||
self.assertEqual(result, "plain rich")
|
||
|
||
|
||
class InklingReasoningEffortTest(unittest.TestCase):
|
||
"""Inkling reasoning-effort mapping and validation."""
|
||
|
||
def test_named_levels(self):
|
||
parse = OpenAIServingChat._parse_inkling_reasoning_effort
|
||
self.assertEqual(parse("none"), 0.0)
|
||
self.assertEqual(parse("minimal"), 0.1)
|
||
self.assertEqual(parse("low"), 0.2)
|
||
self.assertEqual(parse("medium"), 0.7)
|
||
self.assertEqual(parse("high"), 0.9)
|
||
# "xhigh" and "max" are aliases for the same 0.99 ceiling
|
||
self.assertEqual(parse("xhigh"), 0.99)
|
||
self.assertEqual(parse("max"), 0.99)
|
||
self.assertEqual(parse("max"), parse("xhigh"))
|
||
|
||
def test_scalar_range_is_validated(self):
|
||
parse = OpenAIServingChat._parse_inkling_reasoning_effort
|
||
self.assertEqual(parse(0.5), 0.5)
|
||
self.assertEqual(parse(0.99), 0.99)
|
||
for value in (1.0, "1.0", 2.0, "1.5", -1.0, float("nan"), True):
|
||
with self.subTest(value=value), self.assertRaises(ValueError):
|
||
parse(value)
|
||
|
||
def test_invalid_and_none(self):
|
||
parse = OpenAIServingChat._parse_inkling_reasoning_effort
|
||
self.assertIsNone(parse(None))
|
||
with self.assertRaises(ValueError):
|
||
parse("garbage")
|
||
|
||
def test_env_default(self):
|
||
from sglang.srt.environ import envs
|
||
|
||
get = OpenAIServingChat._get_inkling_default_reasoning_effort
|
||
env = envs.SGLANG_INKLING_DEFAULT_REASONING_EFFORT
|
||
try:
|
||
env.clear() # unset -> EnvStr default "0.9"
|
||
self.assertEqual(get(), 0.9)
|
||
env.set("") # explicit empty still uses the protocol default
|
||
self.assertEqual(get(), 0.9)
|
||
env.set("0.7")
|
||
self.assertEqual(get(), 0.7)
|
||
for value in ("1.0", "1.1", "garbage"):
|
||
env.set(value)
|
||
with self.subTest(value=value), self.assertRaises(ValueError):
|
||
get()
|
||
finally:
|
||
env.clear()
|
||
|
||
def test_thinking_disabled_maps_to_no_thinking_effort(self):
|
||
"""Bug regression: Inkling is an always-on parser, so Anthropic
|
||
thinking={"type": "disabled"} was rejected outright even though effort
|
||
"none" (0.0) expresses exactly that."""
|
||
serving = object.__new__(OpenAIServingChat)
|
||
serving.reasoning_parser = "inkling"
|
||
serving.template_manager = Mock(reasoning_config=None)
|
||
serving._reasoning_detector = Mock(reasoning_default="always")
|
||
request = ChatCompletionRequest(
|
||
model="test-model", messages=[{"role": "user", "content": "hi"}]
|
||
)
|
||
|
||
serving.apply_reasoning_enabled(request, False)
|
||
self.assertEqual(request.reasoning_effort, "none")
|
||
|
||
# Enabling leaves an effort set via output_config.effort alone.
|
||
request.reasoning_effort = "low"
|
||
serving.apply_reasoning_enabled(request, True)
|
||
self.assertEqual(request.reasoning_effort, "low")
|
||
|
||
def test_serving_does_not_prefill_model_message(self):
|
||
from sglang.srt.parser.inkling_tokenizer import INKLING_SPECIAL_TOKEN_IDS
|
||
|
||
class Tokenizer:
|
||
def encode(self, text, add_special_tokens=False):
|
||
return list(text.encode())
|
||
|
||
serving = object.__new__(OpenAIServingChat)
|
||
serving.chat_encoding_spec = "inkling"
|
||
serving.tokenizer_manager = Mock(tokenizer=Tokenizer())
|
||
request = ChatCompletionRequest(
|
||
model="test-model",
|
||
messages=[{"role": "user", "content": "hello"}],
|
||
reasoning_effort=0.5,
|
||
)
|
||
prompt_ids = serving._encode_messages(
|
||
[message.model_dump() for message in request.messages],
|
||
request,
|
||
thinking_mode=None,
|
||
)
|
||
self.assertEqual(prompt_ids[-1], INKLING_SPECIAL_TOKEN_IDS["<|end_message|>"])
|
||
|
||
def test_continue_final_message_resumes_open_model_text_block(self):
|
||
"""Bug regression: continue_final_message was silently ignored on the
|
||
inkling path — the trailing assistant message rendered as a CLOSED
|
||
historical turn (<|end_message|> + <|content_model_end_sampling|>), so
|
||
the model started a fresh turn instead of continuing. The prefix must
|
||
render as an OPEN model text block."""
|
||
from sglang.srt.parser.inkling_tokenizer import INKLING_SPECIAL_TOKEN_IDS
|
||
|
||
class Tokenizer:
|
||
def encode(self, text, add_special_tokens=False):
|
||
return list(text.encode())
|
||
|
||
serving = object.__new__(OpenAIServingChat)
|
||
serving.chat_encoding_spec = "inkling"
|
||
serving.tokenizer_manager = Mock(tokenizer=Tokenizer())
|
||
request = ChatCompletionRequest(
|
||
model="test-model",
|
||
messages=[
|
||
{"role": "user", "content": "hello"},
|
||
{"role": "assistant", "content": "The answer"},
|
||
],
|
||
reasoning_effort=0.5,
|
||
continue_final_message=True,
|
||
)
|
||
prompt_ids = serving._encode_messages(
|
||
[message.model_dump() for message in request.messages],
|
||
request,
|
||
thinking_mode=None,
|
||
)
|
||
open_block = [
|
||
INKLING_SPECIAL_TOKEN_IDS["<|message_model|>"],
|
||
INKLING_SPECIAL_TOKEN_IDS["<|content_text|>"],
|
||
*list(b"The answer"),
|
||
]
|
||
self.assertEqual(prompt_ids[-len(open_block) :], open_block)
|
||
self.assertNotIn(
|
||
INKLING_SPECIAL_TOKEN_IDS["<|content_model_end_sampling|>"], prompt_ids
|
||
)
|
||
|
||
def test_continue_final_message_leaves_tool_call_turns_closed(self):
|
||
"""A trailing assistant message with tool_calls cannot be continued —
|
||
it must keep rendering as a closed historical turn."""
|
||
from sglang.srt.parser.inkling_tokenizer import INKLING_SPECIAL_TOKEN_IDS
|
||
|
||
class Tokenizer:
|
||
def encode(self, text, add_special_tokens=False):
|
||
return list(text.encode())
|
||
|
||
serving = object.__new__(OpenAIServingChat)
|
||
serving.chat_encoding_spec = "inkling"
|
||
serving.tokenizer_manager = Mock(tokenizer=Tokenizer())
|
||
request = ChatCompletionRequest(
|
||
model="test-model",
|
||
messages=[
|
||
{"role": "user", "content": "hello"},
|
||
{
|
||
"role": "assistant",
|
||
"content": "calling",
|
||
"tool_calls": [
|
||
{
|
||
"id": "call-1",
|
||
"type": "function",
|
||
"function": {"name": "weather", "arguments": "{}"},
|
||
}
|
||
],
|
||
},
|
||
],
|
||
reasoning_effort=0.5,
|
||
continue_final_message=True,
|
||
)
|
||
prompt_ids = serving._encode_messages(
|
||
[message.model_dump() for message in request.messages],
|
||
request,
|
||
thinking_mode=None,
|
||
)
|
||
self.assertEqual(
|
||
prompt_ids[-1],
|
||
INKLING_SPECIAL_TOKEN_IDS["<|content_model_end_sampling|>"],
|
||
)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main(verbosity=2)
|