Fix invalid escape warnings in tool parsers (#28370)

Co-authored-by: FAN YUCHEN <2994114386@qq.com>
This commit is contained in:
Xinyuan Tong
2026-07-28 23:31:47 +08:00
committed by GitHub
co-authored by FAN YUCHEN
parent 84cdfde5b2
commit ee236086db
12 changed files with 218 additions and 23 deletions
@@ -1,4 +1,3 @@
import ast
import json import json
import logging import logging
import re import re
@@ -17,7 +16,10 @@ from sglang.srt.function_call.core_types import (
ToolCallItem, ToolCallItem,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import infer_type_from_json_schema from sglang.srt.function_call.utils import (
infer_type_from_json_schema,
safe_literal_eval,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -147,9 +149,21 @@ def parse_arguments(
except (json.JSONDecodeError, ValueError, KeyError): except (json.JSONDecodeError, ValueError, KeyError):
pass pass
# Strategy 2.5: string-typed values that are not valid JSON (S1/S2 failed) —
# strip the wrapping quotes and keep the raw bytes, backslashes included.
# Avoids ast.literal_eval so invalid escapes neither warn nor get reinterpreted.
if arg_type == "string":
if (
len(json_value) >= 2
and json_value[0] == json_value[-1]
and json_value[0] in {'"', "'"}
):
return json_value[1:-1], True
return json_value, True
# Strategy 3: ast.literal_eval # Strategy 3: ast.literal_eval
try: try:
parsed_value = ast.literal_eval(json_value) parsed_value = safe_literal_eval(json_value)
return parsed_value, True return parsed_value, True
except (ValueError, SyntaxError): except (ValueError, SyntaxError):
pass pass
@@ -1,4 +1,3 @@
import ast
import json import json
import logging import logging
import re import re
@@ -12,7 +11,10 @@ from sglang.srt.function_call.core_types import (
ToolCallItem, ToolCallItem,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import infer_type_from_json_schema from sglang.srt.function_call.utils import (
infer_type_from_json_schema,
safe_literal_eval,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -116,9 +118,21 @@ def parse_arguments(
except (json.JSONDecodeError, ValueError, KeyError): except (json.JSONDecodeError, ValueError, KeyError):
pass pass
# Strategy 2.5: string-typed values that are not valid JSON (S1/S2 failed) —
# strip the wrapping quotes and keep the raw bytes, backslashes included.
# Avoids ast.literal_eval so invalid escapes neither warn nor get reinterpreted.
if arg_type == "string":
if (
len(json_value) >= 2
and json_value[0] == json_value[-1]
and json_value[0] in {'"', "'"}
):
return json_value[1:-1], True
return json_value, True
# Strategy 3: ast.literal_eval # Strategy 3: ast.literal_eval
try: try:
parsed_value = ast.literal_eval(json_value) parsed_value = safe_literal_eval(json_value)
return parsed_value, True return parsed_value, True
except (ValueError, SyntaxError): except (ValueError, SyntaxError):
pass pass
@@ -32,6 +32,7 @@ from sglang.srt.function_call.core_types import (
ToolCallItem, ToolCallItem,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import safe_ast_parse
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -172,7 +173,7 @@ class Lfm2Detector(BaseFormatDetector):
tool_indices = self._get_tool_indices(tools) tool_indices = self._get_tool_indices(tools)
try: try:
module = ast.parse(content) module = safe_ast_parse(content)
parsed = getattr(module.body[0], "value", None) if module.body else None parsed = getattr(module.body[0], "value", None) if module.body else None
if parsed is None: if parsed is None:
@@ -1,4 +1,3 @@
import ast
import json import json
import logging import logging
import re import re
@@ -11,6 +10,7 @@ from sglang.srt.function_call.core_types import (
StructureInfo, StructureInfo,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import safe_literal_eval
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -37,7 +37,7 @@ class Llama32Detector(BaseFormatDetector):
def _convert_python_dict_to_json(self, text: str) -> str: def _convert_python_dict_to_json(self, text: str) -> str:
"""Convert Python dict strings to JSON format.""" """Convert Python dict strings to JSON format."""
try: try:
parsed = ast.literal_eval(text.strip()) parsed = safe_literal_eval(text.strip())
if isinstance(parsed, dict): if isinstance(parsed, dict):
return json.dumps(parsed, ensure_ascii=False) return json.dumps(parsed, ensure_ascii=False)
except: except:
@@ -12,7 +12,6 @@
# limitations under the License. # limitations under the License.
# ============================================================================== # ==============================================================================
import ast
import html import html
import json import json
import logging import logging
@@ -23,6 +22,7 @@ from sglang.srt.entrypoints.openai.protocol import Tool
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.function_call.base_format_detector import BaseFormatDetector from sglang.srt.function_call.base_format_detector import BaseFormatDetector
from sglang.srt.function_call.core_types import StreamingParseResult, _GetInfoFunc from sglang.srt.function_call.core_types import StreamingParseResult, _GetInfoFunc
from sglang.srt.function_call.utils import safe_literal_eval
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -121,7 +121,7 @@ def _convert_param_value(
func_name, func_name,
) )
try: try:
param_value = ast.literal_eval(param_value) # safer param_value = safe_literal_eval(param_value)
except (ValueError, SyntaxError, TypeError): except (ValueError, SyntaxError, TypeError):
logger.warning( logger.warning(
"Parsed value '%s' of parameter '%s' cannot be " "Parsed value '%s' of parameter '%s' cannot be "
@@ -1,4 +1,3 @@
import ast
import json import json
import logging import logging
import re import re
@@ -10,6 +9,7 @@ from sglang.srt.function_call.core_types import (
StreamingParseResult, StreamingParseResult,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import safe_literal_eval
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -46,7 +46,7 @@ def parse_arguments(json_value):
try: try:
parsed_value = json.loads(json_value) parsed_value = json.loads(json_value)
except (json.JSONDecodeError, TypeError): except (json.JSONDecodeError, TypeError):
parsed_value = ast.literal_eval(json_value) parsed_value = safe_literal_eval(json_value)
return parsed_value, True return parsed_value, True
except (ValueError, SyntaxError, TypeError): except (ValueError, SyntaxError, TypeError):
return json_value, False return json_value, False
@@ -1,4 +1,3 @@
import ast
import json import json
import re import re
from enum import Enum, auto from enum import Enum, auto
@@ -12,6 +11,7 @@ from sglang.srt.function_call.core_types import (
ToolCallItem, ToolCallItem,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import safe_literal_eval
class _ParseState(Enum): class _ParseState(Enum):
@@ -50,7 +50,7 @@ class PoolsideV1Detector(BaseFormatDetector):
String values are emitted as raw text; non-strings are JSON-encoded by String values are emitted as raw text; non-strings are JSON-encoded by
the chat template. The parser does schema-based type coercion to round-trip the chat template. The parser does schema-based type coercion to round-trip
them: schema type `string` keeps the raw value; other types attempt them: schema type `string` keeps the raw value; other types attempt
`json.loads` and fall back to `ast.literal_eval`, then to the raw string. `json.loads` and fall back to `safe_literal_eval`, then to the raw string.
""" """
# Wire-format tag tokens — constants, not per-instance. # Wire-format tag tokens — constants, not per-instance.
@@ -166,7 +166,7 @@ class PoolsideV1Detector(BaseFormatDetector):
- no schema entry → json.loads only (conservative; don't - no schema entry → json.loads only (conservative; don't
ast-eval untyped values) ast-eval untyped values)
- everything else (int, - everything else (int,
number, bool, object, …) → json.loads, then ast.literal_eval number, bool, object, …) → json.loads, then safe_literal_eval
Each decoder result is round-tripped through `json.dumps` before being Each decoder result is round-tripped through `json.dumps` before being
returned; non-JSON-serializable values (sets / complex / bytes from returned; non-JSON-serializable values (sets / complex / bytes from
@@ -179,7 +179,7 @@ class PoolsideV1Detector(BaseFormatDetector):
if param_type in PoolsideV1Detector._STRING_TYPES: if param_type in PoolsideV1Detector._STRING_TYPES:
return raw return raw
decoders = (json.loads,) if not param_type else (json.loads, ast.literal_eval) decoders = (json.loads,) if not param_type else (json.loads, safe_literal_eval)
for decoder in decoders: for decoder in decoders:
try: try:
result = decoder(raw) result = decoder(raw)
@@ -12,6 +12,7 @@ from sglang.srt.function_call.core_types import (
ToolCallItem, ToolCallItem,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import safe_ast_parse
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -72,7 +73,7 @@ class PythonicDetector(BaseFormatDetector):
normal_text = normal_text_before + normal_text_after normal_text = normal_text_before + normal_text_after
try: try:
module = ast.parse(tool_call_text) module = safe_ast_parse(tool_call_text)
parsed = getattr(module.body[0], "value", None) parsed = getattr(module.body[0], "value", None)
if not ( if not (
isinstance(parsed, ast.List) isinstance(parsed, ast.List)
@@ -1,4 +1,3 @@
import ast
import json import json
import logging import logging
import re import re
@@ -11,7 +10,10 @@ from sglang.srt.function_call.core_types import (
ToolCallItem, ToolCallItem,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import infer_type_from_json_schema from sglang.srt.function_call.utils import (
infer_type_from_json_schema,
safe_literal_eval,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -164,7 +166,7 @@ class Qwen3CoderDetector(BaseFormatDetector):
f"'{func_name}', will try other methods to parse it." f"'{func_name}', will try other methods to parse it."
) )
try: try:
param_value = ast.literal_eval(param_value) # safer param_value = safe_literal_eval(param_value)
except Exception: except Exception:
logger.warning( logger.warning(
f"Parsed value '{param_value}' of parameter '{param_name}' cannot be converted via Python `ast.literal_eval()` in tool '{func_name}', degenerating to string." f"Parsed value '{param_value}' of parameter '{param_name}' cannot be converted via Python `ast.literal_eval()` in tool '{func_name}', degenerating to string."
@@ -1,4 +1,3 @@
import ast
import json import json
import logging import logging
import re import re
@@ -11,6 +10,7 @@ from sglang.srt.function_call.core_types import (
ToolCallItem, ToolCallItem,
_GetInfoFunc, _GetInfoFunc,
) )
from sglang.srt.function_call.utils import safe_literal_eval
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -34,7 +34,7 @@ def parse_arguments(value: str) -> tuple[Any, bool]:
try: try:
parsed_value = json.loads(value) parsed_value = json.loads(value)
except: except:
parsed_value = ast.literal_eval(value) parsed_value = safe_literal_eval(value)
return parsed_value, True return parsed_value, True
except: except:
return value, False return value, False
+32
View File
@@ -1,3 +1,6 @@
import ast
import threading
import warnings
from json import JSONDecodeError, JSONDecoder from json import JSONDecodeError, JSONDecoder
from json.decoder import WHITESPACE from json.decoder import WHITESPACE
from typing import Any, Dict, List, Literal, Optional, Tuple, Union from typing import Any, Dict, List, Literal, Optional, Tuple, Union
@@ -228,6 +231,35 @@ def _is_complete_json(input_str: str) -> bool:
return False return False
# ``warnings.catch_warnings`` mutates the *process-global* warning filters and
# is therefore not thread-safe (CPython docs). Tool-call parsing runs on the
# request path and may execute concurrently, so the enter/eval/restore window
# is serialized. These helpers are microsecond-cheap; the lock has no perf impact.
_safe_ast_lock = threading.Lock()
def _run_ast_quiet(fn, *args):
"""Run an ``ast`` function with invalid-escape warnings suppressed.
CPython parses invalid escapes (e.g. ``"\\d+"``) with the backslash kept
and only emits a warning, so the parsed value is already correct —
promoting the warning to an error would drop otherwise-valid tool calls.
Holds ``_safe_ast_lock`` because ``catch_warnings`` touches global state."""
with _safe_ast_lock, warnings.catch_warnings():
warnings.filterwarnings("ignore", category=SyntaxWarning)
warnings.filterwarnings("ignore", category=DeprecationWarning)
return fn(*args)
def safe_literal_eval(value: str) -> Any:
return _run_ast_quiet(ast.literal_eval, value)
def safe_ast_parse(source: str) -> ast.Module:
return _run_ast_quiet(ast.parse, source)
def _get_tool_schema_defs(tools: List[Tool]) -> dict: def _get_tool_schema_defs(tools: List[Tool]) -> dict:
""" """
Get consolidated $defs from all tools, validating for conflicts. Get consolidated $defs from all tools, validating for conflicts.
@@ -1,5 +1,6 @@
import json import json
import unittest import unittest
import warnings
from sglang.srt.entrypoints.openai.protocol import ( from sglang.srt.entrypoints.openai.protocol import (
Function, Function,
@@ -391,6 +392,26 @@ class TestPythonicDetector(unittest.TestCase):
self.assertTrue(self.detector.has_tool_call('[get_weather(location="Tokyo")]')) self.assertTrue(self.detector.has_tool_call('[get_weather(location="Tokyo")]'))
self.assertFalse(self.detector.has_tool_call("plain text only")) self.assertFalse(self.detector.has_tool_call("plain text only"))
def test_invalid_escape_sequence_still_parses(self):
"""An invalid Python escape (e.g. "\\d+") must not drop the tool call.
CPython keeps the backslash and only warns; if the warning were
promoted to an error the whole call would fall out as normal text."""
text = '[search(query="\\d+")]'
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always", SyntaxWarning)
result = self.detector.detect_and_parse(text, self.tools)
self.assertEqual(len(result.calls), 1)
self.assertEqual(result.calls[0].name, "search")
params = json.loads(result.calls[0].parameters)
self.assertEqual(params["query"], "\\d+")
self.assertEqual(result.normal_text, "")
self.assertFalse(
any(isinstance(w.message, SyntaxWarning) for w in caught),
[str(w.message) for w in caught],
)
def test_parse_streaming_no_brackets(self): def test_parse_streaming_no_brackets(self):
"""Test parsing text with no brackets (no tool calls).""" """Test parsing text with no brackets (no tool calls)."""
text = "This is just normal text without any tool calls." text = "This is just normal text without any tool calls."
@@ -3045,6 +3066,55 @@ class TestGlm4MoeDetector(unittest.TestCase):
self.assertEqual(params["old_string"], " indented code") self.assertEqual(params["old_string"], " indented code")
self.assertEqual(params["new_string"], " also indented") self.assertEqual(params["new_string"], " also indented")
def test_quoted_string_invalid_python_escape_no_warning(self):
text = (
'<tool_call>get_weather\n<arg_key>city</arg_key>\n<arg_value>"\\C|\\."</arg_value>\n'
"<arg_key>date</arg_key>\n<arg_value>2024-06-27</arg_value>\n</tool_call>"
)
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always", SyntaxWarning)
result = self.detector.detect_and_parse(text, self.tools)
self.assertEqual(len(result.calls), 1)
params = json.loads(result.calls[0].parameters)
self.assertEqual(params["city"], r"\C|\.")
self.assertFalse(
any(isinstance(w.message, SyntaxWarning) for w in caught),
[str(w.message) for w in caught],
)
def test_parse_arguments_preserves_underscore_in_string_args(self):
"""PEP 515 makes ast.literal_eval strip underscores ("123_456"->123456);
a string-typed arg must keep the raw value. See #30644."""
from sglang.srt.function_call.glm4_moe_detector import parse_arguments
value, is_good = parse_arguments("123_456", arg_type="string")
self.assertTrue(is_good)
self.assertIsInstance(value, str)
self.assertEqual(value, "123_456")
value, is_good = parse_arguments("1_000.5", arg_type="string")
self.assertTrue(is_good)
self.assertIsInstance(value, str)
self.assertEqual(value, "1_000.5")
value, is_good = parse_arguments("123_456")
self.assertTrue(is_good)
self.assertIsInstance(value, int)
self.assertEqual(value, 123456)
def test_parse_arguments_object_with_invalid_escape(self):
"""A dict arg containing an invalid escape ("\\d+") must stay a dict.
If safe_literal_eval raised on the escape warning, Strategy 3 would
fail and Strategy 4 would degrade the whole value to one string."""
from sglang.srt.function_call.glm4_moe_detector import parse_arguments
value, is_good = parse_arguments("{'pattern': '\\d+'}", arg_type="object")
self.assertTrue(is_good)
self.assertIsInstance(value, dict)
self.assertEqual(value, {"pattern": "\\d+"})
class TestGlm47MoeDetector(unittest.TestCase): class TestGlm47MoeDetector(unittest.TestCase):
def setUp(self): def setUp(self):
@@ -3310,6 +3380,51 @@ class TestGlm47MoeDetector(unittest.TestCase):
self.assertEqual(params["old_string"], " indented code") self.assertEqual(params["old_string"], " indented code")
self.assertEqual(params["new_string"], " also indented") self.assertEqual(params["new_string"], " also indented")
def test_quoted_string_invalid_python_escape_no_warning(self):
text = (
'<tool_call>get_weather<arg_key>city</arg_key><arg_value>"\\C|\\."</arg_value>'
"<arg_key>date</arg_key><arg_value>2024-06-27</arg_value></tool_call>"
)
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always", SyntaxWarning)
result = self.detector.detect_and_parse(text, self.tools)
self.assertEqual(len(result.calls), 1)
params = json.loads(result.calls[0].parameters)
self.assertEqual(params["city"], r"\C|\.")
self.assertFalse(
any(isinstance(w.message, SyntaxWarning) for w in caught),
[str(w.message) for w in caught],
)
def test_parse_arguments_preserves_underscore_in_string_args(self):
"""Same PEP 515 guard as the GLM-4 detector, on the GLM-4.7 parser."""
from sglang.srt.function_call.glm47_moe_detector import parse_arguments
value, is_good = parse_arguments("123_456", arg_type="string")
self.assertTrue(is_good)
self.assertIsInstance(value, str)
self.assertEqual(value, "123_456")
value, is_good = parse_arguments("1_000.5", arg_type="string")
self.assertTrue(is_good)
self.assertIsInstance(value, str)
self.assertEqual(value, "1_000.5")
value, is_good = parse_arguments("123_456")
self.assertTrue(is_good)
self.assertIsInstance(value, int)
self.assertEqual(value, 123456)
def test_parse_arguments_object_with_invalid_escape(self):
"""Same object-arg escape guard as the GLM-4 detector."""
from sglang.srt.function_call.glm47_moe_detector import parse_arguments
value, is_good = parse_arguments("{'pattern': '\\d+'}", arg_type="object")
self.assertTrue(is_good)
self.assertIsInstance(value, dict)
self.assertEqual(value, {"pattern": "\\d+"})
def test_get_model_structural_tag(self): def test_get_model_structural_tag(self):
"""GLM-4.7/GLM-5 use xgrammar's native "glm_4_7" structural tag.""" """GLM-4.7/GLM-5 use xgrammar's native "glm_4_7" structural tag."""
import xgrammar as xgr import xgrammar as xgr
@@ -3616,6 +3731,22 @@ class TestLfm2Detector(unittest.TestCase):
params = json.loads(result.calls[0].parameters) params = json.loads(result.calls[0].parameters)
self.assertEqual(params["city"], "Paris") self.assertEqual(params["city"], "Paris")
def test_detect_and_parse_pythonic_invalid_escape(self):
"""An invalid Python escape (e.g. "\\d+") must not drop the tool call."""
text = '<|tool_call_start|>[search(query="\\d+")]<|tool_call_end|>'
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always", SyntaxWarning)
result = self.detector.detect_and_parse(text, self.tools)
self.assertEqual(len(result.calls), 1)
self.assertEqual(result.calls[0].name, "search")
params = json.loads(result.calls[0].parameters)
self.assertEqual(params["query"], "\\d+")
self.assertFalse(
any(isinstance(w.message, SyntaxWarning) for w in caught),
[str(w.message) for w in caught],
)
def test_detect_and_parse_pythonic_multiple_args(self): def test_detect_and_parse_pythonic_multiple_args(self):
"""Test parsing with multiple arguments.""" """Test parsing with multiple arguments."""
text = '<|tool_call_start|>[get_weather(city="London", unit="celsius")]<|tool_call_end|>' text = '<|tool_call_start|>[get_weather(city="London", unit="celsius")]<|tool_call_end|>'