Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 35 additions & 19 deletions python/sglang/srt/function_call/deepseekv32_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import logging
import re

from partial_json_parser.core.exceptions import MalformedJSON
from partial_json_parser.core.options import Allow

from sglang.srt.entrypoints.openai.protocol import Tool
Expand Down Expand Up @@ -179,7 +180,7 @@ def _parse_parameters_from_xml(
parameters[param_name] = _partial_json_loads(
param_value, Allow.ALL
)[0]
except json.JSONDecodeError:
except (json.JSONDecodeError, MalformedJSON, ValueError):
parameters[param_name] = param_value.strip()

return json.dumps(parameters, ensure_ascii=False)
Expand All @@ -199,26 +200,25 @@ def detect_and_parse(self, text: str, tools: list[Tool]) -> StreamingParseResult

calls = []
try:
# Extract content between function_calls tags
function_calls_match = re.search(
self.function_calls_regex,
text,
re.DOTALL,
)
if not function_calls_match:
sections = re.findall(self.function_calls_regex, text, re.DOTALL)
if not sections:
return StreamingParseResult(normal_text=normal_text, calls=[])

function_calls_content = function_calls_match.group(1)

# Find all invoke blocks
for invoke_match in re.finditer(
self.invoke_regex, function_calls_content, re.DOTALL
):
func_name, invoke_content, _ = self._unpack_invoke_match(invoke_match)
func_args = self._parse_parameters_from_xml(invoke_content)
# construct match_result for parse_base_json
match_result = {"name": func_name, "parameters": json.loads(func_args)}
calls.extend(self.parse_base_json(match_result, tools))
for function_calls_content in sections:
for invoke_match in re.finditer(
self.invoke_regex, function_calls_content, re.DOTALL
):
func_name, invoke_content, _ = self._unpack_invoke_match(
invoke_match
)
func_args = self._parse_parameters_from_xml(invoke_content)
# construct match_result for parse_base_json
match_result = {
"name": func_name,
"parameters": json.loads(func_args),
}
calls.extend(self.parse_base_json(match_result, tools))

return StreamingParseResult(normal_text=normal_text, calls=calls)
except Exception as e:
Expand Down Expand Up @@ -259,6 +259,9 @@ def parse_streaming_increment(
return StreamingParseResult(normal_text=current_text)

all_calls: list[ToolCallItem] = []
# Only recovered for the first call: the DSML guard above never releases a
# buffer that still holds a marker, so later prose stays buffered.
preamble = ""
try:
# Loop to handle multiple consecutive invoke blocks
while True:
Expand All @@ -280,6 +283,12 @@ def parse_streaming_increment(
self.current_tool_id = 0
self.prev_tool_call_arr = []
self.streamed_args_for_tool = [""]
call_start = invoke_match.start()
bot_pos = current_text.rfind(self.bot_token, 0, call_start)
if bot_pos != -1:
call_start = bot_pos
# Same trailing-newline trim as detect_and_parse, so both agree.
preamble = current_text[:call_start].removesuffix("\n\n")

# Ensure arrays are large enough for current tool
while len(self.prev_tool_call_arr) <= self.current_tool_id:
Expand Down Expand Up @@ -355,10 +364,17 @@ def parse_streaming_increment(
break

# No more invoke blocks found
return StreamingParseResult(normal_text="", calls=all_calls)
return StreamingParseResult(normal_text=preamble, calls=all_calls)

except Exception as e:
logger.error(f"Error in parse_streaming_increment: {e}")
# Re-emit verbatim rather than swallowing the turn; the preamble is
# still inside current_text unless a completed call advanced past it.
# Calls are dropped on purpose: the failure can land between a tool's
# name and its arguments, and a half-formed call is worse than none.
self._buffer = ""
if not current_text.startswith(preamble):
current_text = preamble + current_text
return StreamingParseResult(normal_text=current_text)

def structure_info(self) -> _GetInfoFunc:
Expand Down
59 changes: 42 additions & 17 deletions python/sglang/srt/parser/reasoning_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -194,6 +194,8 @@ def _parse_streaming_increment_impl(self, new_text: str) -> StreamingParseResult
# Strip `<think>` token if present
if not self.stripped_think_start and think_start_text in current_text:
current_text = current_text.replace(think_start_text, "", 1)
# Write back, or stream_reasoning=False carries the token into finish().
self._buffer = current_text
self.stripped_think_start = True
self._in_reasoning = True

Expand All @@ -213,7 +215,8 @@ def _parse_streaming_increment_impl(self, new_text: str) -> StreamingParseResult

# Continue with reasoning content
if self._in_reasoning:
# Check for tool_start_token interruption
# Check for tool_start_token interruption. Streaming cannot see a
# think_end_token that has not arrived yet; see the chunk_dependent test.
if self.tool_start_token and self.tool_start_token in current_text:
tool_idx = current_text.find(self.tool_start_token)
reasoning_text = current_text[:tool_idx]
Expand All @@ -225,9 +228,21 @@ def _parse_streaming_increment_impl(self, new_text: str) -> StreamingParseResult
normal_text=normal_text, reasoning_text=reasoning_text
)
if self.stream_reasoning:
# Stream the content immediately
self._buffer = ""
return StreamingParseResult(reasoning_text=current_text)
# Minus any trailing slice that could be a token split across chunks.
holdback_tokens = [self.think_end_token]
if self.tool_start_token:
holdback_tokens.append(self.tool_start_token)
if not self.stripped_think_start:
# force_reasoning never saw the opening token; it can still split.
holdback_tokens.append(think_start_text)
holdback = max(
self._ends_with_partial_token(current_text, token)
for token in holdback_tokens
)
self._buffer = current_text[len(current_text) - holdback :]
return StreamingParseResult(
reasoning_text=current_text[: len(current_text) - holdback]
)
else:
return StreamingParseResult()

Expand All @@ -244,15 +259,30 @@ def _strip_leading_think_start(self, text: str) -> str:
return text[len(think_start_text) :]
return text

@staticmethod
def _ends_with_partial_token(buffer: str, token: str) -> int:
"""Length of the longest trailing slice of `buffer` that is a strict prefix
of `token`. Longest, so a token whose prefix repeats inside itself does not
get cut short and leak the rest of the marker."""
for i in range(min(len(buffer), len(token) - 1), 0, -1):
if token.startswith(buffer[-i:]):
return i
return 0

def finish(self) -> StreamingParseResult:
"""Flush reasoning buffered under stream_reasoning=False when the stream ends
before the end token (e.g. max_tokens cut it short), instead of dropping it.
"""Flush reasoning still buffered when the stream ends before the end token
(e.g. max_tokens cut it short), instead of dropping it: the whole block under
stream_reasoning=False, the held-back token suffix under stream_reasoning=True.
force_nonempty_content emits it as normal_text, else as reasoning_text."""
if not self._in_reasoning:
return StreamingParseResult()
# Same as the reasoning-side flush below: a held-back slice that never
# became a token is content.
leftover = self._buffer
self._buffer = ""
return StreamingParseResult(normal_text=leftover)

# stream_reasoning=False never clears _buffer, so the opening think token
# (stripped only from the base class's local view) survives here.
# Defensive: subclasses that fill _buffer themselves may not have stripped
# the opening think token that _parse_streaming_increment_impl removes.
buffer = self._strip_leading_think_start(self._buffer)
self._buffer = ""

Expand All @@ -263,7 +293,7 @@ def finish(self) -> StreamingParseResult:
return StreamingParseResult(normal_text=normal_text)
return StreamingParseResult()

if not self.stream_reasoning and buffer:
if buffer:
return StreamingParseResult(reasoning_text=buffer)

return StreamingParseResult()
Expand Down Expand Up @@ -1116,6 +1146,8 @@ def __init__(
dsv4_thinking_start_token,
dsv4_thinking_end_token,
think_excluded_tokens=[dsv4_eos_token, dsv4_dsml_token],
# Leading "<" included: has_tool_call() matches on it.
tool_start_token=f"<{dsv4_dsml_token}",
force_reasoning=force_reasoning,
stream_reasoning=stream_reasoning,
continue_final_message=continue_final_message,
Expand Down Expand Up @@ -1174,13 +1206,6 @@ def __init__(
self._reasoning_acc: str = ""
self._in_inner_tool: bool = False

@staticmethod
def _ends_with_partial_token(buffer: str, token: str) -> int:
for i in range(1, min(len(buffer) + 1, len(token))):
if token.startswith(buffer[-i:]):
return i
return 0

def detect_and_parse(self, text: str) -> StreamingParseResult:
blocks = self.detect_and_parse_block_sequence(text)
reasoning_parts = [t for k, t in blocks if k == "reasoning"]
Expand Down
129 changes: 129 additions & 0 deletions test/registered/unit/function_call/test_deepseekv4_detector.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,129 @@
"""Unit tests for DeepSeekV4Detector DSML streaming — no server, no model loading."""

from unittest.mock import patch

from sglang.srt.entrypoints.openai.protocol import Function, Tool
from sglang.srt.function_call.deepseekv4_detector import DeepSeekV4Detector
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase

register_cpu_ci(1.0, "base-a-test-cpu")

DSML = "|DSML|"


def _wrapped(invoke: str) -> str:
return f"<{DSML}tool_calls>\n{invoke}\n</{DSML}tool_calls>"


def _invoke(name: str, params: str = "") -> str:
return f'<{DSML}invoke name="{name}">\n{params}\n</{DSML}invoke>'


def _param(name: str, is_string: str, value: str) -> str:
return (
f'<{DSML}parameter name="{name}" string="{is_string}">{value}</{DSML}parameter>'
)


def _weather_call(city: str = "SF") -> str:
return _wrapped(_invoke("get_weather", _param("city", "true", city)))


class TestDeepSeekV4Streaming(CustomTestCase):
def setUp(self):
self.tools = [
Tool(
type="function",
function=Function(
name="get_weather",
description="Get weather information",
parameters={
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
),
)
]

def _feed(self, chunks):
"""Returns (normal_text, calls) accumulated over the chunks."""
detector = DeepSeekV4Detector()
normal, calls = "", []
for chunk in chunks:
result = detector.parse_streaming_increment(chunk, self.tools)
normal += result.normal_text
calls.extend(result.calls)
return normal, calls

def test_preamble_in_same_delta_as_tool_call(self):
"""Prose sharing a delta with the tool call must not be dropped, and the
streaming and one-shot paths must agree on it."""
text = "Let me check.\n" + _weather_call()
normal, calls = self._feed([text])

self.assertEqual([c.name for c in calls if c.name], ["get_weather"])
self.assertEqual(
normal, DeepSeekV4Detector().detect_and_parse(text, self.tools).normal_text
)

def test_preamble_before_bare_invoke_without_wrapper(self):
"""The bare `<|DSML|invoke …>` form has no tool_calls wrapper to walk
back to, so the preamble is computed from the invoke itself."""
text = "Checking.\n" + _invoke("get_weather", _param("city", "true", "SF"))
normal, calls = self._feed([text])

self.assertIn("Checking.", normal)
self.assertEqual([c.name for c in calls if c.name], ["get_weather"])

def test_no_dsml_markers_leak_into_normal_text(self):
text = "Prose.\n" + _weather_call()
normal, _ = self._feed([text[i : i + 4] for i in range(0, len(text), 4)])

self.assertNotIn(DSML, normal)

def test_malformed_partial_json_falls_back_to_raw_value(self):
"""A partial non-string parameter must not escape as MalformedJSON."""
detector = DeepSeekV4Detector()
result = detector.parse_streaming_increment(
f'<{DSML}tool_calls>\n<{DSML}invoke name="get_weather">\n'
f'<{DSML}parameter name="city" string="false">{{"a"',
self.tools,
)

self.assertEqual([c.name for c in result.calls if c.name], ["get_weather"])

def test_non_streaming_parses_every_tool_calls_section(self):
"""A turn with two tool_calls sections must yield both calls."""
result = DeepSeekV4Detector().detect_and_parse(
f"{_weather_call('SF')}\n{_weather_call('NY')}", self.tools
)

self.assertEqual(len(result.calls), 2)

def test_parse_error_neither_swallows_nor_duplicates(self):
"""An unexpected parse error must not empty the turn, and the dropped
buffer must not come back on the next delta."""
detector = DeepSeekV4Detector()

with patch.object(
DeepSeekV4Detector,
"_parse_parameters_from_xml",
side_effect=RuntimeError("boom"),
):
first = detector.parse_streaming_increment(_weather_call(), self.tools)
self.assertEqual(detector._buffer, "")
second = detector.parse_streaming_increment(" tail", self.tools)

self.assertIn("get_weather", first.normal_text)
self.assertNotIn("get_weather", second.normal_text)
# No half-formed call: the failure can land between a tool's name and its
# arguments, so an argument-less named call must not reach the client.
self.assertEqual(first.calls, [])


if __name__ == "__main__":
import unittest

unittest.main()
Loading
Loading