From ebe3a651e061aa2c0e88e1df8f0eac08a5832bd2 Mon Sep 17 00:00:00 2001 From: Xeophon <46377542+xeophon@users.noreply.github.com> Date: Wed, 26 Aug 2026 12:33:16 +0200 Subject: [PATCH 1/7] fix: support root JSON Schema combinators in tool parsers --- .../srt/function_call/function_call_parser.py | 20 +- .../srt/function_call/glm47_moe_detector.py | 2 +- .../srt/function_call/glm4_moe_detector.py | 2 +- python/sglang/srt/function_call/utils.py | 25 ++ .../test_root_combinator_tool_parsers.py | 244 ++++++++++++++++++ 5 files changed, 288 insertions(+), 5 deletions(-) create mode 100644 test/registered/unit/function_call/test_root_combinator_tool_parsers.py diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 535aaa8b183f..f3dc48ac5f15 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -49,6 +49,7 @@ from sglang.srt.function_call.utils import ( _get_tool_schema_defs, get_json_schema_constraint, + get_json_schema_properties, ) logger = logging.getLogger(__name__) @@ -116,6 +117,17 @@ def __init__(self, tools: List[Tool], tool_call_parser: str, tokenizer=None): self.detector = detector self.tools = tools + self.detector_tools: List[Tool] = [] + for tool in tools: + parameters = tool.function.parameters + if isinstance(parameters, dict): + properties = get_json_schema_properties(parameters) + if properties != parameters.get("properties", {}): + function = tool.function.model_copy( + update={"parameters": parameters | {"properties": properties}} + ) + tool = tool.model_copy(update={"function": function}) + self.detector_tools.append(tool) self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get() def has_tool_call(self, text: str) -> bool: @@ -148,7 +160,7 @@ def parse_non_stream(self, full_text: str) -> Tuple[str, list[ToolCallItem]]: if not self.tools: return full_text, [] has_tool_call = self.detector.has_tool_call(full_text) - parsed_result = self.detector.detect_and_parse(full_text, self.tools) + parsed_result = self.detector.detect_and_parse(full_text, self.detector_tools) tool_call_list = parsed_result.calls if tool_call_list or has_tool_call: return parsed_result.normal_text, tool_call_list @@ -172,7 +184,9 @@ def parse_stream_chunk(self, chunk_text: str) -> Tuple[str, list[ToolCallItem]]: final_normal_text = "" final_calls = [] - sp_result = self.detector.parse_streaming_increment(chunk_text, self.tools) + sp_result = self.detector.parse_streaming_increment( + chunk_text, self.detector_tools + ) if sp_result.normal_text: final_normal_text = sp_result.normal_text if sp_result.calls: @@ -189,7 +203,7 @@ def parse_stream_end(self) -> Tuple[str, list[ToolCallItem]]: """ if not self.tools: return "", [] - sp_result = self.detector.finish(self.tools) + sp_result = self.detector.finish(self.detector_tools) return sp_result.normal_text, sp_result.calls def get_legacy_structural_tag( diff --git a/python/sglang/srt/function_call/glm47_moe_detector.py b/python/sglang/srt/function_call/glm47_moe_detector.py index 90bdb2aedcd5..b01c316c5ef2 100644 --- a/python/sglang/srt/function_call/glm47_moe_detector.py +++ b/python/sglang/srt/function_call/glm47_moe_detector.py @@ -613,7 +613,7 @@ def _finalize_tool_call( self._last_arguments += "{}" self.streamed_args_for_tool[self.current_tool_id] += "{}" self._sent_empty_object = True - elif not self._last_arguments.endswith("}") and not self._sent_empty_object: + elif not self._sent_empty_object: # Need to close brace calls.append( ToolCallItem( diff --git a/python/sglang/srt/function_call/glm4_moe_detector.py b/python/sglang/srt/function_call/glm4_moe_detector.py index 0c29a39e7d24..85076316d77f 100644 --- a/python/sglang/srt/function_call/glm4_moe_detector.py +++ b/python/sglang/srt/function_call/glm4_moe_detector.py @@ -566,7 +566,7 @@ def parse_streaming_increment( ) ) self._last_arguments += empty_object - elif not self._last_arguments.endswith("}"): + else: closing_brace = "}" calls.append( ToolCallItem( diff --git a/python/sglang/srt/function_call/utils.py b/python/sglang/srt/function_call/utils.py index 0bd0bef4af79..1d82a1c3cc1d 100644 --- a/python/sglang/srt/function_call/utils.py +++ b/python/sglang/srt/function_call/utils.py @@ -304,6 +304,31 @@ def _get_tool_schema(tool: Tool) -> dict: } +def get_json_schema_properties(schema: Any) -> Dict[str, Any]: + """Collect properties declared directly or in root schema combinators.""" + if not isinstance(schema, dict): + return {} + + direct_properties = schema.get("properties") + if not isinstance(direct_properties, dict): + direct_properties = {} + properties = direct_properties.copy() + + for keyword in ("anyOf", "oneOf", "allOf"): + branches = schema.get(keyword) + if not isinstance(branches, list): + continue + for branch in branches: + for name, property_schema in get_json_schema_properties(branch).items(): + if name in direct_properties: + continue + if name in properties and properties[name] != property_schema: + properties[name] = {keyword: [properties[name], property_schema]} + else: + properties[name] = property_schema + return properties + + def infer_type_from_json_schema(schema: Dict[str, Any]) -> Optional[str]: """ Infer the primary type of a parameter from JSON Schema. diff --git a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py new file mode 100644 index 000000000000..05000ab52ece --- /dev/null +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -0,0 +1,244 @@ +import json +import unittest + +from sglang.srt.entrypoints.openai.protocol import Function, Tool +from sglang.srt.function_call.function_call_parser import FunctionCallParser +from sglang.srt.function_call.minimax_m3 import MINIMAX_NS_TOKEN +from sglang.srt.function_call.utils import get_json_schema_properties +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestRootCombinatorToolParsers(unittest.TestCase): + def setUp(self): + self.tools = [ + Tool( + type="function", + function=Function( + name="acme", + parameters={ + "type": "object", + "oneOf": [ + { + "type": "object", + "properties": { + "kind": {"const": "acme"}, + "payload": { + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + }, + "required": ["kind", "payload"], + }, + { + "type": "object", + "properties": {"kind": {"const": "other"}}, + "required": ["kind"], + }, + ], + }, + ), + ) + ] + self.expected = {"kind": "acme", "payload": {"value": "hello"}} + + def test_property_lookup_supports_root_combinators(self): + payload_schema = { + "type": "object", + "properties": {"value": {"type": "string"}}, + } + for keyword in ("anyOf", "oneOf", "allOf"): + with self.subTest(keyword=keyword): + schema = { + keyword: [ + { + "type": "object", + "properties": {"payload": payload_schema}, + } + ] + } + self.assertEqual( + get_json_schema_properties(schema), {"payload": payload_schema} + ) + + self.assertEqual( + get_json_schema_properties(self.tools[0].function.parameters)["kind"], + {"oneOf": [{"const": "acme"}, {"const": "other"}]}, + ) + + parser = FunctionCallParser(self.tools, "qwen3_coder") + self.assertNotIn("properties", self.tools[0].function.parameters) + self.assertIn( + "payload", parser.detector_tools[0].function.parameters["properties"] + ) + + def test_parsers(self): + ns = MINIMAX_NS_TOKEN + cases = [ + ( + "qwen3_coder", + [ + "", + "", + "acme", + '{"value":"hello"}', + "", + "", + ], + ), + ( + "glm", + [ + "acme\n", + "kind\nacme\n", + ( + 'payload\n{"value":"hello"}' + "\n" + ), + "", + ], + ), + ( + "glm47", + [ + "acme", + "kindacme", + ( + 'payload{"value":"hello"}' + "" + ), + "", + ], + ), + ( + "dots", + [ + '', + 'acme', + '{"value":"hello"}', + "", + ], + ), + ( + "hunyuan", + [ + "acme", + "kindacme", + ( + 'payload{"value":"hello"}' + "" + ), + "", + ], + ), + ( + "mimo", + [ + "", + "acme", + '{"value":"hello"}', + "", + ], + ), + ( + "minicpm5", + [ + '', + 'acme', + '{"value":"hello"}', + "", + ], + ), + ( + "minimax-m2", + [ + '', + 'acme', + '{"value":"hello"}', + "", + ], + ), + ( + "minimax-m3", + [ + ns + segment + for segment in ( + "", + '', + "acme", + "", + "", + "hello", + "", + "", + "", + "", + ) + ], + ), + ( + "poolside_v1", + [ + "acme\n", + "kind\nacme\n", + ( + 'payload\n{"value":"hello"}' + "\n" + ), + "", + ], + ), + ( + "step3", + [ + "<|tool_calls_begin|><|tool_call_begin|>function<|tool_sep|>", + '', + 'acme', + ( + '{"value":"hello"}' + "" + ), + "<|tool_call_end|><|tool_calls_end|>", + ], + ), + ] + + for parser_name, chunks in cases: + with self.subTest(parser=parser_name): + parser = FunctionCallParser(self.tools, parser_name) + _, calls = parser.parse_non_stream("".join(chunks)) + self.assertEqual(json.loads(calls[0].parameters), self.expected) + + parser = FunctionCallParser(self.tools, parser_name) + parameters = "" + for chunk in chunks: + _, calls = parser.parse_stream_chunk(chunk) + parameters += "".join(call.parameters for call in calls) + self.assertEqual(json.loads(parameters), self.expected) + + def test_other_schema_consumers(self): + other_tool = Tool( + type="function", + function=Function( + name="search", + parameters={ + "type": "object", + "properties": {"query": {"type": "string"}}, + }, + ), + ) + parser = FunctionCallParser([self.tools[0], other_tool], "kimi_k2") + _, calls = parser.parse_non_stream( + "<|tool_calls_section_begin|>" + "<|tool_call_begin|>0" + f"<|tool_call_argument_begin|>{json.dumps(self.expected)}" + "<|tool_call_end|>" + "<|tool_calls_section_end|>" + ) + self.assertEqual(calls[0].name, "acme") + + +if __name__ == "__main__": + unittest.main() From c14cd5e650f35819f64b34391766b407c30ba5d5 Mon Sep 17 00:00:00 2001 From: Xeophon <46377542+xeophon@users.noreply.github.com> Date: Wed, 26 Aug 2026 12:59:42 +0200 Subject: [PATCH 2/7] fix: resolve local refs in root schema combinators --- .../function_call/kimik3_structural_tag.py | 21 ++-------- python/sglang/srt/function_call/utils.py | 40 ++++++++++++++++++- .../test_root_combinator_tool_parsers.py | 9 +++-- 3 files changed, 47 insertions(+), 23 deletions(-) diff --git a/python/sglang/srt/function_call/kimik3_structural_tag.py b/python/sglang/srt/function_call/kimik3_structural_tag.py index 04f29a1dfa4a..00300caf0752 100644 --- a/python/sglang/srt/function_call/kimik3_structural_tag.py +++ b/python/sglang/srt/function_call/kimik3_structural_tag.py @@ -29,6 +29,7 @@ TOOLS_CLOSE, TOOLS_OPEN, ) +from sglang.srt.function_call.utils import resolve_local_json_schema_ref _JSON_TYPES = ( "string", @@ -108,22 +109,6 @@ def _matches_json_type(value: Any, json_type: str) -> bool: ) -def _resolve_local_ref( - ref: str, root_schema: Dict[str, Any] -) -> Optional[Union[bool, Dict[str, Any]]]: - if not ref.startswith("#/"): - return None - value: Any = root_schema - for part in ref[2:].split("/"): - key = part.replace("~1", "/").replace("~0", "~") - if not isinstance(value, dict) or key not in value: - return None - value = value[key] - if isinstance(value, (bool, dict)): - return value - return None - - def _schema_types( schema: Union[bool, Dict[str, Any]], root_schema: Dict[str, Any], @@ -138,7 +123,7 @@ def _schema_types( if isinstance(ref, str): seen_refs = set() if seen_refs is None else set(seen_refs) if ref not in seen_refs: - target = _resolve_local_ref(ref, root_schema) + target = resolve_local_json_schema_ref(ref, root_schema) if target is not None: seen_refs.add(ref) return _schema_types(target, root_schema, seen_refs) @@ -219,7 +204,7 @@ def _restrict_schema_type( ref = schema.get("$ref") if isinstance(ref, str): - target = _resolve_local_ref(ref, root_schema) + target = resolve_local_json_schema_ref(ref, root_schema) if target is not None: return _with_root_definitions( _restrict_schema_type(target, json_type, root_schema), root_schema diff --git a/python/sglang/srt/function_call/utils.py b/python/sglang/srt/function_call/utils.py index 1d82a1c3cc1d..c8804f3d30ed 100644 --- a/python/sglang/srt/function_call/utils.py +++ b/python/sglang/srt/function_call/utils.py @@ -304,10 +304,43 @@ def _get_tool_schema(tool: Tool) -> dict: } -def get_json_schema_properties(schema: Any) -> Dict[str, Any]: +def resolve_local_json_schema_ref( + ref: str, root_schema: Dict[str, Any] +) -> Optional[Union[bool, Dict[str, Any]]]: + """Resolve a local JSON Pointer reference against its root schema.""" + if not ref.startswith("#/"): + return None + value: Any = root_schema + for part in ref[2:].split("/"): + key = part.replace("~1", "/").replace("~0", "~") + if not isinstance(value, dict) or key not in value: + return None + value = value[key] + if isinstance(value, (bool, dict)): + return value + return None + + +def get_json_schema_properties( + schema: Any, + root_schema: Optional[Dict[str, Any]] = None, + seen_refs: frozenset[str] = frozenset(), +) -> Dict[str, Any]: """Collect properties declared directly or in root schema combinators.""" if not isinstance(schema, dict): return {} + if root_schema is None: + root_schema = schema + + ref = schema.get("$ref") + if isinstance(ref, str) and ref not in seen_refs: + target = resolve_local_json_schema_ref(ref, root_schema) + if target is not None: + siblings = {key: value for key, value in schema.items() if key != "$ref"} + resolved_schema = {"allOf": [target, siblings]} if siblings else target + return get_json_schema_properties( + resolved_schema, root_schema, seen_refs | {ref} + ) direct_properties = schema.get("properties") if not isinstance(direct_properties, dict): @@ -319,7 +352,10 @@ def get_json_schema_properties(schema: Any) -> Dict[str, Any]: if not isinstance(branches, list): continue for branch in branches: - for name, property_schema in get_json_schema_properties(branch).items(): + branch_properties = get_json_schema_properties( + branch, root_schema, seen_refs + ) + for name, property_schema in branch_properties.items(): if name in direct_properties: continue if name in properties and properties[name] != property_schema: diff --git a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py index 05000ab52ece..d1a5c6cc41cc 100644 --- a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -19,8 +19,8 @@ def setUp(self): name="acme", parameters={ "type": "object", - "oneOf": [ - { + "$defs": { + "AcmeRequest": { "type": "object", "properties": { "kind": {"const": "acme"}, @@ -31,7 +31,10 @@ def setUp(self): }, }, "required": ["kind", "payload"], - }, + } + }, + "oneOf": [ + {"$ref": "#/$defs/AcmeRequest"}, { "type": "object", "properties": {"kind": {"const": "other"}}, From 1ec852ea9a024f29915b900d1afe58c99e47c150 Mon Sep 17 00:00:00 2001 From: Xeophon <46377542+xeophon@users.noreply.github.com> Date: Wed, 26 Aug 2026 13:19:29 +0200 Subject: [PATCH 3/7] fix: resolve refs in projected property schemas --- .../function_call/kimik3_structural_tag.py | 28 ++----- python/sglang/srt/function_call/utils.py | 84 ++++++++++--------- .../test_root_combinator_tool_parsers.py | 17 ++-- 3 files changed, 63 insertions(+), 66 deletions(-) diff --git a/python/sglang/srt/function_call/kimik3_structural_tag.py b/python/sglang/srt/function_call/kimik3_structural_tag.py index 00300caf0752..2cbd2d23efda 100644 --- a/python/sglang/srt/function_call/kimik3_structural_tag.py +++ b/python/sglang/srt/function_call/kimik3_structural_tag.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union +from typing import Any, Dict, List, Literal, Optional, Tuple, Union from xgrammar import StructuralTag from xgrammar.structural_tag import ( @@ -29,7 +29,7 @@ TOOLS_CLOSE, TOOLS_OPEN, ) -from sglang.srt.function_call.utils import resolve_local_json_schema_ref +from sglang.srt.function_call.utils import resolve_local_json_schema_refs _JSON_TYPES = ( "string", @@ -112,22 +112,13 @@ def _matches_json_type(value: Any, json_type: str) -> bool: def _schema_types( schema: Union[bool, Dict[str, Any]], root_schema: Dict[str, Any], - seen_refs: Optional[Set[str]] = None, ) -> List[str]: + schema = resolve_local_json_schema_refs(schema, root_schema) if schema is False: return [] if schema is True: return list(_JSON_TYPES) - ref = schema.get("$ref") - if isinstance(ref, str): - seen_refs = set() if seen_refs is None else set(seen_refs) - if ref not in seen_refs: - target = resolve_local_json_schema_ref(ref, root_schema) - if target is not None: - seen_refs.add(ref) - return _schema_types(target, root_schema, seen_refs) - schema_type = schema.get("type") if isinstance(schema_type, str): return [schema_type] if schema_type in _JSON_TYPES else list(_JSON_TYPES) @@ -141,14 +132,14 @@ def _schema_types( item for option in options if isinstance(option, (bool, dict)) - for item in _schema_types(option, root_schema, seen_refs) + for item in _schema_types(option, root_schema) } return [item for item in _JSON_TYPES if item in option_types] options = schema.get("allOf") if isinstance(options, list): type_sets = [ - set(_schema_types(option, root_schema, seen_refs)) + set(_schema_types(option, root_schema)) for option in options if isinstance(option, (bool, dict)) ] @@ -199,17 +190,10 @@ def _restrict_schema_type( json_type: str, root_schema: Dict[str, Any], ) -> Union[bool, Dict[str, Any]]: + schema = resolve_local_json_schema_refs(schema, root_schema) if not isinstance(schema, dict): return {"type": json_type} if schema else False - ref = schema.get("$ref") - if isinstance(ref, str): - target = resolve_local_json_schema_ref(ref, root_schema) - if target is not None: - return _with_root_definitions( - _restrict_schema_type(target, json_type, root_schema), root_schema - ) - result = dict(schema) schema_type = result.get("type") if isinstance(schema_type, list): diff --git a/python/sglang/srt/function_call/utils.py b/python/sglang/srt/function_call/utils.py index c8804f3d30ed..47ce18ae771e 100644 --- a/python/sglang/srt/function_call/utils.py +++ b/python/sglang/srt/function_call/utils.py @@ -304,57 +304,65 @@ def _get_tool_schema(tool: Tool) -> dict: } -def resolve_local_json_schema_ref( - ref: str, root_schema: Dict[str, Any] -) -> Optional[Union[bool, Dict[str, Any]]]: - """Resolve a local JSON Pointer reference against its root schema.""" - if not ref.startswith("#/"): - return None - value: Any = root_schema - for part in ref[2:].split("/"): - key = part.replace("~1", "/").replace("~0", "~") - if not isinstance(value, dict) or key not in value: - return None - value = value[key] - if isinstance(value, (bool, dict)): - return value - return None +def resolve_local_json_schema_refs( + schema: Any, + root_schema: Dict[str, Any], + seen_refs: frozenset[str] = frozenset(), +) -> Any: + """Resolve local references along a schema's type-inference paths.""" + if not isinstance(schema, dict): + return schema + + ref = schema.get("$ref") + if isinstance(ref, str) and ref in seen_refs: + schema = {key: value for key, value in schema.items() if key != "$ref"} + ref = None + if isinstance(ref, str) and ref.startswith("#/"): + target: Any = root_schema + for part in ref[2:].split("/"): + key = part.replace("~1", "/").replace("~0", "~") + if not isinstance(target, dict) or key not in target: + break + target = target[key] + else: + siblings = {key: value for key, value in schema.items() if key != "$ref"} + schema = {"allOf": [target, siblings]} if siblings else target + return resolve_local_json_schema_refs( + schema, root_schema, seen_refs | {ref} + ) + + return schema | { + keyword: [ + resolve_local_json_schema_refs(branch, root_schema, seen_refs) + for branch in schema[keyword] + ] + for keyword in ("anyOf", "oneOf", "allOf") + if keyword in schema + } def get_json_schema_properties( schema: Any, root_schema: Optional[Dict[str, Any]] = None, - seen_refs: frozenset[str] = frozenset(), ) -> Dict[str, Any]: """Collect properties declared directly or in root schema combinators.""" - if not isinstance(schema, dict): - return {} if root_schema is None: + if not isinstance(schema, dict): + return {} root_schema = schema + schema = resolve_local_json_schema_refs(schema, root_schema) + if not isinstance(schema, dict): + return {} - ref = schema.get("$ref") - if isinstance(ref, str) and ref not in seen_refs: - target = resolve_local_json_schema_ref(ref, root_schema) - if target is not None: - siblings = {key: value for key, value in schema.items() if key != "$ref"} - resolved_schema = {"allOf": [target, siblings]} if siblings else target - return get_json_schema_properties( - resolved_schema, root_schema, seen_refs | {ref} - ) - - direct_properties = schema.get("properties") - if not isinstance(direct_properties, dict): - direct_properties = {} + direct_properties = { + name: resolve_local_json_schema_refs(property_schema, root_schema) + for name, property_schema in schema.get("properties", {}).items() + } properties = direct_properties.copy() for keyword in ("anyOf", "oneOf", "allOf"): - branches = schema.get(keyword) - if not isinstance(branches, list): - continue - for branch in branches: - branch_properties = get_json_schema_properties( - branch, root_schema, seen_refs - ) + for branch in schema.get(keyword, []): + branch_properties = get_json_schema_properties(branch, root_schema) for name, property_schema in branch_properties.items(): if name in direct_properties: continue diff --git a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py index d1a5c6cc41cc..126df7bfe47a 100644 --- a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -20,18 +20,19 @@ def setUp(self): parameters={ "type": "object", "$defs": { + "Payload": { + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, "AcmeRequest": { "type": "object", "properties": { "kind": {"const": "acme"}, - "payload": { - "type": "object", - "properties": {"value": {"type": "string"}}, - "required": ["value"], - }, + "payload": {"$ref": "#/$defs/Payload"}, }, "required": ["kind", "payload"], - } + }, }, "oneOf": [ {"$ref": "#/$defs/AcmeRequest"}, @@ -70,6 +71,10 @@ def test_property_lookup_supports_root_combinators(self): get_json_schema_properties(self.tools[0].function.parameters)["kind"], {"oneOf": [{"const": "acme"}, {"const": "other"}]}, ) + self.assertEqual( + get_json_schema_properties(self.tools[0].function.parameters)["payload"], + payload_schema | {"required": ["value"]}, + ) parser = FunctionCallParser(self.tools, "qwen3_coder") self.assertNotIn("properties", self.tools[0].function.parameters) From 00ca197bd8e10b97a9c783bbc6825a547858287e Mon Sep 17 00:00:00 2001 From: Xeophon <46377542+xeophon@users.noreply.github.com> Date: Wed, 26 Aug 2026 14:07:16 +0200 Subject: [PATCH 4/7] fix: preserve unconstrained tool argument strings --- .../srt/function_call/function_call_parser.py | 66 ++++++++++++++++++- python/sglang/srt/function_call/minimax_m3.py | 7 +- python/sglang/srt/function_call/utils.py | 25 ++++--- .../test_root_combinator_tool_parsers.py | 37 +++++++++++ 4 files changed, 121 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index f3dc48ac5f15..1b8f0bdb92ea 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -1,7 +1,10 @@ import inspect +import json import logging from typing import Dict, List, Literal, Optional, Set, Tuple, Type, Union +from jsonschema import Draft202012Validator + from sglang.srt.entrypoints.openai.protocol import ( LegacyStructuralTagResponseFormat, StructuralTagResponseFormat, @@ -117,12 +120,17 @@ def __init__(self, tools: List[Tool], tool_call_parser: str, tokenizer=None): self.detector = detector self.tools = tools + self.schema_validators = {} + self.streaming_calls: Dict[int, ToolCallItem] = {} self.detector_tools: List[Tool] = [] for tool in tools: parameters = tool.function.parameters if isinstance(parameters, dict): properties = get_json_schema_properties(parameters) if properties != parameters.get("properties", {}): + self.schema_validators[tool.function.name] = Draft202012Validator( + parameters + ) function = tool.function.model_copy( update={"parameters": parameters | {"properties": properties}} ) @@ -130,6 +138,52 @@ def __init__(self, tools: List[Tool], tool_call_parser: str, tokenizer=None): self.detector_tools.append(tool) self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get() + def _normalize_calls( + self, calls: List[ToolCallItem], streaming: bool = False + ) -> List[ToolCallItem]: + """Preserve raw strings whenever they satisfy the original tool schema.""" + result = [] + for call in calls: + pending = self.streaming_calls.get(call.tool_index) + name = call.name or (pending.name if pending else None) + validator = self.schema_validators.get(name) + if validator is None: + result.append(call) + continue + + parameters = (pending.parameters if pending else "") + call.parameters + try: + arguments = json.loads(parameters) + except json.JSONDecodeError: + if streaming: + self.streaming_calls[call.tool_index] = ToolCallItem( + tool_index=call.tool_index, name=name, parameters=parameters + ) + if call.name: + result.append(call.model_copy(update={"parameters": ""})) + else: + result.append(call) + continue + + self.streaming_calls.pop(call.tool_index, None) + for key, value in arguments.items(): + if isinstance(value, str): + continue + string_value = json.dumps( + value, ensure_ascii=False, separators=(",", ":") + ) + if validator.is_valid(arguments | {key: string_value}): + arguments[key] = string_value + result.append( + call.model_copy( + update={ + "name": None if pending else call.name, + "parameters": json.dumps(arguments, ensure_ascii=False), + } + ) + ) + return result + def has_tool_call(self, text: str) -> bool: """ Check if the given text contains a tool call in the format supported by this parser. @@ -161,7 +215,7 @@ def parse_non_stream(self, full_text: str) -> Tuple[str, list[ToolCallItem]]: return full_text, [] has_tool_call = self.detector.has_tool_call(full_text) parsed_result = self.detector.detect_and_parse(full_text, self.detector_tools) - tool_call_list = parsed_result.calls + tool_call_list = self._normalize_calls(parsed_result.calls) if tool_call_list or has_tool_call: return parsed_result.normal_text, tool_call_list else: @@ -190,7 +244,7 @@ def parse_stream_chunk(self, chunk_text: str) -> Tuple[str, list[ToolCallItem]]: if sp_result.normal_text: final_normal_text = sp_result.normal_text if sp_result.calls: - final_calls.extend(sp_result.calls) + final_calls.extend(self._normalize_calls(sp_result.calls, streaming=True)) final_normal_text = sp_result.normal_text return final_normal_text, final_calls @@ -204,7 +258,13 @@ def parse_stream_end(self) -> Tuple[str, list[ToolCallItem]]: if not self.tools: return "", [] sp_result = self.detector.finish(self.detector_tools) - return sp_result.normal_text, sp_result.calls + calls = self._normalize_calls(sp_result.calls, streaming=True) + calls.extend( + call.model_copy(update={"name": None}) + for call in self.streaming_calls.values() + ) + self.streaming_calls.clear() + return sp_result.normal_text, calls def get_legacy_structural_tag( self, at_least_one: bool = False diff --git a/python/sglang/srt/function_call/minimax_m3.py b/python/sglang/srt/function_call/minimax_m3.py index 062fc3ec687d..bc258bf569e6 100644 --- a/python/sglang/srt/function_call/minimax_m3.py +++ b/python/sglang/srt/function_call/minimax_m3.py @@ -278,9 +278,12 @@ def _consume_complex_param(self, calls: List[ToolCallItem]) -> bool: return False self._current_param_buffer += self._buffer[:end] - value = self._parse_parameter( - self._current_param_buffer, self._current_param_schema + parse_value = ( + self._parse_parameter + if self.PARAM_START_PREFIX in self._current_param_buffer + else self._convert_leaf_value ) + value = parse_value(self._current_param_buffer, self._current_param_schema) self._append_stream_call(calls, json.dumps(value, ensure_ascii=False)) self._buffer = self._buffer[end + len(end_token) :] self._clear_current_param() diff --git a/python/sglang/srt/function_call/utils.py b/python/sglang/srt/function_call/utils.py index 47ce18ae771e..777982e5ac1d 100644 --- a/python/sglang/srt/function_call/utils.py +++ b/python/sglang/srt/function_call/utils.py @@ -361,15 +361,22 @@ def get_json_schema_properties( properties = direct_properties.copy() for keyword in ("anyOf", "oneOf", "allOf"): - for branch in schema.get(keyword, []): - branch_properties = get_json_schema_properties(branch, root_schema) - for name, property_schema in branch_properties.items(): - if name in direct_properties: - continue - if name in properties and properties[name] != property_schema: - properties[name] = {keyword: [properties[name], property_schema]} - else: - properties[name] = property_schema + branch_properties = [ + get_json_schema_properties(branch, root_schema) + for branch in schema.get(keyword, []) + ] + for name in dict.fromkeys( + name for branch in branch_properties for name in branch + ): + choices = [branch[name] for branch in branch_properties if name in branch] + property_schema = choices[0] if len(choices) == 1 else {keyword: choices} + current_schema = properties.get(name, {}) + if current_schema in ({}, True): + properties[name] = property_schema + elif ( + property_schema not in ({}, True) and current_schema != property_schema + ): + properties[name] = {"allOf": [current_schema, property_schema]} return properties diff --git a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py index 126df7bfe47a..71cfb9c2adaf 100644 --- a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -66,6 +66,10 @@ def test_property_lookup_supports_root_combinators(self): self.assertEqual( get_json_schema_properties(schema), {"payload": payload_schema} ) + schema["properties"] = {"payload": {}} + self.assertEqual( + get_json_schema_properties(schema), {"payload": payload_schema} + ) self.assertEqual( get_json_schema_properties(self.tools[0].function.parameters)["kind"], @@ -226,6 +230,39 @@ def test_parsers(self): parameters += "".join(call.parameters for call in calls) self.assertEqual(json.loads(parameters), self.expected) + other_chunks = [ + chunk.replace(">acmeother", + '', + "other", + "", + '{"x":1}', + "", + "", + "", + ) + ] + expected = {"kind": "other", "payload": '{"x":1}'} + + parser = FunctionCallParser(self.tools, parser_name) + _, calls = parser.parse_non_stream("".join(other_chunks)) + self.assertEqual(json.loads(calls[0].parameters), expected) + + parser = FunctionCallParser(self.tools, parser_name) + parameters = "" + for chunk in other_chunks: + _, calls = parser.parse_stream_chunk(chunk) + parameters += "".join(call.parameters for call in calls) + self.assertEqual(json.loads(parameters), expected) + def test_other_schema_consumers(self): other_tool = Tool( type="function", From 7224bcaf12466b622aecfb83c8f05cb1a8f9ada1 Mon Sep 17 00:00:00 2001 From: Xeophon <46377542+xeophon@users.noreply.github.com> Date: Wed, 26 Aug 2026 14:41:28 +0200 Subject: [PATCH 5/7] fix: preserve root schema branch streaming --- .../srt/function_call/function_call_parser.py | 178 ++++++++++++------ .../srt/function_call/glm47_moe_detector.py | 28 +-- .../srt/function_call/glm4_moe_detector.py | 28 +-- python/sglang/srt/function_call/minimax_m3.py | 12 ++ python/sglang/srt/function_call/utils.py | 9 +- .../test_root_combinator_tool_parsers.py | 20 +- 6 files changed, 190 insertions(+), 85 deletions(-) diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 1b8f0bdb92ea..23d6eb284e34 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -1,6 +1,7 @@ import inspect import json import logging +import re from typing import Dict, List, Literal, Optional, Set, Tuple, Type, Union from jsonschema import Draft202012Validator @@ -53,6 +54,8 @@ _get_tool_schema_defs, get_json_schema_constraint, get_json_schema_properties, + infer_type_from_json_schema, + resolve_local_json_schema_refs, ) logger = logging.getLogger(__name__) @@ -120,69 +123,131 @@ def __init__(self, tools: List[Tool], tool_call_parser: str, tokenizer=None): self.detector = detector self.tools = tools - self.schema_validators = {} - self.streaming_calls: Dict[int, ToolCallItem] = {} + self.schema_branches = {} + self.streaming_calls: Dict[int, Tuple[Optional[str], str]] = {} self.detector_tools: List[Tool] = [] for tool in tools: parameters = tool.function.parameters if isinstance(parameters, dict): properties = get_json_schema_properties(parameters) if properties != parameters.get("properties", {}): - self.schema_validators[tool.function.name] = Draft202012Validator( - parameters - ) function = tool.function.model_copy( update={"parameters": parameters | {"properties": properties}} ) tool = tool.model_copy(update={"function": function}) + schema = resolve_local_json_schema_refs(parameters, parameters) + for keyword in ("anyOf", "oneOf"): + branches = schema.get(keyword) + if isinstance(branches, list): + if properties == parameters.get("properties", {}): + function = tool.function.model_copy( + update={"parameters": parameters.copy()} + ) + tool = tool.model_copy(update={"function": function}) + base_schema = { + key: value + for key, value in schema.items() + if key != keyword + } + candidates = [] + for branch in branches: + candidate_schema = base_schema | { + "allOf": [*base_schema.get("allOf", []), branch] + } + candidate_properties = get_json_schema_properties( + candidate_schema + ) + candidates.append( + ( + candidate_properties, + { + key: infer_type_from_json_schema(value) + for key, value in candidate_properties.items() + }, + Draft202012Validator(candidate_schema), + ) + ) + self.schema_branches[tool.function.name] = { + "parameters": tool.function.parameters, + "properties": properties, + "candidates": candidates, + } + break self.detector_tools.append(tool) self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get() - def _normalize_calls( + def _update_tool_schemas( self, calls: List[ToolCallItem], streaming: bool = False - ) -> List[ToolCallItem]: - """Preserve raw strings whenever they satisfy the original tool schema.""" - result = [] + ) -> None: + """Select a root schema branch without decoding valid string arguments.""" for call in calls: - pending = self.streaming_calls.get(call.tool_index) - name = call.name or (pending.name if pending else None) - validator = self.schema_validators.get(name) - if validator is None: - result.append(call) + pending_name, pending_parameters = self.streaming_calls.get( + call.tool_index, (None, "") + ) + name = call.name or pending_name + branch_config = self.schema_branches.get(name) + if branch_config is None: continue - parameters = (pending.parameters if pending else "") + call.parameters + tool_parameters = branch_config["parameters"] + if streaming and call.name and not pending_name: + tool_parameters["properties"] = branch_config["properties"] + + parameters = pending_parameters + call.parameters + complete = False try: arguments = json.loads(parameters) + complete = True except json.JSONDecodeError: - if streaming: - self.streaming_calls[call.tool_index] = ToolCallItem( - tool_index=call.tool_index, name=name, parameters=parameters - ) - if call.name: - result.append(call.model_copy(update={"parameters": ""})) - else: - result.append(call) - continue - - self.streaming_calls.pop(call.tool_index, None) - for key, value in arguments.items(): - if isinstance(value, str): + try: + arguments = json.loads(parameters + "}") + except json.JSONDecodeError: + if streaming: + self.streaming_calls[call.tool_index] = (name, parameters) continue - string_value = json.dumps( - value, ensure_ascii=False, separators=(",", ":") - ) - if validator.is_valid(arguments | {key: string_value}): - arguments[key] = string_value - result.append( - call.model_copy( - update={ - "name": None if pending else call.name, - "parameters": json.dumps(arguments, ensure_ascii=False), - } - ) - ) - return result + + if not isinstance(arguments, dict): + continue + if streaming: + self.streaming_calls[call.tool_index] = (name, parameters) + + candidates = [] + for candidate_properties, property_types, validator in branch_config[ + "candidates" + ]: + # Candidate copies can be coerced for validation while the original + # string remains byte-for-byte available for a string branch. + candidate_arguments = arguments.copy() + for key, value in arguments.items(): + expected_type = property_types.get(key) + if isinstance(value, str) and expected_type not in (None, "string"): + try: + candidate_arguments[key] = json.loads(value) + except json.JSONDecodeError: + pass + errors = validator.iter_errors(candidate_arguments) + if all(error.validator == "required" for error in errors): + candidates.append((candidate_properties, candidate_arguments)) + + if len(candidates) == 1: + candidate_properties, selected_arguments = candidates[0] + selected_properties = candidate_properties | { + key: value + for key, value in branch_config["properties"].items() + if candidate_properties.get(key) in (None, {}, True) + } + if streaming and not complete: + tool_parameters["properties"] = selected_properties + if ( + complete + and (not streaming or not pending_parameters) + and selected_arguments != arguments + ): + call.parameters = json.dumps(selected_arguments, ensure_ascii=False) + + if streaming and complete: + self.streaming_calls.pop(call.tool_index, None) + tool_parameters["properties"] = branch_config["properties"] def has_tool_call(self, text: str) -> bool: """ @@ -215,7 +280,8 @@ def parse_non_stream(self, full_text: str) -> Tuple[str, list[ToolCallItem]]: return full_text, [] has_tool_call = self.detector.has_tool_call(full_text) parsed_result = self.detector.detect_and_parse(full_text, self.detector_tools) - tool_call_list = self._normalize_calls(parsed_result.calls) + self._update_tool_schemas(parsed_result.calls) + tool_call_list = parsed_result.calls if tool_call_list or has_tool_call: return parsed_result.normal_text, tool_call_list else: @@ -238,14 +304,18 @@ def parse_stream_chunk(self, chunk_text: str) -> Tuple[str, list[ToolCallItem]]: final_normal_text = "" final_calls = [] - sp_result = self.detector.parse_streaming_increment( - chunk_text, self.detector_tools + chunks = ( + re.findall(r".*?]+>|.+", chunk_text, re.DOTALL) or [""] + if self.schema_branches + else [chunk_text] ) - if sp_result.normal_text: - final_normal_text = sp_result.normal_text - if sp_result.calls: - final_calls.extend(self._normalize_calls(sp_result.calls, streaming=True)) - final_normal_text = sp_result.normal_text + for chunk in chunks: + sp_result = self.detector.parse_streaming_increment( + chunk, self.detector_tools + ) + final_normal_text += sp_result.normal_text + final_calls.extend(sp_result.calls) + self._update_tool_schemas(sp_result.calls, streaming=True) return final_normal_text, final_calls @@ -258,13 +328,9 @@ def parse_stream_end(self) -> Tuple[str, list[ToolCallItem]]: if not self.tools: return "", [] sp_result = self.detector.finish(self.detector_tools) - calls = self._normalize_calls(sp_result.calls, streaming=True) - calls.extend( - call.model_copy(update={"name": None}) - for call in self.streaming_calls.values() - ) + self._update_tool_schemas(sp_result.calls, streaming=True) self.streaming_calls.clear() - return sp_result.normal_text, calls + return sp_result.normal_text, sp_result.calls def get_legacy_structural_tag( self, at_least_one: bool = False diff --git a/python/sglang/srt/function_call/glm47_moe_detector.py b/python/sglang/srt/function_call/glm47_moe_detector.py index b01c316c5ef2..ec50a9e4ef05 100644 --- a/python/sglang/srt/function_call/glm47_moe_detector.py +++ b/python/sglang/srt/function_call/glm47_moe_detector.py @@ -125,6 +125,22 @@ def parse_arguments( Returns: Tuple of (parsed_value, is_valid_json) """ + if arg_type == "string": + try: + parsed_value = json.loads(json_value) + return ( + parsed_value if isinstance(parsed_value, str) else json_value, + True, + ) + except (json.JSONDecodeError, ValueError): + 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 1: Direct JSON parsing try: parsed_value = json.loads(json_value) @@ -149,18 +165,6 @@ def parse_arguments( except (json.JSONDecodeError, ValueError, KeyError): 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 try: parsed_value = safe_literal_eval(json_value) diff --git a/python/sglang/srt/function_call/glm4_moe_detector.py b/python/sglang/srt/function_call/glm4_moe_detector.py index 85076316d77f..35f2e047aad5 100644 --- a/python/sglang/srt/function_call/glm4_moe_detector.py +++ b/python/sglang/srt/function_call/glm4_moe_detector.py @@ -94,6 +94,22 @@ def parse_arguments( Returns: Tuple of (parsed_value, is_valid_json) """ + if arg_type == "string": + try: + parsed_value = json.loads(json_value) + return ( + parsed_value if isinstance(parsed_value, str) else json_value, + True, + ) + except (json.JSONDecodeError, ValueError): + 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 1: Direct JSON parsing try: parsed_value = json.loads(json_value) @@ -118,18 +134,6 @@ def parse_arguments( except (json.JSONDecodeError, ValueError, KeyError): 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 try: parsed_value = safe_literal_eval(json_value) diff --git a/python/sglang/srt/function_call/minimax_m3.py b/python/sglang/srt/function_call/minimax_m3.py index bc258bf569e6..0f4bd6a7a84e 100644 --- a/python/sglang/srt/function_call/minimax_m3.py +++ b/python/sglang/srt/function_call/minimax_m3.py @@ -290,6 +290,14 @@ def _consume_complex_param(self, calls: List[ToolCallItem]) -> bool: return True def _consume_scalar_param(self, calls: List[ToolCallItem]) -> bool: + if ( + not self._current_string_started + and not self._current_param_buffer + and self._buffer.startswith(self.PARAM_START_PREFIX) + ): + self._current_param_is_complex = True + return self._consume_complex_param(calls) + end_token = self._parameter_end_token(self._current_param_name) end = self._buffer.find(end_token) if end == -1: @@ -525,6 +533,10 @@ def _parse_parameter(self, body: str, parameters_schema: Optional[Dict]) -> dict tag = chunk[1:gt].strip() text = chunk[gt + 1 :] parent_frame = stack[-1] + if not isinstance(parent_frame["value"], (dict, list)): + parent_frame["value"] = self._new_container_for_schema( + parent_frame["schema"] + ) child_schema = self._get_child_schema( parent_frame["schema"], tag, parent_frame["value"] ) diff --git a/python/sglang/srt/function_call/utils.py b/python/sglang/srt/function_call/utils.py index 777982e5ac1d..6df93e5788c2 100644 --- a/python/sglang/srt/function_call/utils.py +++ b/python/sglang/srt/function_call/utils.py @@ -369,7 +369,14 @@ def get_json_schema_properties( name for branch in branch_properties for name in branch ): choices = [branch[name] for branch in branch_properties if name in branch] - property_schema = choices[0] if len(choices) == 1 else {keyword: choices} + if keyword in ("anyOf", "oneOf") and len(choices) < len(branch_properties): + # The missing branch leaves this property unconstrained, so preserve + # its raw spelling until sibling arguments select a branch. + property_schema = {"type": "string"} + else: + property_schema = ( + choices[0] if len(choices) == 1 else {keyword: choices} + ) current_schema = properties.get(name, {}) if current_schema in ({}, True): properties[name] = property_schema diff --git a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py index 71cfb9c2adaf..554a22d08959 100644 --- a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -77,7 +77,7 @@ def test_property_lookup_supports_root_combinators(self): ) self.assertEqual( get_json_schema_properties(self.tools[0].function.parameters)["payload"], - payload_schema | {"required": ["value"]}, + {"type": "string"}, ) parser = FunctionCallParser(self.tools, "qwen3_coder") @@ -232,7 +232,7 @@ def test_parsers(self): other_chunks = [ chunk.replace(">acmeother', "other", "", - '{"x":1}', + '{ "x": 1.00 }', "", "", "", ) ] - expected = {"kind": "other", "payload": '{"x":1}'} + expected = {"kind": "other", "payload": '{ "x": 1.00 }'} parser = FunctionCallParser(self.tools, parser_name) _, calls = parser.parse_non_stream("".join(other_chunks)) @@ -284,6 +284,18 @@ def test_other_schema_consumers(self): ) self.assertEqual(calls[0].name, "acme") + def test_streaming_arguments_are_not_buffered(self): + parser = FunctionCallParser(self.tools, "qwen3_coder") + parameters = "" + for chunk in ( + "", + "", + "acme", + ): + _, calls = parser.parse_stream_chunk(chunk) + parameters += "".join(call.parameters for call in calls) + self.assertIn('"kind": "acme"', parameters) + if __name__ == "__main__": unittest.main() From 884beac524410a914f6a32a9125e4a8cddb0d36d Mon Sep 17 00:00:00 2001 From: Xeophon <46377542+xeophon@users.noreply.github.com> Date: Wed, 26 Aug 2026 15:41:12 +0200 Subject: [PATCH 6/7] fix: find alternatives beneath allOf --- .../srt/function_call/function_call_parser.py | 100 ++++++++++++------ .../test_root_combinator_tool_parsers.py | 46 ++++++++ 2 files changed, 112 insertions(+), 34 deletions(-) diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 23d6eb284e34..dd13ff1889c3 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -136,43 +136,75 @@ def __init__(self, tools: List[Tool], tool_call_parser: str, tokenizer=None): ) tool = tool.model_copy(update={"function": function}) schema = resolve_local_json_schema_refs(parameters, parameters) - for keyword in ("anyOf", "oneOf"): - branches = schema.get(keyword) - if isinstance(branches, list): - if properties == parameters.get("properties", {}): - function = tool.function.model_copy( - update={"parameters": parameters.copy()} - ) - tool = tool.model_copy(update={"function": function}) - base_schema = { + conjuncts = [schema] + index = 0 + while index < len(conjuncts): + conjunct = conjuncts[index] + nested = ( + conjunct.get("allOf") if isinstance(conjunct, dict) else None + ) + if isinstance(nested, list): + remainder = { key: value - for key, value in schema.items() - if key != keyword - } - candidates = [] - for branch in branches: - candidate_schema = base_schema | { - "allOf": [*base_schema.get("allOf", []), branch] - } - candidate_properties = get_json_schema_properties( - candidate_schema - ) - candidates.append( - ( - candidate_properties, - { - key: infer_type_from_json_schema(value) - for key, value in candidate_properties.items() - }, - Draft202012Validator(candidate_schema), - ) - ) - self.schema_branches[tool.function.name] = { - "parameters": tool.function.parameters, - "properties": properties, - "candidates": candidates, + for key, value in conjunct.items() + if key != "allOf" } + conjuncts[index : index + 1] = [remainder, *nested] + else: + index += 1 + + alternative = None + for index, conjunct in enumerate(conjuncts): + if not isinstance(conjunct, dict): + continue + for keyword in ("anyOf", "oneOf"): + branches = conjunct.get(keyword) + if isinstance(branches, list): + alternative = index, keyword, branches + break + if alternative: break + + if alternative: + index, keyword, branches = alternative + if properties == parameters.get("properties", {}): + function = tool.function.model_copy( + update={"parameters": parameters.copy()} + ) + tool = tool.model_copy(update={"function": function}) + conjuncts[index] = { + key: value + for key, value in conjuncts[index].items() + if key != keyword + } + base_schema = { + key: parameters[key] + for key in ("$defs", "definitions") + if key in parameters + } | {"allOf": conjuncts} + candidates = [] + for branch in branches: + candidate_schema = base_schema | { + "allOf": [*base_schema["allOf"], branch] + } + candidate_properties = get_json_schema_properties( + candidate_schema + ) + candidates.append( + ( + candidate_properties, + { + key: infer_type_from_json_schema(value) + for key, value in candidate_properties.items() + }, + Draft202012Validator(candidate_schema), + ) + ) + self.schema_branches[tool.function.name] = { + "parameters": tool.function.parameters, + "properties": properties, + "candidates": candidates, + } self.detector_tools.append(tool) self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get() diff --git a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py index 554a22d08959..e583de8d5a49 100644 --- a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -86,6 +86,52 @@ def test_property_lookup_supports_root_combinators(self): "payload", parser.detector_tools[0].function.parameters["properties"] ) + def test_branch_selection_descends_through_all_of(self): + parameters = self.tools[0].function.parameters + variants = { + "explicit": { + "type": "object", + "$defs": parameters["$defs"], + "allOf": [{"oneOf": parameters["oneOf"]}], + }, + "root_ref": { + "$ref": "#/$defs/Root", + "$defs": parameters["$defs"] + | { + "Root": { + "type": "object", + "oneOf": parameters["oneOf"], + } + }, + }, + } + chunks = [ + "", + "", + "acme", + '{"value":"hello"}', + "", + "", + ] + + for name, parameters in variants.items(): + with self.subTest(name=name): + function = self.tools[0].function.model_copy( + update={"parameters": parameters} + ) + tool = self.tools[0].model_copy(update={"function": function}) + + parser = FunctionCallParser([tool], "qwen3_coder") + _, calls = parser.parse_non_stream("".join(chunks)) + self.assertEqual(json.loads(calls[0].parameters), self.expected) + + parser = FunctionCallParser([tool], "qwen3_coder") + streamed = "" + for chunk in chunks: + _, calls = parser.parse_stream_chunk(chunk) + streamed += "".join(call.parameters for call in calls) + self.assertEqual(json.loads(streamed), self.expected) + def test_parsers(self): ns = MINIMAX_NS_TOKEN cases = [ From 3018725ef96a0a7aa015cefc01542adce35d9771 Mon Sep 17 00:00:00 2001 From: Xeophon <46377542+xeophon@users.noreply.github.com> Date: Wed, 26 Aug 2026 19:30:24 +0200 Subject: [PATCH 7/7] fix: project root combinator properties for tool parsers --- .../srt/function_call/function_call_parser.py | 198 ++------ .../srt/function_call/glm47_moe_detector.py | 30 +- .../srt/function_call/glm4_moe_detector.py | 30 +- .../function_call/kimik3_structural_tag.py | 43 +- python/sglang/srt/function_call/minimax_m3.py | 19 +- python/sglang/srt/function_call/utils.py | 84 ++-- .../test_root_combinator_tool_parsers.py | 431 +++++++++--------- 7 files changed, 357 insertions(+), 478 deletions(-) diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index dd13ff1889c3..94e5ddf7fbcc 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -1,11 +1,7 @@ import inspect -import json import logging -import re from typing import Dict, List, Literal, Optional, Set, Tuple, Type, Union -from jsonschema import Draft202012Validator - from sglang.srt.entrypoints.openai.protocol import ( LegacyStructuralTagResponseFormat, StructuralTagResponseFormat, @@ -53,9 +49,7 @@ from sglang.srt.function_call.utils import ( _get_tool_schema_defs, get_json_schema_constraint, - get_json_schema_properties, - infer_type_from_json_schema, - resolve_local_json_schema_refs, + get_tool_parser_property_hints, ) logger = logging.getLogger(__name__) @@ -123,163 +117,32 @@ def __init__(self, tools: List[Tool], tool_call_parser: str, tokenizer=None): self.detector = detector self.tools = tools - self.schema_branches = {} - self.streaming_calls: Dict[int, Tuple[Optional[str], str]] = {} self.detector_tools: List[Tool] = [] + for tool in tools: parameters = tool.function.parameters - if isinstance(parameters, dict): - properties = get_json_schema_properties(parameters) - if properties != parameters.get("properties", {}): - function = tool.function.model_copy( - update={"parameters": parameters | {"properties": properties}} - ) - tool = tool.model_copy(update={"function": function}) - schema = resolve_local_json_schema_refs(parameters, parameters) - conjuncts = [schema] - index = 0 - while index < len(conjuncts): - conjunct = conjuncts[index] - nested = ( - conjunct.get("allOf") if isinstance(conjunct, dict) else None - ) - if isinstance(nested, list): - remainder = { - key: value - for key, value in conjunct.items() - if key != "allOf" - } - conjuncts[index : index + 1] = [remainder, *nested] - else: - index += 1 - - alternative = None - for index, conjunct in enumerate(conjuncts): - if not isinstance(conjunct, dict): - continue - for keyword in ("anyOf", "oneOf"): - branches = conjunct.get(keyword) - if isinstance(branches, list): - alternative = index, keyword, branches - break - if alternative: - break - - if alternative: - index, keyword, branches = alternative - if properties == parameters.get("properties", {}): - function = tool.function.model_copy( - update={"parameters": parameters.copy()} - ) - tool = tool.model_copy(update={"function": function}) - conjuncts[index] = { - key: value - for key, value in conjuncts[index].items() - if key != keyword - } - base_schema = { - key: parameters[key] - for key in ("$defs", "definitions") - if key in parameters - } | {"allOf": conjuncts} - candidates = [] - for branch in branches: - candidate_schema = base_schema | { - "allOf": [*base_schema["allOf"], branch] - } - candidate_properties = get_json_schema_properties( - candidate_schema - ) - candidates.append( - ( - candidate_properties, - { - key: infer_type_from_json_schema(value) - for key, value in candidate_properties.items() - }, - Draft202012Validator(candidate_schema), - ) - ) - self.schema_branches[tool.function.name] = { - "parameters": tool.function.parameters, - "properties": properties, - "candidates": candidates, - } - self.detector_tools.append(tool) - self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get() - - def _update_tool_schemas( - self, calls: List[ToolCallItem], streaming: bool = False - ) -> None: - """Select a root schema branch without decoding valid string arguments.""" - for call in calls: - pending_name, pending_parameters = self.streaming_calls.get( - call.tool_index, (None, "") - ) - name = call.name or pending_name - branch_config = self.schema_branches.get(name) - if branch_config is None: + if not isinstance(parameters, dict): + self.detector_tools.append(tool) continue - tool_parameters = branch_config["parameters"] - if streaming and call.name and not pending_name: - tool_parameters["properties"] = branch_config["properties"] - - parameters = pending_parameters + call.parameters - complete = False - try: - arguments = json.loads(parameters) - complete = True - except json.JSONDecodeError: - try: - arguments = json.loads(parameters + "}") - except json.JSONDecodeError: - if streaming: - self.streaming_calls[call.tool_index] = (name, parameters) - continue - - if not isinstance(arguments, dict): + detector_properties = get_tool_parser_property_hints(parameters) + existing_properties = parameters.get("properties", {}) + if not isinstance(existing_properties, dict): + existing_properties = {} + + if detector_properties == existing_properties: + self.detector_tools.append(tool) continue - if streaming: - self.streaming_calls[call.tool_index] = (name, parameters) - - candidates = [] - for candidate_properties, property_types, validator in branch_config[ - "candidates" - ]: - # Candidate copies can be coerced for validation while the original - # string remains byte-for-byte available for a string branch. - candidate_arguments = arguments.copy() - for key, value in arguments.items(): - expected_type = property_types.get(key) - if isinstance(value, str) and expected_type not in (None, "string"): - try: - candidate_arguments[key] = json.loads(value) - except json.JSONDecodeError: - pass - errors = validator.iter_errors(candidate_arguments) - if all(error.validator == "required" for error in errors): - candidates.append((candidate_properties, candidate_arguments)) - - if len(candidates) == 1: - candidate_properties, selected_arguments = candidates[0] - selected_properties = candidate_properties | { - key: value - for key, value in branch_config["properties"].items() - if candidate_properties.get(key) in (None, {}, True) - } - if streaming and not complete: - tool_parameters["properties"] = selected_properties - if ( - complete - and (not streaming or not pending_parameters) - and selected_arguments != arguments - ): - call.parameters = json.dumps(selected_arguments, ensure_ascii=False) - - if streaming and complete: - self.streaming_calls.pop(call.tool_index, None) - tool_parameters["properties"] = branch_config["properties"] + + detector_parameters = parameters.copy() + detector_parameters["properties"] = detector_properties + detector_function = tool.function.model_copy( + update={"parameters": detector_parameters} + ) + detector_tool = tool.model_copy(update={"function": detector_function}) + self.detector_tools.append(detector_tool) + + self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get() def has_tool_call(self, text: str) -> bool: """ @@ -312,7 +175,6 @@ def parse_non_stream(self, full_text: str) -> Tuple[str, list[ToolCallItem]]: return full_text, [] has_tool_call = self.detector.has_tool_call(full_text) parsed_result = self.detector.detect_and_parse(full_text, self.detector_tools) - self._update_tool_schemas(parsed_result.calls) tool_call_list = parsed_result.calls if tool_call_list or has_tool_call: return parsed_result.normal_text, tool_call_list @@ -336,18 +198,14 @@ def parse_stream_chunk(self, chunk_text: str) -> Tuple[str, list[ToolCallItem]]: final_normal_text = "" final_calls = [] - chunks = ( - re.findall(r".*?]+>|.+", chunk_text, re.DOTALL) or [""] - if self.schema_branches - else [chunk_text] + sp_result = self.detector.parse_streaming_increment( + chunk_text, self.detector_tools ) - for chunk in chunks: - sp_result = self.detector.parse_streaming_increment( - chunk, self.detector_tools - ) - final_normal_text += sp_result.normal_text + if sp_result.normal_text: + final_normal_text = sp_result.normal_text + if sp_result.calls: final_calls.extend(sp_result.calls) - self._update_tool_schemas(sp_result.calls, streaming=True) + final_normal_text = sp_result.normal_text return final_normal_text, final_calls @@ -360,8 +218,6 @@ def parse_stream_end(self) -> Tuple[str, list[ToolCallItem]]: if not self.tools: return "", [] sp_result = self.detector.finish(self.detector_tools) - self._update_tool_schemas(sp_result.calls, streaming=True) - self.streaming_calls.clear() return sp_result.normal_text, sp_result.calls def get_legacy_structural_tag( diff --git a/python/sglang/srt/function_call/glm47_moe_detector.py b/python/sglang/srt/function_call/glm47_moe_detector.py index ec50a9e4ef05..90bdb2aedcd5 100644 --- a/python/sglang/srt/function_call/glm47_moe_detector.py +++ b/python/sglang/srt/function_call/glm47_moe_detector.py @@ -125,22 +125,6 @@ def parse_arguments( Returns: Tuple of (parsed_value, is_valid_json) """ - if arg_type == "string": - try: - parsed_value = json.loads(json_value) - return ( - parsed_value if isinstance(parsed_value, str) else json_value, - True, - ) - except (json.JSONDecodeError, ValueError): - 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 1: Direct JSON parsing try: parsed_value = json.loads(json_value) @@ -165,6 +149,18 @@ def parse_arguments( except (json.JSONDecodeError, ValueError, KeyError): 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 try: parsed_value = safe_literal_eval(json_value) @@ -617,7 +613,7 @@ def _finalize_tool_call( self._last_arguments += "{}" self.streamed_args_for_tool[self.current_tool_id] += "{}" self._sent_empty_object = True - elif not self._sent_empty_object: + elif not self._last_arguments.endswith("}") and not self._sent_empty_object: # Need to close brace calls.append( ToolCallItem( diff --git a/python/sglang/srt/function_call/glm4_moe_detector.py b/python/sglang/srt/function_call/glm4_moe_detector.py index 35f2e047aad5..0c29a39e7d24 100644 --- a/python/sglang/srt/function_call/glm4_moe_detector.py +++ b/python/sglang/srt/function_call/glm4_moe_detector.py @@ -94,22 +94,6 @@ def parse_arguments( Returns: Tuple of (parsed_value, is_valid_json) """ - if arg_type == "string": - try: - parsed_value = json.loads(json_value) - return ( - parsed_value if isinstance(parsed_value, str) else json_value, - True, - ) - except (json.JSONDecodeError, ValueError): - 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 1: Direct JSON parsing try: parsed_value = json.loads(json_value) @@ -134,6 +118,18 @@ def parse_arguments( except (json.JSONDecodeError, ValueError, KeyError): 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 try: parsed_value = safe_literal_eval(json_value) @@ -570,7 +566,7 @@ def parse_streaming_increment( ) ) self._last_arguments += empty_object - else: + elif not self._last_arguments.endswith("}"): closing_brace = "}" calls.append( ToolCallItem( diff --git a/python/sglang/srt/function_call/kimik3_structural_tag.py b/python/sglang/srt/function_call/kimik3_structural_tag.py index 2cbd2d23efda..04f29a1dfa4a 100644 --- a/python/sglang/srt/function_call/kimik3_structural_tag.py +++ b/python/sglang/srt/function_call/kimik3_structural_tag.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Literal, Optional, Tuple, Union +from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union from xgrammar import StructuralTag from xgrammar.structural_tag import ( @@ -29,7 +29,6 @@ TOOLS_CLOSE, TOOLS_OPEN, ) -from sglang.srt.function_call.utils import resolve_local_json_schema_refs _JSON_TYPES = ( "string", @@ -109,16 +108,41 @@ def _matches_json_type(value: Any, json_type: str) -> bool: ) +def _resolve_local_ref( + ref: str, root_schema: Dict[str, Any] +) -> Optional[Union[bool, Dict[str, Any]]]: + if not ref.startswith("#/"): + return None + value: Any = root_schema + for part in ref[2:].split("/"): + key = part.replace("~1", "/").replace("~0", "~") + if not isinstance(value, dict) or key not in value: + return None + value = value[key] + if isinstance(value, (bool, dict)): + return value + return None + + def _schema_types( schema: Union[bool, Dict[str, Any]], root_schema: Dict[str, Any], + seen_refs: Optional[Set[str]] = None, ) -> List[str]: - schema = resolve_local_json_schema_refs(schema, root_schema) if schema is False: return [] if schema is True: return list(_JSON_TYPES) + ref = schema.get("$ref") + if isinstance(ref, str): + seen_refs = set() if seen_refs is None else set(seen_refs) + if ref not in seen_refs: + target = _resolve_local_ref(ref, root_schema) + if target is not None: + seen_refs.add(ref) + return _schema_types(target, root_schema, seen_refs) + schema_type = schema.get("type") if isinstance(schema_type, str): return [schema_type] if schema_type in _JSON_TYPES else list(_JSON_TYPES) @@ -132,14 +156,14 @@ def _schema_types( item for option in options if isinstance(option, (bool, dict)) - for item in _schema_types(option, root_schema) + for item in _schema_types(option, root_schema, seen_refs) } return [item for item in _JSON_TYPES if item in option_types] options = schema.get("allOf") if isinstance(options, list): type_sets = [ - set(_schema_types(option, root_schema)) + set(_schema_types(option, root_schema, seen_refs)) for option in options if isinstance(option, (bool, dict)) ] @@ -190,10 +214,17 @@ def _restrict_schema_type( json_type: str, root_schema: Dict[str, Any], ) -> Union[bool, Dict[str, Any]]: - schema = resolve_local_json_schema_refs(schema, root_schema) if not isinstance(schema, dict): return {"type": json_type} if schema else False + ref = schema.get("$ref") + if isinstance(ref, str): + target = _resolve_local_ref(ref, root_schema) + if target is not None: + return _with_root_definitions( + _restrict_schema_type(target, json_type, root_schema), root_schema + ) + result = dict(schema) schema_type = result.get("type") if isinstance(schema_type, list): diff --git a/python/sglang/srt/function_call/minimax_m3.py b/python/sglang/srt/function_call/minimax_m3.py index 0f4bd6a7a84e..062fc3ec687d 100644 --- a/python/sglang/srt/function_call/minimax_m3.py +++ b/python/sglang/srt/function_call/minimax_m3.py @@ -278,26 +278,15 @@ def _consume_complex_param(self, calls: List[ToolCallItem]) -> bool: return False self._current_param_buffer += self._buffer[:end] - parse_value = ( - self._parse_parameter - if self.PARAM_START_PREFIX in self._current_param_buffer - else self._convert_leaf_value + value = self._parse_parameter( + self._current_param_buffer, self._current_param_schema ) - value = parse_value(self._current_param_buffer, self._current_param_schema) self._append_stream_call(calls, json.dumps(value, ensure_ascii=False)) self._buffer = self._buffer[end + len(end_token) :] self._clear_current_param() return True def _consume_scalar_param(self, calls: List[ToolCallItem]) -> bool: - if ( - not self._current_string_started - and not self._current_param_buffer - and self._buffer.startswith(self.PARAM_START_PREFIX) - ): - self._current_param_is_complex = True - return self._consume_complex_param(calls) - end_token = self._parameter_end_token(self._current_param_name) end = self._buffer.find(end_token) if end == -1: @@ -533,10 +522,6 @@ def _parse_parameter(self, body: str, parameters_schema: Optional[Dict]) -> dict tag = chunk[1:gt].strip() text = chunk[gt + 1 :] parent_frame = stack[-1] - if not isinstance(parent_frame["value"], (dict, list)): - parent_frame["value"] = self._new_container_for_schema( - parent_frame["schema"] - ) child_schema = self._get_child_schema( parent_frame["schema"], tag, parent_frame["value"] ) diff --git a/python/sglang/srt/function_call/utils.py b/python/sglang/srt/function_call/utils.py index 6df93e5788c2..e4232d6348fa 100644 --- a/python/sglang/srt/function_call/utils.py +++ b/python/sglang/srt/function_call/utils.py @@ -341,49 +341,69 @@ def resolve_local_json_schema_refs( } -def get_json_schema_properties( +_ROOT_COMBINATORS = ("allOf", "anyOf", "oneOf") +_MISSING = object() + + +def get_tool_parser_property_hints( schema: Any, root_schema: Optional[Dict[str, Any]] = None, ) -> Dict[str, Any]: - """Collect properties declared directly or in root schema combinators.""" + """Return a lossy property map for tool-format detectors. + + This exposes properties hidden below root-level schema combinators. + + Direct properties are authoritative. Identical branch declarations are + retained. Incompatible declarations collapse to an unconstrained mapping, + causing detectors to use their existing conservative string behavior. + + This result must not be used for validation or constrained decoding. + """ + if not isinstance(schema, dict): + return {} + if root_schema is None: - if not isinstance(schema, dict): - return {} root_schema = schema + schema = resolve_local_json_schema_refs(schema, root_schema) if not isinstance(schema, dict): return {} - direct_properties = { - name: resolve_local_json_schema_refs(property_schema, root_schema) - for name, property_schema in schema.get("properties", {}).items() - } + raw_direct = schema.get("properties", {}) + if not isinstance(raw_direct, dict): + raw_direct = {} + + direct_properties: Dict[str, Any] = {} + for name, property_schema in raw_direct.items(): + property_schema = resolve_local_json_schema_refs(property_schema, root_schema) + # Several detectors assume property schemas are mappings. + direct_properties[name] = ( + property_schema if isinstance(property_schema, dict) else {} + ) + properties = direct_properties.copy() - for keyword in ("anyOf", "oneOf", "allOf"): - branch_properties = [ - get_json_schema_properties(branch, root_schema) - for branch in schema.get(keyword, []) - ] - for name in dict.fromkeys( - name for branch in branch_properties for name in branch - ): - choices = [branch[name] for branch in branch_properties if name in branch] - if keyword in ("anyOf", "oneOf") and len(choices) < len(branch_properties): - # The missing branch leaves this property unconstrained, so preserve - # its raw spelling until sibling arguments select a branch. - property_schema = {"type": "string"} - else: - property_schema = ( - choices[0] if len(choices) == 1 else {keyword: choices} - ) - current_schema = properties.get(name, {}) - if current_schema in ({}, True): - properties[name] = property_schema - elif ( - property_schema not in ({}, True) and current_schema != property_schema - ): - properties[name] = {"allOf": [current_schema, property_schema]} + for keyword in _ROOT_COMBINATORS: + branches = schema.get(keyword) + if not isinstance(branches, list): + continue + + for branch in branches: + branch_properties = get_tool_parser_property_hints( + branch, + root_schema=root_schema, + ) + + for name, candidate_schema in branch_properties.items(): + if name in direct_properties: + continue + + current_schema = properties.get(name, _MISSING) + if current_schema is _MISSING: + properties[name] = candidate_schema + elif current_schema != candidate_schema: + properties[name] = {} + return properties diff --git a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py index e583de8d5a49..9cf71e0df717 100644 --- a/test/registered/unit/function_call/test_root_combinator_tool_parsers.py +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -1,147 +1,157 @@ +import copy import json import unittest from sglang.srt.entrypoints.openai.protocol import Function, Tool from sglang.srt.function_call.function_call_parser import FunctionCallParser from sglang.srt.function_call.minimax_m3 import MINIMAX_NS_TOKEN -from sglang.srt.function_call.utils import get_json_schema_properties +from sglang.srt.function_call.utils import get_tool_parser_property_hints from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=2, suite="base-a-test-cpu") -class TestRootCombinatorToolParsers(unittest.TestCase): - def setUp(self): - self.tools = [ - Tool( - type="function", - function=Function( - name="acme", - parameters={ - "type": "object", - "$defs": { - "Payload": { - "type": "object", - "properties": {"value": {"type": "string"}}, - "required": ["value"], - }, - "AcmeRequest": { - "type": "object", - "properties": { - "kind": {"const": "acme"}, - "payload": {"$ref": "#/$defs/Payload"}, - }, - "required": ["kind", "payload"], - }, - }, - "oneOf": [ - {"$ref": "#/$defs/AcmeRequest"}, - { - "type": "object", - "properties": {"kind": {"const": "other"}}, - "required": ["kind"], - }, - ], - }, - ), - ) - ] - self.expected = {"kind": "acme", "payload": {"value": "hello"}} +PROPERTY_SCHEMAS = { + "number": {"type": "number"}, + "stringList": { + "type": "array", + "items": {"type": "string"}, + }, + "numberList": { + "type": "array", + "items": {"type": "number"}, + }, +} +ARGUMENT_CASES = ( + ("number", "1.5", 1.5), + ("stringList", '["alpha","beta"]', ["alpha", "beta"]), + ("numberList", "[1,2.5]", [1, 2.5]), +) - def test_property_lookup_supports_root_combinators(self): - payload_schema = { - "type": "object", - "properties": {"value": {"type": "string"}}, - } - for keyword in ("anyOf", "oneOf", "allOf"): + +def disjoint_schema(keyword: str = "oneOf") -> dict: + return { + "type": "object", + keyword: [ + { + "type": "object", + "properties": {name: property_schema}, + "required": [name], + "additionalProperties": False, + } + for name, property_schema in PROPERTY_SCHEMAS.items() + ], + } + + +class TestToolParserPropertyHints(unittest.TestCase): + def test_disjoint_combinator_properties_are_projected(self): + for keyword in ("allOf", "anyOf", "oneOf"): with self.subTest(keyword=keyword): - schema = { - keyword: [ - { - "type": "object", - "properties": {"payload": payload_schema}, - } - ] - } - self.assertEqual( - get_json_schema_properties(schema), {"payload": payload_schema} - ) - schema["properties"] = {"payload": {}} + schema = disjoint_schema(keyword) + if keyword == "allOf": + for branch in schema[keyword]: + branch.pop("additionalProperties") + self.assertEqual( - get_json_schema_properties(schema), {"payload": payload_schema} + get_tool_parser_property_hints(schema), + PROPERTY_SCHEMAS, ) + def test_direct_property_is_authoritative(self): + direct_schema = {"type": "object"} + schema = { + "properties": {"payload": direct_schema}, + "oneOf": [ + {"properties": {"payload": {"type": "string"}}}, + {"properties": {"payload": {"type": "array"}}}, + ], + } + self.assertEqual( - get_json_schema_properties(self.tools[0].function.parameters)["kind"], - {"oneOf": [{"const": "acme"}, {"const": "other"}]}, + get_tool_parser_property_hints(schema), + {"payload": direct_schema}, ) + + def test_identical_branch_property_is_retained(self): + property_schema = {"type": "object"} + schema = { + "anyOf": [ + {"properties": {"payload": property_schema}}, + {"properties": {"payload": copy.deepcopy(property_schema)}}, + ] + } + self.assertEqual( - get_json_schema_properties(self.tools[0].function.parameters)["payload"], - {"type": "string"}, + get_tool_parser_property_hints(schema), + {"payload": property_schema}, ) - parser = FunctionCallParser(self.tools, "qwen3_coder") - self.assertNotIn("properties", self.tools[0].function.parameters) - self.assertIn( - "payload", parser.detector_tools[0].function.parameters["properties"] + def test_conflicting_branch_property_is_unconstrained(self): + schema = { + "oneOf": [ + {"properties": {"payload": {"type": "object"}}}, + {"properties": {"payload": {"type": "string"}}}, + ] + } + + self.assertEqual( + get_tool_parser_property_hints(schema), + {"payload": {}}, ) - def test_branch_selection_descends_through_all_of(self): - parameters = self.tools[0].function.parameters - variants = { - "explicit": { - "type": "object", - "$defs": parameters["$defs"], - "allOf": [{"oneOf": parameters["oneOf"]}], - }, - "root_ref": { - "$ref": "#/$defs/Root", - "$defs": parameters["$defs"] - | { - "Root": { - "type": "object", - "oneOf": parameters["oneOf"], - } + def test_local_refs_are_resolved(self): + schema = { + "$defs": { + "Arguments": { + "oneOf": [{"properties": {"payload": {"$ref": "#/$defs/Payload"}}}] }, + "Payload": {"type": "object"}, }, + "$ref": "#/$defs/Arguments", } - chunks = [ - "", - "", - "acme", - '{"value":"hello"}', - "", - "", - ] - for name, parameters in variants.items(): - with self.subTest(name=name): - function = self.tools[0].function.model_copy( - update={"parameters": parameters} - ) - tool = self.tools[0].model_copy(update={"function": function}) + self.assertEqual( + get_tool_parser_property_hints(schema), + {"payload": {"type": "object"}}, + ) + + def test_original_schema_is_not_modified(self): + schema = disjoint_schema() + original = copy.deepcopy(schema) + + get_tool_parser_property_hints(schema) - parser = FunctionCallParser([tool], "qwen3_coder") - _, calls = parser.parse_non_stream("".join(chunks)) - self.assertEqual(json.loads(calls[0].parameters), self.expected) + self.assertEqual(schema, original) - parser = FunctionCallParser([tool], "qwen3_coder") - streamed = "" - for chunk in chunks: - _, calls = parser.parse_stream_chunk(chunk) - streamed += "".join(call.parameters for call in calls) - self.assertEqual(json.loads(streamed), self.expected) - def test_parsers(self): +class TestRootCombinatorToolParsers(unittest.TestCase): + def setUp(self): + parameters = disjoint_schema() + self.original_parameters = copy.deepcopy(parameters) + self.tools = [ + Tool( + type="function", + function=Function(name="convert", parameters=parameters), + ) + ] + + def parser_cases(self, argument: str, raw_value: str): ns = MINIMAX_NS_TOKEN - cases = [ + minimax_value = [f"<{argument}>{raw_value}", f""] + if argument.endswith("List"): + minimax_value = [f"<{argument}>"] + for item in json.loads(raw_value): + minimax_value.extend((f"{item}", "")) + minimax_value.append(f"") + + return [ ( "qwen3_coder", [ "", - "", - "acme", - '{"value":"hello"}', + "", + f"{raw_value}", "", "", ], @@ -149,11 +159,10 @@ def test_parsers(self): ( "glm", [ - "acme\n", - "kind\nacme\n", + "convert\n", ( - 'payload\n{"value":"hello"}' - "\n" + f"{argument}\n" + f"{raw_value}\n" ), "", ], @@ -161,11 +170,10 @@ def test_parsers(self): ( "glm47", [ - "acme", - "kindacme", + "convert", ( - 'payload{"value":"hello"}' - "" + f"{argument}" + f"{raw_value}" ), "", ], @@ -173,20 +181,18 @@ def test_parsers(self): ( "dots", [ - '', - 'acme', - '{"value":"hello"}', + '', + f'{raw_value}', "", ], ), ( "hunyuan", [ - "acme", - "kindacme", + "convert", ( - 'payload{"value":"hello"}' - "" + f"{argument}" + f"{raw_value}" ), "", ], @@ -194,27 +200,24 @@ def test_parsers(self): ( "mimo", [ - "", - "acme", - '{"value":"hello"}', + "", + f"{raw_value}", "", ], ), ( "minicpm5", [ - '', - 'acme', - '{"value":"hello"}', + '', + f'{raw_value}', "", ], ), ( "minimax-m2", [ - '', - 'acme', - '{"value":"hello"}', + '', + f'{raw_value}', "", ], ), @@ -224,13 +227,8 @@ def test_parsers(self): ns + segment for segment in ( "", - '', - "acme", - "", - "", - "hello", - "", - "", + '', + *minimax_value, "", "", ) @@ -239,11 +237,10 @@ def test_parsers(self): ( "poolside_v1", [ - "acme\n", - "kind\nacme\n", + "convert\n", ( - 'payload\n{"value":"hello"}' - "\n" + f"{argument}\n" + f"{raw_value}\n" ), "", ], @@ -252,10 +249,9 @@ def test_parsers(self): "step3", [ "<|tool_calls_begin|><|tool_call_begin|>function<|tool_sep|>", - '', - 'acme', + '', ( - '{"value":"hello"}' + f'{raw_value}' "" ), "<|tool_call_end|><|tool_calls_end|>", @@ -263,84 +259,83 @@ def test_parsers(self): ), ] - for parser_name, chunks in cases: - with self.subTest(parser=parser_name): - parser = FunctionCallParser(self.tools, parser_name) - _, calls = parser.parse_non_stream("".join(chunks)) - self.assertEqual(json.loads(calls[0].parameters), self.expected) + def assert_arguments(self, serialized: str, expected: dict): + self.assertNotIn("<", serialized) + self.assertEqual(json.loads(serialized), expected) - parser = FunctionCallParser(self.tools, parser_name) - parameters = "" - for chunk in chunks: - _, calls = parser.parse_stream_chunk(chunk) - parameters += "".join(call.parameters for call in calls) - self.assertEqual(json.loads(parameters), self.expected) + def test_streaming_and_non_streaming_parsers(self): + parser_names = [ + parser_name for parser_name, _ in self.parser_cases("number", "1.5") + ] + for parser_name in parser_names: + non_streaming_parser = FunctionCallParser(self.tools, parser_name) - other_chunks = [ - chunk.replace(">acmeother", - '', - "other", - "", - '{ "x": 1.00 }', - "", - "", - "", - ) - ] - expected = {"kind": "other", "payload": '{ "x": 1.00 }'} - - parser = FunctionCallParser(self.tools, parser_name) - _, calls = parser.parse_non_stream("".join(other_chunks)) - self.assertEqual(json.loads(calls[0].parameters), expected) - - parser = FunctionCallParser(self.tools, parser_name) - parameters = "" - for chunk in other_chunks: - _, calls = parser.parse_stream_chunk(chunk) + for argument, raw_value, expected_value in ARGUMENT_CASES: + chunks = dict(self.parser_cases(argument, raw_value))[parser_name] + expected = {argument: expected_value} + + with self.subTest( + parser=parser_name, + argument=argument, + mode="non_streaming", + ): + _, calls = non_streaming_parser.parse_non_stream("".join(chunks)) + self.assertEqual(len(calls), 1) + self.assert_arguments(calls[0].parameters, expected) + + with self.subTest( + parser=parser_name, + argument=argument, + mode="streaming", + ): + streaming_parser = FunctionCallParser(self.tools, parser_name) + parameters = "" + for chunk in chunks: + _, calls = streaming_parser.parse_stream_chunk(chunk) + parameters += "".join(call.parameters for call in calls) + _, calls = streaming_parser.parse_stream_end() parameters += "".join(call.parameters for call in calls) - self.assertEqual(json.loads(parameters), expected) - - def test_other_schema_consumers(self): - other_tool = Tool( - type="function", - function=Function( - name="search", - parameters={ - "type": "object", - "properties": {"query": {"type": "string"}}, - }, - ), + self.assert_arguments(parameters, expected) + + self.assertEqual( + self.tools[0].function.parameters, + self.original_parameters, + ) + + def test_detector_tools_use_projected_properties(self): + parser = FunctionCallParser(self.tools, "qwen3_coder") + + self.assertNotIn("properties", self.tools[0].function.parameters) + self.assertEqual( + parser.detector_tools[0].function.parameters["properties"], + PROPERTY_SCHEMAS, ) - parser = FunctionCallParser([self.tools[0], other_tool], "kimi_k2") + + def test_conflicting_property_uses_conservative_string_fallback(self): + parameters = { + "oneOf": [ + {"properties": {"payload": {"type": "object"}}}, + {"properties": {"payload": {"type": "string"}}}, + ] + } + tools = [ + Tool( + type="function", + function=Function(name="convert", parameters=parameters), + ) + ] + parser = FunctionCallParser(tools, "qwen3_coder") + _, calls = parser.parse_non_stream( - "<|tool_calls_section_begin|>" - "<|tool_call_begin|>0" - f"<|tool_call_argument_begin|>{json.dumps(self.expected)}" - "<|tool_call_end|>" - "<|tool_calls_section_end|>" + "" + '{"value":1}' + "" ) - self.assertEqual(calls[0].name, "acme") - def test_streaming_arguments_are_not_buffered(self): - parser = FunctionCallParser(self.tools, "qwen3_coder") - parameters = "" - for chunk in ( - "", - "", - "acme", - ): - _, calls = parser.parse_stream_chunk(chunk) - parameters += "".join(call.parameters for call in calls) - self.assertIn('"kind": "acme"', parameters) + self.assertEqual( + json.loads(calls[0].parameters), + {"payload": '{"value":1}'}, + ) if __name__ == "__main__":