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(">acme", ">other").replace(
+ '{"value":"hello"}', '{"x":1}'
+ )
+ for chunk in chunks
+ ]
+ if parser_name == "minimax-m3":
+ other_chunks = [
+ ns + segment
+ for segment in (
+ "",
+ '',
+ "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(">acme", ">other").replace(
- '{"value":"hello"}', '{"x":1}'
+ '{"value":"hello"}', '{ "x": 1.00 }'
)
for chunk in chunks
]
@@ -244,13 +244,13 @@ def test_parsers(self):
'',
"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"{argument}>"]
+ if argument.endswith("List"):
+ minimax_value = [f"<{argument}>"]
+ for item in json.loads(raw_value):
+ minimax_value.extend((f"- {item}", "
"))
+ minimax_value.append(f"{argument}>")
+
+ 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(">acme", ">other").replace(
- '{"value":"hello"}', '{ "x": 1.00 }'
- )
- for chunk in chunks
- ]
- if parser_name == "minimax-m3":
- other_chunks = [
- ns + segment
- for segment in (
- "",
- '',
- "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__":