diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 535aaa8b183f..94e5ddf7fbcc 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_tool_parser_property_hints, ) logger = logging.getLogger(__name__) @@ -116,6 +117,31 @@ 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 not isinstance(parameters, dict): + self.detector_tools.append(tool) + continue + + 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 + + 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: @@ -148,7 +174,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 +198,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 +217,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/utils.py b/python/sglang/srt/function_call/utils.py index 0bd0bef4af79..e4232d6348fa 100644 --- a/python/sglang/srt/function_call/utils.py +++ b/python/sglang/srt/function_call/utils.py @@ -304,6 +304,109 @@ def _get_tool_schema(tool: Tool) -> dict: } +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 + } + + +_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]: + """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: + root_schema = schema + + schema = resolve_local_json_schema_refs(schema, root_schema) + if not isinstance(schema, dict): + return {} + + 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 _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 + + 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..9cf71e0df717 --- /dev/null +++ b/test/registered/unit/function_call/test_root_combinator_tool_parsers.py @@ -0,0 +1,342 @@ +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_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") + + +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 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 = disjoint_schema(keyword) + if keyword == "allOf": + for branch in schema[keyword]: + branch.pop("additionalProperties") + + self.assertEqual( + 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_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_tool_parser_property_hints(schema), + {"payload": property_schema}, + ) + + 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_local_refs_are_resolved(self): + schema = { + "$defs": { + "Arguments": { + "oneOf": [{"properties": {"payload": {"$ref": "#/$defs/Payload"}}}] + }, + "Payload": {"type": "object"}, + }, + "$ref": "#/$defs/Arguments", + } + + 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) + + self.assertEqual(schema, original) + + +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 + 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", + [ + "", + "", + f"{raw_value}", + "", + "", + ], + ), + ( + "glm", + [ + "convert\n", + ( + f"{argument}\n" + f"{raw_value}\n" + ), + "", + ], + ), + ( + "glm47", + [ + "convert", + ( + f"{argument}" + f"{raw_value}" + ), + "", + ], + ), + ( + "dots", + [ + '', + f'{raw_value}', + "", + ], + ), + ( + "hunyuan", + [ + "convert", + ( + f"{argument}" + f"{raw_value}" + ), + "", + ], + ), + ( + "mimo", + [ + "", + f"{raw_value}", + "", + ], + ), + ( + "minicpm5", + [ + '', + f'{raw_value}', + "", + ], + ), + ( + "minimax-m2", + [ + '', + f'{raw_value}', + "", + ], + ), + ( + "minimax-m3", + [ + ns + segment + for segment in ( + "", + '', + *minimax_value, + "", + "", + ) + ], + ), + ( + "poolside_v1", + [ + "convert\n", + ( + f"{argument}\n" + f"{raw_value}\n" + ), + "", + ], + ), + ( + "step3", + [ + "<|tool_calls_begin|><|tool_call_begin|>function<|tool_sep|>", + '', + ( + f'{raw_value}' + "" + ), + "<|tool_call_end|><|tool_calls_end|>", + ], + ), + ] + + def assert_arguments(self, serialized: str, expected: dict): + self.assertNotIn("<", serialized) + self.assertEqual(json.loads(serialized), 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) + + 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.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, + ) + + 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( + "" + '{"value":1}' + "" + ) + + self.assertEqual( + json.loads(calls[0].parameters), + {"payload": '{"value":1}'}, + ) + + +if __name__ == "__main__": + unittest.main()