feat(dsv32): better error handling for DeepSeek-v3.2 encoder (#14353)

This commit is contained in:
Jimmy
2025-12-18 16:34:27 -08:00
committed by GitHub
parent e72b02db28
commit 216067c0cb
2 changed files with 53 additions and 32 deletions
@@ -4,6 +4,11 @@ import json
import re import re
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any, Dict, List, Optional, Tuple, Union
class DS32EncodingError(Exception):
pass
TOOLS_SYSTEM_TEMPLATE = """## Tools TOOLS_SYSTEM_TEMPLATE = """## Tools
You have access to a set of tools you can use to answer the user's question. You have access to a set of tools you can use to answer the user's question.
You can invoke functions by writing a "<{dsml_token}function_calls>" block like the following as part of your reply to the user: You can invoke functions by writing a "<{dsml_token}function_calls>" block like the following as part of your reply to the user:
@@ -148,11 +153,12 @@ def find_last_user_index(messages: List[Dict[str, Any]]) -> int:
def render_message( def render_message(
index: int, messages: List[Dict[str, Any]], thinking_mode: str index: int, messages: List[Dict[str, Any]], thinking_mode: str
) -> str: ) -> str:
assert 0 <= index < len(messages) if not (0 <= index < len(messages)):
assert thinking_mode in [ raise DS32EncodingError(
"chat", f"Index {index} out of range for messages list of length {len(messages)}"
"thinking", )
], f"Invalid thinking_mode `{thinking_mode}`" if thinking_mode not in ["chat", "thinking"]:
raise DS32EncodingError(f"Invalid thinking_mode `{thinking_mode}`")
prompt = "" prompt = ""
msg = messages[index] msg = messages[index]
@@ -181,7 +187,8 @@ def render_message(
) )
elif role == "developer": elif role == "developer":
assert content, f"Invalid message for role `{role}`: {msg}" if not content:
raise DS32EncodingError(f"Invalid message for role `{role}`: {msg}")
content_developer = "" content_developer = ""
if tools: if tools:
content_developer += "\n\n" + render_tools(tools) content_developer += "\n\n" + render_tools(tools)
@@ -214,17 +221,16 @@ def render_message(
prev_assistant_idx -= 1 prev_assistant_idx -= 1
assistant_msg = messages[prev_assistant_idx] assistant_msg = messages[prev_assistant_idx]
assert ( if not (
index == 0 index == 0
or prev_assistant_idx >= 0 or (prev_assistant_idx >= 0 and assistant_msg.get("role") == "assistant")
and assistant_msg.get("role") == "assistant" ):
), f"Invalid messages at {index}:\n{assistant_msg}" raise DS32EncodingError(f"Invalid messages at {index}:\n{assistant_msg}")
tool_call_order = index - prev_assistant_idx tool_call_order = index - prev_assistant_idx
assistant_tool_calls = assistant_msg.get("tool_calls") assistant_tool_calls = assistant_msg.get("tool_calls")
assert ( if not (assistant_tool_calls and len(assistant_tool_calls) >= tool_call_order):
assistant_tool_calls and len(assistant_tool_calls) >= tool_call_order raise DS32EncodingError("No tool calls but found tool output")
), "No tool calls but found tool output"
if tool_call_order == 1: if tool_call_order == 1:
prompt += "\n\n<function_results>" prompt += "\n\n<function_results>"
@@ -260,9 +266,10 @@ def render_message(
summary_content = content or "" summary_content = content or ""
if thinking_mode == "thinking" and index > last_user_idx: if thinking_mode == "thinking" and index > last_user_idx:
assert ( if not (reasoning_content or tool_calls):
reasoning_content or tool_calls raise DS32EncodingError(
), f"ThinkingMode: {thinking_mode}, invalid message without reasoning_content/tool_calls `{msg}` after last user message" f"ThinkingMode: {thinking_mode}, invalid message without reasoning_content/tool_calls `{msg}` after last user message"
)
thinking_part = ( thinking_part = (
thinking_template.format(reasoning_content=reasoning_content or "") thinking_template.format(reasoning_content=reasoning_content or "")
+ thinking_end_token + thinking_end_token
@@ -352,12 +359,14 @@ def parse_tool_calls(index: int, text: str):
index, _, stop_token = _read_until_stop( index, _, stop_token = _read_until_stop(
index, text, [f"<{dsml_token}invoke", tool_calls_end_token] index, text, [f"<{dsml_token}invoke", tool_calls_end_token]
) )
assert _ == ">\n", "Tool call format error" if _ != ">\n":
raise DS32EncodingError("Tool call format error")
if stop_token == tool_calls_end_token: if stop_token == tool_calls_end_token:
break break
assert stop_token is not None, "Missing special token" if stop_token is None:
raise DS32EncodingError("Missing special token")
index, tool_name_content, stop_token = _read_until_stop( index, tool_name_content, stop_token = _read_until_stop(
index, text, [f"<{dsml_token}parameter", f"</{dsml_token}invoke"] index, text, [f"<{dsml_token}parameter", f"</{dsml_token}invoke"]
@@ -366,7 +375,8 @@ def parse_tool_calls(index: int, text: str):
p_tool_name = re.findall( p_tool_name = re.findall(
r'^\s*name="(.*?)">\n$', tool_name_content, flags=re.DOTALL r'^\s*name="(.*?)">\n$', tool_name_content, flags=re.DOTALL
) )
assert len(p_tool_name) == 1, "Tool name format error" if len(p_tool_name) != 1:
raise DS32EncodingError("Tool name format error")
tool_name = p_tool_name[0] tool_name = p_tool_name[0]
tool_args: Dict[str, Tuple[str, str]] = {} tool_args: Dict[str, Tuple[str, str]] = {}
@@ -380,16 +390,19 @@ def parse_tool_calls(index: int, text: str):
param_content, param_content,
flags=re.DOTALL, flags=re.DOTALL,
) )
assert len(param_kv) == 1, "Parameter format error" if len(param_kv) != 1:
raise DS32EncodingError("Parameter format error")
param_name, string, param_value = param_kv[0] param_name, string, param_value = param_kv[0]
assert param_name not in tool_args, "Duplicate parameter name" if param_name in tool_args:
raise DS32EncodingError("Duplicate parameter name")
tool_args[param_name] = (param_value, string) tool_args[param_name] = (param_value, string)
index, content, stop_token = _read_until_stop( index, content, stop_token = _read_until_stop(
index, text, [f"<{dsml_token}parameter", f"</{dsml_token}invoke"] index, text, [f"<{dsml_token}parameter", f"</{dsml_token}invoke"]
) )
assert content == ">\n", "Parameter format error" if content != ">\n":
raise DS32EncodingError("Parameter format error")
tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args) tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args)
tool_calls.append(tool_call) tool_calls.append(tool_call)
@@ -410,7 +423,8 @@ def parse_message_from_completion_text(text: str, thinking_mode: str):
index, text, [thinking_end_token, tool_calls_start_token] index, text, [thinking_end_token, tool_calls_start_token]
) )
reasoning_content = content_delta reasoning_content = content_delta
assert stop_token == thinking_end_token, "Invalid thinking format" if stop_token != thinking_end_token:
raise DS32EncodingError("Invalid thinking format")
index, content_delta, stop_token = _read_until_stop( index, content_delta, stop_token = _read_until_stop(
index, text, [eos_token, tool_calls_start_token] index, text, [eos_token, tool_calls_start_token]
@@ -419,18 +433,18 @@ def parse_message_from_completion_text(text: str, thinking_mode: str):
if stop_token == tool_calls_start_token: if stop_token == tool_calls_start_token:
is_tool_calling = True is_tool_calling = True
else: else:
assert stop_token == eos_token, "Invalid summary format" if stop_token != eos_token:
raise DS32EncodingError("Invalid summary format")
if is_tool_calling: if is_tool_calling:
index, stop_token, tool_calls = parse_tool_calls(index, text) index, stop_token, tool_calls = parse_tool_calls(index, text)
index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token]) index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token])
assert not tool_ends_text, "Unexpected content after tool calls" if tool_ends_text:
raise DS32EncodingError("Unexpected content after tool calls")
assert len(text) == index and stop_token in [ if not (len(text) == index and stop_token in [eos_token, None]):
eos_token, raise DS32EncodingError("Unexpected content at end")
None,
], "Unexpected content at end"
for sp_token in [ for sp_token in [
bos_token, bos_token,
@@ -439,9 +453,8 @@ def parse_message_from_completion_text(text: str, thinking_mode: str):
thinking_end_token, thinking_end_token,
dsml_token, dsml_token,
]: ]:
assert ( if sp_token in summary_content or sp_token in reasoning_content:
sp_token not in summary_content and sp_token not in reasoning_content raise DS32EncodingError("Unexpected special token in content")
), "Unexpected special token in content"
return { return {
"role": "assistant", "role": "assistant",
@@ -11,6 +11,7 @@ import orjson
from fastapi import HTTPException, Request from fastapi import HTTPException, Request
from fastapi.responses import ORJSONResponse, StreamingResponse from fastapi.responses import ORJSONResponse, StreamingResponse
from sglang.srt.entrypoints.openai.encoding_dsv32 import DS32EncodingError
from sglang.srt.entrypoints.openai.protocol import ErrorResponse, OpenAIServingRequest from sglang.srt.entrypoints.openai.protocol import ErrorResponse, OpenAIServingRequest
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -129,6 +130,13 @@ class OpenAIServingBase(ABC):
err_type="BadRequest", err_type="BadRequest",
status_code=400, status_code=400,
) )
except DS32EncodingError as e:
logger.info(f"DS32EncodingError: {e}")
return self.create_error_response(
message=str(e),
err_type="BadRequest",
status_code=400,
)
except Exception as e: except Exception as e:
logger.exception(f"Error in request: {e}") logger.exception(f"Error in request: {e}")
return self.create_error_response( return self.create_error_response(