Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 31 additions & 3 deletions python/sglang/srt/function_call/function_call_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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(
Expand Down
103 changes: 103 additions & 0 deletions python/sglang/srt/function_call/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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] = {}
Comment on lines +404 to +405

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Preserve common types across non-identical branch schemas

When every anyOf/oneOf branch declares the same argument as an object or array but applies different nested constraints, the schemas are unequal and this collapses the property hint to {}. Detectors such as Qwen then infer string, so a valid nested argument remains JSON-encoded text even though every alternative agrees on its container type—a common case for discriminated unions with branch-specific payload fields. Combine the candidate schemas under the relevant combinator, or at least preserve their common inferred type, instead of requiring exact schema equality.

Useful? React with 👍 / 👎.


return properties


def infer_type_from_json_schema(schema: Dict[str, Any]) -> Optional[str]:
"""
Infer the primary type of a parameter from JSON Schema.
Expand Down
Loading
Loading