Files
sglang/python/sglang/srt/function_call/deepseekv32_detector.py
2025-12-26 11:35:30 -08:00

355 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
import logging
import re
from partial_json_parser.core.options import Allow
from sglang.srt.entrypoints.openai.protocol import Tool
from sglang.srt.function_call.base_format_detector import BaseFormatDetector
from sglang.srt.function_call.core_types import (
StreamingParseResult,
StructureInfo,
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import _find_common_prefix, _partial_json_loads
logger = logging.getLogger(__name__)
class DeepSeekV32Detector(BaseFormatDetector):
"""
Detector for DeepSeek V3.2 model function call format.
The DeepSeek V3.2 format uses XML-like DSML tags to delimit function calls.
Supports two parameter formats:
Format 1 - XML Parameter Tags:
```
<DSMLfunction_calls>
<DSMLinvoke name="function_name">
<DSMLparameter name="param_name" string="true">value</DSMLparameter>
...
</DSMLinvoke>
</DSMLfunction_calls>
```
Format 2 - Direct JSON:
```
<DSMLfunction_calls>
<DSMLinvoke name="function_name">
{
"param_name": "value"
}
</DSMLinvoke>
</DSMLfunction_calls>
```
Examples:
```
<DSMLfunction_calls>
<DSMLinvoke name="get_favorite_tourist_spot">
<DSMLparameter name="city" string="true">San Francisco</DSMLparameter>
</DSMLinvoke>
</DSMLfunction_calls>
<DSMLfunction_calls>
<DSMLinvoke name="get_favorite_tourist_spot">
{ "city": "San Francisco" }
</DSMLinvoke>
</DSMLfunction_calls>
```
Key Components:
- Tool Calls Section: Wrapped between `<DSMLfunction_calls>` and `</DSMLfunction_calls>`
- Individual Tool Call: Wrapped between `<DSMLinvoke name="...">` and `</DSMLinvoke>`
- Parameters: Either XML tags or direct JSON format
- Supports multiple tool calls
Reference: DeepSeek V3.2 format specification
"""
def __init__(self):
super().__init__()
self.bot_token = "<DSMLfunction_calls>"
self.eot_token = "</DSMLfunction_calls>"
self.invoke_end_token = "</DSMLinvoke>"
self.parameter_regex = r'<DSMLparameter\s+name="([^"]+)"\s+string="([^"]+)"\s*>(.*?)</DSMLparameter>'
self.partial_parameter_regex = (
r'<DSMLparameter\s+name="([^"]+)"\s+string="([^"]+)"\s*>(.*)$'
)
self.function_calls_regex = (
r"<DSMLfunction_calls>(.*?)</DSMLfunction_calls>"
)
self.invoke_regex = (
r'<DSMLinvoke\s+name="([^"]+)"\s*>(.*?)(</DSMLinvoke>|$)'
)
self.prefix_parameter_end_call = ["</", "DSML", "parameter"]
self.current_tool_id = -1
def has_tool_call(self, text: str) -> bool:
"""Check if the text contains a deepseek v32 format tool call."""
return self.bot_token in text or "<DSMLinvoke" in text
def _parse_parameters_from_xml(
self, invoke_content: str, allow_partial: bool = False
) -> dict:
"""
Parse parameters from either XML-like format or JSON format to dict.
Supports two formats:
1. XML parameter tags: <DSMLparameter name="..." string="...">value</DSMLparameter>
2. Direct JSON: { "key": "value" }
"""
# First, try to parse as direct JSON (new format)
invoke_content_stripped = invoke_content.strip()
if invoke_content_stripped.startswith("{") and invoke_content_stripped.endswith(
"}"
):
try:
parameters = json.loads(invoke_content_stripped)
if isinstance(parameters, dict):
return parameters
except (json.JSONDecodeError, ValueError):
# If JSON parsing fails, fall through to XML parsing
pass
# Fall back to XML parameter tag parsing (original format)
parameters = {}
# Find all complete parameter matches
param_matches = list(
re.finditer(self.parameter_regex, invoke_content, re.DOTALL)
)
last_match_end = 0
for match in param_matches:
param_name = match.group(1)
param_type = match.group(2)
param_value = match.group(3)
last_match_end = match.end()
# Convert value based on type
if param_type == "true": # string type
parameters[param_name] = param_value.strip()
else:
# Try to parse as JSON for other types
try:
parameters[param_name] = json.loads(param_value.strip())
except (json.JSONDecodeError, ValueError):
parameters[param_name] = param_value.strip()
# If allowed, try to parse a partial parameter at the end
if allow_partial:
remaining_content = invoke_content[last_match_end:]
# Remove incomplete parameter_end_call prefix in case they are captured by param
for token in reversed(self.prefix_parameter_end_call):
remaining_content = remaining_content.rstrip(token)
# Match start of a parameter tag + value (potentially incomplete)
# Regex: <tag name="..." string="...">VALUE... (no end tag)
partial_match = re.search(
self.partial_parameter_regex, remaining_content, re.DOTALL
)
if partial_match and (param_value := partial_match.group(3)):
param_name = partial_match.group(1)
if partial_match.group(2) == "true":
parameters[param_name] = param_value.strip()
else:
parameters[param_name] = _partial_json_loads(
param_value, Allow.ALL
)[0]
return parameters
def detect_and_parse(self, text: str, tools: list[Tool]) -> StreamingParseResult:
"""
One-time parsing: Detects and parses tool calls in the provided text.
:param text: The complete text to parse.
:param tools: List of available tools.
:return: ParseResult indicating success or failure, consumed text, leftover text, and parsed calls.
"""
idx = text.find(self.bot_token)
normal_text = text[:idx].strip() if idx != -1 else text
if self.bot_token not in text:
return StreamingParseResult(normal_text=normal_text, calls=[])
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:
return StreamingParseResult(normal_text=normal_text, calls=[])
function_calls_content = function_calls_match.group(1)
# Find all invoke blocks
invoke_matches = re.findall(
self.invoke_regex, function_calls_content, re.DOTALL
)
for func_name, invoke_content, _ in invoke_matches:
# Parse parameters from XML format
func_args = self._parse_parameters_from_xml(invoke_content)
# construct match_result for parse_base_json
match_result = {"name": func_name, "parameters": func_args}
calls.extend(self.parse_base_json(match_result, tools))
return StreamingParseResult(normal_text=normal_text, calls=calls)
except Exception as e:
logger.error(f"Error in detect_and_parse: {e}")
# return the normal text if parsing fails
return StreamingParseResult(normal_text=text)
def parse_streaming_increment(
self, new_text: str, tools: list[Tool]
) -> StreamingParseResult:
"""
Streaming incremental parsing tool calls for DeepSeekV32 format.
Supports multiple consecutive invoke blocks and argument streaming.
"""
self._buffer += new_text
current_text = self._buffer
# Check if buffer contains any DSML markers or ends with potential tag prefix
# This handles partial/streaming DSML content
dsml_markers = ["DSML", "<", "</"]
potentially_dsml = any(marker in current_text for marker in dsml_markers)
# Also check if text ends with start of a tag (to handle "<" arriving separately)
dsml_prefixes = ["<", "<", "</", "</"]
ends_with_prefix = any(
current_text.rstrip().endswith(prefix) for prefix in dsml_prefixes
)
if (
not self.has_tool_call(current_text)
and not potentially_dsml
and not ends_with_prefix
):
self._buffer = ""
for e_token in [self.eot_token, self.invoke_end_token]:
if e_token in current_text:
current_text = current_text.replace(e_token, "")
return StreamingParseResult(normal_text=current_text)
all_calls: list[ToolCallItem] = []
try:
# Loop to handle multiple consecutive invoke blocks
while True:
# Try to match an invoke block (may be partial)
invoke_match = re.search(
pattern=self.invoke_regex,
string=current_text,
flags=re.DOTALL,
)
if not invoke_match:
break
func_name = invoke_match.group(1).strip()
invoke_content = invoke_match.group(2)
# group(3) is either "</DSMLinvoke>" (complete) or "" (incomplete, matched with $)
is_tool_end = bool(invoke_match.group(3))
# Initialize state if this is the first tool call
if self.current_tool_id == -1:
self.current_tool_id = 0
self.prev_tool_call_arr = []
self.streamed_args_for_tool = [""]
# Ensure arrays are large enough for current tool
while len(self.prev_tool_call_arr) <= self.current_tool_id:
self.prev_tool_call_arr.append({})
while len(self.streamed_args_for_tool) <= self.current_tool_id:
self.streamed_args_for_tool.append("")
# 1. Send tool name if not sent yet
if not self.current_tool_name_sent:
all_calls.append(
ToolCallItem(
tool_index=self.current_tool_id,
name=func_name,
parameters="",
)
)
self.current_tool_name_sent = True
# 2. Parse current parameters (partial or complete)
current_params = self._parse_parameters_from_xml(
invoke_content, allow_partial=not is_tool_end
)
current_args_json = json.dumps(current_params, ensure_ascii=False)
# 3. Calculate and send incremental arguments
sent_len = len(self.streamed_args_for_tool[self.current_tool_id])
prev_params = self.prev_tool_call_arr[self.current_tool_id].get(
"arguments"
)
argument_diff = None
if is_tool_end:
# If complete, send everything remaining
argument_diff = current_args_json[sent_len:]
elif prev_params is not None:
# If partial, send stable prefix diff
prev_args_json = json.dumps(prev_params, ensure_ascii=False)
if current_args_json != prev_args_json:
prefix = _find_common_prefix(prev_args_json, current_args_json)
if len(prefix) > sent_len:
argument_diff = prefix[sent_len:]
if argument_diff:
all_calls.append(
ToolCallItem(
tool_index=self.current_tool_id,
name=None,
parameters=argument_diff,
)
)
self.streamed_args_for_tool[self.current_tool_id] += argument_diff
# Update the stored arguments
self.prev_tool_call_arr[self.current_tool_id] = {
"name": func_name,
"arguments": current_params,
}
# Check if tool call is complete (has closing tag)
if is_tool_end:
# Remove the completed tool call from buffer
self._buffer = current_text[invoke_match.end() :]
current_text = self._buffer # Update for next iteration
# Move to next tool call
self.current_tool_id += 1
self.current_tool_name_sent = False
# Continue loop to check for more invoke blocks
continue
else:
# Tool call not complete yet, don't return anything
# Wait for more chunks until we see </DSMLinvoke>
break
# No more invoke blocks found
return StreamingParseResult(normal_text="", calls=all_calls)
except Exception as e:
logger.error(f"Error in parse_streaming_increment: {e}")
return StreamingParseResult(normal_text=current_text)
def structure_info(self) -> _GetInfoFunc:
return lambda name: StructureInfo(
begin=f'<DSMLinvoke name="{name}">',
end="</DSMLinvoke>",
trigger=f"<DSMLinvoke",
)