Fix invalid escape warnings in tool parsers (#28370)
Co-authored-by: FAN YUCHEN <2994114386@qq.com>
This commit is contained in:
co-authored by
FAN YUCHEN
parent
84cdfde5b2
commit
ee236086db
@@ -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
|
||||||
|
|||||||
@@ -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|>'
|
||||||
|
|||||||
Reference in New Issue
Block a user