Fix GLM-4.7 MoE Detector complex JSON Schema type parsing (#15753)
This commit is contained in:
@@ -12,6 +12,7 @@ from sglang.srt.function_call.core_types import (
|
|||||||
ToolCallItem,
|
ToolCallItem,
|
||||||
_GetInfoFunc,
|
_GetInfoFunc,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.function_call.utils import infer_type_from_json_schema
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -31,6 +32,14 @@ def get_argument_type(
|
|||||||
) -> Optional[str]:
|
) -> Optional[str]:
|
||||||
"""Get the expected type of a function argument from tool definitions.
|
"""Get the expected type of a function argument from tool definitions.
|
||||||
|
|
||||||
|
Supports complex JSON Schema definitions including:
|
||||||
|
- Direct type field (including type arrays)
|
||||||
|
- anyOf/oneOf: parameter can be any of multiple types
|
||||||
|
- enum: parameter must be one of enum values
|
||||||
|
- allOf: parameter must satisfy all type definitions
|
||||||
|
- properties: inferred as object type
|
||||||
|
- items: inferred as array type
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
func_name: Name of the function/tool
|
func_name: Name of the function/tool
|
||||||
arg_key: Name of the argument
|
arg_key: Name of the argument
|
||||||
@@ -58,7 +67,8 @@ def get_argument_type(
|
|||||||
|
|
||||||
arg_spec = properties.get(arg_key)
|
arg_spec = properties.get(arg_key)
|
||||||
if isinstance(arg_spec, dict):
|
if isinstance(arg_spec, dict):
|
||||||
return arg_spec.get("type")
|
# Use the new type inference function for complex JSON Schema support
|
||||||
|
return infer_type_from_json_schema(arg_spec)
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -243,25 +253,48 @@ class Glm47MoeDetector(BaseFormatDetector):
|
|||||||
tools: List of available tools
|
tools: List of available tools
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Type string: 'string', 'number', or 'object'
|
Type string: 'string', 'number', 'object', 'array', or 'boolean'
|
||||||
"""
|
"""
|
||||||
arg_type = get_argument_type(func_name, key, tools)
|
arg_type = get_argument_type(func_name, key, tools)
|
||||||
if arg_type:
|
if arg_type:
|
||||||
return arg_type
|
return arg_type
|
||||||
|
|
||||||
# Auto-detect type from value (best effort)
|
# Improved auto-detection type from value (best effort)
|
||||||
first_chars = (
|
value_content = self._current_value.strip() if self._current_value else ""
|
||||||
self._current_value.strip()[:10]
|
|
||||||
if self._current_value and self._current_value.strip()
|
if not value_content:
|
||||||
else ""
|
return "string"
|
||||||
)
|
|
||||||
if first_chars:
|
# Try to parse as valid JSON first
|
||||||
first_char = first_chars[0]
|
try:
|
||||||
|
parsed = json.loads(value_content)
|
||||||
|
if isinstance(parsed, dict):
|
||||||
|
return "object"
|
||||||
|
elif isinstance(parsed, list):
|
||||||
|
return "array"
|
||||||
|
elif isinstance(parsed, bool):
|
||||||
|
return "boolean"
|
||||||
|
elif isinstance(parsed, (int, float)):
|
||||||
|
return "number"
|
||||||
|
# For string values, check if they look like numbers
|
||||||
|
elif isinstance(parsed, str):
|
||||||
|
if parsed.isdigit() or (
|
||||||
|
parsed.startswith("-") and parsed[1:].isdigit()
|
||||||
|
):
|
||||||
|
return "number"
|
||||||
|
return "string"
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
# Not valid JSON, try heuristic detection
|
||||||
|
first_char = value_content[0] if value_content else ""
|
||||||
|
|
||||||
if first_char.isdigit() or first_char in ["-", "."]:
|
if first_char.isdigit() or first_char in ["-", "."]:
|
||||||
return "number"
|
return "number"
|
||||||
elif first_char in ["{", "["]:
|
elif first_char in ["{", "["]:
|
||||||
return "object"
|
return "object"
|
||||||
|
elif first_char in ['"', "'"]:
|
||||||
|
return "string"
|
||||||
|
|
||||||
|
# Default to string (safest fallback)
|
||||||
return "string"
|
return "string"
|
||||||
|
|
||||||
def _format_value_complete(self, value: str, value_type: str) -> str:
|
def _format_value_complete(self, value: str, value_type: str) -> str:
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from sglang.srt.function_call.core_types import (
|
|||||||
ToolCallItem,
|
ToolCallItem,
|
||||||
_GetInfoFunc,
|
_GetInfoFunc,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.function_call.utils import infer_type_from_json_schema
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -31,6 +32,14 @@ def get_argument_type(
|
|||||||
) -> Optional[str]:
|
) -> Optional[str]:
|
||||||
"""Get the expected type of a function argument from tool definitions.
|
"""Get the expected type of a function argument from tool definitions.
|
||||||
|
|
||||||
|
Supports complex JSON Schema definitions including:
|
||||||
|
- Direct type field (including type arrays)
|
||||||
|
- anyOf/oneOf: parameter can be any of multiple types
|
||||||
|
- enum: parameter must be one of enum values
|
||||||
|
- allOf: parameter must satisfy all type definitions
|
||||||
|
- properties: inferred as object type
|
||||||
|
- items: inferred as array type
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
func_name: Name of the function/tool
|
func_name: Name of the function/tool
|
||||||
arg_key: Name of the argument
|
arg_key: Name of the argument
|
||||||
@@ -48,7 +57,9 @@ def get_argument_type(
|
|||||||
properties = {}
|
properties = {}
|
||||||
if arg_key not in properties:
|
if arg_key not in properties:
|
||||||
return None
|
return None
|
||||||
return properties[arg_key].get("type", None)
|
|
||||||
|
# Use new type inference function for complex JSON Schema support
|
||||||
|
return infer_type_from_json_schema(properties[arg_key])
|
||||||
|
|
||||||
|
|
||||||
def _convert_to_number(value: str) -> Any:
|
def _convert_to_number(value: str) -> Any:
|
||||||
@@ -216,21 +227,48 @@ class Glm4MoeDetector(BaseFormatDetector):
|
|||||||
tools: List of available tools
|
tools: List of available tools
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Type string: 'string', 'number', or 'object'
|
Type string: 'string', 'number', 'object', 'array', or 'boolean'
|
||||||
"""
|
"""
|
||||||
arg_type = get_argument_type(func_name, key, tools)
|
arg_type = get_argument_type(func_name, key, tools)
|
||||||
if arg_type:
|
if arg_type:
|
||||||
return arg_type
|
return arg_type
|
||||||
|
|
||||||
# Auto-detect type from value (best effort)
|
# Improved auto-detection type from value (best effort)
|
||||||
first_chars = self._current_value.strip()[:10] if self._current_value else ""
|
value_content = self._current_value.strip() if self._current_value else ""
|
||||||
if first_chars:
|
|
||||||
first_char = first_chars[0]
|
if not value_content:
|
||||||
|
return "string"
|
||||||
|
|
||||||
|
# Try to parse as valid JSON first
|
||||||
|
try:
|
||||||
|
parsed = json.loads(value_content)
|
||||||
|
if isinstance(parsed, dict):
|
||||||
|
return "object"
|
||||||
|
elif isinstance(parsed, list):
|
||||||
|
return "array"
|
||||||
|
elif isinstance(parsed, bool):
|
||||||
|
return "boolean"
|
||||||
|
elif isinstance(parsed, (int, float)):
|
||||||
|
return "number"
|
||||||
|
# For string values, check if they look like numbers
|
||||||
|
elif isinstance(parsed, str):
|
||||||
|
if parsed.isdigit() or (
|
||||||
|
parsed.startswith("-") and parsed[1:].isdigit()
|
||||||
|
):
|
||||||
|
return "number"
|
||||||
|
return "string"
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
# Not valid JSON, try heuristic detection
|
||||||
|
first_char = value_content[0] if value_content else ""
|
||||||
|
|
||||||
if first_char.isdigit() or first_char in ["-", "."]:
|
if first_char.isdigit() or first_char in ["-", "."]:
|
||||||
return "number"
|
return "number"
|
||||||
elif first_char in ["{", "["]:
|
elif first_char in ["{", "["]:
|
||||||
return "object"
|
return "object"
|
||||||
|
elif first_char in ['"', "'"]:
|
||||||
|
return "string"
|
||||||
|
|
||||||
|
# Default to string (safest fallback)
|
||||||
return "string"
|
return "string"
|
||||||
|
|
||||||
def _format_value_complete(self, value: str, value_type: str) -> str:
|
def _format_value_complete(self, value: str, value_type: str) -> str:
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from json import JSONDecodeError, JSONDecoder
|
from json import JSONDecodeError, JSONDecoder
|
||||||
from json.decoder import WHITESPACE
|
from json.decoder import WHITESPACE
|
||||||
from typing import Any, List, Literal, Optional, Tuple, Union
|
from typing import Any, Dict, List, Literal, Optional, Tuple, Union
|
||||||
|
|
||||||
import orjson
|
import orjson
|
||||||
import partial_json_parser
|
import partial_json_parser
|
||||||
@@ -101,6 +101,109 @@ def _get_tool_schema(tool: Tool) -> dict:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def infer_type_from_json_schema(schema: Dict[str, Any]) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
Infer the primary type of a parameter from JSON Schema.
|
||||||
|
|
||||||
|
Supports complex JSON Schema structures including:
|
||||||
|
- Direct type field (including type arrays)
|
||||||
|
- anyOf/oneOf: parameter can be any of multiple types
|
||||||
|
- enum: parameter must be one of enum values
|
||||||
|
- allOf: parameter must satisfy all type definitions
|
||||||
|
- properties: inferred as object type
|
||||||
|
- items: inferred as array type
|
||||||
|
|
||||||
|
Args:
|
||||||
|
schema: JSON Schema definition
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Inferred type ('string', 'number', 'object', 'array', etc.) or None
|
||||||
|
"""
|
||||||
|
if not isinstance(schema, dict):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Priority 1: Direct type field (including type arrays)
|
||||||
|
if "type" in schema:
|
||||||
|
type_value = schema["type"]
|
||||||
|
if isinstance(type_value, str):
|
||||||
|
return type_value
|
||||||
|
elif isinstance(type_value, list) and type_value:
|
||||||
|
# Handle type arrays: return first non-null type
|
||||||
|
non_null_types = [t for t in type_value if t != "null"]
|
||||||
|
if non_null_types:
|
||||||
|
return non_null_types[0]
|
||||||
|
return "string" # If only null, default to string
|
||||||
|
|
||||||
|
# Priority 2: Handle anyOf/oneOf
|
||||||
|
if "anyOf" in schema or "oneOf" in schema:
|
||||||
|
schemas = schema.get("anyOf") or schema.get("oneOf")
|
||||||
|
types = []
|
||||||
|
|
||||||
|
if isinstance(schemas, list):
|
||||||
|
for sub_schema in schemas:
|
||||||
|
inferred_type = infer_type_from_json_schema(sub_schema)
|
||||||
|
if inferred_type:
|
||||||
|
types.append(inferred_type)
|
||||||
|
|
||||||
|
if types:
|
||||||
|
# If all types are the same, return unified type
|
||||||
|
if len(set(types)) == 1:
|
||||||
|
return types[0]
|
||||||
|
# When types differ, prioritize string (safest)
|
||||||
|
if "string" in types:
|
||||||
|
return "string"
|
||||||
|
# Otherwise return first type
|
||||||
|
return types[0]
|
||||||
|
|
||||||
|
# Priority 3: Handle enum (infer type from enum values)
|
||||||
|
if "enum" in schema and isinstance(schema["enum"], list):
|
||||||
|
if not schema["enum"]:
|
||||||
|
return "string"
|
||||||
|
|
||||||
|
# Infer type from enum values
|
||||||
|
enum_types = set()
|
||||||
|
for value in schema["enum"]:
|
||||||
|
if value is None:
|
||||||
|
enum_types.add("null")
|
||||||
|
elif isinstance(value, bool):
|
||||||
|
enum_types.add("boolean")
|
||||||
|
elif isinstance(value, int):
|
||||||
|
enum_types.add("integer")
|
||||||
|
elif isinstance(value, float):
|
||||||
|
enum_types.add("number")
|
||||||
|
elif isinstance(value, str):
|
||||||
|
enum_types.add("string")
|
||||||
|
elif isinstance(value, list):
|
||||||
|
enum_types.add("array")
|
||||||
|
elif isinstance(value, dict):
|
||||||
|
enum_types.add("object")
|
||||||
|
|
||||||
|
# If type is uniform, return that type
|
||||||
|
if len(enum_types) == 1:
|
||||||
|
return enum_types.pop()
|
||||||
|
# Mixed types, prioritize string
|
||||||
|
return "string"
|
||||||
|
|
||||||
|
# Priority 4: Handle allOf (must satisfy all types)
|
||||||
|
if "allOf" in schema and isinstance(schema["allOf"], list):
|
||||||
|
schemas = schema["allOf"]
|
||||||
|
for sub_schema in schemas:
|
||||||
|
inferred_type = infer_type_from_json_schema(sub_schema)
|
||||||
|
if inferred_type and inferred_type != "string":
|
||||||
|
return inferred_type
|
||||||
|
return "string"
|
||||||
|
|
||||||
|
# Priority 5: Infer object type
|
||||||
|
if "properties" in schema:
|
||||||
|
return "object"
|
||||||
|
|
||||||
|
# Priority 6: Infer array type
|
||||||
|
if "items" in schema:
|
||||||
|
return "array"
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def get_json_schema_constraint(
|
def get_json_schema_constraint(
|
||||||
tools: List[Tool], tool_choice: Union[ToolChoice, Literal["required"]]
|
tools: List[Tool], tool_choice: Union[ToolChoice, Literal["required"]]
|
||||||
) -> Optional[dict]:
|
) -> Optional[dict]:
|
||||||
|
|||||||
@@ -3,7 +3,11 @@ import unittest
|
|||||||
|
|
||||||
from sglang.srt.entrypoints.openai.protocol import Function, Tool
|
from sglang.srt.entrypoints.openai.protocol import Function, Tool
|
||||||
from sglang.srt.function_call.core_types import StreamingParseResult
|
from sglang.srt.function_call.core_types import StreamingParseResult
|
||||||
from sglang.srt.function_call.glm47_moe_detector import Glm47MoeDetector
|
from sglang.srt.function_call.glm4_moe_detector import Glm4MoeDetector
|
||||||
|
from sglang.srt.function_call.glm47_moe_detector import (
|
||||||
|
Glm47MoeDetector,
|
||||||
|
get_argument_type,
|
||||||
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(1.0, "default")
|
register_cpu_ci(1.0, "default")
|
||||||
@@ -1168,5 +1172,676 @@ class TestGlm47MoeDetector(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
class TestGlm4ComplexJsonSchema(unittest.TestCase):
|
||||||
unittest.main()
|
"""Test complex JSON Schema type inference for GLM function call parsers."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test tools with complex JSON schemas."""
|
||||||
|
self.tools_with_complex_schema = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="search",
|
||||||
|
description="Search for information",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"query": {
|
||||||
|
"description": "Search query, can be a string or a complex object",
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"text": {"type": "string"},
|
||||||
|
"filters": {"type": "object"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
"priority": {"enum": ["low", "medium", "high"]},
|
||||||
|
"options": {
|
||||||
|
"oneOf": [{"type": "string"}, {"type": "number"}]
|
||||||
|
},
|
||||||
|
"config": {
|
||||||
|
"allOf": [
|
||||||
|
{"type": "object"},
|
||||||
|
{"properties": {"timeout": {"type": "number"}}},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"tags": {"type": ["string", "null"]},
|
||||||
|
"data": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"nested": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"value": {"type": "string"}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["query"],
|
||||||
|
},
|
||||||
|
),
|
||||||
|
),
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="get_weather",
|
||||||
|
description="Get weather information",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"location": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "Location to get weather for",
|
||||||
|
},
|
||||||
|
"unit": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "Temperature unit",
|
||||||
|
"enum": ["celsius", "fahrenheit"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["location"],
|
||||||
|
},
|
||||||
|
),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
self.glm4_detector = Glm4MoeDetector()
|
||||||
|
self.glm47_detector = Glm47MoeDetector()
|
||||||
|
|
||||||
|
def test_get_argument_type_simple_type(self):
|
||||||
|
"""Test that get_argument_type correctly handles simple type fields."""
|
||||||
|
result = get_argument_type(
|
||||||
|
"get_weather", "location", self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
self.assertEqual(result, "string")
|
||||||
|
|
||||||
|
def test_get_argument_type_enum_type(self):
|
||||||
|
"""Test that get_argument_type correctly identifies enum as string type."""
|
||||||
|
result = get_argument_type(
|
||||||
|
"get_weather", "unit", self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
# Current implementation returns the direct type field, which is "string" for the enum parameter
|
||||||
|
# But it doesn't handle enum-only schemas properly (without type field)
|
||||||
|
self.assertEqual(result, "string")
|
||||||
|
|
||||||
|
def test_get_argument_type_anyof_type(self):
|
||||||
|
"""Test that get_argument_type correctly handles anyOf type fields."""
|
||||||
|
result = get_argument_type("search", "query", self.tools_with_complex_schema)
|
||||||
|
# anyOf with [{"type": "string"}, {"type": "object", ...}] should return "string"
|
||||||
|
self.assertEqual(result, "string") # Returns first common type
|
||||||
|
|
||||||
|
def test_get_argument_type_oneof_type(self):
|
||||||
|
"""Test that get_argument_type correctly handles oneOf type fields."""
|
||||||
|
result = get_argument_type("search", "options", self.tools_with_complex_schema)
|
||||||
|
# oneOf with [{"type": "string"}, {"type": "number"}] should return "string" (prioritizes string)
|
||||||
|
self.assertEqual(result, "string")
|
||||||
|
|
||||||
|
def test_get_argument_type_allof_type(self):
|
||||||
|
"""Test that get_argument_type correctly handles allOf type fields."""
|
||||||
|
result = get_argument_type("search", "config", self.tools_with_complex_schema)
|
||||||
|
# allOf with [{"type": "object"}, ...] should return "object"
|
||||||
|
self.assertEqual(result, "object")
|
||||||
|
|
||||||
|
def test_get_argument_type_type_array(self):
|
||||||
|
"""Test that get_argument_type correctly handles type arrays."""
|
||||||
|
result = get_argument_type("search", "tags", self.tools_with_complex_schema)
|
||||||
|
# Type arrays should return the first non-null type
|
||||||
|
self.assertEqual(
|
||||||
|
result, "string"
|
||||||
|
) # ["string", "null"] -> "string" (non-null type)
|
||||||
|
|
||||||
|
def test_glm4_detector_with_complex_schema_anyof(self):
|
||||||
|
"""Test GLM4 detector with anyOf schema - should demonstrate current issues."""
|
||||||
|
# This test shows the current behavior with complex schemas
|
||||||
|
text = (
|
||||||
|
"<tool_call>search\n"
|
||||||
|
"<arg_key>query</arg_key>\n<arg_value>Hello world</arg_value>\n"
|
||||||
|
"<arg_key>priority</arg_key>\n<arg_value>medium</arg_value>\n"
|
||||||
|
"</tool_call>"
|
||||||
|
)
|
||||||
|
result = self.glm4_detector.detect_and_parse(
|
||||||
|
text, self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(len(result.calls), 1)
|
||||||
|
self.assertEqual(result.calls[0].name, "search")
|
||||||
|
|
||||||
|
# Parse parameters to check if they are correctly handled
|
||||||
|
params = json.loads(result.calls[0].parameters)
|
||||||
|
self.assertEqual(params["query"], "Hello world")
|
||||||
|
self.assertEqual(params["priority"], "medium")
|
||||||
|
|
||||||
|
def test_glm47_detector_with_complex_schema_anyof(self):
|
||||||
|
"""Test GLM47 detector with anyOf schema - should demonstrate current issues."""
|
||||||
|
# This test shows the current behavior with complex schemas
|
||||||
|
text = (
|
||||||
|
"<tool_call>search"
|
||||||
|
"<arg_key>query</arg_key><arg_value>Hello world</arg_value>"
|
||||||
|
"<arg_key>priority</arg_key><arg_value>medium</arg_value>"
|
||||||
|
"</tool_call>"
|
||||||
|
)
|
||||||
|
result = self.glm47_detector.detect_and_parse(
|
||||||
|
text, self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(len(result.calls), 1)
|
||||||
|
self.assertEqual(result.calls[0].name, "search")
|
||||||
|
|
||||||
|
# Parse parameters to check if they are correctly handled
|
||||||
|
params = json.loads(result.calls[0].parameters)
|
||||||
|
self.assertEqual(params["query"], "Hello world")
|
||||||
|
self.assertEqual(params["priority"], "medium")
|
||||||
|
|
||||||
|
def test_glm4_detector_with_enum_values(self):
|
||||||
|
"""Test GLM4 detector with enum values in complex schema."""
|
||||||
|
text = (
|
||||||
|
"<tool_call>search\n"
|
||||||
|
"<arg_key>query</arg_key>\n<arg_value>test query</arg_value>\n"
|
||||||
|
"<arg_key>priority</arg_key>\n<arg_value>high</arg_value>\n"
|
||||||
|
"</tool_call>"
|
||||||
|
)
|
||||||
|
result = self.glm4_detector.detect_and_parse(
|
||||||
|
text, self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(len(result.calls), 1)
|
||||||
|
self.assertEqual(result.calls[0].name, "search")
|
||||||
|
|
||||||
|
params = json.loads(result.calls[0].parameters)
|
||||||
|
self.assertEqual(params["query"], "test query")
|
||||||
|
self.assertEqual(params["priority"], "high")
|
||||||
|
|
||||||
|
def test_glm47_detector_with_enum_values(self):
|
||||||
|
"""Test GLM47 detector with enum values in complex schema."""
|
||||||
|
text = (
|
||||||
|
"<tool_call>search"
|
||||||
|
"<arg_key>query</arg_key><arg_value>test query</arg_value>"
|
||||||
|
"<arg_key>priority</arg_key><arg_value>high</arg_value>"
|
||||||
|
"</tool_call>"
|
||||||
|
)
|
||||||
|
result = self.glm47_detector.detect_and_parse(
|
||||||
|
text, self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(len(result.calls), 1)
|
||||||
|
self.assertEqual(result.calls[0].name, "search")
|
||||||
|
|
||||||
|
params = json.loads(result.calls[0].parameters)
|
||||||
|
self.assertEqual(params["query"], "test query")
|
||||||
|
self.assertEqual(params["priority"], "high")
|
||||||
|
|
||||||
|
def test_glm4_detector_streaming_with_complex_schema(self):
|
||||||
|
"""Test GLM4 detector streaming with complex schema."""
|
||||||
|
chunks = [
|
||||||
|
"<tool_call>search\n",
|
||||||
|
"<arg_key>query</arg_key>\n<arg_value>nested object</arg_value>\n",
|
||||||
|
"<arg_key>priority</arg_key>\n<arg_value>low</arg_value>\n",
|
||||||
|
"</tool_call>",
|
||||||
|
]
|
||||||
|
tool_calls = []
|
||||||
|
for chunk in chunks:
|
||||||
|
result = self.glm4_detector.parse_streaming_increment(
|
||||||
|
chunk, self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
for tool_call_chunk in result.calls:
|
||||||
|
if (
|
||||||
|
hasattr(tool_call_chunk, "tool_index")
|
||||||
|
and tool_call_chunk.tool_index is not None
|
||||||
|
):
|
||||||
|
while len(tool_calls) <= tool_call_chunk.tool_index:
|
||||||
|
tool_calls.append({"name": "", "parameters": ""})
|
||||||
|
tc = tool_calls[tool_call_chunk.tool_index]
|
||||||
|
if tool_call_chunk.name:
|
||||||
|
tc["name"] = tool_call_chunk.name
|
||||||
|
if tool_call_chunk.parameters:
|
||||||
|
tc["parameters"] += tool_call_chunk.parameters
|
||||||
|
|
||||||
|
self.assertEqual(len(tool_calls), 1)
|
||||||
|
self.assertEqual(tool_calls[0]["name"], "search")
|
||||||
|
|
||||||
|
params = json.loads(tool_calls[0]["parameters"])
|
||||||
|
self.assertEqual(params["query"], "nested object")
|
||||||
|
self.assertEqual(params["priority"], "low")
|
||||||
|
|
||||||
|
def test_glm47_detector_streaming_with_complex_schema(self):
|
||||||
|
"""Test GLM47 detector streaming with complex schema."""
|
||||||
|
chunks = [
|
||||||
|
"<tool_call>search",
|
||||||
|
"<arg_key>query</arg_key><arg_value>nested object</arg_value>",
|
||||||
|
"<arg_key>priority</arg_key><arg_value>low</arg_value>",
|
||||||
|
"</tool_call>",
|
||||||
|
]
|
||||||
|
tool_calls = []
|
||||||
|
for chunk in chunks:
|
||||||
|
result = self.glm47_detector.parse_streaming_increment(
|
||||||
|
chunk, self.tools_with_complex_schema
|
||||||
|
)
|
||||||
|
for tool_call_chunk in result.calls:
|
||||||
|
if (
|
||||||
|
hasattr(tool_call_chunk, "tool_index")
|
||||||
|
and tool_call_chunk.tool_index is not None
|
||||||
|
):
|
||||||
|
while len(tool_calls) <= tool_call_chunk.tool_index:
|
||||||
|
tool_calls.append({"name": "", "parameters": ""})
|
||||||
|
tc = tool_calls[tool_call_chunk.tool_index]
|
||||||
|
if tool_call_chunk.name:
|
||||||
|
tc["name"] = tool_call_chunk.name
|
||||||
|
if tool_call_chunk.parameters:
|
||||||
|
tc["parameters"] += tool_call_chunk.parameters
|
||||||
|
|
||||||
|
self.assertEqual(len(tool_calls), 1)
|
||||||
|
self.assertEqual(tool_calls[0]["name"], "search")
|
||||||
|
|
||||||
|
params = json.loads(tool_calls[0]["parameters"])
|
||||||
|
self.assertEqual(params["query"], "nested object")
|
||||||
|
self.assertEqual(params["priority"], "low")
|
||||||
|
|
||||||
|
def test_type_inference_issue_reproduction(self):
|
||||||
|
"""Reproduce the issue where complex JSON schemas are not properly handled."""
|
||||||
|
# This test demonstrates the current limitations
|
||||||
|
complex_tools = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="complex_function",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"complex_param": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"value": {"type": "string"}},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"enum_param": {"enum": ["option1", "option2", "option3"]},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Test that get_argument_type returns appropriate types for complex schemas
|
||||||
|
anyof_result = get_argument_type(
|
||||||
|
"complex_function", "complex_param", complex_tools
|
||||||
|
)
|
||||||
|
enum_result = get_argument_type("complex_function", "enum_param", complex_tools)
|
||||||
|
|
||||||
|
# Verify complex schema types are correctly inferred
|
||||||
|
self.assertEqual(anyof_result, "string") # anyOf prioritizes string type
|
||||||
|
self.assertEqual(enum_result, "string") # enum values are strings
|
||||||
|
|
||||||
|
def test_expected_behavior_for_complex_schemas(self):
|
||||||
|
"""Test cases that should work but currently fail - demonstrating the issue."""
|
||||||
|
# This test shows what the behavior SHOULD be after the fix
|
||||||
|
complex_tools = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="complex_function",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"complex_param": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"value": {"type": "string"}},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"enum_param": {"enum": ["option1", "option2", "option3"]},
|
||||||
|
"oneof_param": {
|
||||||
|
"oneOf": [{"type": "string"}, {"type": "number"}]
|
||||||
|
},
|
||||||
|
"allof_param": {
|
||||||
|
"allOf": [
|
||||||
|
{"type": "object"},
|
||||||
|
{"properties": {"timeout": {"type": "number"}}},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# These assertions represent the EXPECTED behavior after implementing RFC improvements
|
||||||
|
# Currently they will fail, demonstrating the issue
|
||||||
|
anyof_result = get_argument_type(
|
||||||
|
"complex_function", "complex_param", complex_tools
|
||||||
|
)
|
||||||
|
enum_result = get_argument_type("complex_function", "enum_param", complex_tools)
|
||||||
|
oneof_result = get_argument_type(
|
||||||
|
"complex_function", "oneof_param", complex_tools
|
||||||
|
)
|
||||||
|
allof_result = get_argument_type(
|
||||||
|
"complex_function", "allof_param", complex_tools
|
||||||
|
)
|
||||||
|
|
||||||
|
# These should pass after implementing the RFC improvements, but will currently fail
|
||||||
|
# This demonstrates the issue exists
|
||||||
|
self.assertIsNotNone(
|
||||||
|
anyof_result, "anyOf should return a type after RFC implementation"
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
enum_result,
|
||||||
|
"string",
|
||||||
|
"enum should return 'string' type after RFC implementation",
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(
|
||||||
|
oneof_result, "oneOf should return a type after RFC implementation"
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(
|
||||||
|
allof_result, "allOf should return a type after RFC implementation"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_complex_schema_type_inference_scenarios(self):
|
||||||
|
"""Test various complex schema scenarios mentioned in the RFC."""
|
||||||
|
# Create tools with different complex schema structures
|
||||||
|
complex_schema_tools = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="search_complex",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
# anyOf example - parameter can be string or object
|
||||||
|
"query": {
|
||||||
|
"description": "Search query, can be a string or a complex object",
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"text": {"type": "string"},
|
||||||
|
"filters": {"type": "object"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
# oneOf example - parameter must be one of the specified types
|
||||||
|
"priority": {
|
||||||
|
"oneOf": [{"type": "string"}, {"type": "integer"}]
|
||||||
|
},
|
||||||
|
# enum example - parameter must be one of the enum values
|
||||||
|
"category": {"enum": ["news", "sports", "tech"]},
|
||||||
|
# allOf example - parameter must satisfy all schemas
|
||||||
|
"config": {
|
||||||
|
"allOf": [
|
||||||
|
{"type": "object"},
|
||||||
|
{"properties": {"timeout": {"type": "number"}}},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
# Type array example
|
||||||
|
"tags": {"type": ["string", "null"]},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
),
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="get_data",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
# Complex nested anyOf
|
||||||
|
"input": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{"type": "number"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"type": {"type": "string"},
|
||||||
|
"value": {},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
# Test each complex type scenario
|
||||||
|
query_type = get_argument_type("search_complex", "query", complex_schema_tools)
|
||||||
|
priority_type = get_argument_type(
|
||||||
|
"search_complex", "priority", complex_schema_tools
|
||||||
|
)
|
||||||
|
category_type = get_argument_type(
|
||||||
|
"search_complex", "category", complex_schema_tools
|
||||||
|
)
|
||||||
|
config_type = get_argument_type(
|
||||||
|
"search_complex", "config", complex_schema_tools
|
||||||
|
)
|
||||||
|
tags_type = get_argument_type("search_complex", "tags", complex_schema_tools)
|
||||||
|
input_type = get_argument_type("get_data", "input", complex_schema_tools)
|
||||||
|
|
||||||
|
# All of these should return appropriate types according to RFC
|
||||||
|
self.assertEqual(query_type, "string") # anyOf: string | object -> string
|
||||||
|
self.assertEqual(priority_type, "string") # oneOf: string | integer -> string
|
||||||
|
self.assertEqual(
|
||||||
|
category_type, "string"
|
||||||
|
) # enum: ["news", "sports", "tech"] -> string
|
||||||
|
self.assertEqual(config_type, "object") # allOf with object -> object
|
||||||
|
self.assertEqual(
|
||||||
|
tags_type, "string"
|
||||||
|
) # type array: ["string", "null"] -> string
|
||||||
|
self.assertEqual(
|
||||||
|
input_type, "string"
|
||||||
|
) # nested anyOf: string | number | object -> string
|
||||||
|
|
||||||
|
def test_glm4_detector_type_handling_with_complex_schema(self):
|
||||||
|
"""Test how GLM4 detector handles type inference for complex schemas in practice."""
|
||||||
|
complex_tools = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="complex_search",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"query": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"text": {"type": "string"}},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"category": {"enum": ["tech", "news", "sports"]},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Test with string value for anyOf parameter
|
||||||
|
text = (
|
||||||
|
"<tool_call>complex_search\n"
|
||||||
|
"<arg_key>query</arg_key>\n<arg_value>test search</arg_value>\n"
|
||||||
|
"<arg_key>category</arg_key>\n<arg_value>tech</arg_value>\n"
|
||||||
|
"</tool_call>"
|
||||||
|
)
|
||||||
|
result = self.glm4_detector.detect_and_parse(text, complex_tools)
|
||||||
|
|
||||||
|
self.assertEqual(len(result.calls), 1)
|
||||||
|
self.assertEqual(result.calls[0].name, "complex_search")
|
||||||
|
|
||||||
|
params = json.loads(result.calls[0].parameters)
|
||||||
|
self.assertEqual(params["query"], "test search")
|
||||||
|
self.assertEqual(params["category"], "tech")
|
||||||
|
|
||||||
|
def test_glm47_detector_type_handling_with_complex_schema(self):
|
||||||
|
"""Test how GLM47 detector handles type inference for complex schemas in practice."""
|
||||||
|
complex_tools = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="complex_search",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"query": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"text": {"type": "string"}},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"category": {"enum": ["tech", "news", "sports"]},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Test with string value for anyOf parameter
|
||||||
|
text = (
|
||||||
|
"<tool_call>complex_search"
|
||||||
|
"<arg_key>query</arg_key><arg_value>test search</arg_value>"
|
||||||
|
"<arg_key>category</arg_key><arg_value>tech</arg_value>"
|
||||||
|
"</tool_call>"
|
||||||
|
)
|
||||||
|
result = self.glm47_detector.detect_and_parse(text, complex_tools)
|
||||||
|
|
||||||
|
self.assertEqual(len(result.calls), 1)
|
||||||
|
self.assertEqual(result.calls[0].name, "complex_search")
|
||||||
|
|
||||||
|
params = json.loads(result.calls[0].parameters)
|
||||||
|
self.assertEqual(params["query"], "test search")
|
||||||
|
self.assertEqual(params["category"], "tech")
|
||||||
|
|
||||||
|
def test_streaming_with_complex_schema_type_inference(self):
|
||||||
|
"""Test streaming behavior with complex schema type inference."""
|
||||||
|
complex_tools = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="stream_test",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"data": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"value": {"type": "string"}},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"status": {"enum": ["active", "inactive"]},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Test GLM4 detector streaming
|
||||||
|
chunks = [
|
||||||
|
"<tool_call>stream_test\n",
|
||||||
|
"<arg_key>data</arg_key>\n<arg_value>nested data</arg_value>\n",
|
||||||
|
"<arg_key>status</arg_key>\n<arg_value>active</arg_value>\n",
|
||||||
|
"</tool_call>",
|
||||||
|
]
|
||||||
|
tool_calls = []
|
||||||
|
for chunk in chunks:
|
||||||
|
result = self.glm4_detector.parse_streaming_increment(chunk, complex_tools)
|
||||||
|
for tool_call_chunk in result.calls:
|
||||||
|
if (
|
||||||
|
hasattr(tool_call_chunk, "tool_index")
|
||||||
|
and tool_call_chunk.tool_index is not None
|
||||||
|
):
|
||||||
|
while len(tool_calls) <= tool_call_chunk.tool_index:
|
||||||
|
tool_calls.append({"name": "", "parameters": ""})
|
||||||
|
tc = tool_calls[tool_call_chunk.tool_index]
|
||||||
|
if tool_call_chunk.name:
|
||||||
|
tc["name"] = tool_call_chunk.name
|
||||||
|
if tool_call_chunk.parameters:
|
||||||
|
tc["parameters"] += tool_call_chunk.parameters
|
||||||
|
|
||||||
|
self.assertEqual(len(tool_calls), 1)
|
||||||
|
self.assertEqual(tool_calls[0]["name"], "stream_test")
|
||||||
|
|
||||||
|
params = json.loads(tool_calls[0]["parameters"])
|
||||||
|
self.assertEqual(params["data"], "nested data")
|
||||||
|
self.assertEqual(params["status"], "active")
|
||||||
|
|
||||||
|
def test_streaming_with_complex_schema_type_inference_glm47(self):
|
||||||
|
"""Test GLM47 streaming behavior with complex schema type inference."""
|
||||||
|
complex_tools = [
|
||||||
|
Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="stream_test",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"data": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string"},
|
||||||
|
{
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"value": {"type": "string"}},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"status": {"enum": ["active", "inactive"]},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Test GLM47 detector streaming
|
||||||
|
chunks = [
|
||||||
|
"<tool_call>stream_test",
|
||||||
|
"<arg_key>data</arg_key><arg_value>nested data</arg_value>",
|
||||||
|
"<arg_key>status</arg_key><arg_value>active</arg_value>",
|
||||||
|
"</tool_call>",
|
||||||
|
]
|
||||||
|
tool_calls = []
|
||||||
|
for chunk in chunks:
|
||||||
|
result = self.glm47_detector.parse_streaming_increment(chunk, complex_tools)
|
||||||
|
for tool_call_chunk in result.calls:
|
||||||
|
if (
|
||||||
|
hasattr(tool_call_chunk, "tool_index")
|
||||||
|
and tool_call_chunk.tool_index is not None
|
||||||
|
):
|
||||||
|
while len(tool_calls) <= tool_call_chunk.tool_index:
|
||||||
|
tool_calls.append({"name": "", "parameters": ""})
|
||||||
|
tc = tool_calls[tool_call_chunk.tool_index]
|
||||||
|
if tool_call_chunk.name:
|
||||||
|
tc["name"] = tool_call_chunk.name
|
||||||
|
if tool_call_chunk.parameters:
|
||||||
|
tc["parameters"] += tool_call_chunk.parameters
|
||||||
|
|
||||||
|
self.assertEqual(len(tool_calls), 1)
|
||||||
|
self.assertEqual(tool_calls[0]["name"], "stream_test")
|
||||||
|
|
||||||
|
params = json.loads(tool_calls[0]["parameters"])
|
||||||
|
self.assertEqual(params["data"], "nested data")
|
||||||
|
self.assertEqual(params["status"], "active")
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user