Support non-strict GLM47 tool calls with EBNF constraints (#38890)

Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
Yuxuan Zhang
2026-09-15 16:41:48 +08:00
committed by GitHub
co-authored by Xinyuan Tong
parent 2c0a70960c
commit 17ba2c2e7c
14 changed files with 834 additions and 22 deletions
@@ -269,7 +269,7 @@ class BaseGrammarBackend:
grammar = self.dispatch_json(key_string) grammar = self.dispatch_json(key_string)
elif key_type == "regex": elif key_type == "regex":
grammar = self.dispatch_regex(key_string) grammar = self.dispatch_regex(key_string)
elif key_type == "ebnf": elif key_type in ("ebnf", "full_assistant_ebnf"):
grammar = self.dispatch_ebnf(key_string) grammar = self.dispatch_ebnf(key_string)
elif key_type == "structural_tag": elif key_type == "structural_tag":
grammar = self.dispatch_structural_tag(key_string) grammar = self.dispatch_structural_tag(key_string)
@@ -160,7 +160,12 @@ class GrammarManager:
elif req.sampling_params.regex is not None: elif req.sampling_params.regex is not None:
key = ("regex", req.sampling_params.regex) key = ("regex", req.sampling_params.regex)
elif req.sampling_params.ebnf is not None: elif req.sampling_params.ebnf is not None:
key = ("ebnf", req.sampling_params.ebnf) key_type = (
"full_assistant_ebnf"
if req.sampling_params.ebnf_full_assistant
else "ebnf"
)
key = (key_type, req.sampling_params.ebnf)
elif req.sampling_params.structural_tag is not None: elif req.sampling_params.structural_tag is not None:
key = ("structural_tag", req.sampling_params.structural_tag) key = ("structural_tag", req.sampling_params.structural_tag)
@@ -352,4 +352,6 @@ class ReasonerGrammarBackend(BaseGrammarBackend):
ret = self.grammar_backend._init_value_dispatch(key, reasoning) ret = self.grammar_backend._init_value_dispatch(key, reasoning)
if ret is None or isinstance(ret, InvalidGrammarObject): if ret is None or isinstance(ret, InvalidGrammarObject):
return ret return ret
if key[0] == "full_assistant_ebnf":
return ret
return self._make_grammar_object(ret, reasoning) return self._make_grammar_object(ret, reasoning)
@@ -256,6 +256,8 @@ StructuralTagResponseFormat: TypeAlias = Union[
ToolCallConstraint: TypeAlias = Union[ ToolCallConstraint: TypeAlias = Union[
Tuple[Literal["structural_tag"], StructuralTagResponseFormat], Tuple[Literal["structural_tag"], StructuralTagResponseFormat],
Tuple[Literal["json_schema"], Any], # json_schema can be dict/str/None Tuple[Literal["json_schema"], Any], # json_schema can be dict/str/None
Tuple[Literal["ebnf"], str],
Tuple[Literal["full_assistant_ebnf"], str],
] ]
@@ -1174,8 +1176,9 @@ class ChatCompletionRequest(BaseModel):
) )
if tool_call_constraint and has_existing_constraints: if tool_call_constraint and has_existing_constraints:
if self.tool_choice == "required" or isinstance( if tool_call_constraint[0] != "full_assistant_ebnf" and (
self.tool_choice, ToolChoice self.tool_choice == "required"
or isinstance(self.tool_choice, ToolChoice)
): ):
raise ValueError( raise ValueError(
"tool_choice 'required' or a named tool cannot be combined with " "tool_choice 'required' or a named tool cannot be combined with "
@@ -1193,6 +1196,9 @@ class ChatCompletionRequest(BaseModel):
sampling_params[constraint_type] = convert_json_schema_to_str( sampling_params[constraint_type] = convert_json_schema_to_str(
constraint_value # type: ignore constraint_value # type: ignore
) )
elif constraint_type == "full_assistant_ebnf":
sampling_params["ebnf"] = constraint_value
sampling_params["ebnf_full_assistant"] = True
else: else:
sampling_params[constraint_type] = constraint_value sampling_params[constraint_type] = constraint_value
@@ -1883,6 +1889,12 @@ class ResponsesRequest(BaseModel):
or params.get("json_schema") or params.get("json_schema")
) )
if tool_call_constraint and has_existing_constraints: if tool_call_constraint and has_existing_constraints:
if tool_call_constraint[0] == "full_assistant_ebnf":
# Explicit output constraints take precedence over the default EBNF.
logger.warning(
"Constrained decoding is not compatible with tool calls."
)
return params
# Refuse rather than silently drop the tool-call grammar. # Refuse rather than silently drop the tool-call grammar.
raise ValueError( raise ValueError(
"Cannot combine tool calls with constrained decoding " "Cannot combine tool calls with constrained decoding "
@@ -1897,6 +1909,9 @@ class ResponsesRequest(BaseModel):
if hasattr(constraint_value, "model_dump") if hasattr(constraint_value, "model_dump")
else constraint_value else constraint_value
) )
elif constraint_type == "full_assistant_ebnf":
params["ebnf"] = constraint_value
params["ebnf_full_assistant"] = True
else: else:
params[constraint_type] = constraint_value params[constraint_type] = constraint_value
@@ -1233,11 +1233,27 @@ class OpenAIServingChat(OpenAIServingBase):
xgrammar_reasoning = thinking_mode and (self.reasoning_parser is None) xgrammar_reasoning = thinking_mode and (self.reasoning_parser is None)
tool_call_constraint = None tool_call_constraint = None
effective_tools = self._effective_tools(request)
glm_constraint = self.tool_call_parser == "glm47" and not any(
tool.function.strict for tool in effective_tools
)
if glm_constraint:
enable_thinking = (request.chat_template_kwargs or {}).get(
"enable_thinking"
)
parser = FunctionCallParser(request.tools or [], self.tool_call_parser)
tool_call_constraint = parser.get_structure_constraint(
request.tool_choice,
parallel_tool_calls=request.parallel_tool_calls,
thinking_mode=True
if enable_thinking is None
else bool(enable_thinking),
)
# Apply chat template and its stop strings # Apply chat template and its stop strings
tools = None tools = None
tool_call_stop = None tool_call_stop = None
required_parsed_natively = False required_parsed_natively = glm_constraint
effective_tools = self._effective_tools(request)
if effective_tools and request.tool_choice != "none": if effective_tools and request.tool_choice != "none":
request.skip_special_tokens = False request.skip_special_tokens = False
if not isinstance(request.tool_choice, str): if not isinstance(request.tool_choice, str):
@@ -1248,7 +1264,7 @@ class OpenAIServingChat(OpenAIServingBase):
] or None ] or None
elif request.tools: elif request.tools:
tools = [item.model_dump() for item in request.tools] tools = [item.model_dump() for item in request.tools]
if self.tool_call_parser: if self.tool_call_parser and not glm_constraint:
parser = FunctionCallParser( parser = FunctionCallParser(
effective_tools, effective_tools,
self.tool_call_parser, self.tool_call_parser,
@@ -22,7 +22,11 @@ from sglang.srt.function_call.deepseekv32_detector import DeepSeekV32Detector
from sglang.srt.function_call.dots_detector import DotsToolDetector from sglang.srt.function_call.dots_detector import DotsToolDetector
from sglang.srt.function_call.gemma4_detector import Gemma4Detector from sglang.srt.function_call.gemma4_detector import Gemma4Detector
from sglang.srt.function_call.gigachat3_detector import GigaChat3Detector from sglang.srt.function_call.gigachat3_detector import GigaChat3Detector
from sglang.srt.function_call.glm4_moe_detector import Glm4MoeDetector from sglang.srt.function_call.glm4_moe_detector import (
Glm4MoeDetector,
GlmSpecialTokenConfig,
generate_glm_grammar,
)
from sglang.srt.function_call.glm47_moe_detector import Glm47MoeDetector from sglang.srt.function_call.glm47_moe_detector import Glm47MoeDetector
from sglang.srt.function_call.gpt_oss_detector import GptOssDetector from sglang.srt.function_call.gpt_oss_detector import GptOssDetector
from sglang.srt.function_call.hermes_detector import HermesDetector from sglang.srt.function_call.hermes_detector import HermesDetector
@@ -119,6 +123,10 @@ class FunctionCallParser:
else: else:
raise ValueError(f"Unsupported tool_call_parser: {tool_call_parser}") raise ValueError(f"Unsupported tool_call_parser: {tool_call_parser}")
if isinstance(detector, Glm47MoeDetector):
detector.use_full_assistant_constraint = not any(
tool.function.strict for tool in tools
)
self.detector = detector self.detector = detector
self.tools = tools self.tools = tools
self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get() self.tool_strict_level = envs.SGLANG_TOOL_STRICT_LEVEL.get()
@@ -272,8 +280,35 @@ class FunctionCallParser:
or self.tool_strict_level >= ToolStrictLevel.FUNCTION or self.tool_strict_level >= ToolStrictLevel.FUNCTION
) )
# Highest priority: model-native structural_tag when available.
try: try:
if (
isinstance(self.detector, Glm47MoeDetector)
and self.detector.use_full_assistant_constraint
):
functions = (
[
tool.function
for tool in self.tools
if not isinstance(tool_choice, ToolChoice)
or tool.function.name == tool_choice.function.name
]
if self.tools and tool_choice != "none"
else None
)
return (
"full_assistant_ebnf",
generate_glm_grammar(
enable_thinking=thinking_mode,
functions=functions,
special_tokens=GlmSpecialTokenConfig(),
chat_template_version="glm47",
accommodate_chat_template=True,
allow_multiple_assistant_turns=False,
required=is_required,
parallel_tool_calls=parallel_tool_calls,
),
)
# Highest priority: model-native structural_tag when available.
if tool_choice == "auto" and not should_constrain_auto: if tool_choice == "auto" and not should_constrain_auto:
structural_tag = self.detector.get_auto_tool_call_structural_tag( structural_tag = self.detector.get_auto_tool_call_structural_tag(
tools=self.tools, tools=self.tools,
@@ -307,6 +307,7 @@ class Glm47MoeDetector(BaseFormatDetector):
def __init__(self): def __init__(self):
super().__init__() super().__init__()
self.use_full_assistant_constraint = False
self.bot_token = "<tool_call>" self.bot_token = "<tool_call>"
self.eot_token = "</tool_call>" self.eot_token = "</tool_call>"
self.func_call_regex = r"<tool_call>.*?</tool_call>" self.func_call_regex = r"<tool_call>.*?</tool_call>"
@@ -911,6 +912,9 @@ class Glm47MoeDetector(BaseFormatDetector):
return arguments return arguments
def parses_required_natively(self) -> bool:
return self.use_full_assistant_constraint
def supports_structural_tag(self) -> bool: def supports_structural_tag(self) -> bool:
return _glm47_native_structural_tag_available() return _glm47_native_structural_tag_available()
@@ -1,8 +1,12 @@
import hashlib
import json import json
import logging import logging
import re import re
from collections import deque
from dataclasses import dataclass
from enum import Enum from enum import Enum
from typing import Any, Dict, List, Optional, Tuple from functools import lru_cache
from typing import Any, Dict, List, Literal, Optional, Set, Tuple
from sglang.srt.entrypoints.openai.protocol import Tool from sglang.srt.entrypoints.openai.protocol import Tool
from sglang.srt.function_call.base_format_detector import BaseFormatDetector from sglang.srt.function_call.base_format_detector import BaseFormatDetector
@@ -660,3 +664,493 @@ class Glm4MoeDetector(BaseFormatDetector):
def structure_info(self) -> _GetInfoFunc: def structure_info(self) -> _GetInfoFunc:
raise NotImplementedError() raise NotImplementedError()
class _GlmTrieNode:
"""Trie node with Aho-Corasick failure link."""
def __init__(self, node_id: int):
self.id = node_id
self.children: Dict[str, _GlmTrieNode] = {}
self.is_end = False
self.fail: _GlmTrieNode = None
def _glm_build_trie_with_failure_links(
patterns: List[str],
) -> tuple["_GlmTrieNode", List["_GlmTrieNode"]]:
"""Build Trie and compute Aho-Corasick failure links."""
root = _GlmTrieNode(0)
all_nodes = [root]
next_id = 1
for pattern in patterns:
node = root
for char in pattern:
if char not in node.children:
new_node = _GlmTrieNode(next_id)
next_id += 1
all_nodes.append(new_node)
node.children[char] = new_node
node = node.children[char]
node.is_end = True
root.fail = root
queue = deque()
for child in root.children.values():
child.fail = root
queue.append(child)
while queue:
node = queue.popleft()
for char, child in node.children.items():
queue.append(child)
fail_node = node.fail
while fail_node != root and char not in fail_node.children:
fail_node = fail_node.fail
if char in fail_node.children and fail_node.children[char] != child:
child.fail = fail_node.children[char]
else:
child.fail = root
# A suffix match also completes a forbidden pattern.
if child.fail.is_end:
child.is_end = True
return root, all_nodes
def _glm_get_transition(
node: "_GlmTrieNode", char: str, root: "_GlmTrieNode"
) -> "_GlmTrieNode":
"""Follow Aho-Corasick failure links to the next state."""
current = node
while True:
if char in current.children:
return current.children[char]
if current == root:
return root
current = current.fail
def _glm_escape_char_class(s: str) -> str:
"""Escape special characters for use in EBNF character class [...]."""
result = []
for c in s:
if c in r"\]^-":
result.append("\\" + c)
elif c == "\n":
result.append("\\n")
elif c == "\t":
result.append("\\t")
elif c == "\r":
result.append("\\r")
elif ord(c) < 32 or ord(c) > 126:
result.append(f"\\x{ord(c):02X}")
else:
result.append(c)
return "".join(result)
def _glm_escape_string(c: str) -> str:
"""Escape a character for use in EBNF string literal "..."."""
if c == '"':
return '\\"'
elif c == "\\":
return "\\\\"
elif c == "\n":
return "\\n"
elif c == "\t":
return "\\t"
elif c == "\r":
return "\\r"
elif ord(c) < 32 or ord(c) > 126:
return f"\\x{ord(c):02X}"
return c
def _glm_any_string_exclude(rule_name: str, negative_strings) -> List[str]:
return list(_glm_cached_string_exclude(rule_name, tuple(negative_strings)))
@lru_cache(maxsize=32)
def _glm_cached_string_exclude(
rule_name: str, negative_strings: tuple[str, ...]
) -> tuple[str, ...]:
"""Build EBNF that excludes forbidden substrings using Aho-Corasick states."""
if not negative_strings:
return (f"{rule_name} ::= [^]*",)
sorted_strings = sorted(set(s for s in negative_strings if s))
if not sorted_strings:
return (f"{rule_name} ::= [^]*",)
hash_input = "\x00".join(sorted_strings)
hash_prefix = hashlib.sha256(hash_input.encode("utf-8")).hexdigest()[:16]
root, all_nodes = _glm_build_trie_with_failure_links(sorted_strings)
all_pattern_chars: Set[str] = set()
for pattern in sorted_strings:
all_pattern_chars.update(pattern)
def state_name(node: "_GlmTrieNode") -> str:
return f"s_{hash_prefix}_{node.id}"
rules = []
rules.append(f"{rule_name} ::= {state_name(root)}")
for node in all_nodes:
if node.is_end:
continue
excluded_chars: List[str] = []
transitions_by_target: Dict[int, List[str]] = {}
for char in all_pattern_chars:
target = _glm_get_transition(node, char, root)
if target.is_end:
excluded_chars.append(char)
else:
if target.id not in transitions_by_target:
transitions_by_target[target.id] = []
transitions_by_target[target.id].append(char)
alternatives = []
all_explicit_chars = set(excluded_chars)
for chars in transitions_by_target.values():
all_explicit_chars.update(chars)
if all_explicit_chars:
escaped = _glm_escape_char_class("".join(sorted(all_explicit_chars)))
alternatives.append(f"[^{escaped}] {state_name(root)}")
else:
alternatives.append(f"[^] {state_name(root)}")
for target_id in sorted(transitions_by_target.keys()):
chars = transitions_by_target[target_id]
target_node = next(n for n in all_nodes if n.id == target_id)
for char in sorted(chars):
alternatives.append(
f'"{_glm_escape_string(char)}" {state_name(target_node)}'
)
alternatives.append('""')
rules.append(f"{state_name(node)} ::= {' | '.join(alternatives)}")
return tuple(rules)
_GLM_XML_GRAMMAR_RULES = [
'basic_string ::= (([\\"] basic_string_1 [\\"]))',
'basic_string_1 ::= "" | [^"\\\\\\x00-\\x1F] basic_string_1 | "\\\\" escape basic_string_1',
'escape ::= ["\\\\//bfnrt] | "u" [A-Fa-f0-9]{4}',
'basic_integer ::= "-"? ("0" | [1-9] [0-9]*) ".0"?',
'basic_number ::= "-"? ("0" | [1-9] [0-9]*) ("." [0-9]+)? ([eE] [+-]? [0-9]+)?',
'basic_array ::= "[" ("" | ws basic_any (ws "," ws basic_any)*) ws "]"',
'basic_object ::= "{" ("" | ws basic_string ws ":" ws basic_any ( ws "," ws basic_string ws ":" ws basic_any)*) ws "}"',
"ws ::= [ \\n\\t]*",
"basic_any ::= basic_number | basic_string | basic_boolean | basic_null | basic_array | basic_object",
'basic_boolean ::= "true" | "false"',
'basic_null ::= "null"',
]
_GLM_TYPE_MAPPING = {
"string": "text_without_special_tokens",
"number": "basic_number",
"integer": "basic_number",
"boolean": "basic_boolean",
"null": "basic_null",
"array": "basic_array",
"object": "basic_object",
}
def _glm_hash_name(name: str) -> str:
return hashlib.sha256(name.encode("utf-8")).hexdigest()[:16]
def _glm_get_value_rule(prop: Any) -> str:
if not isinstance(prop, dict):
return "text_without_special_tokens"
if "enum" in prop:
return _glm_handle_enum(prop)
if "type" in prop:
return _glm_handle_type(prop)
return "text_without_special_tokens"
def _glm_escape_ebnf_string(s: str) -> str:
return json.dumps(s, ensure_ascii=False)[1:-1]
def _glm_handle_enum(prop: dict) -> str:
enum_values = prop["enum"]
def format_enum_val(v: Any) -> str:
value = v if isinstance(v, str) else json.dumps(v, ensure_ascii=False)
return f'"{_glm_escape_ebnf_string(value)}"'
formatted_values = [format_enum_val(v) for v in enum_values]
if not formatted_values:
return "text_without_special_tokens"
enum_rule = " | ".join(formatted_values)
return f"({enum_rule})" if len(formatted_values) > 1 else enum_rule
def _glm_handle_type(prop: dict) -> str:
prop_type = prop["type"]
if isinstance(prop_type, list):
type_rules = [
_GLM_TYPE_MAPPING.get(t, "text_without_special_tokens") for t in prop_type
]
return " | ".join(type_rules) if type_rules else "text_without_special_tokens"
return _GLM_TYPE_MAPPING.get(prop_type, "text_without_special_tokens")
def _glm_has_complete_properties(schema: Any) -> bool:
if not isinstance(schema, dict):
return schema is False
if any(
keyword in schema
for keyword in (
"$ref",
"$dynamicRef",
"patternProperties",
"dependentSchemas",
"if",
"then",
"else",
)
) or any(
schema.get(keyword, False) is not False
for keyword in ("additionalProperties", "unevaluatedProperties")
):
return False
branches = [
branch
for keyword in ("allOf", "anyOf", "oneOf")
for branch in schema.get(keyword, [])
]
properties = schema.get("properties")
if isinstance(properties, dict):
return not branches and (
bool(properties) or schema.get("additionalProperties") is False
)
if branches:
return all(_glm_has_complete_properties(branch) for branch in branches)
return schema.get("additionalProperties") is False
def _glm_build_tool_call_rules(
non_terminal_name: str,
functions: list[Any],
special_tokens: "GlmSpecialTokenConfig",
chat_template_version: Literal["glm45", "glm47"],
required: bool = False,
parallel_tool_calls: bool = True,
) -> list[str]:
"""Build non-strict XML tool-call rules with shallow value constraints."""
if chat_template_version == "glm45":
extra_seperator = '"\\n"'
elif chat_template_version == "glm47":
extra_seperator = ""
else:
raise NotImplementedError(
f"Unsupported chat_template_version: {chat_template_version}"
)
repetition = (
("+" if required else "*") if parallel_tool_calls else ("" if required else "?")
)
rules = [
f"{non_terminal_name} ::= ( {extra_seperator} tool_call_unit ){repetition}",
f'tool_call_unit ::= "{special_tokens.begin_of_tool_call}" single_tool_call "{special_tokens.end_of_tool_call}"',
]
# Include the index to distinguish duplicate function names.
tool_alternatives = " | ".join(
f"call_{_glm_hash_name(func.name + str(function_index))}"
for function_index, func in enumerate(functions)
)
rules.append(f"single_tool_call ::= {tool_alternatives}")
# Group alternatives so nullable values remain inside the argument tags.
kv_template = f'"{special_tokens.begin_of_key}{{key}}{special_tokens.end_of_key}" {extra_seperator} "{special_tokens.begin_of_value}" ({{valrule}}) "{special_tokens.end_of_value}"'
kv_separator = extra_seperator
for function_index, func in enumerate(functions):
tool_name = _glm_escape_ebnf_string(func.name)
namehash = _glm_hash_name(func.name + str(function_index))
params = func.parameters or {}
properties = get_schema_properties(params)
if not _glm_has_complete_properties(params):
properties = {}
prop_kv_pairs = {}
for prop_name, prop_schema in properties.items():
# Composition branches can disagree on a property's value schema.
value_rule = (
"text_without_special_tokens"
if any(keyword in params for keyword in ("allOf", "anyOf", "oneOf"))
else _glm_get_value_rule(prop_schema)
)
pair = kv_template.format(
key=_glm_escape_ebnf_string(prop_name), valrule=value_rule
)
prop_kv_pairs[prop_name] = pair
# Non-strict arguments may be omitted, repeated, or emitted in any order.
all_props = list(properties.keys())
if all_props:
all_choices = " | ".join(prop_kv_pairs[k] for k in all_props)
arguments_rule = (
f"( ( {all_choices} ) ( {kv_separator} ( {all_choices} ) )* )?"
)
else:
arguments_rule = (
f'( "{special_tokens.begin_of_key}" text_without_special_tokens '
f'"{special_tokens.end_of_key}" {extra_seperator} '
f'"{special_tokens.begin_of_value}" text_without_special_tokens '
f'"{special_tokens.end_of_value}" {kv_separator} )*'
)
rules.append(
f'call_{namehash} ::= "{tool_name}" {extra_seperator} ( arguments_{namehash} {extra_seperator} )?'
)
rules.append(f"arguments_{namehash} ::= {arguments_rule}")
rules.extend(_GLM_XML_GRAMMAR_RULES)
return rules
@dataclass
class GlmSpecialTokenConfig:
begin_of_thinking: str = "<think>"
end_of_thinking: str = "</think>"
begin_of_tool_call: str = "<tool_call>"
end_of_tool_call: str = "</tool_call>"
begin_of_key: str = "<arg_key>"
end_of_key: str = "</arg_key>"
begin_of_value: str = "<arg_value>"
end_of_value: str = "</arg_value>"
assistant_token: str = "<|assistant|>"
def all_special_tokens(self) -> list[str]:
return vars(self).values()
def generate_glm_grammar(
enable_thinking: bool,
functions: list[Any] | None,
special_tokens: GlmSpecialTokenConfig,
chat_template_version: Literal["glm45", "glm47"],
accommodate_chat_template: bool,
allow_multiple_assistant_turns: bool,
root_name: str = "root",
required: bool = False,
parallel_tool_calls: bool = True,
) -> str:
ebnf_lines = [
f'{root_name} ::= assistant_turn ( "{special_tokens.assistant_token}" assistant_turn )*'
if allow_multiple_assistant_turns
else f"{root_name} ::= assistant_turn",
"assistant_turn ::= thinking_block text_block tool_call_blocks",
]
thinking_exclusions = [
special_tokens.begin_of_tool_call,
special_tokens.end_of_tool_call,
special_tokens.begin_of_key,
special_tokens.end_of_key,
special_tokens.begin_of_value,
special_tokens.end_of_value,
special_tokens.end_of_thinking,
]
if chat_template_version == "glm45":
extra_seperator = '"\\n"'
elif chat_template_version == "glm47":
extra_seperator = ""
else:
raise NotImplementedError(
f"Unsupported chat_template_version: {chat_template_version}"
)
if chat_template_version == "glm45":
if enable_thinking:
ebnf_lines.append(
rf'thinking_block ::= "\n{special_tokens.begin_of_thinking}" thinking_block_content "{special_tokens.end_of_thinking}"'
)
ebnf_lines.extend(
_glm_any_string_exclude("thinking_block_content", thinking_exclusions)
)
else:
if accommodate_chat_template:
ebnf_lines.append('thinking_block ::= ""')
else:
ebnf_lines.append(
rf'thinking_block ::= "\n{special_tokens.begin_of_thinking}" "{special_tokens.end_of_thinking}"'
)
elif chat_template_version == "glm47":
if enable_thinking:
if accommodate_chat_template:
ebnf_lines.append(
rf'thinking_block ::= thinking_block_content "{special_tokens.end_of_thinking}"'
)
else:
ebnf_lines.append(
rf'thinking_block ::= "{special_tokens.begin_of_thinking}" thinking_block_content "{special_tokens.end_of_thinking}"'
)
ebnf_lines.extend(
_glm_any_string_exclude("thinking_block_content", thinking_exclusions)
)
else:
if accommodate_chat_template:
ebnf_lines.append('thinking_block ::= ""')
else:
ebnf_lines.append(
rf'thinking_block ::= "{special_tokens.end_of_thinking}"'
)
else:
raise NotImplementedError(
f"Unsupported chat_template_version: {chat_template_version}"
)
ebnf_lines.extend(
_glm_any_string_exclude(
"text_without_special_tokens", special_tokens.all_special_tokens()
)
)
ebnf_lines.append(
f"text_block ::= ( {extra_seperator} text_without_special_tokens )?"
)
if functions:
ebnf_lines.extend(
_glm_build_tool_call_rules(
non_terminal_name="tool_call_blocks",
functions=functions,
special_tokens=special_tokens,
chat_template_version=chat_template_version,
required=required,
parallel_tool_calls=parallel_tool_calls,
)
)
else:
ebnf_lines.append('tool_call_blocks ::= ""')
non_terminals = {}
deduped_lines = []
for line in ebnf_lines:
assert "\n" not in line, "Each EBNF rule should be in a single line."
lhs = line.split("::=")[0].strip()
if lhs in non_terminals:
if non_terminals[lhs] == line:
continue
raise ValueError(f"Duplicate non-terminal found: {lhs}")
non_terminals[lhs] = line
deduped_lines.append(line)
return "\n".join(deduped_lines)
@@ -159,6 +159,7 @@ class SamplingParams(msgspec.Struct, kw_only=True, array_like=True):
stop_str_max_len: int = 0 # set by normalize() stop_str_max_len: int = 0 # set by normalize()
stop_regex_max_len: int = 0 # set by normalize() stop_regex_max_len: int = 0 # set by normalize()
is_normalized: bool = False # set by normalize() is_normalized: bool = False # set by normalize()
ebnf_full_assistant: bool = False
def __post_init__(self): def __post_init__(self):
# For non-optional params, treat None as "use default" so that callers # For non-optional params, treat None as "use default" so that callers
@@ -235,6 +235,12 @@ pub struct SamplingParams {
/// Set by `normalize`; tells the scheduler its own pass can early-return. /// Set by `normalize`; tells the scheduler its own pass can early-return.
#[serde(skip_deserializing)] #[serde(skip_deserializing)]
pub is_normalized: bool, pub is_normalized: bool,
/// Set by the OpenAI serving layer for generated full-assistant EBNF
/// constraints, which already cover reasoning; the scheduler skips the
/// reasoner grammar wrapper for them. Client-settable would let a request
/// strip that wrapper from its own grammar, so it is a pipeline output only.
#[serde(skip_deserializing)]
pub ebnf_full_assistant: bool,
/// API fields present in the request object. Serde defaults erase this /// API fields present in the request object. Serde defaults erase this
/// distinction, but preferred sampling parameters must not overwrite an /// distinction, but preferred sampling parameters must not overwrite an
/// explicit request value, including an explicit default or null. /// explicit request value, including an explicit default or null.
@@ -380,6 +386,7 @@ impl Default for SamplingParams {
stop_str_max_len: 0, stop_str_max_len: 0,
stop_regex_max_len: 0, stop_regex_max_len: 0,
is_normalized: false, is_normalized: false,
ebnf_full_assistant: false,
explicit_fields: BTreeSet::new(), explicit_fields: BTreeSet::new(),
} }
} }
@@ -831,6 +838,7 @@ mod tests {
"stop_str_max_len", "stop_str_max_len",
"stop_regex_max_len", "stop_regex_max_len",
"is_normalized", "is_normalized",
"ebnf_full_assistant",
]; ];
/// Every field reaches the wire, at the position Python expects. /// Every field reaches the wire, at the position Python expects.
@@ -907,6 +915,7 @@ mod tests {
// `normalize` outputs occupy the tail. // `normalize` outputs occupy the tail.
assert!(arr[at("stop_strs")].is_array()); assert!(arr[at("stop_strs")].is_array());
assert_eq!(arr[at("is_normalized")].as_bool(), Some(false)); assert_eq!(arr[at("is_normalized")].as_bool(), Some(false));
assert_eq!(arr[at("ebnf_full_assistant")].as_bool(), Some(false));
} }
#[test] #[test]
@@ -87,6 +87,7 @@ def _make_req(
req.sampling_params.json_schema = json_schema req.sampling_params.json_schema = json_schema
req.sampling_params.regex = regex req.sampling_params.regex = regex
req.sampling_params.ebnf = ebnf req.sampling_params.ebnf = ebnf
req.sampling_params.ebnf_full_assistant = False
req.sampling_params.structural_tag = structural_tag req.sampling_params.structural_tag = structural_tag
req.sampling_params.custom_params = custom_params req.sampling_params.custom_params = custom_params
req.require_reasoning = False req.require_reasoning = False
@@ -201,11 +202,16 @@ class TestProcessReqWithGrammar(unittest.TestCase):
future = Future() future = Future()
mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False) mgr.grammar_backend.get_cached_or_future_value.return_value = (future, False)
req = _make_req(ebnf="root ::= 'hello'") for full_assistant, key_type in (
result = mgr.process_req_with_grammar(req) (False, "ebnf"),
(True, "full_assistant_ebnf"),
self.assertTrue(result) ):
self.assertEqual(req.grammar_key, ("ebnf", "root ::= 'hello'")) with self.subTest(full_assistant=full_assistant):
req = _make_req(ebnf='root ::= "hello"')
req.sampling_params.ebnf_full_assistant = full_assistant
result = mgr.process_req_with_grammar(req)
self.assertTrue(result)
self.assertEqual(req.grammar_key, (key_type, 'root ::= "hello"'))
def test_structural_tag_cache_miss(self): def test_structural_tag_cache_miss(self):
mgr = self._make_mgr() mgr = self._make_mgr()
@@ -244,13 +244,21 @@ class TestReasonerGrammarBackend(unittest.TestCase):
enable_strict_thinking=True, enable_strict_thinking=True,
) )
wrapped = reasoner._init_value_dispatch(("json", "{}"), reasoning=True) for key in (("json", "{}"), ("ebnf", 'root ::= "OK"')):
self.assertIsInstance(wrapped, ReasonerGrammarObject) with self.subTest(key=key):
wrapped.accept_token(10) inner_grammar.reset_mock()
inner_grammar.accept_token.assert_not_called() wrapped = reasoner._init_value_dispatch(key, reasoning=True)
wrapped.accept_token(2) self.assertIsInstance(wrapped, ReasonerGrammarObject)
wrapped.accept_token(42) wrapped.accept_token(10)
inner_grammar.accept_token.assert_called_once_with(42) inner_grammar.accept_token.assert_not_called()
wrapped.accept_token(2)
wrapped.accept_token(42)
inner_grammar.accept_token.assert_called_once_with(42)
bare = reasoner._init_value_dispatch(
("full_assistant_ebnf", 'root ::= "OK"'), reasoning=True
)
self.assertIs(bare, inner_grammar)
def test_accepts_multi_token_think_start_marker(self): def test_accepts_multi_token_think_start_marker(self):
"""think_start_token can be multi-token (e.g., GPT-OSS) since it's not used.""" """think_start_token can be multi-token (e.g., GPT-OSS) since it's not used."""
@@ -13,6 +13,7 @@
# ============================================================================== # ==============================================================================
"""Tests for OpenAI API protocol models""" """Tests for OpenAI API protocol models"""
import json
import unittest import unittest
from typing import List, Optional from typing import List, Optional
@@ -116,6 +117,31 @@ class TestCompletionRequest(unittest.TestCase):
class TestChatCompletionRequest(unittest.TestCase): class TestChatCompletionRequest(unittest.TestCase):
"""Test ChatCompletionRequest protocol model""" """Test ChatCompletionRequest protocol model"""
def test_full_assistant_ebnf_preserves_explicit_output_constraints(self):
constraint = ("full_assistant_ebnf", 'root ::= "generated"')
for explicit in (
{},
{"ebnf": 'root ::= "OK"'},
{"response_format": {"type": "json_object"}},
):
with self.subTest(explicit=explicit):
request = ChatCompletionRequest(
model="test",
messages=[{"role": "user", "content": "Hi"}],
tool_choice="required",
**explicit,
)
params = request.to_sampling_params([], {}, constraint)
self.assertEqual(params.get("ebnf_full_assistant", False), not explicit)
if "ebnf" in explicit:
self.assertEqual(params["ebnf"], explicit["ebnf"])
elif "response_format" in explicit:
self.assertEqual(
json.loads(params["json_schema"]), {"type": "object"}
)
else:
self.assertEqual(params["ebnf"], constraint[1])
def test_json_schema_strict_requires_json_boolean(self): def test_json_schema_strict_requires_json_boolean(self):
base_request = { base_request = {
"model": "test-model", "model": "test-model",
@@ -3,6 +3,8 @@ import json
import unittest import unittest
import warnings import warnings
import xgrammar as xgr
from sglang.srt.entrypoints.openai.protocol import ( from sglang.srt.entrypoints.openai.protocol import (
Function, Function,
Tool, Tool,
@@ -14,6 +16,7 @@ from sglang.srt.function_call.core_types import StreamingParseResult
from sglang.srt.function_call.deepseekv3_detector import DeepSeekV3Detector from sglang.srt.function_call.deepseekv3_detector import DeepSeekV3Detector
from sglang.srt.function_call.deepseekv4_detector import DeepSeekV4Detector from sglang.srt.function_call.deepseekv4_detector import DeepSeekV4Detector
from sglang.srt.function_call.deepseekv32_detector import DeepSeekV32Detector from sglang.srt.function_call.deepseekv32_detector import DeepSeekV32Detector
from sglang.srt.function_call.function_call_parser import FunctionCallParser
from sglang.srt.function_call.gemma4_detector import ( from sglang.srt.function_call.gemma4_detector import (
Gemma4Detector, Gemma4Detector,
_parse_gemma4_args, _parse_gemma4_args,
@@ -3671,6 +3674,19 @@ class TestGlm47MoeDetector(unittest.TestCase):
self.assertIsNone(self.detector.get_structural_tag(self.tools)) self.assertIsNone(self.detector.get_structural_tag(self.tools))
parser = FunctionCallParser(self.tools, "glm47") parser = FunctionCallParser(self.tools, "glm47")
self.assertEqual(
"full_assistant_ebnf",
parser.get_structure_constraint("required")[0],
)
strict_tools = [
tool.model_copy(
update={
"function": tool.function.model_copy(update={"strict": True})
}
)
for tool in self.tools
]
parser = FunctionCallParser(strict_tools, "glm47")
constraint = parser.get_structure_constraint("required") constraint = parser.get_structure_constraint("required")
self.assertIsNotNone(constraint) self.assertIsNotNone(constraint)
@@ -3678,6 +3694,181 @@ class TestGlm47MoeDetector(unittest.TestCase):
_glm47_native_structural_tag_available.cache_clear() _glm47_native_structural_tag_available.cache_clear()
class TestGlm47FullAssistantGrammar(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.compiler = xgr.GrammarCompiler(
xgr.TokenizerInfo(
[bytes([i]) for i in range(256)], vocab_type=xgr.VocabType.RAW
),
max_threads=1,
)
def _compile(self, parameters=None, choice="auto", parallel=True, thinking=False):
tools = [
Tool(type="function", function=Function(name=name, parameters=parameters))
for name in ("alpha", "beta")
]
parser = FunctionCallParser(tools, "glm47")
constraint = parser.get_structure_constraint(
choice, parallel_tool_calls=parallel, thinking_mode=thinking
)
self.assertIsNotNone(constraint)
return self.compiler.compile_grammar(xgr.Grammar.from_ebnf(constraint[1]))
def _accepts(self, grammar, text):
matcher = xgr.GrammarMatcher(grammar)
return matcher.accept_string(text) and matcher.is_completed()
def test_tool_choice_and_parallel_calls(self):
alpha = "<tool_call>alpha</tool_call>"
beta = "<tool_call>beta</tool_call>"
named = ToolChoice(function=ToolChoiceFuncName(name="alpha"))
for thinking in (False, True):
prefix = "analysis</think>" if thinking else ""
for parallel in (False, True):
for choice in ("auto", "required", named, "none"):
with self.subTest(
thinking=thinking, parallel=parallel, choice=choice
):
grammar = self._compile(
choice=choice, parallel=parallel, thinking=thinking
)
self.assertEqual(
self._accepts(grammar, prefix + "Hello"),
choice in ("auto", "none"),
)
self.assertEqual(
self._accepts(grammar, prefix + alpha), choice != "none"
)
self.assertEqual(
self._accepts(grammar, prefix + beta),
choice in ("auto", "required"),
)
self.assertEqual(
self._accepts(grammar, prefix + alpha * 2),
parallel and choice != "none",
)
def test_enum_json_types_and_boolean_schemas(self):
cases = [
({"enum": [1, 2]}, ["1", "2"], ["3"]),
({"enum": [True, False]}, ["true", "false"], ["True", "1"]),
(
{"type": ["string", "null"], "enum": ["ok", None]},
["ok", "null"],
["None", "bad"],
),
({"enum": [{"x": 1}, [True, None]]}, ['{"x": 1}', "[true, null]"], ["{}"]),
(True, ["anything"], []),
(False, ["anything"], []),
]
for schema, accepted, rejected in cases:
with self.subTest(schema=schema):
grammar = self._compile({"properties": {"p": schema}})
for values, expected in ((accepted, True), (rejected, False)):
for value in values:
text = f"<tool_call>alpha<arg_key>p</arg_key><arg_value>{value}</arg_value></tool_call>"
self.assertEqual(self._accepts(grammar, text), expected, text)
def test_composed_and_unresolved_schemas_allow_arguments(self):
schemas = [
{keyword: [{"properties": {"city": {"type": "string"}}}]}
for keyword in ("allOf", "anyOf", "oneOf")
]
schemas += [
{
"properties": {"country": {"type": "string"}},
"allOf": [{"properties": {"city": {"type": "string"}}}],
},
{
"$ref": "#/$defs/args",
"$defs": {"args": {"properties": {"city": {"type": "string"}}}},
},
{
"anyOf": [
{"properties": {"city": {"enum": [1]}}},
{"properties": {"city": {"enum": ["Paris"]}}},
]
},
]
for schema in schemas:
with self.subTest(schema=schema):
grammar = self._compile(schema)
arg = "<arg_key>city</arg_key><arg_value>Paris</arg_value>"
self.assertTrue(
self._accepts(grammar, f"<tool_call>alpha{arg}</tool_call>")
)
self.assertTrue(
self._accepts(grammar, f"<tool_call>alpha{arg}{arg}</tool_call>")
)
self.assertTrue(self._accepts(grammar, "<tool_call>alpha</tool_call>"))
def test_incomplete_composition_branches_allow_arguments(self):
city = {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
"additionalProperties": False,
}
country = {
"type": "object",
"properties": {"country": {"type": "string"}},
"required": ["country"],
"additionalProperties": False,
}
branches = [
{"$ref": "#/$defs/by_country"},
{"patternProperties": {"^country$": {"type": "string"}}},
{"properties": {"region": {"type": "string"}}, "allOf": [country]},
{"additionalProperties": {"type": "string"}},
True,
{},
{"properties": {}},
]
text = "<tool_call>alpha<arg_key>country</arg_key><arg_value>France</arg_value></tool_call>"
for branch in branches:
for nested in (False, True):
with self.subTest(branch=branch, nested=nested):
schema = {
"type": "object",
"anyOf": [
city,
{"allOf": [{"oneOf": [branch]}]} if nested else branch,
],
"$defs": {"by_country": country},
}
grammar = self._compile(schema, choice="required", parallel=False)
self.assertTrue(self._accepts(grammar, text))
self.assertFalse(
self._accepts(grammar, text.replace("</arg_key>", ""))
)
self.assertFalse(self._accepts(grammar, text + text))
def test_complete_compositions_restrict_argument_names(self):
schema = {
"allOf": [
{"properties": {"city": {"type": "string"}}},
{
"anyOf": [
{"oneOf": [{"properties": {"country": {"type": "string"}}}]}
]
},
]
}
grammar = self._compile(schema, choice="required", parallel=False)
for key, accepted in (("city", True), ("country", True), ("unknown", False)):
text = f"<tool_call>alpha<arg_key>{key}</arg_key><arg_value>Paris</arg_value></tool_call>"
self.assertEqual(self._accepts(grammar, text), accepted)
def test_escaped_property_names(self):
for key in ['a"b', "path\\name", "line\nbreak", "tab\tkey", "control\x01key"]:
with self.subTest(key=key):
grammar = self._compile({"properties": {key: {"type": "string"}}})
text = f"<tool_call>alpha<arg_key>{key}</arg_key><arg_value>v</arg_value></tool_call>"
self.assertTrue(self._accepts(grammar, text))
class TestLing3Detector(unittest.TestCase): class TestLing3Detector(unittest.TestCase):
def setUp(self): def setUp(self):
self.tools = [ self.tools = [