Skip to content
Merged
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
8 changes: 5 additions & 3 deletions python/sglang/srt/function_call/dots_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,10 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import _is_complete_json
from sglang.srt.function_call.utils import (
_is_complete_json,
get_schema_properties,
)

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -137,10 +140,9 @@ def _tool_schema(name: str, tools: list[Tool]) -> tuple[dict, dict]:
schema = tool.function.parameters
if not isinstance(schema, dict):
break
properties = schema.get("properties", {})
defs = schema.get("$defs", {})
return (
properties if isinstance(properties, dict) else {},
get_schema_properties(schema),
defs if isinstance(defs, dict) else {},
)
return {}, {}
Expand Down
15 changes: 5 additions & 10 deletions python/sglang/srt/function_call/glm47_moe_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
_GetInfoFunc,
)
from sglang.srt.function_call.utils import (
get_schema_properties,
infer_type_from_json_schema,
safe_literal_eval,
)
Expand Down Expand Up @@ -79,15 +80,8 @@ def get_argument_type(

# Get parameters safely using getattr
params = getattr(tool.function, "parameters", None)
if not isinstance(params, dict):
return None

# Navigate to the type using dict.get() for safe access
properties = params.get("properties")
if not isinstance(properties, dict):
return None

arg_spec = properties.get(arg_key)
arg_spec = get_schema_properties(params).get(arg_key)
if isinstance(arg_spec, dict):
# Use the new type inference function for complex JSON Schema support
return infer_type_from_json_schema(arg_spec)
Expand Down Expand Up @@ -613,8 +607,9 @@ 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:
# Need to close brace
elif not self._is_first_param and not self._sent_empty_object:
# The streamed outer `{` is only closed here; a trailing "}" in
# _last_arguments may belong to a nested object value.
calls.append(
ToolCallItem(
tool_index=self.current_tool_id,
Expand Down
9 changes: 5 additions & 4 deletions python/sglang/srt/function_call/glm4_moe_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
_GetInfoFunc,
)
from sglang.srt.function_call.utils import (
get_schema_properties,
infer_type_from_json_schema,
safe_literal_eval,
)
Expand Down Expand Up @@ -54,9 +55,7 @@ def get_argument_type(
if func_name not in name2tool:
return None
tool = name2tool[func_name]
properties = (tool.function.parameters or {}).get("properties", {})
if not isinstance(properties, dict):
properties = {}
properties = get_schema_properties(tool.function.parameters)
if arg_key not in properties:
return None

Expand Down Expand Up @@ -572,7 +571,9 @@ def parse_streaming_increment(
self.streamed_args_for_tool[
self.current_tool_id
] += empty_object
elif not self._last_arguments.endswith("}"):
else:
# The streamed outer `{` is only closed here; a
# trailing "}" may belong to a nested object value.
closing_brace = "}"
calls.append(
ToolCallItem(
Expand Down
3 changes: 2 additions & 1 deletion python/sglang/srt/function_call/hunyuan_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import get_schema_properties

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -176,7 +177,7 @@ def _get_arg_schema(
if tool.function.name == function_name:
if tool.function.parameters is None:
return {}
return tool.function.parameters.get("properties", {}).get(arg_key, {})
return get_schema_properties(tool.function.parameters).get(arg_key, {})
return {}

@staticmethod
Expand Down
3 changes: 2 additions & 1 deletion python/sglang/srt/function_call/kimik2_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import get_schema_properties

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -144,7 +145,7 @@ def _infer_tool_name(self, tools: List[Tool], function_args: str = None):
best_score = -1
for tool in tools:
params = tool.function.parameters or {}
props = set(params.get("properties", {}).keys())
props = set(get_schema_properties(params).keys())
if not props:
continue
overlap = len(arg_keys & props)
Expand Down
7 changes: 5 additions & 2 deletions python/sglang/srt/function_call/mimo_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,10 @@
from sglang.srt.environ import envs
from sglang.srt.function_call.base_format_detector import BaseFormatDetector
from sglang.srt.function_call.core_types import StreamingParseResult, _GetInfoFunc
from sglang.srt.function_call.utils import safe_literal_eval
from sglang.srt.function_call.utils import (
get_schema_properties,
safe_literal_eval,
)

logger = logging.getLogger(__name__)

Expand All @@ -31,7 +34,7 @@ def _get_param_type(func_name: str, param_name: str, tools: List[Tool]) -> str:
"""Get parameter type from tool schema."""
for tool in tools:
if tool.function.name == func_name:
props = tool.function.parameters.get("properties", {})
props = get_schema_properties(tool.function.parameters)
if param_name in props:
return props[param_name].get("type", "string")
return "string"
Expand Down
12 changes: 6 additions & 6 deletions python/sglang/srt/function_call/minicpm5_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,10 @@
StreamingParseResult,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import safe_literal_eval
from sglang.srt.function_call.utils import (
get_schema_properties,
safe_literal_eval,
)

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -38,7 +41,7 @@ def get_argument_type(
params = tool.function.parameters or {}
if not isinstance(params, dict):
return None
return params.get("properties", {}).get(arg_key, {}).get("type")
return get_schema_properties(params).get(arg_key, {}).get("type")


def parse_arguments(json_value):
Expand Down Expand Up @@ -84,10 +87,7 @@ def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult
name_to_required = {}
for name, t in name_to_tool.items():
params = t.function.parameters or {}
props = (
(params.get("properties", {}) or {}) if isinstance(params, dict) else {}
)
name_to_allowed_props[name] = set(props.keys())
name_to_allowed_props[name] = set(get_schema_properties(params).keys())
req = params.get("required", []) if isinstance(params, dict) else []
try:
name_to_required[name] = set(req)
Expand Down
5 changes: 3 additions & 2 deletions python/sglang/srt/function_call/minimax_m2.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import get_schema_properties

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -508,8 +509,8 @@ def _parse_parameter(
for tool in tools:
if tool.function.name == fname and tool.function.parameters is not None:
parameters = tool.function.parameters
if isinstance(parameters, dict) and "properties" in parameters:
param_config = parameters["properties"]
if isinstance(parameters, dict):
param_config = get_schema_properties(parameters)
break

param_type = self._get_param_types_from_config(pname, param_config)
Expand Down
5 changes: 3 additions & 2 deletions python/sglang/srt/function_call/minimax_m3.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import get_schema_properties

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -374,8 +375,8 @@ def _get_child_schema(
if self._schema_has_type(parent_schema, ("array",)) and child_tag == "item":
return self._get_array_item_schema(parent_schema, parent_value)

properties = parent_schema.get("properties")
if isinstance(properties, dict) and child_tag in properties:
properties = get_schema_properties(parent_schema)
if child_tag in properties:
child_schema = properties[child_tag]
return child_schema if isinstance(child_schema, dict) else None

Expand Down
7 changes: 5 additions & 2 deletions python/sglang/srt/function_call/poolside_v1_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,10 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import safe_literal_eval
from sglang.srt.function_call.utils import (
get_schema_properties,
safe_literal_eval,
)


class _ParseState(Enum):
Expand Down Expand Up @@ -150,7 +153,7 @@ def _get_param_schema(
and tool.function.name == func_name
and isinstance(tool.function.parameters, dict)
):
return tool.function.parameters.get("properties", {})
return get_schema_properties(tool.function.parameters)
except AttributeError:
continue
return {}
Expand Down
8 changes: 5 additions & 3 deletions python/sglang/srt/function_call/qwen3_coder_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
_GetInfoFunc,
)
from sglang.srt.function_call.utils import (
get_schema_properties,
infer_type_from_json_schema,
safe_literal_eval,
)
Expand Down Expand Up @@ -80,9 +81,10 @@ def _get_arguments_config(
except AttributeError:
return {}

if isinstance(params, dict) and "properties" in params:
return params["properties"]
elif isinstance(params, dict):
if isinstance(params, dict):
properties = get_schema_properties(params)
if properties or "properties" in params:
return properties
return params
else:
return {}
Expand Down
6 changes: 2 additions & 4 deletions python/sglang/srt/function_call/spark25_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import get_schema_properties

TOOL_CALL_BEGIN = "<tool_call>"
TOOL_CALL_END = "</tool_call>"
Expand Down Expand Up @@ -47,10 +48,7 @@ def _get_param_type(tools: list[Tool], function_name: str, param_name: str) -> s
parameters = getattr(function, "parameters", None)
if not isinstance(parameters, dict):
continue
properties = parameters.get("properties")
if not isinstance(properties, dict):
continue
definition = properties.get(param_name)
definition = get_schema_properties(parameters).get(param_name)
if isinstance(definition, dict) and isinstance(definition.get("type"), str):
return definition["type"]
return "string"
Expand Down
7 changes: 5 additions & 2 deletions python/sglang/srt/function_call/step3_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,10 @@
ToolCallItem,
_GetInfoFunc,
)
from sglang.srt.function_call.utils import safe_literal_eval
from sglang.srt.function_call.utils import (
get_schema_properties,
safe_literal_eval,
)

logger = logging.getLogger(__name__)

Expand All @@ -22,7 +25,7 @@ def get_argument_type(func_name: str, arg_key: str, defined_tools: List[Tool]) -
return None
tool = name2tool[func_name]
parameters = tool.function.parameters or {}
properties = parameters.get("properties", {})
properties = get_schema_properties(parameters)
if arg_key not in properties:
return None
return properties[arg_key].get("type", None)
Expand Down
19 changes: 19 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,25 @@ def _get_tool_schema(tool: Tool) -> dict:
}


def get_schema_properties(schema: Any) -> Dict[str, Any]:
"""Top-level ``properties`` of a tool ``parameters`` schema, descending
into ``anyOf``/``oneOf``/``allOf`` branches when the top level declares
none (legal JSON Schema, e.g. discriminated-union arguments)."""
if not isinstance(schema, dict):
return {}
properties = schema.get("properties")
if isinstance(properties, dict):
return properties
merged: Dict[str, Any] = {}
for keyword in ("anyOf", "oneOf", "allOf"):
branches = schema.get(keyword)
if isinstance(branches, list):
for branch in branches:
for key, value in get_schema_properties(branch).items():
merged.setdefault(key, value)
return merged


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