From 265202cda251bd466930f132ce7ca3b45cf33697 Mon Sep 17 00:00:00 2001 From: xianzhiT Date: Tue, 16 Jun 2026 20:21:23 +0800 Subject: [PATCH] fix(openai): validate assistant tool call arguments before chat template (#28035) Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> Co-authored-by: Xinyuan Tong --- .../srt/entrypoints/openai/encoding_dsv32.py | 9 +- .../srt/entrypoints/openai/encoding_dsv4.py | 14 +- .../srt/entrypoints/openai/serving_chat.py | 65 +++++--- .../entrypoints/openai/test_serving_chat.py | 151 ++++++++++++++++++ 4 files changed, 209 insertions(+), 30 deletions(-) diff --git a/python/sglang/srt/entrypoints/openai/encoding_dsv32.py b/python/sglang/srt/entrypoints/openai/encoding_dsv32.py index 8840173cb..f20f07451 100644 --- a/python/sglang/srt/entrypoints/openai/encoding_dsv32.py +++ b/python/sglang/srt/entrypoints/openai/encoding_dsv32.py @@ -97,7 +97,14 @@ def encode_arguments_to_dsml(tool_call: Dict[str, str]) -> str: p_dsml_template = """<{dsml_token}parameter name="{key}" string="{is_str}">{value}""" P_dsml_strs = [] - arguments = json.loads(tool_call["arguments"]) + raw_arguments = tool_call["arguments"] + arguments = ( + json.loads(raw_arguments) if isinstance(raw_arguments, str) else raw_arguments + ) + if not isinstance(arguments, dict): + raise ValueError( + "Assistant tool call function.arguments must be a JSON object." + ) for k, v in arguments.items(): p_dsml_str = p_dsml_template.format( diff --git a/python/sglang/srt/entrypoints/openai/encoding_dsv4.py b/python/sglang/srt/entrypoints/openai/encoding_dsv4.py index 0822c9ad1..a19d1f4b1 100644 --- a/python/sglang/srt/entrypoints/openai/encoding_dsv4.py +++ b/python/sglang/srt/entrypoints/openai/encoding_dsv4.py @@ -141,7 +141,7 @@ def encode_arguments_to_dsml(tool_call: Dict[str, str]) -> str: Encode tool call arguments into DSML parameter format. Args: - tool_call: Dict with "name" and "arguments" (JSON string) keys. + tool_call: Dict with "name" and "arguments" keys. Returns: DSML-formatted parameter string. @@ -149,10 +149,14 @@ def encode_arguments_to_dsml(tool_call: Dict[str, str]) -> str: p_dsml_template = '<{dsml_token}parameter name="{key}" string="{is_str}">{value}' P_dsml_strs = [] - try: - arguments = json.loads(tool_call["arguments"]) - except Exception as err: - arguments = {"arguments": tool_call["arguments"]} + raw_arguments = tool_call["arguments"] + arguments = ( + json.loads(raw_arguments) if isinstance(raw_arguments, str) else raw_arguments + ) + if not isinstance(arguments, dict): + raise ValueError( + "Assistant tool call function.arguments must be a JSON object." + ) for k, v in arguments.items(): p_dsml_str = p_dsml_template.format( diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 3ac1301cd..48d43831e 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -98,6 +98,38 @@ def normalize_tool_content(role: str, content): return content +def parse_tool_call_arguments(arguments: str) -> Dict[str, Any]: + """Parse OpenAI tool call arguments for chat templates.""" + try: + parsed_arguments = orjson.loads(arguments) + except orjson.JSONDecodeError as exc: + raise ValueError( + "Assistant tool call function.arguments must be valid JSON." + ) from exc + + if not isinstance(parsed_arguments, dict): + raise ValueError( + "Assistant tool call function.arguments must be a JSON object." + ) + + return parsed_arguments + + +def normalize_assistant_tool_call_arguments(message: Dict[str, Any]) -> None: + """Normalize assistant history tool call arguments in-place.""" + if message.get("role") != "assistant" or not isinstance( + message.get("tool_calls"), list + ): + return + + for item in message["tool_calls"]: + function = item.get("function") if isinstance(item, dict) else None + if not isinstance(function, dict): + continue + if "arguments" in function and isinstance(function["arguments"], str): + function["arguments"] = parse_tool_call_arguments(function["arguments"]) + + def _extract_max_dynamic_patch(request: ChatCompletionRequest): img_vals = [] vid_vals = [] @@ -651,8 +683,12 @@ class OpenAIServingChat(OpenAIServingBase): thinking_mode = ( ThinkingMode.THINKING if thinking_requested else ThinkingMode.CHAT ) + messages = [msg.model_dump() for msg in request.messages] + for message in messages: + normalize_assistant_tool_call_arguments(message) + prompt_ids = self._encode_messages( - [msg.model_dump() for msg in request.messages], request, thinking_mode + copy.deepcopy(messages), request, thinking_mode ) if prompt_ids is not None: @@ -660,7 +696,7 @@ class OpenAIServingChat(OpenAIServingBase): pass elif self.chat_encoding_spec is not None: # dsv4/dsv32 encoding path - messages = [msg.model_dump() for msg in request.messages] + messages = copy.deepcopy(messages) # dsv4/dsv32 are text-only and consume string content; flatten # OpenAI parts-list content here so the encoder sees a plain string. @@ -730,10 +766,9 @@ class OpenAIServingChat(OpenAIServingBase): prompt_ids, assistant_prefix ) else: - for message in request.messages: - if message.content is None: - message.content = "" - msg_dict = message.model_dump() + for msg_dict in copy.deepcopy(messages): + if msg_dict.get("content") is None: + msg_dict["content"] = "" # Process content based on detected template format processed_msg = process_content_for_template_format( @@ -749,24 +784,6 @@ class OpenAIServingChat(OpenAIServingBase): processed_msg["role"], processed_msg.get("content") ) - # per the Transformers docs & maintainers, tool call arguments in - # assistant-role messages with tool_calls need to be dicts not JSON str - - # this is how tool-use chat templates will expect them moving forwards - # so, for messages that have tool_calls, parse the string (which we get - # from openAI format) to dict - if ( - processed_msg["role"] == "assistant" - and "tool_calls" in processed_msg - and isinstance(processed_msg["tool_calls"], list) - ): - for item in processed_msg["tool_calls"]: - if "arguments" in item["function"] and isinstance( - item["function"]["arguments"], str - ): - item["function"]["arguments"] = orjson.loads( - item["function"]["arguments"] - ) - openai_compatible_messages.append(processed_msg) # Handle continue_final_message: separate final assistant message diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 33df9a611..3ba70cc9f 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -558,6 +558,157 @@ class ServingChatTestCase(unittest.TestCase): parser.get_structure_constraint.call_args.kwargs["thinking_mode"] ) + 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" + + 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" + + 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.