[Kimi K3] Add reasoning, tool-call, and OpenAI serving support (#33025)
Co-authored-by: hnyls2002 <lsyincs@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com> Co-authored-by: A-transformer <cl5743590921@gmail.com>
This commit is contained in:
co-authored by
hnyls2002
Liangsheng Yin
A-transformer
parent
f1b41a5b3d
commit
e2cf21b9e5
@@ -502,7 +502,7 @@ class ModelConfig:
|
|||||||
# Cache attributes
|
# Cache attributes
|
||||||
self.hf_eos_token_id = self._get_hf_eos_token_id()
|
self.hf_eos_token_id = self._get_hf_eos_token_id()
|
||||||
# Set by scheduler when reasoning_parser is enabled
|
# Set by scheduler when reasoning_parser is enabled
|
||||||
self.think_end_id: Optional[int] = None
|
self.think_end_ids: Optional[List[int]] = None
|
||||||
|
|
||||||
# multimodal
|
# multimodal
|
||||||
self.image_token_id = getattr(
|
self.image_token_id = getattr(
|
||||||
|
|||||||
@@ -313,7 +313,7 @@ def create_grammar_backend(
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
vocab_size: int,
|
vocab_size: int,
|
||||||
eos_token_ids: Optional[set] = None,
|
eos_token_ids: Optional[set] = None,
|
||||||
think_end_id: Optional[int] = None,
|
think_end_ids: Optional[List[int]] = None,
|
||||||
) -> Optional[BaseGrammarBackend]:
|
) -> Optional[BaseGrammarBackend]:
|
||||||
name = server_args.grammar_backend
|
name = server_args.grammar_backend
|
||||||
|
|
||||||
@@ -384,7 +384,7 @@ def create_grammar_backend(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid grammar backend: {name}")
|
raise ValueError(f"Invalid grammar backend: {name}")
|
||||||
|
|
||||||
if server_args.reasoning_parser and think_end_id is not None:
|
if server_args.reasoning_parser and think_end_ids:
|
||||||
from sglang.srt.constrained.reasoner_grammar_backend import (
|
from sglang.srt.constrained.reasoner_grammar_backend import (
|
||||||
ReasonerGrammarBackend,
|
ReasonerGrammarBackend,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ class GrammarManager:
|
|||||||
scheduler.tokenizer,
|
scheduler.tokenizer,
|
||||||
scheduler.model_config.vocab_size,
|
scheduler.model_config.vocab_size,
|
||||||
scheduler.model_config.hf_eos_token_id,
|
scheduler.model_config.hf_eos_token_id,
|
||||||
think_end_id=scheduler.model_config.think_end_id,
|
think_end_ids=scheduler.model_config.think_end_ids,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.grammar_backend = None
|
self.grammar_backend = None
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast
|
|||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||||
|
from sglang.srt.utils.token_sequence_matcher import TokenSequenceMatcher
|
||||||
|
|
||||||
from .base_grammar_backend import (
|
from .base_grammar_backend import (
|
||||||
BaseGrammarBackend,
|
BaseGrammarBackend,
|
||||||
@@ -32,26 +33,12 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
class ReasonerGrammarObject(BaseGrammarObject):
|
class ReasonerGrammarObject(BaseGrammarObject):
|
||||||
"""Wraps a grammar object to handle reasoning (think/generation) phases.
|
"""Defers grammar constraints until the reasoning end sequence is complete."""
|
||||||
|
|
||||||
State machine (must call maybe_init_reasoning before use):
|
|
||||||
THINKING (tokens_in_think >= 0, tokens_after_end == -1)
|
|
||||||
-> grammar not consulted, optional token filtering
|
|
||||||
GENERATION (tokens_after_end >= 0)
|
|
||||||
-> grammar consulted for accept/fill/rollback
|
|
||||||
|
|
||||||
When enable_token_filter=True (strict mode), fill_vocab_mask filters
|
|
||||||
excluded tokens during THINKING and enforces max_think_tokens budget.
|
|
||||||
When the budget is exhausted, only think_end_id is allowed, forcing the
|
|
||||||
model to exit the thinking phase.
|
|
||||||
When enable_token_filter=False (non-strict mode), fill_vocab_mask is
|
|
||||||
a no-op during THINKING.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
grammar: Optional[BaseGrammarObject],
|
grammar: Optional[BaseGrammarObject],
|
||||||
think_end_id: int,
|
think_end_ids: List[int],
|
||||||
think_excluded_token_ids: Optional[List[int]] = None,
|
think_excluded_token_ids: Optional[List[int]] = None,
|
||||||
max_think_tokens: int = -1,
|
max_think_tokens: int = -1,
|
||||||
enable_token_filter: bool = False,
|
enable_token_filter: bool = False,
|
||||||
@@ -62,7 +49,8 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.grammar = grammar
|
self.grammar = grammar
|
||||||
self.think_end_id = think_end_id
|
self.think_end_ids = tuple(think_end_ids)
|
||||||
|
self._think_end_matcher = TokenSequenceMatcher(self.think_end_ids)
|
||||||
self.think_excluded_token_ids = think_excluded_token_ids
|
self.think_excluded_token_ids = think_excluded_token_ids
|
||||||
self.max_think_tokens = max_think_tokens
|
self.max_think_tokens = max_think_tokens
|
||||||
self.enable_token_filter = enable_token_filter
|
self.enable_token_filter = enable_token_filter
|
||||||
@@ -70,12 +58,15 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
self.allocate_vocab_mask_fn = allocate_vocab_mask_fn
|
self.allocate_vocab_mask_fn = allocate_vocab_mask_fn
|
||||||
self.move_vocab_mask_fn = move_vocab_mask_fn
|
self.move_vocab_mask_fn = move_vocab_mask_fn
|
||||||
self.apply_vocab_mask_fn = apply_vocab_mask_fn
|
self.apply_vocab_mask_fn = apply_vocab_mask_fn
|
||||||
self._think_end_id_list = [think_end_id]
|
|
||||||
|
|
||||||
self.tokens_in_think = -1
|
self.tokens_in_think = -1
|
||||||
self.tokens_after_end = -1
|
self.tokens_after_end = -1
|
||||||
|
self._matched_think_end_tokens = 0
|
||||||
|
self._thinking_match_history: List[int] = []
|
||||||
|
|
||||||
def maybe_init_reasoning(self, reasoning: bool):
|
def maybe_init_reasoning(self, reasoning: bool):
|
||||||
|
self._matched_think_end_tokens = 0
|
||||||
|
self._thinking_match_history.clear()
|
||||||
if reasoning:
|
if reasoning:
|
||||||
self.tokens_in_think = 0
|
self.tokens_in_think = 0
|
||||||
self.tokens_after_end = -1
|
self.tokens_after_end = -1
|
||||||
@@ -91,20 +82,31 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
|
|
||||||
def transfer_state(self, token: int) -> None:
|
def transfer_state(self, token: int) -> None:
|
||||||
if self._is_thinking():
|
if self._is_thinking():
|
||||||
if token == self.think_end_id:
|
previous_match = self._matched_think_end_tokens
|
||||||
|
self._thinking_match_history.append(previous_match)
|
||||||
|
matched = self._think_end_matcher.advance(previous_match, token)
|
||||||
|
if matched == len(self._think_end_matcher):
|
||||||
|
self._matched_think_end_tokens = 0
|
||||||
self.tokens_after_end = 0
|
self.tokens_after_end = 0
|
||||||
else:
|
else:
|
||||||
self.tokens_in_think += 1
|
self.tokens_in_think += previous_match + 1 - matched
|
||||||
|
self._matched_think_end_tokens = matched
|
||||||
elif self._is_generation():
|
elif self._is_generation():
|
||||||
self.tokens_after_end += 1
|
self.tokens_after_end += 1
|
||||||
|
|
||||||
def rollback_state(self):
|
def rollback_state(self):
|
||||||
if self._is_thinking():
|
if self._is_thinking():
|
||||||
if self.tokens_in_think > 0:
|
if self._thinking_match_history:
|
||||||
self.tokens_in_think -= 1
|
previous_match = self._thinking_match_history.pop()
|
||||||
|
self.tokens_in_think -= (
|
||||||
|
previous_match + 1 - self._matched_think_end_tokens
|
||||||
|
)
|
||||||
|
self._matched_think_end_tokens = previous_match
|
||||||
elif self._is_generation():
|
elif self._is_generation():
|
||||||
if self.tokens_after_end == 0:
|
if self.tokens_after_end == 0:
|
||||||
|
if self._thinking_match_history:
|
||||||
self.tokens_after_end = -1
|
self.tokens_after_end = -1
|
||||||
|
self._matched_think_end_tokens = self._thinking_match_history.pop()
|
||||||
elif self.tokens_after_end > 0:
|
elif self.tokens_after_end > 0:
|
||||||
self.tokens_after_end -= 1
|
self.tokens_after_end -= 1
|
||||||
|
|
||||||
@@ -122,7 +124,7 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
self.transfer_state(token)
|
self.transfer_state(token)
|
||||||
|
|
||||||
def is_terminated(self):
|
def is_terminated(self):
|
||||||
if self.grammar is not None:
|
if self._is_generation() and self.grammar is not None:
|
||||||
return self.grammar.is_terminated()
|
return self.grammar.is_terminated()
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -135,7 +137,11 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
self.rollback_state()
|
self.rollback_state()
|
||||||
|
|
||||||
def _can_think_more(self):
|
def _can_think_more(self):
|
||||||
return self.max_think_tokens < 0 or self.tokens_in_think < self.max_think_tokens
|
return (
|
||||||
|
self.max_think_tokens < 0
|
||||||
|
or self.tokens_in_think + self._matched_think_end_tokens
|
||||||
|
< self.max_think_tokens
|
||||||
|
)
|
||||||
|
|
||||||
def _do_token_filter(self, vocab_mask, token_ids, idx, is_allowed=True):
|
def _do_token_filter(self, vocab_mask, token_ids, idx, is_allowed=True):
|
||||||
if self.token_filter_fn is not None:
|
if self.token_filter_fn is not None:
|
||||||
@@ -146,12 +152,19 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
if not self.enable_token_filter:
|
if not self.enable_token_filter:
|
||||||
return
|
return
|
||||||
if self._can_think_more():
|
if self._can_think_more():
|
||||||
|
if self.think_excluded_token_ids is not None:
|
||||||
self._do_token_filter(
|
self._do_token_filter(
|
||||||
vocab_mask, self.think_excluded_token_ids, idx, is_allowed=False
|
vocab_mask,
|
||||||
|
self.think_excluded_token_ids,
|
||||||
|
idx,
|
||||||
|
is_allowed=False,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self._do_token_filter(
|
self._do_token_filter(
|
||||||
vocab_mask, self._think_end_id_list, idx, is_allowed=True
|
vocab_mask,
|
||||||
|
[self.think_end_ids[self._matched_think_end_tokens]],
|
||||||
|
idx,
|
||||||
|
is_allowed=True,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
if self._is_generation() and self.grammar is not None:
|
if self._is_generation() and self.grammar is not None:
|
||||||
@@ -179,19 +192,22 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
|
|
||||||
def copy(self):
|
def copy(self):
|
||||||
new_obj = ReasonerGrammarObject(
|
new_obj = ReasonerGrammarObject(
|
||||||
self.grammar.copy() if self.grammar is not None else None,
|
grammar=self.grammar.copy() if self.grammar is not None else None,
|
||||||
self.think_end_id,
|
think_end_ids=self.think_end_ids,
|
||||||
self.think_excluded_token_ids,
|
think_excluded_token_ids=self.think_excluded_token_ids,
|
||||||
self.max_think_tokens,
|
max_think_tokens=self.max_think_tokens,
|
||||||
self.enable_token_filter,
|
enable_token_filter=self.enable_token_filter,
|
||||||
self.token_filter_fn,
|
token_filter_fn=self.token_filter_fn,
|
||||||
self.allocate_vocab_mask_fn,
|
allocate_vocab_mask_fn=self.allocate_vocab_mask_fn,
|
||||||
self.move_vocab_mask_fn,
|
move_vocab_mask_fn=self.move_vocab_mask_fn,
|
||||||
self.apply_vocab_mask_fn,
|
apply_vocab_mask_fn=self.apply_vocab_mask_fn,
|
||||||
)
|
)
|
||||||
new_obj.tokens_in_think = self.tokens_in_think
|
new_obj.tokens_in_think = self.tokens_in_think
|
||||||
new_obj.tokens_after_end = self.tokens_after_end
|
new_obj.tokens_after_end = self.tokens_after_end
|
||||||
|
new_obj._matched_think_end_tokens = self._matched_think_end_tokens
|
||||||
|
new_obj._thinking_match_history = list(self._thinking_match_history)
|
||||||
new_obj._finished = self._finished
|
new_obj._finished = self._finished
|
||||||
|
new_obj.current_token = self.current_token
|
||||||
return new_obj
|
return new_obj
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -208,17 +224,17 @@ class ReasonerGrammarObject(BaseGrammarObject):
|
|||||||
self._finished = finished
|
self._finished = finished
|
||||||
|
|
||||||
def try_jump_forward(self, tokenizer):
|
def try_jump_forward(self, tokenizer):
|
||||||
if self.grammar is not None:
|
if self._is_generation() and self.grammar is not None:
|
||||||
return self.grammar.try_jump_forward(tokenizer)
|
return self.grammar.try_jump_forward(tokenizer)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def jump_forward_str_state(self, helper):
|
def jump_forward_str_state(self, helper):
|
||||||
if self.grammar is not None:
|
if self._is_generation() and self.grammar is not None:
|
||||||
return self.grammar.jump_forward_str_state(helper)
|
return self.grammar.jump_forward_str_state(helper)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def jump_and_retokenize(self, old_output_ids, new_output_ids, next_state):
|
def jump_and_retokenize(self, old_output_ids, new_output_ids, next_state):
|
||||||
if self.grammar is not None:
|
if self._is_generation() and self.grammar is not None:
|
||||||
return self.grammar.jump_and_retokenize(
|
return self.grammar.jump_and_retokenize(
|
||||||
old_output_ids, new_output_ids, next_state
|
old_output_ids, new_output_ids, next_state
|
||||||
)
|
)
|
||||||
@@ -242,20 +258,17 @@ class ReasonerGrammarBackend(BaseGrammarBackend):
|
|||||||
f"think_end_token '{reasoning_parser.detector.think_end_token}' "
|
f"think_end_token '{reasoning_parser.detector.think_end_token}' "
|
||||||
f"could not be encoded by the tokenizer."
|
f"could not be encoded by the tokenizer."
|
||||||
)
|
)
|
||||||
if len(think_end_ids) != 1:
|
self.think_end_ids = think_end_ids
|
||||||
raise ValueError(
|
|
||||||
f"think_end_token '{reasoning_parser.detector.think_end_token}' "
|
|
||||||
"must encode to exactly one token for constrained reasoning."
|
|
||||||
)
|
|
||||||
self.think_end_id = think_end_ids[0]
|
|
||||||
self._enable_strict_thinking = enable_strict_thinking
|
self._enable_strict_thinking = enable_strict_thinking
|
||||||
self.think_excluded_token_ids = self._get_think_excluded_token_ids(
|
self.think_excluded_token_ids = self._get_think_excluded_token_ids(
|
||||||
reasoning_parser, tokenizer
|
reasoning_parser, tokenizer
|
||||||
)
|
)
|
||||||
self.max_think_tokens = envs.SGLANG_MAX_THINK_TOKENS.get()
|
self.max_think_tokens = envs.SGLANG_MAX_THINK_TOKENS.get()
|
||||||
|
self.enable_token_filter = self.enable_strict_thinking and (
|
||||||
|
self.think_excluded_token_ids is not None or self.max_think_tokens >= 0
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
self.enable_strict_thinking
|
self.enable_token_filter
|
||||||
and self.think_excluded_token_ids is not None
|
|
||||||
and not self.grammar_backend.is_support_token_filter
|
and not self.grammar_backend.is_support_token_filter
|
||||||
):
|
):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -263,11 +276,6 @@ class ReasonerGrammarBackend(BaseGrammarBackend):
|
|||||||
"support token filtering. Use a grammar backend that supports token "
|
"support token filtering. Use a grammar backend that supports token "
|
||||||
"filtering (e.g., xgrammar) or disable strict reasoning mode."
|
"filtering (e.g., xgrammar) or disable strict reasoning mode."
|
||||||
)
|
)
|
||||||
self.enable_token_filter = (
|
|
||||||
self.enable_strict_thinking
|
|
||||||
and self.think_excluded_token_ids is not None
|
|
||||||
and self.grammar_backend.is_support_token_filter
|
|
||||||
)
|
|
||||||
self._token_filter_fn = (
|
self._token_filter_fn = (
|
||||||
self.grammar_backend.set_token_filter if self.enable_token_filter else None
|
self.grammar_backend.set_token_filter if self.enable_token_filter else None
|
||||||
)
|
)
|
||||||
@@ -298,7 +306,7 @@ class ReasonerGrammarBackend(BaseGrammarBackend):
|
|||||||
) -> ReasonerGrammarObject:
|
) -> ReasonerGrammarObject:
|
||||||
obj = ReasonerGrammarObject(
|
obj = ReasonerGrammarObject(
|
||||||
grammar=grammar,
|
grammar=grammar,
|
||||||
think_end_id=self.think_end_id,
|
think_end_ids=self.think_end_ids,
|
||||||
think_excluded_token_ids=self.think_excluded_token_ids,
|
think_excluded_token_ids=self.think_excluded_token_ids,
|
||||||
max_think_tokens=self.max_think_tokens,
|
max_think_tokens=self.max_think_tokens,
|
||||||
enable_token_filter=self.enable_token_filter,
|
enable_token_filter=self.enable_token_filter,
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ def resolve_chat_encoding_spec(
|
|||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
tool_call_parser: Optional[str] = None,
|
tool_call_parser: Optional[str] = None,
|
||||||
) -> Optional[str]:
|
) -> Optional[str]:
|
||||||
"""Return the chat encoding spec for a model: "dsv4", "dsv32", "inkling", or None.
|
"""Return the chat encoding spec for a model.
|
||||||
|
|
||||||
None means the default path (HF chat template).
|
None means the default path (HF chat template).
|
||||||
"""
|
"""
|
||||||
@@ -24,12 +24,16 @@ def resolve_chat_encoding_spec(
|
|||||||
return "dsv4"
|
return "dsv4"
|
||||||
if tool_call_parser == "deepseekv32":
|
if tool_call_parser == "deepseekv32":
|
||||||
return "dsv32"
|
return "dsv32"
|
||||||
|
if tool_call_parser == "kimi_k3":
|
||||||
|
return "kimi_k3"
|
||||||
|
|
||||||
architectures = hf_config.architectures
|
architectures = hf_config.architectures
|
||||||
arch = architectures[0] if architectures else ""
|
arch = architectures[0] if architectures else ""
|
||||||
|
|
||||||
if "DeepseekV4" in arch:
|
if "DeepseekV4" in arch:
|
||||||
return "dsv4"
|
return "dsv4"
|
||||||
|
if "KimiK3" in arch:
|
||||||
|
return "kimi_k3"
|
||||||
|
|
||||||
# Inkling has no Jinja chat_template and uses a tiktoken base + a special-token
|
# Inkling has no Jinja chat_template and uses a tiktoken base + a special-token
|
||||||
# overlay + negative MM placeholders, so it can't go through apply_chat_template;
|
# overlay + negative MM placeholders, so it can't go through apply_chat_template;
|
||||||
|
|||||||
@@ -43,7 +43,6 @@ from openai.types.responses import (
|
|||||||
ResponseReasoningItem,
|
ResponseReasoningItem,
|
||||||
)
|
)
|
||||||
from openai.types.responses.response import ToolChoice
|
from openai.types.responses.response import ToolChoice
|
||||||
from openai.types.responses.tool import Tool
|
|
||||||
from pydantic import (
|
from pydantic import (
|
||||||
BaseModel,
|
BaseModel,
|
||||||
ConfigDict,
|
ConfigDict,
|
||||||
@@ -213,7 +212,7 @@ class JsonSchemaResponseFormat(BaseModel):
|
|||||||
description: Optional[str] = None
|
description: Optional[str] = None
|
||||||
# use alias to workaround pydantic conflict
|
# use alias to workaround pydantic conflict
|
||||||
schema_: Optional[Dict[str, object]] = Field(alias="schema", default=None)
|
schema_: Optional[Dict[str, object]] = Field(alias="schema", default=None)
|
||||||
strict: Optional[bool] = False
|
strict: Optional[bool] = None
|
||||||
|
|
||||||
|
|
||||||
class ResponseFormat(BaseModel):
|
class ResponseFormat(BaseModel):
|
||||||
@@ -687,6 +686,11 @@ class Tool(BaseModel):
|
|||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
# Tool is defined after the message params that reference it, so the forward
|
||||||
|
# reference has to be resolved explicitly.
|
||||||
|
ChatCompletionMessageGenericParam.model_rebuild()
|
||||||
|
|
||||||
|
|
||||||
class ToolChoiceFuncName(BaseModel):
|
class ToolChoiceFuncName(BaseModel):
|
||||||
"""The name of tool choice function."""
|
"""The name of tool choice function."""
|
||||||
|
|
||||||
@@ -717,6 +721,18 @@ ReasoningEffortType = Optional[
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _has_message_level_tools(messages: Any) -> bool:
|
||||||
|
if not isinstance(messages, list):
|
||||||
|
return False
|
||||||
|
return any(
|
||||||
|
isinstance(message, dict)
|
||||||
|
and isinstance(message.get("role"), str)
|
||||||
|
and message["role"].lower() in ("system", "developer")
|
||||||
|
and bool(message.get("tools"))
|
||||||
|
for message in messages
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ChatCompletionRequest(BaseModel):
|
class ChatCompletionRequest(BaseModel):
|
||||||
# Ordered by official OpenAI API documentation
|
# Ordered by official OpenAI API documentation
|
||||||
# https://platform.openai.com/docs/api-reference/chat/create
|
# https://platform.openai.com/docs/api-reference/chat/create
|
||||||
@@ -859,7 +875,9 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def set_tool_choice_default(cls, values):
|
def set_tool_choice_default(cls, values):
|
||||||
if values.get("tool_choice") is None:
|
if values.get("tool_choice") is None:
|
||||||
if values.get("tools") is None:
|
if values.get("tools") is None and not _has_message_level_tools(
|
||||||
|
values.get("messages")
|
||||||
|
):
|
||||||
values["tool_choice"] = "none"
|
values["tool_choice"] = "none"
|
||||||
else:
|
else:
|
||||||
values["tool_choice"] = "auto"
|
values["tool_choice"] = "auto"
|
||||||
@@ -949,11 +967,10 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
|
|
||||||
if schema:
|
if schema:
|
||||||
name_ = schema.get("title", "Schema")
|
name_ = schema.get("title", "Schema")
|
||||||
strict_ = False
|
strict_ = None
|
||||||
if "properties" in schema and "strict" in schema["properties"]:
|
if "properties" in schema and "strict" in schema["properties"]:
|
||||||
item = schema["properties"].pop("strict", None)
|
item = schema["properties"].pop("strict", None)
|
||||||
if item and item.get("default", False):
|
strict_ = bool(item and item.get("default", False))
|
||||||
strict_ = True
|
|
||||||
|
|
||||||
response_format["json_schema"] = {
|
response_format["json_schema"] = {
|
||||||
"name": name_,
|
"name": name_,
|
||||||
@@ -968,6 +985,7 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
stop: List[str],
|
stop: List[str],
|
||||||
model_generation_config: Dict[str, Any],
|
model_generation_config: Dict[str, Any],
|
||||||
tool_call_constraint: Optional[ToolCallConstraint] = None,
|
tool_call_constraint: Optional[ToolCallConstraint] = None,
|
||||||
|
renderer_handles_response_format: bool = False,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Convert request to sampling parameters.
|
Convert request to sampling parameters.
|
||||||
@@ -1015,6 +1033,12 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
}
|
}
|
||||||
|
|
||||||
if self.response_format and self.response_format.type == "json_schema":
|
if self.response_format and self.response_format.type == "json_schema":
|
||||||
|
# strict=false may only go unconstrained when the renderer forwards
|
||||||
|
# response_format to the model; plain chat templates never see it.
|
||||||
|
if (
|
||||||
|
self.response_format.json_schema.strict is not False
|
||||||
|
or not renderer_handles_response_format
|
||||||
|
):
|
||||||
sampling_params["json_schema"] = convert_json_schema_to_str(
|
sampling_params["json_schema"] = convert_json_schema_to_str(
|
||||||
self.response_format.json_schema.schema_
|
self.response_format.json_schema.schema_
|
||||||
)
|
)
|
||||||
@@ -1781,22 +1805,6 @@ class RequestResponseMetadata(BaseModel):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class MessageProcessingResult:
|
class MessageProcessingResult:
|
||||||
"""Result of processing chat messages and applying templates.
|
|
||||||
|
|
||||||
This dataclass encapsulates all the outputs from message processing including
|
|
||||||
prompt generation, multimodal data extraction, and constraint preparation.
|
|
||||||
Used internally by OpenAIServingChat to pass processed data between methods.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
prompt: The final text prompt after applying chat template
|
|
||||||
prompt_ids: Either the text prompt (str) or tokenized IDs (List[int])
|
|
||||||
image_data: Extracted image data from messages, if any
|
|
||||||
audio_data: Extracted audio data from messages, if any
|
|
||||||
modalities: List of modality types present in the messages
|
|
||||||
stop: Combined stop strings from template and request
|
|
||||||
tool_call_constraint: Optional constraint for structured tool calls
|
|
||||||
"""
|
|
||||||
|
|
||||||
prompt: str
|
prompt: str
|
||||||
prompt_ids: Union[str, List[int]]
|
prompt_ids: Union[str, List[int]]
|
||||||
image_data: Optional[Any]
|
image_data: Optional[Any]
|
||||||
@@ -1805,6 +1813,7 @@ class MessageProcessingResult:
|
|||||||
modalities: List[str]
|
modalities: List[str]
|
||||||
stop: List[str]
|
stop: List[str]
|
||||||
tool_call_constraint: Optional[ToolCallConstraint] = None
|
tool_call_constraint: Optional[ToolCallConstraint] = None
|
||||||
|
require_reasoning: bool = False
|
||||||
|
|
||||||
|
|
||||||
class ToolCallProcessingResult(NamedTuple):
|
class ToolCallProcessingResult(NamedTuple):
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ from jsonschema import Draft202012Validator, SchemaError
|
|||||||
|
|
||||||
from sglang.srt.entrypoints.openai import encoding_dsv4, encoding_dsv32
|
from sglang.srt.entrypoints.openai import encoding_dsv4, encoding_dsv32
|
||||||
from sglang.srt.entrypoints.openai.protocol import (
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
|
ChatCompletionMessageGenericParam,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
ChatCompletionResponse,
|
ChatCompletionResponse,
|
||||||
ChatCompletionResponseChoice,
|
ChatCompletionResponseChoice,
|
||||||
@@ -42,6 +43,7 @@ from sglang.srt.entrypoints.openai.protocol import (
|
|||||||
PromptTokensDetails,
|
PromptTokensDetails,
|
||||||
ResponseParserProtocol,
|
ResponseParserProtocol,
|
||||||
SglExt,
|
SglExt,
|
||||||
|
Tool,
|
||||||
ToolCall,
|
ToolCall,
|
||||||
ToolCallProcessingResult,
|
ToolCallProcessingResult,
|
||||||
ToolChoice,
|
ToolChoice,
|
||||||
@@ -119,7 +121,9 @@ def parse_tool_call_arguments(arguments: str) -> Dict[str, Any]:
|
|||||||
return parsed_arguments
|
return parsed_arguments
|
||||||
|
|
||||||
|
|
||||||
def normalize_assistant_tool_call_arguments(message: Dict[str, Any]) -> None:
|
def normalize_assistant_tool_call_arguments(
|
||||||
|
message: Dict[str, Any], *, strict: bool = True
|
||||||
|
) -> None:
|
||||||
"""Normalize assistant history tool call arguments in-place."""
|
"""Normalize assistant history tool call arguments in-place."""
|
||||||
if message.get("role") != "assistant" or not isinstance(
|
if message.get("role") != "assistant" or not isinstance(
|
||||||
message.get("tool_calls"), list
|
message.get("tool_calls"), list
|
||||||
@@ -131,7 +135,11 @@ def normalize_assistant_tool_call_arguments(message: Dict[str, Any]) -> None:
|
|||||||
if not isinstance(function, dict):
|
if not isinstance(function, dict):
|
||||||
continue
|
continue
|
||||||
if "arguments" in function and isinstance(function["arguments"], str):
|
if "arguments" in function and isinstance(function["arguments"], str):
|
||||||
|
try:
|
||||||
function["arguments"] = parse_tool_call_arguments(function["arguments"])
|
function["arguments"] = parse_tool_call_arguments(function["arguments"])
|
||||||
|
except ValueError:
|
||||||
|
if strict:
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
def _extract_max_dynamic_patch(request: ChatCompletionRequest):
|
def _extract_max_dynamic_patch(request: ChatCompletionRequest):
|
||||||
@@ -160,6 +168,27 @@ def _extract_max_dynamic_patch(request: ChatCompletionRequest):
|
|||||||
return img_max_dynamic_patch, vid_max_dynamic_patch
|
return img_max_dynamic_patch, vid_max_dynamic_patch
|
||||||
|
|
||||||
|
|
||||||
|
KIMI_K3_IMAGE_PLACEHOLDER = "<|kimi_image_placeholder|>"
|
||||||
|
KIMI_K3_IMAGE_PLACEHOLDER_ESCAPED = "<| kimi_image_placeholder |>"
|
||||||
|
|
||||||
|
|
||||||
|
def neutralize_kimi_k3_image_placeholder(text: str) -> str:
|
||||||
|
return text.replace(KIMI_K3_IMAGE_PLACEHOLDER, KIMI_K3_IMAGE_PLACEHOLDER_ESCAPED)
|
||||||
|
|
||||||
|
|
||||||
|
def neutralize_kimi_k3_image_placeholder_value(value: Any) -> Any:
|
||||||
|
if isinstance(value, str):
|
||||||
|
return neutralize_kimi_k3_image_placeholder(value)
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [neutralize_kimi_k3_image_placeholder_value(item) for item in value]
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {
|
||||||
|
key: neutralize_kimi_k3_image_placeholder_value(item)
|
||||||
|
for key, item in value.items()
|
||||||
|
}
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
class OpenAIServingChat(OpenAIServingBase):
|
class OpenAIServingChat(OpenAIServingBase):
|
||||||
"""Handler for /v1/chat/completions requests"""
|
"""Handler for /v1/chat/completions requests"""
|
||||||
|
|
||||||
@@ -321,6 +350,91 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
def _request_id_prefix(self) -> str:
|
def _request_id_prefix(self) -> str:
|
||||||
return "chatcmpl-"
|
return "chatcmpl-"
|
||||||
|
|
||||||
|
def _effective_tools(self, request: ChatCompletionRequest) -> List[Tool]:
|
||||||
|
tools = list(request.tools or [])
|
||||||
|
for message in request.messages:
|
||||||
|
if (
|
||||||
|
isinstance(message, ChatCompletionMessageGenericParam)
|
||||||
|
and message.role in ("system", "developer")
|
||||||
|
and message.tools
|
||||||
|
):
|
||||||
|
tools.extend(message.tools)
|
||||||
|
return tools
|
||||||
|
|
||||||
|
def _prepare_kimi_k3_messages(
|
||||||
|
self,
|
||||||
|
messages: List[Dict[str, Any]],
|
||||||
|
request: ChatCompletionRequest,
|
||||||
|
) -> tuple[List[Dict[str, Any]], int, Optional[str]]:
|
||||||
|
image_count = 0
|
||||||
|
for index, message in enumerate(messages):
|
||||||
|
content = message.get("content")
|
||||||
|
if isinstance(content, list):
|
||||||
|
parts = []
|
||||||
|
for part in content:
|
||||||
|
if not isinstance(part, dict):
|
||||||
|
continue
|
||||||
|
part_type = part.get("type")
|
||||||
|
if part_type in ("text", "input_text"):
|
||||||
|
parts.append(
|
||||||
|
{
|
||||||
|
"type": "text",
|
||||||
|
"text": neutralize_kimi_k3_image_placeholder(
|
||||||
|
part["text"]
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
elif part_type in ("image_url", "input_image"):
|
||||||
|
image = part.get("image_url") or {}
|
||||||
|
if isinstance(image, str):
|
||||||
|
image = {"url": image, "detail": part.get("detail")}
|
||||||
|
parts.append({"type": "image_url", "image_url": image})
|
||||||
|
image_count += 1
|
||||||
|
message["content"] = parts
|
||||||
|
elif isinstance(content, str):
|
||||||
|
message["content"] = neutralize_kimi_k3_image_placeholder(content)
|
||||||
|
elif content is None:
|
||||||
|
message["content"] = ""
|
||||||
|
|
||||||
|
if message.get("role") == "assistant":
|
||||||
|
for key in ("reasoning_content", "reasoning"):
|
||||||
|
if key in message:
|
||||||
|
message[key] = neutralize_kimi_k3_image_placeholder_value(
|
||||||
|
message[key]
|
||||||
|
)
|
||||||
|
for tool_call in message.get("tool_calls") or []:
|
||||||
|
function = (
|
||||||
|
tool_call.get("function")
|
||||||
|
if isinstance(tool_call, dict)
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if isinstance(function, dict) and "arguments" in function:
|
||||||
|
function["arguments"] = (
|
||||||
|
neutralize_kimi_k3_image_placeholder_value(
|
||||||
|
function["arguments"]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
source = request.messages[index]
|
||||||
|
if (
|
||||||
|
isinstance(source, ChatCompletionMessageGenericParam)
|
||||||
|
and source.role in ("system", "developer")
|
||||||
|
and source.tools
|
||||||
|
):
|
||||||
|
message["tools"] = [
|
||||||
|
tool.model_dump(exclude_unset=True, by_alias=True)
|
||||||
|
for tool in source.tools
|
||||||
|
]
|
||||||
|
if message.get("role") == "developer":
|
||||||
|
message["role"] = "system"
|
||||||
|
|
||||||
|
assistant_prefix = None
|
||||||
|
if request.continue_final_message:
|
||||||
|
messages, assistant_prefix = self._handle_last_assistant_message(
|
||||||
|
messages, request
|
||||||
|
)
|
||||||
|
return messages, image_count, assistant_prefix
|
||||||
|
|
||||||
def _encode_messages(
|
def _encode_messages(
|
||||||
self,
|
self,
|
||||||
messages: List[Dict[str, Any]],
|
messages: List[Dict[str, Any]],
|
||||||
@@ -370,6 +484,71 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
*inkling_tokenizer.encode_text(assistant_prefix),
|
*inkling_tokenizer.encode_text(assistant_prefix),
|
||||||
]
|
]
|
||||||
return prompt_ids
|
return prompt_ids
|
||||||
|
if self.chat_encoding_spec == "kimi_k3":
|
||||||
|
messages, image_count, assistant_prefix = self._prepare_kimi_k3_messages(
|
||||||
|
messages, request
|
||||||
|
)
|
||||||
|
template_kwargs = dict(request.chat_template_kwargs or {})
|
||||||
|
template_kwargs.pop("tokenize", None)
|
||||||
|
template_kwargs.pop("return_dict", None)
|
||||||
|
template_kwargs.pop("image_prompts", None)
|
||||||
|
if image_count:
|
||||||
|
template_kwargs["image_prompts"] = ["<|media_pad|>"] * image_count
|
||||||
|
|
||||||
|
if (
|
||||||
|
request.reasoning_effort in ("low", "high", "max")
|
||||||
|
and "thinking_effort" not in template_kwargs
|
||||||
|
):
|
||||||
|
template_kwargs["thinking_effort"] = request.reasoning_effort
|
||||||
|
elif request.reasoning_effort not in (
|
||||||
|
None,
|
||||||
|
"none",
|
||||||
|
"low",
|
||||||
|
"high",
|
||||||
|
"max",
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"Kimi K3 does not support reasoning_effort=%r; using the "
|
||||||
|
"encoder default.",
|
||||||
|
request.reasoning_effort,
|
||||||
|
)
|
||||||
|
|
||||||
|
effective_tools = self._effective_tools(request)
|
||||||
|
if (
|
||||||
|
effective_tools
|
||||||
|
and isinstance(request.tool_choice, str)
|
||||||
|
and request.tool_choice in ("required", "none")
|
||||||
|
):
|
||||||
|
template_kwargs.setdefault("tool_choice", request.tool_choice)
|
||||||
|
if request.response_format is not None:
|
||||||
|
template_kwargs.setdefault(
|
||||||
|
"response_format",
|
||||||
|
request.response_format.model_dump(
|
||||||
|
exclude_unset=True, by_alias=True
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
request_tools = (
|
||||||
|
[
|
||||||
|
tool.model_dump(exclude_unset=True, by_alias=True)
|
||||||
|
for tool in request.tools
|
||||||
|
]
|
||||||
|
if request.tools
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
prompt_ids = self.tokenizer_manager.tokenizer.apply_chat_template(
|
||||||
|
messages,
|
||||||
|
tokenize=True,
|
||||||
|
add_generation_prompt=True,
|
||||||
|
tools=request_tools,
|
||||||
|
return_dict=False,
|
||||||
|
**template_kwargs,
|
||||||
|
)
|
||||||
|
if assistant_prefix:
|
||||||
|
prompt_ids = self._append_assistant_prefix_to_prompt_ids(
|
||||||
|
prompt_ids, assistant_prefix
|
||||||
|
)
|
||||||
|
return prompt_ids
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -531,7 +710,11 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
remaining_logprobs = None
|
remaining_logprobs = None
|
||||||
|
|
||||||
# Handle tool calls
|
# Handle tool calls
|
||||||
if request.tool_choice != "none" and request.tools and self.tool_call_parser:
|
if (
|
||||||
|
request.tool_choice != "none"
|
||||||
|
and self._effective_tools(request)
|
||||||
|
and self.tool_call_parser
|
||||||
|
):
|
||||||
async for chunk in self._process_tool_call_stream(
|
async for chunk in self._process_tool_call_stream(
|
||||||
index,
|
index,
|
||||||
delta,
|
delta,
|
||||||
@@ -611,23 +794,37 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
if media_error:
|
if media_error:
|
||||||
return media_error
|
return media_error
|
||||||
|
|
||||||
|
effective_tools = self._effective_tools(request)
|
||||||
|
has_message_tools = any(
|
||||||
|
isinstance(message, ChatCompletionMessageGenericParam)
|
||||||
|
and message.role in ("system", "developer")
|
||||||
|
and message.tools
|
||||||
|
for message in request.messages
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
isinstance(request.tool_choice, str)
|
isinstance(request.tool_choice, str)
|
||||||
and request.tool_choice.lower() == "required"
|
and request.tool_choice.lower() == "required"
|
||||||
and not request.tools
|
and not effective_tools
|
||||||
):
|
):
|
||||||
return "Tools cannot be empty if tool choice is set to required."
|
return "Tools cannot be empty if tool choice is set to required."
|
||||||
|
|
||||||
if request.tool_choice is not None and not isinstance(request.tool_choice, str):
|
if request.tool_choice is not None and not isinstance(request.tool_choice, str):
|
||||||
if not request.tools:
|
if not effective_tools:
|
||||||
return "Tools cannot be empty if tool choice is set to a specific tool."
|
return "Tools cannot be empty if tool choice is set to a specific tool."
|
||||||
tool_name = request.tool_choice.function.name
|
tool_name = request.tool_choice.function.name
|
||||||
tool_exists = any(tool.function.name == tool_name for tool in request.tools)
|
tool_exists = any(
|
||||||
|
tool.function.name == tool_name for tool in effective_tools
|
||||||
|
)
|
||||||
if not tool_exists:
|
if not tool_exists:
|
||||||
return f"Tool '{tool_name}' not found in tools list."
|
return f"Tool '{tool_name}' not found in tools list."
|
||||||
|
|
||||||
|
if has_message_tools:
|
||||||
|
names = [tool.function.name for tool in effective_tools]
|
||||||
|
if len(names) != len(set(names)):
|
||||||
|
return "Tool names must be unique across request and message tools."
|
||||||
|
|
||||||
# Validate tool definitions
|
# Validate tool definitions
|
||||||
for i, tool in enumerate(request.tools or []):
|
for i, tool in enumerate(effective_tools):
|
||||||
if tool.function.parameters is None:
|
if tool.function.parameters is None:
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
@@ -731,11 +928,14 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
stop=processed_messages.stop,
|
stop=processed_messages.stop,
|
||||||
model_generation_config=self.default_sampling_params,
|
model_generation_config=self.default_sampling_params,
|
||||||
tool_call_constraint=processed_messages.tool_call_constraint,
|
tool_call_constraint=processed_messages.tool_call_constraint,
|
||||||
|
renderer_handles_response_format=self.chat_encoding_spec == "kimi_k3",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Handle single vs multiple requests
|
# Handle single vs multiple requests
|
||||||
if request.input_ids is not None:
|
if request.input_ids is not None:
|
||||||
prompt_kwargs = {"input_ids": processed_messages.prompt_ids}
|
prompt_kwargs = {"input_ids": processed_messages.prompt_ids}
|
||||||
|
elif is_multimodal and self.chat_encoding_spec == "kimi_k3":
|
||||||
|
prompt_kwargs = {"input_ids": processed_messages.prompt_ids}
|
||||||
elif is_multimodal:
|
elif is_multimodal:
|
||||||
# Standard VLMs render a text prompt (with placeholder strings) for the MM
|
# Standard VLMs render a text prompt (with placeholder strings) for the MM
|
||||||
# processor to tokenize. Inkling's custom encoder instead produces pre-rendered
|
# processor to tokenize. Inkling's custom encoder instead produces pre-rendered
|
||||||
@@ -769,8 +969,6 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
img_max_dynamic_patch, vid_max_dynamic_patch = _extract_max_dynamic_patch(
|
img_max_dynamic_patch, vid_max_dynamic_patch = _extract_max_dynamic_patch(
|
||||||
request
|
request
|
||||||
)
|
)
|
||||||
require_reasoning = self._get_reasoning_from_request(request)
|
|
||||||
|
|
||||||
adapted_request = GenerateReqInput(
|
adapted_request = GenerateReqInput(
|
||||||
**prompt_kwargs,
|
**prompt_kwargs,
|
||||||
image_data=processed_messages.image_data,
|
image_data=processed_messages.image_data,
|
||||||
@@ -795,7 +993,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
rid=request.rid,
|
rid=request.rid,
|
||||||
session_id=request.session_id,
|
session_id=request.session_id,
|
||||||
extra_key=self._compute_extra_key(request),
|
extra_key=self._compute_extra_key(request),
|
||||||
require_reasoning=require_reasoning,
|
require_reasoning=processed_messages.require_reasoning,
|
||||||
priority=request.priority,
|
priority=request.priority,
|
||||||
routing_key=self.extract_routing_key(raw_request),
|
routing_key=self.extract_routing_key(raw_request),
|
||||||
custom_labels=custom_labels,
|
custom_labels=custom_labels,
|
||||||
@@ -841,19 +1039,22 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
|
|
||||||
# Apply chat template and its stop strings
|
# Apply chat template and its stop strings
|
||||||
tools = None
|
tools = None
|
||||||
if request.tools and request.tool_choice != "none":
|
tool_call_stop = None
|
||||||
|
required_parsed_natively = False
|
||||||
|
effective_tools = self._effective_tools(request)
|
||||||
|
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):
|
||||||
tools = [
|
tools = [
|
||||||
item.model_dump()
|
item.model_dump()
|
||||||
for item in request.tools
|
for item in request.tools or []
|
||||||
if item.function.name == request.tool_choice.function.name
|
if item.function.name == request.tool_choice.function.name
|
||||||
]
|
] or None
|
||||||
else:
|
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:
|
||||||
parser = FunctionCallParser(
|
parser = FunctionCallParser(
|
||||||
request.tools,
|
effective_tools,
|
||||||
self.tool_call_parser,
|
self.tool_call_parser,
|
||||||
tokenizer=self.tokenizer_manager.tokenizer,
|
tokenizer=self.tokenizer_manager.tokenizer,
|
||||||
)
|
)
|
||||||
@@ -862,14 +1063,23 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
parallel_tool_calls=request.parallel_tool_calls,
|
parallel_tool_calls=request.parallel_tool_calls,
|
||||||
thinking_mode=xgrammar_reasoning,
|
thinking_mode=xgrammar_reasoning,
|
||||||
)
|
)
|
||||||
# Fallback: use generic JSON schema for required/named tool choice
|
required_parsed_natively = parser.detector.parses_required_natively()
|
||||||
# only when no parser-specific constraint was set
|
if self.chat_encoding_spec == "kimi_k3":
|
||||||
if tool_call_constraint is None and (
|
tool_call_stop = parser.detector.eot_token
|
||||||
|
if (
|
||||||
|
tool_call_constraint is None
|
||||||
|
and not required_parsed_natively
|
||||||
|
and not (
|
||||||
|
self.chat_encoding_spec == "kimi_k3"
|
||||||
|
and self.tool_call_parser == "kimi_k3"
|
||||||
|
)
|
||||||
|
and (
|
||||||
request.tool_choice == "required"
|
request.tool_choice == "required"
|
||||||
or isinstance(request.tool_choice, ToolChoice)
|
or isinstance(request.tool_choice, ToolChoice)
|
||||||
|
)
|
||||||
):
|
):
|
||||||
json_schema = get_json_schema_constraint(
|
json_schema = get_json_schema_constraint(
|
||||||
request.tools,
|
effective_tools,
|
||||||
request.tool_choice,
|
request.tool_choice,
|
||||||
parallel_tool_calls=request.parallel_tool_calls,
|
parallel_tool_calls=request.parallel_tool_calls,
|
||||||
)
|
)
|
||||||
@@ -892,7 +1102,18 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
else:
|
else:
|
||||||
result = self._apply_conversation_template(request, is_multimodal)
|
result = self._apply_conversation_template(request, is_multimodal)
|
||||||
|
|
||||||
|
if tool_call_stop is not None:
|
||||||
|
if isinstance(result.stop, str):
|
||||||
|
result.stop = [result.stop]
|
||||||
|
elif result.stop is None:
|
||||||
|
result.stop = []
|
||||||
|
else:
|
||||||
|
result.stop = list(result.stop)
|
||||||
|
if tool_call_stop not in result.stop:
|
||||||
|
result.stop.append(tool_call_stop)
|
||||||
|
|
||||||
result.tool_call_constraint = tool_call_constraint
|
result.tool_call_constraint = tool_call_constraint
|
||||||
|
result.require_reasoning = thinking_mode
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def _apply_jinja_template(
|
def _apply_jinja_template(
|
||||||
@@ -921,7 +1142,9 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
)
|
)
|
||||||
messages = [msg.model_dump() for msg in request.messages]
|
messages = [msg.model_dump() for msg in request.messages]
|
||||||
for message in messages:
|
for message in messages:
|
||||||
normalize_assistant_tool_call_arguments(message)
|
normalize_assistant_tool_call_arguments(
|
||||||
|
message, strict=self.chat_encoding_spec != "kimi_k3"
|
||||||
|
)
|
||||||
|
|
||||||
prompt_ids = self._encode_messages(
|
prompt_ids = self._encode_messages(
|
||||||
copy.deepcopy(messages),
|
copy.deepcopy(messages),
|
||||||
@@ -931,10 +1154,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if prompt_ids is not None:
|
if prompt_ids is not None:
|
||||||
# Custom encoding produced prompt_ids. Text-only encoders (dsv4/dsv32) need
|
if self.chat_encoding_spec in ("inkling", "kimi_k3"):
|
||||||
# nothing more; Inkling is the only multimodal custom encoder and still needs the
|
|
||||||
# image/audio media harvested from the messages for the MM processor.
|
|
||||||
if self.chat_encoding_spec == "inkling":
|
|
||||||
for message in request.messages:
|
for message in request.messages:
|
||||||
msg_dict = message.model_dump()
|
msg_dict = message.model_dump()
|
||||||
if msg_dict.get("content") is None:
|
if msg_dict.get("content") is None:
|
||||||
@@ -1557,15 +1777,16 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
|
|
||||||
# Handle tool calls
|
# Handle tool calls
|
||||||
tool_calls = None
|
tool_calls = None
|
||||||
|
effective_tools = self._effective_tools(request)
|
||||||
if (
|
if (
|
||||||
request.tool_choice != "none"
|
request.tool_choice != "none"
|
||||||
and request.tools
|
and effective_tools
|
||||||
and self.tool_call_parser
|
and self.tool_call_parser
|
||||||
):
|
):
|
||||||
history_tool_calls_cnt = self._get_history_tool_calls_cnt(request)
|
history_tool_calls_cnt = self._get_history_tool_calls_cnt(request)
|
||||||
tool_calls, text, finish_reason = self._process_tool_calls(
|
tool_calls, text, finish_reason = self._process_tool_calls(
|
||||||
text,
|
text,
|
||||||
request.tools,
|
effective_tools,
|
||||||
finish_reason,
|
finish_reason,
|
||||||
request.tool_choice,
|
request.tool_choice,
|
||||||
history_tool_calls_cnt,
|
history_tool_calls_cnt,
|
||||||
@@ -1702,15 +1923,16 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
history_tool_calls_cnt: int,
|
history_tool_calls_cnt: int,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Process for generating a new and unique `tool_call_id`"""
|
"""Process for generating a new and unique `tool_call_id`"""
|
||||||
|
if self.tool_call_parser == "kimi_k3":
|
||||||
|
return f"{call_item.name}:{history_tool_calls_cnt + call_item.tool_index}"
|
||||||
if self.tool_call_parser != "kimi_k2":
|
if self.tool_call_parser != "kimi_k2":
|
||||||
# A simple uuid is sufficient for all models except for Kimi-K2.
|
# A simple uuid is sufficient for all models except for Kimi-K2.
|
||||||
tool_call_id = f"call_{uuid.uuid4().hex[:24]}"
|
tool_call_id = f"call_{uuid.uuid4().hex[:24]}"
|
||||||
return tool_call_id
|
return tool_call_id
|
||||||
else:
|
tool_call_id = (
|
||||||
# Align with Kimi-K2 format: functions.{name}:{index}
|
f"functions.{call_item.name}:"
|
||||||
# Kimi-K2 allows multiple tool_calls in one message; SGLang sets call_item.tool_index to the *local* position inside that message.
|
f"{history_tool_calls_cnt + call_item.tool_index}"
|
||||||
# Therefore, the index must be corrected by using `history_tool_calls_cnt + call_item.tool_index` to ensure globally unique and properly ordered.
|
)
|
||||||
tool_call_id = f"functions.{call_item.name}:{history_tool_calls_cnt+call_item.tool_index}"
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Process tool call idx, parser: {self.tool_call_parser}, tool_call_id: {tool_call_id}, history_cnt: {history_tool_calls_cnt}"
|
f"Process tool call idx, parser: {self.tool_call_parser}, tool_call_id: {tool_call_id}, history_cnt: {history_tool_calls_cnt}"
|
||||||
)
|
)
|
||||||
@@ -1736,7 +1958,9 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
tools, self.tool_call_parser, tokenizer=self.tokenizer_manager.tokenizer
|
tools, self.tool_call_parser, tokenizer=self.tokenizer_manager.tokenizer
|
||||||
)
|
)
|
||||||
should_try_parser = (
|
should_try_parser = (
|
||||||
not is_required or parser.detector.supports_structural_tag()
|
not is_required
|
||||||
|
or parser.detector.supports_structural_tag()
|
||||||
|
or parser.detector.parses_required_natively()
|
||||||
)
|
)
|
||||||
if should_try_parser and parser.has_tool_call(text):
|
if should_try_parser and parser.has_tool_call(text):
|
||||||
try:
|
try:
|
||||||
@@ -1891,6 +2115,8 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
"""
|
"""
|
||||||
if self.reasoning_parser == "apertus2509":
|
if self.reasoning_parser == "apertus2509":
|
||||||
request.skip_special_tokens = False
|
request.skip_special_tokens = False
|
||||||
|
if self.reasoning_parser == "kimi_k3" or self.chat_encoding_spec == "kimi_k3":
|
||||||
|
request.skip_special_tokens = False
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.reasoning_parser in ["mistral"]
|
self.reasoning_parser in ["mistral"]
|
||||||
@@ -2095,6 +2321,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
continuous_usage_stats: bool = False,
|
continuous_usage_stats: bool = False,
|
||||||
):
|
):
|
||||||
"""Process tool calls in streaming response"""
|
"""Process tool calls in streaming response"""
|
||||||
|
effective_tools = self._effective_tools(request)
|
||||||
if index not in parser_dict:
|
if index not in parser_dict:
|
||||||
is_required = request.tool_choice == "required" or isinstance(
|
is_required = request.tool_choice == "required" or isinstance(
|
||||||
request.tool_choice, ToolChoice
|
request.tool_choice, ToolChoice
|
||||||
@@ -2108,18 +2335,21 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
use_native_parser = False
|
use_native_parser = False
|
||||||
if self.tool_call_parser:
|
if self.tool_call_parser:
|
||||||
probe = FunctionCallParser(
|
probe = FunctionCallParser(
|
||||||
tools=request.tools,
|
tools=effective_tools,
|
||||||
tool_call_parser=self.tool_call_parser,
|
tool_call_parser=self.tool_call_parser,
|
||||||
tokenizer=self.tokenizer_manager.tokenizer,
|
tokenizer=self.tokenizer_manager.tokenizer,
|
||||||
)
|
)
|
||||||
use_native_parser = probe.detector.supports_structural_tag()
|
use_native_parser = (
|
||||||
|
probe.detector.supports_structural_tag()
|
||||||
|
or probe.detector.parses_required_natively()
|
||||||
|
)
|
||||||
if use_native_parser:
|
if use_native_parser:
|
||||||
parser_dict[index] = probe
|
parser_dict[index] = probe
|
||||||
else:
|
else:
|
||||||
parser_dict[index] = JsonArrayParser()
|
parser_dict[index] = JsonArrayParser()
|
||||||
else:
|
else:
|
||||||
parser_dict[index] = FunctionCallParser(
|
parser_dict[index] = FunctionCallParser(
|
||||||
tools=request.tools,
|
tools=effective_tools,
|
||||||
tool_call_parser=self.tool_call_parser,
|
tool_call_parser=self.tool_call_parser,
|
||||||
tokenizer=self.tokenizer_manager.tokenizer,
|
tokenizer=self.tokenizer_manager.tokenizer,
|
||||||
)
|
)
|
||||||
@@ -2128,7 +2358,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
|
|
||||||
# Handle both FunctionCallParser and JsonArrayParser
|
# Handle both FunctionCallParser and JsonArrayParser
|
||||||
if isinstance(parser, JsonArrayParser):
|
if isinstance(parser, JsonArrayParser):
|
||||||
result = parser.parse_streaming_increment(delta, request.tools)
|
result = parser.parse_streaming_increment(delta, effective_tools)
|
||||||
normal_text, calls = result.normal_text, result.calls
|
normal_text, calls = result.normal_text, result.calls
|
||||||
else:
|
else:
|
||||||
normal_text, calls = parser.parse_stream_chunk(delta)
|
normal_text, calls = parser.parse_stream_chunk(delta)
|
||||||
|
|||||||
@@ -235,6 +235,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
messages, request_prompts, engine_prompts = (
|
messages, request_prompts, engine_prompts = (
|
||||||
self._make_request_with_harmony(request, prev_response)
|
self._make_request_with_harmony(request, prev_response)
|
||||||
)
|
)
|
||||||
|
require_reasoning = self._is_thinking_enabled_for_request(request)
|
||||||
else:
|
else:
|
||||||
(
|
(
|
||||||
messages,
|
messages,
|
||||||
@@ -242,6 +243,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
engine_prompts,
|
engine_prompts,
|
||||||
processed_messages,
|
processed_messages,
|
||||||
) = await self._make_request(request, prev_response, tokenizer)
|
) = await self._make_request(request, prev_response, tokenizer)
|
||||||
|
require_reasoning = processed_messages.require_reasoning
|
||||||
|
|
||||||
except _MediaInputValidationError as e:
|
except _MediaInputValidationError as e:
|
||||||
return self.create_error_response(str(e))
|
return self.create_error_response(str(e))
|
||||||
@@ -369,6 +371,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
session_id=request.session_id,
|
session_id=request.session_id,
|
||||||
extra_key=self._compute_extra_key(request),
|
extra_key=self._compute_extra_key(request),
|
||||||
background=request.background,
|
background=request.background,
|
||||||
|
require_reasoning=require_reasoning,
|
||||||
)
|
)
|
||||||
|
|
||||||
generator = self._generate_with_builtin_tools(
|
generator = self._generate_with_builtin_tools(
|
||||||
@@ -416,6 +419,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
request_metadata,
|
request_metadata,
|
||||||
created_time,
|
created_time,
|
||||||
|
require_reasoning=require_reasoning,
|
||||||
),
|
),
|
||||||
name=f"create_{response.id}",
|
name=f"create_{response.id}",
|
||||||
)
|
)
|
||||||
@@ -437,6 +441,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
model_name,
|
model_name,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
request_metadata,
|
request_metadata,
|
||||||
|
require_reasoning=require_reasoning,
|
||||||
)
|
)
|
||||||
return self.responses_stream_generator_non_harmony(
|
return self.responses_stream_generator_non_harmony(
|
||||||
request,
|
request,
|
||||||
@@ -445,6 +450,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
model_name,
|
model_name,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
request_metadata,
|
request_metadata,
|
||||||
|
require_reasoning=require_reasoning,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
result: Union[ORJSONResponse, ResponsesResponse] = (
|
result: Union[ORJSONResponse, ResponsesResponse] = (
|
||||||
@@ -456,6 +462,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
model_name,
|
model_name,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
request_metadata,
|
request_metadata,
|
||||||
|
require_reasoning=require_reasoning,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
@@ -527,6 +534,8 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
request_metadata: RequestResponseMetadata,
|
request_metadata: RequestResponseMetadata,
|
||||||
created_time: Optional[int] = None,
|
created_time: Optional[int] = None,
|
||||||
|
*,
|
||||||
|
require_reasoning: bool,
|
||||||
) -> Union[ResponsesResponse, ORJSONResponse]:
|
) -> Union[ResponsesResponse, ORJSONResponse]:
|
||||||
if created_time is None:
|
if created_time is None:
|
||||||
created_time = int(time.time())
|
created_time = int(time.time())
|
||||||
@@ -553,7 +562,10 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
assert final_res is not None
|
assert final_res is not None
|
||||||
|
|
||||||
output = self._make_response_output_items(
|
output = self._make_response_output_items(
|
||||||
request, final_res["text"], tokenizer
|
request,
|
||||||
|
final_res["text"],
|
||||||
|
tokenizer,
|
||||||
|
require_reasoning=require_reasoning,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Calculate usage from actual output
|
# Calculate usage from actual output
|
||||||
@@ -636,7 +648,6 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
return request.reasoning is not None and request.reasoning.summary is not None
|
return request.reasoning is not None and request.reasoning.summary is not None
|
||||||
|
|
||||||
def _is_thinking_enabled_for_request(self, request: ResponsesRequest) -> bool:
|
def _is_thinking_enabled_for_request(self, request: ResponsesRequest) -> bool:
|
||||||
"""Whether to start the reasoning detector in thinking mode."""
|
|
||||||
if not self.reasoning_parser:
|
if not self.reasoning_parser:
|
||||||
return False
|
return False
|
||||||
effort = request.reasoning.effort if request.reasoning is not None else None
|
effort = request.reasoning.effort if request.reasoning is not None else None
|
||||||
@@ -674,14 +685,14 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
request: ResponsesRequest,
|
request: ResponsesRequest,
|
||||||
final_output: Any,
|
final_output: Any,
|
||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
|
*,
|
||||||
|
require_reasoning: bool,
|
||||||
):
|
):
|
||||||
if self.reasoning_parser:
|
if self.reasoning_parser:
|
||||||
# Templates that prefill ``<think>`` only emit the close tag, so
|
|
||||||
# start the detector in thinking mode.
|
|
||||||
reasoning_parser = ReasoningParser(
|
reasoning_parser = ReasoningParser(
|
||||||
model_type=self.reasoning_parser,
|
model_type=self.reasoning_parser,
|
||||||
stream_reasoning=False,
|
stream_reasoning=False,
|
||||||
force_reasoning=self._is_thinking_enabled_for_request(request),
|
force_reasoning=require_reasoning,
|
||||||
request=request,
|
request=request,
|
||||||
tokenizer=self.tokenizer_manager.tokenizer,
|
tokenizer=self.tokenizer_manager.tokenizer,
|
||||||
)
|
)
|
||||||
@@ -1201,8 +1212,8 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
request_metadata: RequestResponseMetadata,
|
request_metadata: RequestResponseMetadata,
|
||||||
created_time: Optional[int] = None,
|
created_time: Optional[int] = None,
|
||||||
*args,
|
*,
|
||||||
**kwargs,
|
require_reasoning: bool,
|
||||||
):
|
):
|
||||||
try:
|
try:
|
||||||
# Update the status to "in_progress"
|
# Update the status to "in_progress"
|
||||||
@@ -1220,8 +1231,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
request_metadata,
|
request_metadata,
|
||||||
created_time,
|
created_time,
|
||||||
*args,
|
require_reasoning=require_reasoning,
|
||||||
**kwargs,
|
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception("Background request failed for %s", request.request_id)
|
logger.exception("Background request failed for %s", request.request_id)
|
||||||
@@ -1311,6 +1321,8 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
request_metadata: RequestResponseMetadata,
|
request_metadata: RequestResponseMetadata,
|
||||||
created_time: Optional[int] = None,
|
created_time: Optional[int] = None,
|
||||||
|
*,
|
||||||
|
require_reasoning: bool,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
# TODO:
|
# TODO:
|
||||||
# 1. Handle disconnect
|
# 1. Handle disconnect
|
||||||
@@ -1717,6 +1729,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
request_metadata,
|
request_metadata,
|
||||||
created_time=created_time,
|
created_time=created_time,
|
||||||
|
require_reasoning=require_reasoning,
|
||||||
)
|
)
|
||||||
# Convert final_response to the format expected by ResponseCompletedEvent
|
# Convert final_response to the format expected by ResponseCompletedEvent
|
||||||
response_dict = final_response.model_dump()
|
response_dict = final_response.model_dump()
|
||||||
@@ -1755,6 +1768,8 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
request_metadata: RequestResponseMetadata,
|
request_metadata: RequestResponseMetadata,
|
||||||
created_time: Optional[int] = None,
|
created_time: Optional[int] = None,
|
||||||
|
*,
|
||||||
|
require_reasoning: bool,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
"""Stream a /v1/responses response as typed OpenAI SSE events for
|
"""Stream a /v1/responses response as typed OpenAI SSE events for
|
||||||
non-harmony models. Each engine chunk is run through the reasoning
|
non-harmony models. Each engine chunk is run through the reasoning
|
||||||
@@ -1836,7 +1851,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
reasoning_parser_obj = ReasoningParser(
|
reasoning_parser_obj = ReasoningParser(
|
||||||
model_type=self.reasoning_parser,
|
model_type=self.reasoning_parser,
|
||||||
stream_reasoning=True,
|
stream_reasoning=True,
|
||||||
force_reasoning=self._is_thinking_enabled_for_request(request),
|
force_reasoning=require_reasoning,
|
||||||
request=request,
|
request=request,
|
||||||
tokenizer=self.tokenizer_manager.tokenizer,
|
tokenizer=self.tokenizer_manager.tokenizer,
|
||||||
)
|
)
|
||||||
@@ -2396,6 +2411,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
|||||||
return_text_in_logprobs=adapted_request.return_text_in_logprobs,
|
return_text_in_logprobs=adapted_request.return_text_in_logprobs,
|
||||||
return_hidden_states=adapted_request.return_hidden_states,
|
return_hidden_states=adapted_request.return_hidden_states,
|
||||||
background=adapted_request.background,
|
background=adapted_request.background,
|
||||||
|
require_reasoning=adapted_request.require_reasoning,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Update sampling params with reduced max_tokens
|
# Update sampling params with reduced max_tokens
|
||||||
|
|||||||
@@ -354,6 +354,11 @@ class BaseFormatDetector(ABC):
|
|||||||
"""Return True if this detector supports structural tag format."""
|
"""Return True if this detector supports structural tag format."""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
def parses_required_natively(self) -> bool:
|
||||||
|
"""Return True if ``tool_choice="required"`` must skip grammar
|
||||||
|
constraints and parse the model's native output format instead."""
|
||||||
|
return False
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def structure_info(self) -> _GetInfoFunc:
|
def structure_info(self) -> _GetInfoFunc:
|
||||||
"""
|
"""
|
||||||
@@ -377,6 +382,7 @@ class BaseFormatDetector(ABC):
|
|||||||
tools: Union[List[Tool], None] = None,
|
tools: Union[List[Tool], None] = None,
|
||||||
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
||||||
thinking_mode: bool = False,
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
) -> Optional[StructuralTag]:
|
) -> Optional[StructuralTag]:
|
||||||
"""
|
"""
|
||||||
Return a model-native XGrammar structural tag when supported.
|
Return a model-native XGrammar structural tag when supported.
|
||||||
@@ -389,6 +395,11 @@ class BaseFormatDetector(ABC):
|
|||||||
ReasonerGrammarBackend will own the <think>...</think> prefix
|
ReasonerGrammarBackend will own the <think>...</think> prefix
|
||||||
(the typical case when --reasoning-parser is configured) so
|
(the typical case when --reasoning-parser is configured) so
|
||||||
only one layer constrains the reasoning section.
|
only one layer constrains the reasoning section.
|
||||||
|
parallel_tool_calls: Whether multiple tool calls may appear in one
|
||||||
|
assistant response. xgrammar's get_model_structural_tag does
|
||||||
|
not expose this knob, so this base implementation ignores it;
|
||||||
|
only detectors that build their own tags (e.g. Kimi K3)
|
||||||
|
honor it.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
StructuralTag if this detector supports model-native tags, otherwise None
|
StructuralTag if this detector supports model-native tags, otherwise None
|
||||||
@@ -411,7 +422,10 @@ class BaseFormatDetector(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_auto_tool_call_structural_tag(
|
def get_auto_tool_call_structural_tag(
|
||||||
self, tools: Union[List[Tool], None] = None
|
self,
|
||||||
|
tools: Union[List[Tool], None] = None,
|
||||||
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
) -> Optional[StructuralTag]:
|
) -> Optional[StructuralTag]:
|
||||||
"""Return an always-on structural tag for automatic tool choice.
|
"""Return an always-on structural tag for automatic tool choice.
|
||||||
|
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ from sglang.srt.function_call.hunyuan_detector import HunyuanDetector
|
|||||||
from sglang.srt.function_call.inkling_detector import InklingDetector
|
from sglang.srt.function_call.inkling_detector import InklingDetector
|
||||||
from sglang.srt.function_call.internlm_detector import InternlmDetector
|
from sglang.srt.function_call.internlm_detector import InternlmDetector
|
||||||
from sglang.srt.function_call.kimik2_detector import KimiK2Detector
|
from sglang.srt.function_call.kimik2_detector import KimiK2Detector
|
||||||
|
from sglang.srt.function_call.kimik3_detector import KimiK3Detector
|
||||||
from sglang.srt.function_call.lfm2_detector import Lfm2Detector
|
from sglang.srt.function_call.lfm2_detector import Lfm2Detector
|
||||||
from sglang.srt.function_call.llama32_detector import Llama32Detector
|
from sglang.srt.function_call.llama32_detector import Llama32Detector
|
||||||
from sglang.srt.function_call.mimo_detector import MiMoDetector
|
from sglang.srt.function_call.mimo_detector import MiMoDetector
|
||||||
@@ -71,6 +72,7 @@ class FunctionCallParser:
|
|||||||
"glm47": Glm47MoeDetector,
|
"glm47": Glm47MoeDetector,
|
||||||
"gpt-oss": GptOssDetector,
|
"gpt-oss": GptOssDetector,
|
||||||
"kimi_k2": KimiK2Detector,
|
"kimi_k2": KimiK2Detector,
|
||||||
|
"kimi_k3": KimiK3Detector,
|
||||||
"lfm2": Lfm2Detector,
|
"lfm2": Lfm2Detector,
|
||||||
"llama3": Llama32Detector,
|
"llama3": Llama32Detector,
|
||||||
"mimo": MiMoDetector,
|
"mimo": MiMoDetector,
|
||||||
@@ -252,16 +254,31 @@ class FunctionCallParser:
|
|||||||
try:
|
try:
|
||||||
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,
|
||||||
|
thinking_mode=thinking_mode,
|
||||||
|
parallel_tool_calls=parallel_tool_calls,
|
||||||
)
|
)
|
||||||
if structural_tag is not None:
|
if structural_tag is not None:
|
||||||
return ("structural_tag", structural_tag)
|
return ("structural_tag", structural_tag)
|
||||||
|
|
||||||
if is_required or should_constrain_auto:
|
if is_required or should_constrain_auto:
|
||||||
|
structural_tag_tools = self.tools
|
||||||
|
if self.tool_strict_level >= ToolStrictLevel.PARAMETER:
|
||||||
|
structural_tag_tools = [
|
||||||
|
tool.model_copy(
|
||||||
|
update={
|
||||||
|
"function": tool.function.model_copy(
|
||||||
|
update={"strict": True}
|
||||||
|
)
|
||||||
|
}
|
||||||
|
)
|
||||||
|
for tool in self.tools
|
||||||
|
]
|
||||||
structural_tag = self.detector.get_structural_tag(
|
structural_tag = self.detector.get_structural_tag(
|
||||||
tools=self.tools,
|
tools=structural_tag_tools,
|
||||||
thinking_mode=thinking_mode,
|
thinking_mode=thinking_mode,
|
||||||
tool_choice=tool_choice,
|
tool_choice=tool_choice,
|
||||||
|
parallel_tool_calls=parallel_tool_calls,
|
||||||
)
|
)
|
||||||
if structural_tag is not None:
|
if structural_tag is not None:
|
||||||
return ("structural_tag", structural_tag)
|
return ("structural_tag", structural_tag)
|
||||||
@@ -275,7 +292,9 @@ class FunctionCallParser:
|
|||||||
tag = self.get_legacy_structural_tag(at_least_one=is_required)
|
tag = self.get_legacy_structural_tag(at_least_one=is_required)
|
||||||
return ("structural_tag", tag)
|
return ("structural_tag", tag)
|
||||||
|
|
||||||
if tool_choice == "required" or isinstance(tool_choice, ToolChoice):
|
if (
|
||||||
|
tool_choice == "required" or isinstance(tool_choice, ToolChoice)
|
||||||
|
) and not self.detector.parses_required_natively():
|
||||||
json_schema = get_json_schema_constraint(
|
json_schema = get_json_schema_constraint(
|
||||||
self.tools, tool_choice, parallel_tool_calls=parallel_tool_calls
|
self.tools, tool_choice, parallel_tool_calls=parallel_tool_calls
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -822,11 +822,15 @@ class Glm47MoeDetector(BaseFormatDetector):
|
|||||||
tools: Union[List[Tool], None] = None,
|
tools: Union[List[Tool], None] = None,
|
||||||
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
||||||
thinking_mode: bool = False,
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
) -> Optional[StructuralTag]:
|
) -> Optional[StructuralTag]:
|
||||||
if not self.supports_structural_tag():
|
if not self.supports_structural_tag():
|
||||||
return None
|
return None
|
||||||
return super().get_structural_tag(
|
return super().get_structural_tag(
|
||||||
tools=tools, tool_choice=tool_choice, thinking_mode=thinking_mode
|
tools=tools,
|
||||||
|
tool_choice=tool_choice,
|
||||||
|
thinking_mode=thinking_mode,
|
||||||
|
parallel_tool_calls=parallel_tool_calls,
|
||||||
)
|
)
|
||||||
|
|
||||||
def structure_info(self) -> _GetInfoFunc:
|
def structure_info(self) -> _GetInfoFunc:
|
||||||
|
|||||||
@@ -238,7 +238,10 @@ class InklingDetector(BaseFormatDetector):
|
|||||||
return info
|
return info
|
||||||
|
|
||||||
def get_auto_tool_call_structural_tag(
|
def get_auto_tool_call_structural_tag(
|
||||||
self, tools: Optional[List[Tool]] = None
|
self,
|
||||||
|
tools: Optional[List[Tool]] = None,
|
||||||
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
) -> StructuralTag:
|
) -> StructuralTag:
|
||||||
"""Constrain JSON after Inkling's tool-payload trigger token.
|
"""Constrain JSON after Inkling's tool-payload trigger token.
|
||||||
|
|
||||||
@@ -248,7 +251,7 @@ class InklingDetector(BaseFormatDetector):
|
|||||||
``END_MESSAGE``. This mirrors the TML sampling default used by the OAI
|
``END_MESSAGE``. This mirrors the TML sampling default used by the OAI
|
||||||
API and intentionally does not restrict names to the request's tools.
|
API and intentionally does not restrict names to the request's tools.
|
||||||
"""
|
"""
|
||||||
del tools
|
del tools, thinking_mode, parallel_tool_calls
|
||||||
return StructuralTag.model_validate(
|
return StructuralTag.model_validate(
|
||||||
{
|
{
|
||||||
"type": "structural_tag",
|
"type": "structural_tag",
|
||||||
|
|||||||
@@ -430,12 +430,16 @@ class KimiK2Detector(BaseFormatDetector):
|
|||||||
tools: Union[List[Tool], None] = None,
|
tools: Union[List[Tool], None] = None,
|
||||||
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
||||||
thinking_mode: bool = False,
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
) -> Optional[StructuralTag]:
|
) -> Optional[StructuralTag]:
|
||||||
if not (
|
if not (
|
||||||
tools and (tool_choice == "required" or isinstance(tool_choice, ToolChoice))
|
tools and (tool_choice == "required" or isinstance(tool_choice, ToolChoice))
|
||||||
):
|
):
|
||||||
return super().get_structural_tag(
|
return super().get_structural_tag(
|
||||||
tools=tools, tool_choice=tool_choice, thinking_mode=thinking_mode
|
tools=tools,
|
||||||
|
tool_choice=tool_choice,
|
||||||
|
thinking_mode=thinking_mode,
|
||||||
|
parallel_tool_calls=parallel_tool_calls,
|
||||||
)
|
)
|
||||||
if get_model_structural_tag is None:
|
if get_model_structural_tag is None:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -0,0 +1,234 @@
|
|||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from typing import List, Literal, Optional, Union
|
||||||
|
|
||||||
|
from xgrammar import StructuralTag
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import Tool, ToolChoice
|
||||||
|
from sglang.srt.function_call.base_format_detector import BaseFormatDetector
|
||||||
|
from sglang.srt.function_call.core_types import (
|
||||||
|
StreamingParseResult,
|
||||||
|
ToolCallItem,
|
||||||
|
_GetInfoFunc,
|
||||||
|
)
|
||||||
|
from sglang.srt.function_call.kimik3_format import (
|
||||||
|
MESSAGE_CLOSE,
|
||||||
|
RESPONSE_CLOSE,
|
||||||
|
RESPONSE_OPEN,
|
||||||
|
TOOLS_CLOSE,
|
||||||
|
TOOLS_OPEN,
|
||||||
|
partial_suffix_len,
|
||||||
|
strip_response_wrappers,
|
||||||
|
)
|
||||||
|
from sglang.srt.function_call.kimik3_structural_tag import (
|
||||||
|
get_kimik3_auto_tool_call_structural_tag,
|
||||||
|
get_kimik3_structural_tag,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_CALL_RE = re.compile(
|
||||||
|
r"<\|open\|>call\s+(?P<attrs>(?:(?!<\|sep\|>).)*?)<\|sep\|>"
|
||||||
|
r"(?P<body>.*?)<\|close\|>call<\|sep\|>",
|
||||||
|
re.DOTALL,
|
||||||
|
)
|
||||||
|
_ARG_RE = re.compile(
|
||||||
|
r"<\|open\|>argument\s+(?P<attrs>(?:(?!<\|sep\|>).)*?)<\|sep\|>"
|
||||||
|
r"(?P<val>.*?)<\|close\|>argument<\|sep\|>",
|
||||||
|
re.DOTALL,
|
||||||
|
)
|
||||||
|
_ATTR_RE = re.compile(r'(?P<k>\w+)="(?P<v>[^"]*)"')
|
||||||
|
|
||||||
|
|
||||||
|
def _unescape_attr(value: str) -> str:
|
||||||
|
return value.replace(""", '"').replace("&", "&")
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_attrs(attrs: str) -> dict:
|
||||||
|
return {m["k"]: _unescape_attr(m["v"]) for m in _ATTR_RE.finditer(attrs)}
|
||||||
|
|
||||||
|
|
||||||
|
class KimiK3Detector(BaseFormatDetector):
|
||||||
|
"""Detector for the Kimi K3 XTML tool-call format.
|
||||||
|
|
||||||
|
K3 emits tool calls in a ``tools`` channel built from dedicated special
|
||||||
|
tokens; the plain reply lives in a preceding ``response`` channel:
|
||||||
|
|
||||||
|
```
|
||||||
|
<|open|>response<|sep|>text<|close|>response<|sep|>
|
||||||
|
<|open|>tools<|sep|>
|
||||||
|
<|open|>call tool="name" index="1"<|sep|>
|
||||||
|
<|open|>argument key="k" type="string"<|sep|>raw text<|close|>argument<|sep|>
|
||||||
|
<|close|>call<|sep|>
|
||||||
|
<|close|>tools<|sep|>
|
||||||
|
```
|
||||||
|
|
||||||
|
``type="string"`` argument values are raw text; other types are
|
||||||
|
JSON-decoded. Attribute values reverse the template's ``&``/``"``
|
||||||
|
escaping.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.bot_token = TOOLS_OPEN
|
||||||
|
self.eot_token = TOOLS_CLOSE
|
||||||
|
self._sent_normal_idx = 0
|
||||||
|
|
||||||
|
def has_tool_call(self, text: str) -> bool:
|
||||||
|
return self.bot_token in text
|
||||||
|
|
||||||
|
def supports_structural_tag(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def parses_required_natively(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def structure_info(self) -> _GetInfoFunc:
|
||||||
|
raise NotImplementedError(
|
||||||
|
"Kimi K3 uses its model-native structural tag implementation"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_auto_tool_call_structural_tag(
|
||||||
|
self,
|
||||||
|
tools: Union[List[Tool], None] = None,
|
||||||
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
|
) -> Optional[StructuralTag]:
|
||||||
|
return get_kimik3_auto_tool_call_structural_tag(
|
||||||
|
tools or [],
|
||||||
|
thinking_mode=thinking_mode,
|
||||||
|
parallel_tool_calls=parallel_tool_calls,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_structural_tag(
|
||||||
|
self,
|
||||||
|
tools: Union[List[Tool], None] = None,
|
||||||
|
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
||||||
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
|
) -> StructuralTag:
|
||||||
|
return get_kimik3_structural_tag(
|
||||||
|
tools=tools or [],
|
||||||
|
tool_choice=tool_choice,
|
||||||
|
thinking_mode=thinking_mode,
|
||||||
|
parallel_tool_calls=parallel_tool_calls,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _decode_call(self, attrs: str, body: str) -> dict | None:
|
||||||
|
call_attrs = _parse_attrs(attrs)
|
||||||
|
tool_name = call_attrs.get("tool", "")
|
||||||
|
if not tool_name:
|
||||||
|
return None
|
||||||
|
arguments = {}
|
||||||
|
for arg in _ARG_RE.finditer(body):
|
||||||
|
arg_attrs = _parse_attrs(arg["attrs"])
|
||||||
|
key = arg_attrs.get("key", "")
|
||||||
|
arg_type = arg_attrs.get("type", "string")
|
||||||
|
raw_value = arg["val"]
|
||||||
|
if arg_type == "string":
|
||||||
|
arguments[key] = raw_value
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
arguments[key] = json.loads(raw_value)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
arguments[key] = raw_value
|
||||||
|
return {
|
||||||
|
"name": tool_name,
|
||||||
|
"arguments": json.dumps(arguments, ensure_ascii=False),
|
||||||
|
}
|
||||||
|
|
||||||
|
def _parse_calls(self, section: str) -> List[dict]:
|
||||||
|
return [
|
||||||
|
call
|
||||||
|
for m in _CALL_RE.finditer(section)
|
||||||
|
if (call := self._decode_call(m["attrs"], m["body"])) is not None
|
||||||
|
]
|
||||||
|
|
||||||
|
def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult:
|
||||||
|
open_idx = text.find(self.bot_token)
|
||||||
|
if open_idx == -1:
|
||||||
|
return StreamingParseResult(normal_text=strip_response_wrappers(text))
|
||||||
|
# Computed outside the try so the error path can reuse it instead of
|
||||||
|
# falling back to raw text, which would ship the XTML tools markup to
|
||||||
|
# the client.
|
||||||
|
before = strip_response_wrappers(text[:open_idx])
|
||||||
|
try:
|
||||||
|
section_start = open_idx + len(self.bot_token)
|
||||||
|
close_idx = text.find(self.eot_token, section_start)
|
||||||
|
section = (
|
||||||
|
text[section_start:]
|
||||||
|
if close_idx == -1
|
||||||
|
else text[section_start:close_idx]
|
||||||
|
)
|
||||||
|
calls = [
|
||||||
|
ToolCallItem(
|
||||||
|
tool_index=i,
|
||||||
|
name=call["name"],
|
||||||
|
parameters=call["arguments"],
|
||||||
|
)
|
||||||
|
for i, call in enumerate(self._parse_calls(section))
|
||||||
|
]
|
||||||
|
return StreamingParseResult(normal_text=before, calls=calls)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Error in Kimi K3 detect_and_parse: %s", e, exc_info=True)
|
||||||
|
return StreamingParseResult(normal_text=before)
|
||||||
|
|
||||||
|
def parse_streaming_increment(
|
||||||
|
self, new_text: str, tools: List[Tool]
|
||||||
|
) -> StreamingParseResult:
|
||||||
|
self._buffer += new_text
|
||||||
|
try:
|
||||||
|
open_idx = self._buffer.find(self.bot_token)
|
||||||
|
if open_idx == -1:
|
||||||
|
return StreamingParseResult(normal_text=self._emit_normal_text())
|
||||||
|
|
||||||
|
normal_text = self._emit_normal_text(limit=open_idx)
|
||||||
|
section = self._buffer[open_idx + len(self.bot_token) :]
|
||||||
|
calls = []
|
||||||
|
parsed = self._parse_calls(section)
|
||||||
|
for call in parsed[self.current_tool_id + 1 :]:
|
||||||
|
self.current_tool_id += 1
|
||||||
|
while len(self.prev_tool_call_arr) <= self.current_tool_id:
|
||||||
|
self.prev_tool_call_arr.append({})
|
||||||
|
while len(self.streamed_args_for_tool) <= self.current_tool_id:
|
||||||
|
self.streamed_args_for_tool.append("")
|
||||||
|
self.prev_tool_call_arr[self.current_tool_id] = {
|
||||||
|
"name": call["name"],
|
||||||
|
"arguments": json.loads(call["arguments"]),
|
||||||
|
}
|
||||||
|
self.streamed_args_for_tool[self.current_tool_id] = call["arguments"]
|
||||||
|
calls.append(
|
||||||
|
ToolCallItem(
|
||||||
|
tool_index=self.current_tool_id,
|
||||||
|
name=call["name"],
|
||||||
|
parameters=call["arguments"],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return StreamingParseResult(normal_text=normal_text, calls=calls)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
"Error in Kimi K3 parse_streaming_increment: %s", e, exc_info=True
|
||||||
|
)
|
||||||
|
# _sent_normal_idx indexes into _buffer, so it must be reset with it;
|
||||||
|
# otherwise every later _emit_normal_text sees limit <= _sent_normal_idx
|
||||||
|
# and silently drops the rest of the response.
|
||||||
|
self._buffer = ""
|
||||||
|
self._sent_normal_idx = 0
|
||||||
|
return StreamingParseResult()
|
||||||
|
|
||||||
|
def _emit_normal_text(self, limit: int | None = None) -> str:
|
||||||
|
if limit is None:
|
||||||
|
holdback = partial_suffix_len(
|
||||||
|
self._buffer,
|
||||||
|
[self.bot_token, RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE],
|
||||||
|
)
|
||||||
|
limit = len(self._buffer) - holdback
|
||||||
|
if limit <= self._sent_normal_idx:
|
||||||
|
return ""
|
||||||
|
pending = self._buffer[self._sent_normal_idx : limit]
|
||||||
|
for marker in (RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE):
|
||||||
|
if marker in pending:
|
||||||
|
pending = pending.replace(marker, "")
|
||||||
|
self._sent_normal_idx = limit
|
||||||
|
return pending
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
from typing import List
|
||||||
|
|
||||||
|
THINK_OPEN = "<|open|>think<|sep|>"
|
||||||
|
THINK_CLOSE = "<|close|>think<|sep|>"
|
||||||
|
RESPONSE_OPEN = "<|open|>response<|sep|>"
|
||||||
|
RESPONSE_CLOSE = "<|close|>response<|sep|>"
|
||||||
|
TOOLS_OPEN = "<|open|>tools<|sep|>"
|
||||||
|
TOOLS_CLOSE = "<|close|>tools<|sep|>"
|
||||||
|
MESSAGE_CLOSE = "<|close|>message<|sep|>"
|
||||||
|
CALL_OPEN = "<|open|>call"
|
||||||
|
CALL_CLOSE = "<|close|>call<|sep|>"
|
||||||
|
ARGUMENT_CLOSE = "<|close|>argument<|sep|>"
|
||||||
|
|
||||||
|
# max_tokens can stop after an XTML control token or channel name, before <|sep|>.
|
||||||
|
_PARTIAL_MARKER_SUFFIXES = (
|
||||||
|
"<|open|>",
|
||||||
|
"<|close|>",
|
||||||
|
THINK_CLOSE.removesuffix("<|sep|>"),
|
||||||
|
RESPONSE_OPEN.removesuffix("<|sep|>"),
|
||||||
|
RESPONSE_CLOSE.removesuffix("<|sep|>"),
|
||||||
|
TOOLS_OPEN.removesuffix("<|sep|>"),
|
||||||
|
TOOLS_CLOSE.removesuffix("<|sep|>"),
|
||||||
|
MESSAGE_CLOSE.removesuffix("<|sep|>"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def partial_suffix_len(text: str, markers: List[str]) -> int:
|
||||||
|
best = 0
|
||||||
|
for marker in markers:
|
||||||
|
for length in range(min(len(marker) - 1, len(text)), best, -1):
|
||||||
|
if text.endswith(marker[:length]):
|
||||||
|
best = length
|
||||||
|
break
|
||||||
|
return best
|
||||||
|
|
||||||
|
|
||||||
|
def strip_partial_marker_suffix(text: str) -> str:
|
||||||
|
for suffix in _PARTIAL_MARKER_SUFFIXES:
|
||||||
|
if text.endswith(suffix):
|
||||||
|
return text[: -len(suffix)]
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def strip_response_wrappers(text: str) -> str:
|
||||||
|
open_idx = text.find(RESPONSE_OPEN)
|
||||||
|
if open_idx != -1:
|
||||||
|
close_idx = text.find(RESPONSE_CLOSE, open_idx + len(RESPONSE_OPEN))
|
||||||
|
if close_idx != -1:
|
||||||
|
text = text[open_idx + len(RESPONSE_OPEN) : close_idx]
|
||||||
|
else:
|
||||||
|
text = text[open_idx + len(RESPONSE_OPEN) :]
|
||||||
|
else:
|
||||||
|
text = text.replace(RESPONSE_CLOSE, "")
|
||||||
|
text = text.replace(MESSAGE_CLOSE, "")
|
||||||
|
return strip_partial_marker_suffix(text)
|
||||||
@@ -0,0 +1,598 @@
|
|||||||
|
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union
|
||||||
|
|
||||||
|
from xgrammar import StructuralTag
|
||||||
|
from xgrammar.structural_tag import (
|
||||||
|
AnyTextFormat,
|
||||||
|
AnyTokensFormat,
|
||||||
|
ConstStringFormat,
|
||||||
|
ExcludeTokenFormat,
|
||||||
|
Format,
|
||||||
|
JSONSchemaFormat,
|
||||||
|
OptionalFormat,
|
||||||
|
OrFormat,
|
||||||
|
RegexFormat,
|
||||||
|
SequenceFormat,
|
||||||
|
StarFormat,
|
||||||
|
TagFormat,
|
||||||
|
TagsWithSeparatorFormat,
|
||||||
|
TokenFormat,
|
||||||
|
TriggeredTagsFormat,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import Tool, ToolChoice
|
||||||
|
from sglang.srt.function_call.kimik3_format import (
|
||||||
|
ARGUMENT_CLOSE,
|
||||||
|
CALL_CLOSE,
|
||||||
|
CALL_OPEN,
|
||||||
|
THINK_CLOSE,
|
||||||
|
THINK_OPEN,
|
||||||
|
TOOLS_CLOSE,
|
||||||
|
TOOLS_OPEN,
|
||||||
|
)
|
||||||
|
|
||||||
|
_JSON_TYPES = (
|
||||||
|
"string",
|
||||||
|
"number",
|
||||||
|
"integer",
|
||||||
|
"boolean",
|
||||||
|
"array",
|
||||||
|
"object",
|
||||||
|
"null",
|
||||||
|
)
|
||||||
|
_CLOSE_TOKEN = "<|close|>"
|
||||||
|
_ARGUMENT_CLOSE_SUFFIX = ARGUMENT_CLOSE.removeprefix(_CLOSE_TOKEN)
|
||||||
|
_JSON_TO_XTML_TYPE = {
|
||||||
|
"string": "string",
|
||||||
|
"number": "number",
|
||||||
|
"integer": "number",
|
||||||
|
"boolean": "boolean",
|
||||||
|
"array": "array",
|
||||||
|
"object": "object",
|
||||||
|
"null": "null",
|
||||||
|
}
|
||||||
|
_STRING_KEYWORDS = {"format", "maxLength", "minLength", "pattern"}
|
||||||
|
_NUMBER_KEYWORDS = {
|
||||||
|
"exclusiveMaximum",
|
||||||
|
"exclusiveMinimum",
|
||||||
|
"maximum",
|
||||||
|
"minimum",
|
||||||
|
"multipleOf",
|
||||||
|
}
|
||||||
|
_ARRAY_KEYWORDS = {
|
||||||
|
"contains",
|
||||||
|
"items",
|
||||||
|
"maxContains",
|
||||||
|
"maxItems",
|
||||||
|
"minContains",
|
||||||
|
"minItems",
|
||||||
|
"prefixItems",
|
||||||
|
"uniqueItems",
|
||||||
|
}
|
||||||
|
_OBJECT_KEYWORDS = {
|
||||||
|
"additionalProperties",
|
||||||
|
"dependentRequired",
|
||||||
|
"dependentSchemas",
|
||||||
|
"maxProperties",
|
||||||
|
"minProperties",
|
||||||
|
"patternProperties",
|
||||||
|
"properties",
|
||||||
|
"propertyNames",
|
||||||
|
"required",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _escape_attr(value: str) -> str:
|
||||||
|
return value.replace("&", "&").replace('"', """)
|
||||||
|
|
||||||
|
|
||||||
|
def _json_type(value: Any) -> str:
|
||||||
|
if value is None:
|
||||||
|
return "null"
|
||||||
|
if isinstance(value, bool):
|
||||||
|
return "boolean"
|
||||||
|
if isinstance(value, int):
|
||||||
|
return "integer"
|
||||||
|
if isinstance(value, float):
|
||||||
|
return "number"
|
||||||
|
if isinstance(value, str):
|
||||||
|
return "string"
|
||||||
|
if isinstance(value, list):
|
||||||
|
return "array"
|
||||||
|
return "object"
|
||||||
|
|
||||||
|
|
||||||
|
def _matches_json_type(value: Any, json_type: str) -> bool:
|
||||||
|
value_type = _json_type(value)
|
||||||
|
return value_type == json_type or (
|
||||||
|
json_type == "number" and value_type == "integer"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_local_ref(
|
||||||
|
ref: str, root_schema: Dict[str, Any]
|
||||||
|
) -> Optional[Union[bool, Dict[str, Any]]]:
|
||||||
|
if not ref.startswith("#/"):
|
||||||
|
return None
|
||||||
|
value: Any = root_schema
|
||||||
|
for part in ref[2:].split("/"):
|
||||||
|
key = part.replace("~1", "/").replace("~0", "~")
|
||||||
|
if not isinstance(value, dict) or key not in value:
|
||||||
|
return None
|
||||||
|
value = value[key]
|
||||||
|
if isinstance(value, (bool, dict)):
|
||||||
|
return value
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _schema_types(
|
||||||
|
schema: Union[bool, Dict[str, Any]],
|
||||||
|
root_schema: Dict[str, Any],
|
||||||
|
seen_refs: Optional[Set[str]] = None,
|
||||||
|
) -> List[str]:
|
||||||
|
if schema is False:
|
||||||
|
return []
|
||||||
|
if schema is True:
|
||||||
|
return list(_JSON_TYPES)
|
||||||
|
|
||||||
|
ref = schema.get("$ref")
|
||||||
|
if isinstance(ref, str):
|
||||||
|
seen_refs = set() if seen_refs is None else set(seen_refs)
|
||||||
|
if ref not in seen_refs:
|
||||||
|
target = _resolve_local_ref(ref, root_schema)
|
||||||
|
if target is not None:
|
||||||
|
seen_refs.add(ref)
|
||||||
|
return _schema_types(target, root_schema, seen_refs)
|
||||||
|
|
||||||
|
schema_type = schema.get("type")
|
||||||
|
if isinstance(schema_type, str):
|
||||||
|
return [schema_type] if schema_type in _JSON_TYPES else list(_JSON_TYPES)
|
||||||
|
if isinstance(schema_type, list):
|
||||||
|
return [item for item in _JSON_TYPES if item in schema_type]
|
||||||
|
|
||||||
|
for keyword in ("anyOf", "oneOf"):
|
||||||
|
options = schema.get(keyword)
|
||||||
|
if isinstance(options, list):
|
||||||
|
option_types = {
|
||||||
|
item
|
||||||
|
for option in options
|
||||||
|
if isinstance(option, (bool, dict))
|
||||||
|
for item in _schema_types(option, root_schema, seen_refs)
|
||||||
|
}
|
||||||
|
return [item for item in _JSON_TYPES if item in option_types]
|
||||||
|
|
||||||
|
options = schema.get("allOf")
|
||||||
|
if isinstance(options, list):
|
||||||
|
type_sets = [
|
||||||
|
set(_schema_types(option, root_schema, seen_refs))
|
||||||
|
for option in options
|
||||||
|
if isinstance(option, (bool, dict))
|
||||||
|
]
|
||||||
|
constrained = [types for types in type_sets if types != set(_JSON_TYPES)]
|
||||||
|
if constrained:
|
||||||
|
result = constrained[0] | (
|
||||||
|
{"integer"} if "number" in constrained[0] else set()
|
||||||
|
)
|
||||||
|
for types in constrained[1:]:
|
||||||
|
result &= types | ({"integer"} if "number" in types else set())
|
||||||
|
# number survives the intersection only if every branch allows it,
|
||||||
|
# making the widened integer redundant rather than the other way.
|
||||||
|
if "number" in result and "integer" in result:
|
||||||
|
result.remove("integer")
|
||||||
|
return [item for item in _JSON_TYPES if item in result]
|
||||||
|
|
||||||
|
if "const" in schema:
|
||||||
|
return [_json_type(schema["const"])]
|
||||||
|
enum = schema.get("enum")
|
||||||
|
if isinstance(enum, list):
|
||||||
|
enum_types = {_json_type(value) for value in enum}
|
||||||
|
return [item for item in _JSON_TYPES if item in enum_types]
|
||||||
|
if _OBJECT_KEYWORDS.intersection(schema):
|
||||||
|
return ["object"]
|
||||||
|
if _ARRAY_KEYWORDS.intersection(schema):
|
||||||
|
return ["array"]
|
||||||
|
if _STRING_KEYWORDS.intersection(schema):
|
||||||
|
return ["string"]
|
||||||
|
if _NUMBER_KEYWORDS.intersection(schema):
|
||||||
|
return ["number"]
|
||||||
|
return list(_JSON_TYPES)
|
||||||
|
|
||||||
|
|
||||||
|
def _with_root_definitions(
|
||||||
|
schema: Union[bool, Dict[str, Any]], root_schema: Dict[str, Any]
|
||||||
|
) -> Union[bool, Dict[str, Any]]:
|
||||||
|
if not isinstance(schema, dict):
|
||||||
|
return schema
|
||||||
|
result = dict(schema)
|
||||||
|
for key in ("$defs", "definitions"):
|
||||||
|
if key in root_schema and key not in result:
|
||||||
|
result[key] = root_schema[key]
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _restrict_schema_type(
|
||||||
|
schema: Union[bool, Dict[str, Any]],
|
||||||
|
json_type: str,
|
||||||
|
root_schema: Dict[str, Any],
|
||||||
|
) -> Union[bool, Dict[str, Any]]:
|
||||||
|
if not isinstance(schema, dict):
|
||||||
|
return {"type": json_type} if schema else False
|
||||||
|
|
||||||
|
ref = schema.get("$ref")
|
||||||
|
if isinstance(ref, str):
|
||||||
|
target = _resolve_local_ref(ref, root_schema)
|
||||||
|
if target is not None:
|
||||||
|
return _with_root_definitions(
|
||||||
|
_restrict_schema_type(target, json_type, root_schema), root_schema
|
||||||
|
)
|
||||||
|
|
||||||
|
result = dict(schema)
|
||||||
|
schema_type = result.get("type")
|
||||||
|
if isinstance(schema_type, list):
|
||||||
|
if json_type not in schema_type:
|
||||||
|
return False
|
||||||
|
result["type"] = json_type
|
||||||
|
elif isinstance(schema_type, str):
|
||||||
|
if schema_type != json_type:
|
||||||
|
return False
|
||||||
|
else:
|
||||||
|
result["type"] = json_type
|
||||||
|
|
||||||
|
for keyword in ("anyOf", "oneOf"):
|
||||||
|
options = result.get(keyword)
|
||||||
|
if not isinstance(options, list):
|
||||||
|
continue
|
||||||
|
restricted = [
|
||||||
|
_restrict_schema_type(option, json_type, root_schema)
|
||||||
|
for option in options
|
||||||
|
if isinstance(option, (bool, dict))
|
||||||
|
and json_type in _schema_types(option, root_schema)
|
||||||
|
]
|
||||||
|
if not restricted:
|
||||||
|
return False
|
||||||
|
if len(restricted) == 1 and isinstance(restricted[0], dict):
|
||||||
|
result.pop(keyword)
|
||||||
|
result.update(restricted[0])
|
||||||
|
else:
|
||||||
|
result[keyword] = restricted
|
||||||
|
|
||||||
|
enum = result.get("enum")
|
||||||
|
if isinstance(enum, list):
|
||||||
|
result["enum"] = [
|
||||||
|
value for value in enum if _matches_json_type(value, json_type)
|
||||||
|
]
|
||||||
|
if not result["enum"]:
|
||||||
|
return False
|
||||||
|
if "const" in result and not _matches_json_type(result["const"], json_type):
|
||||||
|
return False
|
||||||
|
return _with_root_definitions(result, root_schema)
|
||||||
|
|
||||||
|
|
||||||
|
def _value_format(
|
||||||
|
schema: Union[bool, Dict[str, Any]],
|
||||||
|
json_type: str,
|
||||||
|
loose_string: bool = False,
|
||||||
|
) -> Format:
|
||||||
|
if loose_string and json_type == "string":
|
||||||
|
return AnyTextFormat()
|
||||||
|
return JSONSchemaFormat(
|
||||||
|
json_schema=schema,
|
||||||
|
style="qwen_xml" if json_type == "string" else "json",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _argument_value_variants(
|
||||||
|
schema: Union[bool, Dict[str, Any]],
|
||||||
|
root_schema: Dict[str, Any],
|
||||||
|
loose_strings: bool = False,
|
||||||
|
) -> List[Tuple[str, Format]]:
|
||||||
|
return [
|
||||||
|
(
|
||||||
|
json_type,
|
||||||
|
_value_format(restricted, json_type, loose_string=loose_strings),
|
||||||
|
)
|
||||||
|
for json_type in _schema_types(schema, root_schema)
|
||||||
|
if (restricted := _restrict_schema_type(schema, json_type, root_schema))
|
||||||
|
is not False
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _known_argument_format(
|
||||||
|
key: str,
|
||||||
|
schema: Union[bool, Dict[str, Any]],
|
||||||
|
root_schema: Dict[str, Any],
|
||||||
|
) -> Optional[Format]:
|
||||||
|
escaped_key = _escape_attr(key)
|
||||||
|
variants = [
|
||||||
|
TagFormat(
|
||||||
|
begin=(
|
||||||
|
f'<|open|>argument key="{escaped_key}" '
|
||||||
|
f'type="{_JSON_TO_XTML_TYPE[json_type]}"<|sep|>'
|
||||||
|
),
|
||||||
|
content=value_format,
|
||||||
|
end=ARGUMENT_CLOSE,
|
||||||
|
)
|
||||||
|
for json_type, value_format in _argument_value_variants(schema, root_schema)
|
||||||
|
]
|
||||||
|
if not variants:
|
||||||
|
return None
|
||||||
|
if len(variants) == 1:
|
||||||
|
return variants[0]
|
||||||
|
return OrFormat(elements=variants)
|
||||||
|
|
||||||
|
|
||||||
|
def _dynamic_argument_format(
|
||||||
|
schema: Union[bool, Dict[str, Any]],
|
||||||
|
root_schema: Dict[str, Any],
|
||||||
|
loose_strings: bool = False,
|
||||||
|
) -> Format:
|
||||||
|
variants = [
|
||||||
|
SequenceFormat(
|
||||||
|
elements=[
|
||||||
|
RegexFormat(pattern=r'[^"& \t\r\n\f\v=<>]+'),
|
||||||
|
ConstStringFormat(
|
||||||
|
value=(f'" type="{_JSON_TO_XTML_TYPE[json_type]}"<|sep|>')
|
||||||
|
),
|
||||||
|
value_format,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
for json_type, value_format in _argument_value_variants(
|
||||||
|
schema, root_schema, loose_strings=loose_strings
|
||||||
|
)
|
||||||
|
]
|
||||||
|
if not variants:
|
||||||
|
raise ValueError("Kimi K3 additional parameter schema accepts no values")
|
||||||
|
content = variants[0] if len(variants) == 1 else OrFormat(elements=variants)
|
||||||
|
return TagFormat(
|
||||||
|
begin='<|open|>argument key="',
|
||||||
|
content=content,
|
||||||
|
end=ARGUMENT_CLOSE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _strict_arguments_format(parameters: Dict[str, Any]) -> Format:
|
||||||
|
properties = parameters.get("properties", {})
|
||||||
|
if not isinstance(properties, dict):
|
||||||
|
raise ValueError("Kimi K3 tool parameters 'properties' must be an object")
|
||||||
|
required = parameters.get("required", [])
|
||||||
|
if not isinstance(required, list) or not all(
|
||||||
|
isinstance(item, str) for item in required
|
||||||
|
):
|
||||||
|
raise ValueError("Kimi K3 tool parameters 'required' must be a string list")
|
||||||
|
|
||||||
|
required_set = set(required)
|
||||||
|
missing = required_set.difference(properties)
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f"Kimi K3 required parameters are missing schemas: {sorted(missing)!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
elements: List[Format] = []
|
||||||
|
for key, schema in properties.items():
|
||||||
|
if not isinstance(key, str) or not isinstance(schema, (bool, dict)):
|
||||||
|
raise ValueError("Kimi K3 tool property schemas must be JSON schemas")
|
||||||
|
argument = _known_argument_format(key, schema, parameters)
|
||||||
|
if argument is None:
|
||||||
|
if key in required_set:
|
||||||
|
raise ValueError(
|
||||||
|
f"Kimi K3 required parameter {key!r} accepts no values"
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
elements.append(
|
||||||
|
argument if key in required_set else OptionalFormat(content=argument)
|
||||||
|
)
|
||||||
|
|
||||||
|
additional = parameters.get("additionalProperties", True)
|
||||||
|
if additional is True:
|
||||||
|
elements.append(StarFormat(content=_dynamic_argument_format(True, parameters)))
|
||||||
|
elif isinstance(additional, dict):
|
||||||
|
elements.append(
|
||||||
|
StarFormat(content=_dynamic_argument_format(additional, parameters))
|
||||||
|
)
|
||||||
|
elif additional is not False:
|
||||||
|
raise ValueError(
|
||||||
|
"Kimi K3 tool parameters 'additionalProperties' must be a schema"
|
||||||
|
)
|
||||||
|
if not elements:
|
||||||
|
return ConstStringFormat(value="")
|
||||||
|
return SequenceFormat(elements=elements)
|
||||||
|
|
||||||
|
|
||||||
|
def _tool_arguments_format(tool: Tool) -> Format:
|
||||||
|
parameters = tool.function.parameters
|
||||||
|
if not tool.function.strict:
|
||||||
|
root_schema = parameters if isinstance(parameters, dict) else {}
|
||||||
|
return StarFormat(
|
||||||
|
content=_dynamic_argument_format(True, root_schema, loose_strings=True)
|
||||||
|
)
|
||||||
|
if parameters is None:
|
||||||
|
# Server-side strict levels mark tools without parameters strict too;
|
||||||
|
# they take no arguments rather than failing the whole constraint.
|
||||||
|
return ConstStringFormat(value="")
|
||||||
|
if not isinstance(parameters, dict):
|
||||||
|
raise ValueError(
|
||||||
|
f"Kimi K3 strict tool {tool.function.name!r} must define parameters"
|
||||||
|
)
|
||||||
|
schema_types = _schema_types(parameters, parameters)
|
||||||
|
if "object" not in schema_types:
|
||||||
|
raise ValueError(
|
||||||
|
f"Kimi K3 tool {tool.function.name!r} parameters must be an object schema"
|
||||||
|
)
|
||||||
|
return _strict_arguments_format(parameters)
|
||||||
|
|
||||||
|
|
||||||
|
def _tool_call_tag(tool: Tool, arguments_format: Optional[Format] = None) -> TagFormat:
|
||||||
|
name = _escape_attr(tool.function.name)
|
||||||
|
if arguments_format is None:
|
||||||
|
arguments_format = _tool_arguments_format(tool)
|
||||||
|
return TagFormat(
|
||||||
|
begin=f'{CALL_OPEN} tool="{name}" index="',
|
||||||
|
content=SequenceFormat(
|
||||||
|
elements=[
|
||||||
|
RegexFormat(pattern=r"[1-9][0-9]*"),
|
||||||
|
ConstStringFormat(value='"<|sep|>'),
|
||||||
|
arguments_format,
|
||||||
|
]
|
||||||
|
),
|
||||||
|
end=CALL_CLOSE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _tool_calls_tag(call_tags: List[TagFormat], parallel_tool_calls: bool) -> TagFormat:
|
||||||
|
if parallel_tool_calls:
|
||||||
|
content: Format = TagsWithSeparatorFormat(
|
||||||
|
tags=call_tags, separator="", at_least_one=True
|
||||||
|
)
|
||||||
|
elif len(call_tags) == 1:
|
||||||
|
content = call_tags[0]
|
||||||
|
else:
|
||||||
|
content = OrFormat(elements=call_tags)
|
||||||
|
return TagFormat(begin=TOOLS_OPEN, content=content, end=TOOLS_CLOSE)
|
||||||
|
|
||||||
|
|
||||||
|
def _auto_suffix(
|
||||||
|
tools_tag: TagFormat, parallel_tool_calls: bool
|
||||||
|
) -> TriggeredTagsFormat:
|
||||||
|
# A retriggered second tools section would evade the single-call limit.
|
||||||
|
return TriggeredTagsFormat(
|
||||||
|
triggers=[TOOLS_OPEN],
|
||||||
|
tags=[tools_tag],
|
||||||
|
excludes=[THINK_OPEN, THINK_CLOSE, CALL_OPEN],
|
||||||
|
stop_after_first=not parallel_tool_calls,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _with_reasoning(suffix: Format, thinking_mode: bool) -> Format:
|
||||||
|
if not thinking_mode:
|
||||||
|
return suffix
|
||||||
|
return SequenceFormat(
|
||||||
|
elements=[
|
||||||
|
TagFormat(begin="", content=AnyTextFormat(), end=THINK_CLOSE),
|
||||||
|
suffix,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _single_xtml_type(
|
||||||
|
schema: Union[bool, Dict[str, Any]], root_schema: Dict[str, Any]
|
||||||
|
) -> Optional[str]:
|
||||||
|
schema_types = _schema_types(schema, root_schema)
|
||||||
|
if len(schema_types) != 1:
|
||||||
|
return None
|
||||||
|
return _JSON_TO_XTML_TYPE[schema_types[0]]
|
||||||
|
|
||||||
|
|
||||||
|
def _nonempty_argument_format(key: str, xtml_type: str) -> Format:
|
||||||
|
# A token-based end keeps the first close token out of both content formats.
|
||||||
|
argument = TagFormat(
|
||||||
|
begin=(
|
||||||
|
f'<|open|>argument key="{_escape_attr(key)}" ' f'type="{xtml_type}"<|sep|>'
|
||||||
|
),
|
||||||
|
content=SequenceFormat(elements=[ExcludeTokenFormat(), AnyTokensFormat()]),
|
||||||
|
end=TokenFormat(token=_CLOSE_TOKEN),
|
||||||
|
)
|
||||||
|
return SequenceFormat(
|
||||||
|
elements=[
|
||||||
|
argument,
|
||||||
|
ConstStringFormat(value=_ARGUMENT_CLOSE_SUFFIX),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _auto_tool_arguments_format(tool: Tool) -> Format:
|
||||||
|
parameters = tool.function.parameters
|
||||||
|
if not isinstance(parameters, dict):
|
||||||
|
return AnyTextFormat()
|
||||||
|
properties = parameters.get("properties", {})
|
||||||
|
required = parameters.get("required", [])
|
||||||
|
if not isinstance(properties, dict) or not isinstance(required, list):
|
||||||
|
return AnyTextFormat()
|
||||||
|
|
||||||
|
required_tags: List[Format] = []
|
||||||
|
for key in required:
|
||||||
|
schema = properties.get(key)
|
||||||
|
if not isinstance(key, str) or not isinstance(schema, (bool, dict)):
|
||||||
|
return AnyTextFormat()
|
||||||
|
xtml_type = _single_xtml_type(schema, parameters)
|
||||||
|
if xtml_type is None:
|
||||||
|
return AnyTextFormat()
|
||||||
|
required_tags.append(_nonempty_argument_format(key, xtml_type))
|
||||||
|
|
||||||
|
if required_tags:
|
||||||
|
elements = required_tags
|
||||||
|
else:
|
||||||
|
alternatives = [
|
||||||
|
_nonempty_argument_format(key, xtml_type)
|
||||||
|
for key, schema in properties.items()
|
||||||
|
if isinstance(key, str)
|
||||||
|
and isinstance(schema, (bool, dict))
|
||||||
|
and (xtml_type := _single_xtml_type(schema, parameters)) is not None
|
||||||
|
]
|
||||||
|
if not alternatives:
|
||||||
|
return AnyTextFormat()
|
||||||
|
elements = [
|
||||||
|
(
|
||||||
|
alternatives[0]
|
||||||
|
if len(alternatives) == 1
|
||||||
|
else OrFormat(elements=alternatives)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
return SequenceFormat(
|
||||||
|
elements=[*elements, AnyTextFormat(excludes=[CALL_OPEN, CALL_CLOSE])]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_kimik3_auto_tool_call_structural_tag(
|
||||||
|
tools: List[Tool],
|
||||||
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
|
) -> Optional[StructuralTag]:
|
||||||
|
if not tools:
|
||||||
|
return None
|
||||||
|
|
||||||
|
call_tags = [
|
||||||
|
_tool_call_tag(tool, _auto_tool_arguments_format(tool)) for tool in tools
|
||||||
|
]
|
||||||
|
suffix = _auto_suffix(
|
||||||
|
_tool_calls_tag(call_tags, parallel_tool_calls), parallel_tool_calls
|
||||||
|
)
|
||||||
|
return StructuralTag(format=_with_reasoning(suffix, thinking_mode))
|
||||||
|
|
||||||
|
|
||||||
|
def _select_tools(
|
||||||
|
tools: List[Tool],
|
||||||
|
tool_choice: Union[ToolChoice, Literal["auto", "required"]],
|
||||||
|
) -> Tuple[List[Tool], bool]:
|
||||||
|
if not isinstance(tool_choice, ToolChoice):
|
||||||
|
return tools, tool_choice == "required"
|
||||||
|
name = tool_choice.function.name
|
||||||
|
selected = [tool for tool in tools if tool.function.name == name]
|
||||||
|
if not selected:
|
||||||
|
raise ValueError(f"Kimi K3 tool choice {name!r} is not in the tools list")
|
||||||
|
return selected, True
|
||||||
|
|
||||||
|
|
||||||
|
def get_kimik3_structural_tag(
|
||||||
|
tools: List[Tool],
|
||||||
|
tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto",
|
||||||
|
thinking_mode: bool = False,
|
||||||
|
parallel_tool_calls: bool = True,
|
||||||
|
) -> StructuralTag:
|
||||||
|
selected_tools, at_least_one = _select_tools(tools, tool_choice)
|
||||||
|
if not selected_tools:
|
||||||
|
raise ValueError("Kimi K3 structural tags require at least one tool")
|
||||||
|
|
||||||
|
call_tags = [_tool_call_tag(tool) for tool in selected_tools]
|
||||||
|
tools_tag = _tool_calls_tag(call_tags, parallel_tool_calls)
|
||||||
|
if at_least_one:
|
||||||
|
suffix: Format = SequenceFormat(
|
||||||
|
elements=[
|
||||||
|
AnyTextFormat(
|
||||||
|
excludes=[TOOLS_OPEN, THINK_OPEN, THINK_CLOSE, CALL_OPEN]
|
||||||
|
),
|
||||||
|
tools_tag,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
suffix = _auto_suffix(tools_tag, parallel_tool_calls)
|
||||||
|
return StructuralTag(format=_with_reasoning(suffix, thinking_mode))
|
||||||
@@ -115,6 +115,7 @@ from sglang.srt.utils.cuda_ipc_transport_utils import (
|
|||||||
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
||||||
CudaIpcTensorTransportProxy,
|
CudaIpcTensorTransportProxy,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.token_sequence_matcher import TokenSequenceMatcher
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from typing import Any, Dict
|
from typing import Any, Dict
|
||||||
@@ -816,6 +817,8 @@ class Req(ReqDllmMixin):
|
|||||||
# State indicating whether the reasoning phase has finished (only meaningful when require_reasoning is True)
|
# State indicating whether the reasoning phase has finished (only meaningful when require_reasoning is True)
|
||||||
self._is_reasoning_over = False
|
self._is_reasoning_over = False
|
||||||
self.reasoning_tokens = 0
|
self.reasoning_tokens = 0
|
||||||
|
self._think_end_matcher: Optional[TokenSequenceMatcher] = None
|
||||||
|
self._think_end_match_len = 0
|
||||||
|
|
||||||
# Sampling info
|
# Sampling info
|
||||||
if isinstance(sampling_params.custom_params, dict):
|
if isinstance(sampling_params.custom_params, dict):
|
||||||
@@ -1690,18 +1693,25 @@ class Req(ReqDllmMixin):
|
|||||||
error_msg, HTTPStatus.BAD_REQUEST, "BadRequestError"
|
error_msg, HTTPStatus.BAD_REQUEST, "BadRequestError"
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_reasoning_tokens(self, token_id, think_end_id):
|
def update_reasoning_tokens(self, token_id, think_end_ids):
|
||||||
if self._is_reasoning_over:
|
if self._is_reasoning_over:
|
||||||
return
|
return
|
||||||
|
|
||||||
if not isinstance(token_id, list):
|
if not isinstance(token_id, list):
|
||||||
token_id = [token_id]
|
token_id = [token_id]
|
||||||
|
|
||||||
try:
|
if self._think_end_matcher is None:
|
||||||
end_pos = token_id.index(think_end_id)
|
self._think_end_matcher = TokenSequenceMatcher(think_end_ids)
|
||||||
self.reasoning_tokens += end_pos + 1
|
|
||||||
|
matched = self._think_end_match_len
|
||||||
|
for position, token in enumerate(token_id):
|
||||||
|
matched = self._think_end_matcher.advance(matched, token)
|
||||||
|
if matched == len(self._think_end_matcher):
|
||||||
|
self.reasoning_tokens += position + 1
|
||||||
self._is_reasoning_over = True
|
self._is_reasoning_over = True
|
||||||
except ValueError:
|
return
|
||||||
|
|
||||||
|
self._think_end_match_len = matched
|
||||||
self.reasoning_tokens += len(token_id)
|
self.reasoning_tokens += len(token_id)
|
||||||
|
|
||||||
def __repr__(self):
|
def __repr__(self):
|
||||||
|
|||||||
@@ -789,16 +789,23 @@ class Scheduler(
|
|||||||
"M-RoPE fallback will not be available."
|
"M-RoPE fallback will not be available."
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set reasoning_parser and think_end_id if --reasoning_parser is enabled
|
|
||||||
if get_serving().reasoning_parser and self.tokenizer:
|
if get_serving().reasoning_parser and self.tokenizer:
|
||||||
reasoning_parser = ReasoningParser(
|
reasoning_parser = ReasoningParser(
|
||||||
model_type=get_serving().reasoning_parser,
|
model_type=get_serving().reasoning_parser,
|
||||||
stream_reasoning=False,
|
stream_reasoning=False,
|
||||||
tokenizer=self.tokenizer,
|
tokenizer=self.tokenizer,
|
||||||
)
|
)
|
||||||
self.model_config.think_end_id = self.tokenizer.encode(
|
think_end_ids = self.tokenizer.encode(
|
||||||
reasoning_parser.detector.think_end_token, add_special_tokens=False
|
reasoning_parser.detector.think_end_token, add_special_tokens=False
|
||||||
)[0]
|
)
|
||||||
|
if think_end_ids:
|
||||||
|
self.model_config.think_end_ids = think_end_ids
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Reasoning parser think_end_token %r could not be encoded; "
|
||||||
|
"grammar-gated reasoning is disabled.",
|
||||||
|
reasoning_parser.detector.think_end_token,
|
||||||
|
)
|
||||||
|
|
||||||
def init_mamba_backend(self) -> None:
|
def init_mamba_backend(self) -> None:
|
||||||
initialize_mamba_selective_state_update_backend(self.server_args)
|
initialize_mamba_selective_state_update_backend(self.server_args)
|
||||||
|
|||||||
@@ -993,9 +993,9 @@ class SchedulerBatchResultProcessor:
|
|||||||
req: Req,
|
req: Req,
|
||||||
next_token_id: Union[int, List[int]],
|
next_token_id: Union[int, List[int]],
|
||||||
):
|
):
|
||||||
think_end_id = self.model_config.think_end_id
|
think_end_ids = self.model_config.think_end_ids
|
||||||
if req.require_reasoning and think_end_id is not None:
|
if req.require_reasoning and think_end_ids:
|
||||||
req.update_reasoning_tokens(next_token_id, think_end_id)
|
req.update_reasoning_tokens(next_token_id, think_end_ids)
|
||||||
|
|
||||||
def _mamba_prefix_cache_update(
|
def _mamba_prefix_cache_update(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -12,6 +12,17 @@ from sglang.srt.entrypoints.openai.encoding_dsv4 import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.entrypoints.openai.protocol import ChatCompletionRequest
|
from sglang.srt.entrypoints.openai.protocol import ChatCompletionRequest
|
||||||
from sglang.srt.function_call.hunyuan_detector import resolve_hunyuan_tokens
|
from sglang.srt.function_call.hunyuan_detector import resolve_hunyuan_tokens
|
||||||
|
from sglang.srt.function_call.kimik3_format import (
|
||||||
|
MESSAGE_CLOSE,
|
||||||
|
RESPONSE_CLOSE,
|
||||||
|
RESPONSE_OPEN,
|
||||||
|
THINK_CLOSE,
|
||||||
|
THINK_OPEN,
|
||||||
|
TOOLS_OPEN,
|
||||||
|
partial_suffix_len,
|
||||||
|
strip_partial_marker_suffix,
|
||||||
|
strip_response_wrappers,
|
||||||
|
)
|
||||||
from sglang.srt.parser.harmony_parser import HarmonyParser
|
from sglang.srt.parser.harmony_parser import HarmonyParser
|
||||||
from sglang.srt.parser.inkling_tokenizer import (
|
from sglang.srt.parser.inkling_tokenizer import (
|
||||||
CONTENT_INVOKE_TOOL_JSON,
|
CONTENT_INVOKE_TOOL_JSON,
|
||||||
@@ -419,6 +430,182 @@ class KimiK2Detector(BaseReasoningFormatDetector):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class KimiK3Detector(BaseReasoningFormatDetector):
|
||||||
|
"""Detector for the Kimi K3 XTML think channel.
|
||||||
|
|
||||||
|
K3 wraps reasoning as ``<|open|>think<|sep|>...<|close|>think<|sep|>``
|
||||||
|
where each marker is a multi-token special sequence, so partial markers
|
||||||
|
can straddle streaming chunks and must be held back. In thinking mode
|
||||||
|
the serving layer may feed the open marker as the generation prefix, so
|
||||||
|
output can begin inside the think channel with no open marker
|
||||||
|
(``force_reasoning=True`` covers this).
|
||||||
|
|
||||||
|
Post-reasoning content is unwrapped from the XTML ``response`` /
|
||||||
|
``message`` markers; a ``tools`` channel is passed through raw for the
|
||||||
|
kimi_k3 tool-call detector.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
stream_reasoning: bool = True,
|
||||||
|
force_reasoning: bool = True,
|
||||||
|
continue_final_message: bool = False,
|
||||||
|
previous_content: str = "",
|
||||||
|
):
|
||||||
|
# strict-thinking flattens these to single token ids, so the full marker
|
||||||
|
# "<|open|>response<|sep|>" is inexpressible. The bare name works: it
|
||||||
|
# follows <|open|> unspaced, so it tokenizes to the no-space variant, not
|
||||||
|
# the " response"/" message" tokens prose uses -- at the cost of not being
|
||||||
|
# able to start those words unspaced mid-reasoning. tools is left out on
|
||||||
|
# purpose: the model may jump from think straight into that channel.
|
||||||
|
think_excluded_tokens = [
|
||||||
|
"response",
|
||||||
|
"message",
|
||||||
|
"<|end_of_msg|>",
|
||||||
|
"[EOS]",
|
||||||
|
"[EOT]",
|
||||||
|
]
|
||||||
|
super().__init__(
|
||||||
|
THINK_OPEN,
|
||||||
|
THINK_CLOSE,
|
||||||
|
think_excluded_tokens=think_excluded_tokens,
|
||||||
|
force_reasoning=force_reasoning,
|
||||||
|
stream_reasoning=stream_reasoning,
|
||||||
|
tool_start_token=TOOLS_OPEN,
|
||||||
|
continue_final_message=continue_final_message,
|
||||||
|
previous_content=previous_content,
|
||||||
|
reasoning_default="thinking",
|
||||||
|
)
|
||||||
|
self._reasoning_done = False
|
||||||
|
self._tools_passthrough = False
|
||||||
|
|
||||||
|
def _clean_content(self, text: str) -> str:
|
||||||
|
tools_idx = text.find(TOOLS_OPEN)
|
||||||
|
if tools_idx != -1:
|
||||||
|
return strip_response_wrappers(text[:tools_idx]) + text[tools_idx:]
|
||||||
|
return strip_response_wrappers(text)
|
||||||
|
|
||||||
|
def _next_channel_idx(self, text: str, start: int = 0) -> int:
|
||||||
|
found = [
|
||||||
|
idx
|
||||||
|
for token in (RESPONSE_OPEN, self.tool_start_token)
|
||||||
|
if (idx := text.find(token, start)) != -1
|
||||||
|
]
|
||||||
|
return min(found) if found else -1
|
||||||
|
|
||||||
|
def detect_and_parse(self, text: str) -> StreamingParseResult:
|
||||||
|
in_reasoning = self._in_reasoning or self.think_start_token in text
|
||||||
|
if not in_reasoning and self.think_end_token not in text:
|
||||||
|
return StreamingParseResult(normal_text=self._clean_content(text))
|
||||||
|
|
||||||
|
open_idx = text.find(self.think_start_token)
|
||||||
|
start = open_idx + len(self.think_start_token) if open_idx != -1 else 0
|
||||||
|
close_idx = text.find(self.think_end_token, start)
|
||||||
|
if close_idx == -1:
|
||||||
|
channel_idx = self._next_channel_idx(text, start)
|
||||||
|
if channel_idx != -1:
|
||||||
|
return StreamingParseResult(
|
||||||
|
reasoning_text=strip_partial_marker_suffix(text[start:channel_idx]),
|
||||||
|
normal_text=self._clean_content(text[channel_idx:]),
|
||||||
|
)
|
||||||
|
return StreamingParseResult(
|
||||||
|
reasoning_text=strip_partial_marker_suffix(text[start:])
|
||||||
|
)
|
||||||
|
|
||||||
|
reasoning_text = text[start:close_idx]
|
||||||
|
rest = text[close_idx + len(self.think_end_token) :]
|
||||||
|
return StreamingParseResult(
|
||||||
|
reasoning_text=reasoning_text, normal_text=self._clean_content(rest)
|
||||||
|
)
|
||||||
|
|
||||||
|
def parse_streaming_increment(self, new_text: str) -> StreamingParseResult:
|
||||||
|
self._buffer += new_text
|
||||||
|
|
||||||
|
if not self._in_reasoning and not self._reasoning_done:
|
||||||
|
open_idx = self._buffer.find(self.think_start_token)
|
||||||
|
if open_idx != -1:
|
||||||
|
self._buffer = self._buffer[open_idx + len(self.think_start_token) :]
|
||||||
|
self._in_reasoning = True
|
||||||
|
self.stripped_think_start = True
|
||||||
|
elif self.think_start_token.startswith(self._buffer):
|
||||||
|
return StreamingParseResult()
|
||||||
|
else:
|
||||||
|
self._reasoning_done = True
|
||||||
|
|
||||||
|
if self._in_reasoning:
|
||||||
|
buf = self._buffer
|
||||||
|
if not self.stripped_think_start:
|
||||||
|
open_idx = buf.find(self.think_start_token)
|
||||||
|
if open_idx != -1:
|
||||||
|
buf = buf[open_idx + len(self.think_start_token) :]
|
||||||
|
self._buffer = buf
|
||||||
|
self.stripped_think_start = True
|
||||||
|
|
||||||
|
close_idx = buf.find(self.think_end_token)
|
||||||
|
if close_idx != -1:
|
||||||
|
reasoning_text = buf[:close_idx]
|
||||||
|
self._buffer = buf[close_idx + len(self.think_end_token) :]
|
||||||
|
self._in_reasoning = False
|
||||||
|
self._reasoning_done = True
|
||||||
|
return StreamingParseResult(
|
||||||
|
reasoning_text=reasoning_text or None,
|
||||||
|
normal_text=self._drain_content() or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
channel_idx = self._next_channel_idx(buf)
|
||||||
|
if channel_idx != -1:
|
||||||
|
reasoning_text = strip_partial_marker_suffix(buf[:channel_idx])
|
||||||
|
self._buffer = buf[channel_idx:]
|
||||||
|
self._in_reasoning = False
|
||||||
|
self._reasoning_done = True
|
||||||
|
self._tools_passthrough = buf.startswith(
|
||||||
|
self.tool_start_token, channel_idx
|
||||||
|
)
|
||||||
|
return StreamingParseResult(
|
||||||
|
reasoning_text=reasoning_text or None,
|
||||||
|
normal_text=self._drain_content() or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not self.stream_reasoning:
|
||||||
|
return StreamingParseResult()
|
||||||
|
markers = [self.think_end_token, self.tool_start_token, RESPONSE_OPEN]
|
||||||
|
if not self.stripped_think_start:
|
||||||
|
markers.append(self.think_start_token)
|
||||||
|
holdback = partial_suffix_len(buf, markers)
|
||||||
|
emit = buf[: len(buf) - holdback] if holdback else buf
|
||||||
|
emit = strip_partial_marker_suffix(emit)
|
||||||
|
self._buffer = buf[len(emit) :]
|
||||||
|
return StreamingParseResult(reasoning_text=emit)
|
||||||
|
|
||||||
|
return StreamingParseResult(normal_text=self._drain_content())
|
||||||
|
|
||||||
|
def _drain_content(self) -> str:
|
||||||
|
buf = self._buffer
|
||||||
|
if not buf:
|
||||||
|
return ""
|
||||||
|
if self._tools_passthrough:
|
||||||
|
self._buffer = ""
|
||||||
|
return buf
|
||||||
|
|
||||||
|
tools_idx = buf.find(TOOLS_OPEN)
|
||||||
|
if tools_idx != -1:
|
||||||
|
head = buf[:tools_idx]
|
||||||
|
for marker in (RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE):
|
||||||
|
head = head.replace(marker, "")
|
||||||
|
self._tools_passthrough = True
|
||||||
|
self._buffer = ""
|
||||||
|
return head + buf[tools_idx:]
|
||||||
|
|
||||||
|
holdback = partial_suffix_len(
|
||||||
|
buf, [RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE, TOOLS_OPEN]
|
||||||
|
)
|
||||||
|
emit = buf[: len(buf) - holdback] if holdback else buf
|
||||||
|
self._buffer = buf[len(emit) :]
|
||||||
|
for marker in (RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE):
|
||||||
|
emit = emit.replace(marker, "")
|
||||||
|
return emit
|
||||||
|
|
||||||
|
|
||||||
class Glm45Detector(BaseReasoningFormatDetector):
|
class Glm45Detector(BaseReasoningFormatDetector):
|
||||||
"""
|
"""
|
||||||
Detector for GLM-4.5 models.
|
Detector for GLM-4.5 models.
|
||||||
@@ -1448,6 +1635,7 @@ class ReasoningParser:
|
|||||||
"gpt-oss": GptOssDetector,
|
"gpt-oss": GptOssDetector,
|
||||||
"kimi": KimiDetector,
|
"kimi": KimiDetector,
|
||||||
"kimi_k2": KimiK2Detector,
|
"kimi_k2": KimiK2Detector,
|
||||||
|
"kimi_k3": KimiK3Detector,
|
||||||
"mimo": _MimoDetector,
|
"mimo": _MimoDetector,
|
||||||
"poolside_v1": _PoolsideV1Detector,
|
"poolside_v1": _PoolsideV1Detector,
|
||||||
"qwen3": Qwen3Detector,
|
"qwen3": Qwen3Detector,
|
||||||
|
|||||||
@@ -677,8 +677,11 @@ def _resolve_architecture_auto_parsers(server_args) -> None:
|
|||||||
)
|
)
|
||||||
architectures = getattr(config, "architectures", None) or []
|
architectures = getattr(config, "architectures", None) or []
|
||||||
arch = architectures[0] if architectures else ""
|
arch = architectures[0] if architectures else ""
|
||||||
|
model_type = getattr(config, "model_type", "")
|
||||||
|
|
||||||
if "DeepseekV4" in arch:
|
if "KimiK3" in arch or model_type == "kimi_k3":
|
||||||
|
reasoning_parser, tool_call_parser = "kimi_k3", "kimi_k3"
|
||||||
|
elif "DeepseekV4" in arch:
|
||||||
reasoning_parser, tool_call_parser = "deepseek-v4", "deepseekv4"
|
reasoning_parser, tool_call_parser = "deepseek-v4", "deepseekv4"
|
||||||
elif "DeepseekV3" in arch:
|
elif "DeepseekV3" in arch:
|
||||||
reasoning_parser, tool_call_parser = "deepseek-v3", "deepseekv32"
|
reasoning_parser, tool_call_parser = "deepseek-v3", "deepseekv32"
|
||||||
|
|||||||
@@ -0,0 +1,44 @@
|
|||||||
|
# Copyright 2023-2024 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
from typing import Sequence
|
||||||
|
|
||||||
|
|
||||||
|
class TokenSequenceMatcher:
|
||||||
|
def __init__(self, pattern: Sequence[int]):
|
||||||
|
if not pattern:
|
||||||
|
raise ValueError("pattern must contain at least one token")
|
||||||
|
self.pattern = tuple(pattern)
|
||||||
|
self.prefix_lengths = self._build_prefix_lengths()
|
||||||
|
|
||||||
|
def _build_prefix_lengths(self) -> tuple[int, ...]:
|
||||||
|
prefix_lengths = [0] * len(self.pattern)
|
||||||
|
matched = 0
|
||||||
|
for index in range(1, len(self.pattern)):
|
||||||
|
while matched > 0 and self.pattern[index] != self.pattern[matched]:
|
||||||
|
matched = prefix_lengths[matched - 1]
|
||||||
|
if self.pattern[index] == self.pattern[matched]:
|
||||||
|
matched += 1
|
||||||
|
prefix_lengths[index] = matched
|
||||||
|
return tuple(prefix_lengths)
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self.pattern)
|
||||||
|
|
||||||
|
def advance(self, matched: int, token: int) -> int:
|
||||||
|
while matched > 0 and token != self.pattern[matched]:
|
||||||
|
matched = self.prefix_lengths[matched - 1]
|
||||||
|
if token == self.pattern[matched]:
|
||||||
|
matched += 1
|
||||||
|
return matched
|
||||||
@@ -0,0 +1,221 @@
|
|||||||
|
import json
|
||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import Function, Tool
|
||||||
|
from sglang.srt.function_call.core_types import ToolCallItem
|
||||||
|
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
||||||
|
from sglang.srt.function_call.kimik3_detector import KimiK3Detector
|
||||||
|
from sglang.srt.function_call.kimik3_format import (
|
||||||
|
MESSAGE_CLOSE,
|
||||||
|
RESPONSE_CLOSE,
|
||||||
|
RESPONSE_OPEN,
|
||||||
|
TOOLS_CLOSE,
|
||||||
|
TOOLS_OPEN,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_tool(name: str) -> Tool:
|
||||||
|
return Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name=name,
|
||||||
|
description=f"{name} tool",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"code": {"type": "string"}},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _call_block(tool: str, index: int, args: dict[str, tuple[str, str]]) -> str:
|
||||||
|
parts = [f'<|open|>call tool="{tool}" index="{index}"<|sep|>']
|
||||||
|
for key, (arg_type, value) in args.items():
|
||||||
|
parts.append(
|
||||||
|
f'<|open|>argument key="{key}" type="{arg_type}"<|sep|>'
|
||||||
|
f"{value}<|close|>argument<|sep|>"
|
||||||
|
)
|
||||||
|
parts.append("<|close|>call<|sep|>")
|
||||||
|
return "".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _chunks(text: str, size: int) -> list[str]:
|
||||||
|
return [text[index : index + size] for index in range(0, len(text), size)]
|
||||||
|
|
||||||
|
|
||||||
|
def _stream(
|
||||||
|
detector: KimiK3Detector, chunks: list[str], tools: list[Tool]
|
||||||
|
) -> tuple[str, list[ToolCallItem]]:
|
||||||
|
text = ""
|
||||||
|
calls = []
|
||||||
|
for chunk in chunks:
|
||||||
|
result = detector.parse_streaming_increment(chunk, tools)
|
||||||
|
text += result.normal_text
|
||||||
|
calls.extend(result.calls)
|
||||||
|
return text, calls
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_and_parse_single_call() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
tools = [_make_tool("python")]
|
||||||
|
text = (
|
||||||
|
f"{RESPONSE_OPEN}Let me run it.{RESPONSE_CLOSE}{TOOLS_OPEN}"
|
||||||
|
+ _call_block(
|
||||||
|
"python",
|
||||||
|
1,
|
||||||
|
{"code": ("string", "print(1)"), "opts": ("object", '{"a": 1}')},
|
||||||
|
)
|
||||||
|
+ TOOLS_CLOSE
|
||||||
|
)
|
||||||
|
result = detector.detect_and_parse(text, tools)
|
||||||
|
assert result.normal_text == "Let me run it."
|
||||||
|
assert len(result.calls) == 1
|
||||||
|
assert result.calls[0].name == "python"
|
||||||
|
assert json.loads(result.calls[0].parameters) == {
|
||||||
|
"code": "print(1)",
|
||||||
|
"opts": {"a": 1},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_and_parse_no_tools_channel() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
result = detector.detect_and_parse(
|
||||||
|
f"{RESPONSE_OPEN}hi there{RESPONSE_CLOSE}{MESSAGE_CLOSE}",
|
||||||
|
[_make_tool("python")],
|
||||||
|
)
|
||||||
|
assert result.normal_text == "hi there"
|
||||||
|
assert result.calls == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_and_parse_multiple_calls() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
text = (
|
||||||
|
TOOLS_OPEN
|
||||||
|
+ _call_block("python", 1, {"code": ("string", "a")})
|
||||||
|
+ _call_block("python", 2, {"code": ("string", "b")})
|
||||||
|
+ TOOLS_CLOSE
|
||||||
|
)
|
||||||
|
result = detector.detect_and_parse(text, [_make_tool("python")])
|
||||||
|
assert [call.tool_index for call in result.calls] == [0, 1]
|
||||||
|
assert json.loads(result.calls[1].parameters) == {"code": "b"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_and_parse_unclosed_tools_section() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
text = TOOLS_OPEN + _call_block("python", 1, {"code": ("string", "x")})
|
||||||
|
result = detector.detect_and_parse(text, [_make_tool("python")])
|
||||||
|
assert len(result.calls) == 1
|
||||||
|
assert json.loads(result.calls[0].parameters) == {"code": "x"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_attr_unescaping_and_raw_string_args() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
text = (
|
||||||
|
f"{TOOLS_OPEN}"
|
||||||
|
'<|open|>call tool="a&b" index="1"<|sep|>'
|
||||||
|
'<|open|>argument key="q" type="string"<|sep|>'
|
||||||
|
"say "hi"<|close|>argument<|sep|>"
|
||||||
|
"<|close|>call<|sep|>"
|
||||||
|
f"{TOOLS_CLOSE}"
|
||||||
|
)
|
||||||
|
result = detector.detect_and_parse(text, [_make_tool("python")])
|
||||||
|
assert result.calls[0].name == "a&b"
|
||||||
|
assert json.loads(result.calls[0].parameters) == {"q": "say "hi""}
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_string_arg_json_decoding() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
text = (
|
||||||
|
TOOLS_OPEN
|
||||||
|
+ _call_block(
|
||||||
|
"python",
|
||||||
|
1,
|
||||||
|
{
|
||||||
|
"n": ("number", "42"),
|
||||||
|
"flag": ("boolean", "true"),
|
||||||
|
"bad": ("object", "{not json"),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
+ TOOLS_CLOSE
|
||||||
|
)
|
||||||
|
result = detector.detect_and_parse(text, [_make_tool("python")])
|
||||||
|
assert json.loads(result.calls[0].parameters) == {
|
||||||
|
"n": 42,
|
||||||
|
"flag": True,
|
||||||
|
"bad": "{not json",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("chunk_size", [1, 7, 23])
|
||||||
|
def test_streaming_split_markers(chunk_size: int) -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
tools = [_make_tool("python")]
|
||||||
|
text = (
|
||||||
|
f"{RESPONSE_OPEN}Hello!{RESPONSE_CLOSE}{TOOLS_OPEN}"
|
||||||
|
+ _call_block("python", 1, {"code": ("string", "print(2)")})
|
||||||
|
+ TOOLS_CLOSE
|
||||||
|
)
|
||||||
|
normal_text, calls = _stream(detector, _chunks(text, chunk_size), tools)
|
||||||
|
assert normal_text == "Hello!"
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert calls[0].name == "python"
|
||||||
|
assert json.loads(calls[0].parameters) == {"code": "print(2)"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_streaming_two_calls() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
tools = [_make_tool("python")]
|
||||||
|
text = (
|
||||||
|
TOOLS_OPEN
|
||||||
|
+ _call_block("python", 1, {"code": ("string", "a")})
|
||||||
|
+ _call_block("python", 2, {"code": ("string", "b")})
|
||||||
|
+ TOOLS_CLOSE
|
||||||
|
)
|
||||||
|
_, calls = _stream(detector, _chunks(text, 7), tools)
|
||||||
|
assert [call.tool_index for call in calls] == [0, 1]
|
||||||
|
assert [json.loads(call.parameters) for call in calls] == [
|
||||||
|
{"code": "a"},
|
||||||
|
{"code": "b"},
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_streaming_plain_text_only() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
text, calls = _stream(
|
||||||
|
detector, ["just a ", "plain ", "reply"], [_make_tool("python")]
|
||||||
|
)
|
||||||
|
assert text == "just a plain reply"
|
||||||
|
assert calls == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_streaming_bookkeeping_for_serving_layer() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
tools = [_make_tool("python")]
|
||||||
|
text = (
|
||||||
|
TOOLS_OPEN + _call_block("python", 1, {"code": ("string", "a")}) + TOOLS_CLOSE
|
||||||
|
)
|
||||||
|
_stream(detector, _chunks(text, 9), tools)
|
||||||
|
assert detector.current_tool_id == 0
|
||||||
|
assert detector.prev_tool_call_arr[0] == {
|
||||||
|
"name": "python",
|
||||||
|
"arguments": {"code": "a"},
|
||||||
|
}
|
||||||
|
assert json.loads(detector.streamed_args_for_tool[0]) == {"code": "a"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_detector_capabilities_and_registration() -> None:
|
||||||
|
detector = KimiK3Detector()
|
||||||
|
assert detector.supports_structural_tag()
|
||||||
|
assert not detector.parses_required_natively()
|
||||||
|
parser = FunctionCallParser([_make_tool("python")], "kimi_k3")
|
||||||
|
assert isinstance(parser.detector, KimiK3Detector)
|
||||||
|
assert parser.get_structure_constraint("required") is not None
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__]))
|
||||||
@@ -359,30 +359,28 @@ class TestCreateGrammarBackend(unittest.TestCase):
|
|||||||
# encode must return a single-token list for think_start/end tokens
|
# encode must return a single-token list for think_start/end tokens
|
||||||
tokenizer.encode.return_value = [42]
|
tokenizer.encode.return_value = [42]
|
||||||
|
|
||||||
result = create_grammar_backend(args, tokenizer, 32000, think_end_id=42)
|
result = create_grammar_backend(args, tokenizer, 32000, think_end_ids=[42])
|
||||||
self.assertIsInstance(result, ReasonerGrammarBackend)
|
self.assertIsInstance(result, ReasonerGrammarBackend)
|
||||||
self.assertIs(result.grammar_backend, mock_backend)
|
self.assertIs(result.grammar_backend, mock_backend)
|
||||||
|
|
||||||
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
||||||
def test_no_reasoner_wrapping_without_think_end_id(self, mock_outlines_cls):
|
def test_no_reasoner_wrapping_without_think_end_ids(self, mock_outlines_cls):
|
||||||
"""Without think_end_id passed in, no reasoner wrapping."""
|
|
||||||
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
mock_outlines_cls.return_value = mock_backend
|
mock_outlines_cls.return_value = mock_backend
|
||||||
args = self._make_server_args("outlines", reasoning_parser="deepseek-r1")
|
args = self._make_server_args("outlines", reasoning_parser="deepseek-r1")
|
||||||
tokenizer = MagicMock(spec=[]) # No think_end_id attribute
|
tokenizer = MagicMock(spec=[])
|
||||||
|
|
||||||
result = create_grammar_backend(args, tokenizer, 32000, think_end_id=None)
|
result = create_grammar_backend(args, tokenizer, 32000, think_end_ids=None)
|
||||||
self.assertIs(result, mock_backend)
|
self.assertIs(result, mock_backend)
|
||||||
|
|
||||||
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
@patch("sglang.srt.constrained.outlines_backend.OutlinesGrammarBackend")
|
||||||
def test_no_reasoner_wrapping_without_reasoning_parser(self, mock_outlines_cls):
|
def test_no_reasoner_wrapping_without_reasoning_parser(self, mock_outlines_cls):
|
||||||
"""Without reasoning_parser, no reasoner wrapping even with think_end_id."""
|
|
||||||
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
mock_outlines_cls.return_value = mock_backend
|
mock_outlines_cls.return_value = mock_backend
|
||||||
args = self._make_server_args("outlines", reasoning_parser=None)
|
args = self._make_server_args("outlines", reasoning_parser=None)
|
||||||
tokenizer = MagicMock()
|
tokenizer = MagicMock()
|
||||||
|
|
||||||
result = create_grammar_backend(args, tokenizer, 32000, think_end_id=42)
|
result = create_grammar_backend(args, tokenizer, 32000, think_end_ids=[42])
|
||||||
self.assertIs(result, mock_backend)
|
self.assertIs(result, mock_backend)
|
||||||
|
|
||||||
@patch("sglang.srt.constrained.xgrammar_backend.XGrammarGrammarBackend")
|
@patch("sglang.srt.constrained.xgrammar_backend.XGrammarGrammarBackend")
|
||||||
|
|||||||
@@ -273,7 +273,7 @@ class TestProcessReqWithGrammar(unittest.TestCase):
|
|||||||
def test_cache_hit_applies_request_thinking_budget(self):
|
def test_cache_hit_applies_request_thinking_budget(self):
|
||||||
mgr = self._make_mgr()
|
mgr = self._make_mgr()
|
||||||
grammar_obj = ReasonerGrammarObject(
|
grammar_obj = ReasonerGrammarObject(
|
||||||
grammar=None, think_end_id=0, max_think_tokens=99
|
grammar=None, think_end_ids=[0], max_think_tokens=99
|
||||||
)
|
)
|
||||||
mgr.grammar_backend.get_cached_or_future_value.return_value = (
|
mgr.grammar_backend.get_cached_or_future_value.return_value = (
|
||||||
grammar_obj,
|
grammar_obj,
|
||||||
@@ -292,7 +292,7 @@ class TestProcessReqWithGrammar(unittest.TestCase):
|
|||||||
mgr = self._make_mgr()
|
mgr = self._make_mgr()
|
||||||
mgr._enable_strict_thinking = True
|
mgr._enable_strict_thinking = True
|
||||||
grammar_obj = ReasonerGrammarObject(
|
grammar_obj = ReasonerGrammarObject(
|
||||||
grammar=None, think_end_id=0, max_think_tokens=99
|
grammar=None, think_end_ids=[0], max_think_tokens=99
|
||||||
)
|
)
|
||||||
mgr.grammar_backend.init_strict_reasoning_grammar.return_value = grammar_obj
|
mgr.grammar_backend.init_strict_reasoning_grammar.return_value = grammar_obj
|
||||||
|
|
||||||
@@ -545,7 +545,7 @@ class TestGetReadyGrammarRequests(unittest.TestCase):
|
|||||||
mgr = self._make_mgr()
|
mgr = self._make_mgr()
|
||||||
|
|
||||||
grammar_obj = ReasonerGrammarObject(
|
grammar_obj = ReasonerGrammarObject(
|
||||||
grammar=None, think_end_id=0, max_think_tokens=99
|
grammar=None, think_end_ids=[0], max_think_tokens=99
|
||||||
)
|
)
|
||||||
future = Future()
|
future = Future()
|
||||||
future.set_result(grammar_obj)
|
future.set_result(grammar_obj)
|
||||||
|
|||||||
@@ -13,6 +13,8 @@ from sglang.srt.constrained.reasoner_grammar_backend import (
|
|||||||
from sglang.srt.constrained.torch_ops.token_filter_torch_ops import (
|
from sglang.srt.constrained.torch_ops.token_filter_torch_ops import (
|
||||||
set_token_filter_torch,
|
set_token_filter_torch,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.function_call.kimik3_format import THINK_CLOSE
|
||||||
|
from sglang.srt.parser.reasoning_parser import KimiK3Detector as KimiK3ReasoningDetector
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(2.0, "base-a-test-cpu")
|
register_cpu_ci(2.0, "base-a-test-cpu")
|
||||||
@@ -75,7 +77,7 @@ class TestReasonerGrammarObject(unittest.TestCase):
|
|||||||
def _make_strict_object(self):
|
def _make_strict_object(self):
|
||||||
return ReasonerGrammarObject(
|
return ReasonerGrammarObject(
|
||||||
grammar=None,
|
grammar=None,
|
||||||
think_end_id=7,
|
think_end_ids=[7],
|
||||||
think_excluded_token_ids=[3, 5],
|
think_excluded_token_ids=[3, 5],
|
||||||
max_think_tokens=2,
|
max_think_tokens=2,
|
||||||
enable_token_filter=True,
|
enable_token_filter=True,
|
||||||
@@ -117,6 +119,26 @@ class TestReasonerGrammarObject(unittest.TestCase):
|
|||||||
self.assertIs(obj.move_vocab_mask(mask, "cpu"), mask)
|
self.assertIs(obj.move_vocab_mask(mask, "cpu"), mask)
|
||||||
self.assertIsNotNone(obj.apply_vocab_mask)
|
self.assertIsNotNone(obj.apply_vocab_mask)
|
||||||
|
|
||||||
|
def test_budget_exhaustion_walks_multi_token_end(self):
|
||||||
|
obj = ReasonerGrammarObject(
|
||||||
|
grammar=None,
|
||||||
|
think_end_ids=[7, 8],
|
||||||
|
max_think_tokens=1,
|
||||||
|
enable_token_filter=True,
|
||||||
|
token_filter_fn=set_token_filter_torch,
|
||||||
|
)
|
||||||
|
obj.maybe_init_reasoning(True)
|
||||||
|
obj.accept_token(10)
|
||||||
|
|
||||||
|
first_mask = torch.zeros((1, 2), dtype=torch.int32)
|
||||||
|
obj.fill_vocab_mask(first_mask, 0)
|
||||||
|
self.assertEqual(_allowed_token_ids(first_mask, [7, 8, 10]), [7])
|
||||||
|
|
||||||
|
obj.accept_token(7)
|
||||||
|
second_mask = torch.zeros((1, 2), dtype=torch.int32)
|
||||||
|
obj.fill_vocab_mask(second_mask, 0)
|
||||||
|
self.assertEqual(_allowed_token_ids(second_mask, [7, 8, 10]), [8])
|
||||||
|
|
||||||
|
|
||||||
class TestReasonerGrammarBackend(unittest.TestCase):
|
class TestReasonerGrammarBackend(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
@@ -163,6 +185,42 @@ class TestReasonerGrammarBackend(unittest.TestCase):
|
|||||||
self.assertEqual(obj.max_think_tokens, 2)
|
self.assertEqual(obj.max_think_tokens, 2)
|
||||||
self.assertEqual(obj.think_excluded_token_ids, [3, 4])
|
self.assertEqual(obj.think_excluded_token_ids, [3, 4])
|
||||||
|
|
||||||
|
def test_kimi_k3_excluded_tokens_spare_the_xtml_control_tokens(self):
|
||||||
|
"""Kimi K3 bans bare channel names, never the marker-composing tokens.
|
||||||
|
|
||||||
|
The excluded list is flattened into single token ids, so listing a whole
|
||||||
|
marker such as "<|open|>response<|sep|>" would ban <|open|> and <|sep|>
|
||||||
|
individually -- which also blocks the think-end sequence and the jump
|
||||||
|
into the tools channel, leaving the model unable to stop thinking.
|
||||||
|
"""
|
||||||
|
control_ids = {"<|open|>": [1], "<|close|>": [2], "<|sep|>": [3]}
|
||||||
|
think_end_ids = [2, 4, 3]
|
||||||
|
tokenizer = _DummyTokenizer(
|
||||||
|
{
|
||||||
|
THINK_CLOSE: think_end_ids,
|
||||||
|
"response": [10],
|
||||||
|
"message": [11],
|
||||||
|
"<|end_of_msg|>": [12],
|
||||||
|
"[EOS]": [13],
|
||||||
|
"[EOT]": [14],
|
||||||
|
**control_ids,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
reasoner = ReasonerGrammarBackend(
|
||||||
|
_DummyGrammarBackend(support_token_filter=True),
|
||||||
|
SimpleNamespace(detector=KimiK3ReasoningDetector()),
|
||||||
|
tokenizer,
|
||||||
|
enable_strict_thinking=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
excluded = reasoner.think_excluded_token_ids
|
||||||
|
|
||||||
|
self.assertEqual(excluded, [10, 11, 12, 13, 14])
|
||||||
|
for token, ids in control_ids.items():
|
||||||
|
for token_id in ids:
|
||||||
|
self.assertNotIn(token_id, excluded, f"{token} must stay generatable")
|
||||||
|
self.assertEqual(set(think_end_ids) & set(excluded), set())
|
||||||
|
|
||||||
def test_init_strict_reasoning_grammar_none_when_strict_disabled(self):
|
def test_init_strict_reasoning_grammar_none_when_strict_disabled(self):
|
||||||
backend = _DummyGrammarBackend(support_token_filter=True)
|
backend = _DummyGrammarBackend(support_token_filter=True)
|
||||||
reasoner = ReasonerGrammarBackend(
|
reasoner = ReasonerGrammarBackend(
|
||||||
@@ -205,16 +263,15 @@ class TestReasonerGrammarBackend(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
self.assertIsNotNone(reasoner)
|
self.assertIsNotNone(reasoner)
|
||||||
|
|
||||||
def test_rejects_multi_token_think_end_marker(self):
|
def test_accepts_multi_token_think_end_marker(self):
|
||||||
backend = _DummyGrammarBackend(support_token_filter=True)
|
backend = _DummyGrammarBackend(support_token_filter=True)
|
||||||
|
reasoner = ReasonerGrammarBackend(
|
||||||
with self.assertRaisesRegex(ValueError, "must encode to exactly one token"):
|
|
||||||
ReasonerGrammarBackend(
|
|
||||||
backend,
|
backend,
|
||||||
self._make_parser(),
|
self._make_parser(),
|
||||||
self._make_tokenizer(end_ids=[2, 3]),
|
self._make_tokenizer(end_ids=[2, 3]),
|
||||||
enable_strict_thinking=True,
|
enable_strict_thinking=True,
|
||||||
)
|
)
|
||||||
|
self.assertEqual(reasoner.think_end_ids, [2, 3])
|
||||||
|
|
||||||
def test_rejects_unencodable_excluded_token(self):
|
def test_rejects_unencodable_excluded_token(self):
|
||||||
backend = _DummyGrammarBackend(support_token_filter=True)
|
backend = _DummyGrammarBackend(support_token_filter=True)
|
||||||
@@ -255,7 +312,7 @@ class TestReasonerGrammarObjectRollback(unittest.TestCase):
|
|||||||
inner_grammar.is_terminated.return_value = False
|
inner_grammar.is_terminated.return_value = False
|
||||||
obj = ReasonerGrammarObject(
|
obj = ReasonerGrammarObject(
|
||||||
grammar=inner_grammar,
|
grammar=inner_grammar,
|
||||||
think_end_id=7,
|
think_end_ids=[7],
|
||||||
think_excluded_token_ids=[3, 5],
|
think_excluded_token_ids=[3, 5],
|
||||||
max_think_tokens=-1,
|
max_think_tokens=-1,
|
||||||
enable_token_filter=True,
|
enable_token_filter=True,
|
||||||
@@ -272,11 +329,10 @@ class TestReasonerGrammarObjectRollback(unittest.TestCase):
|
|||||||
obj, inner_grammar = self._make_object_with_mock_grammar()
|
obj, inner_grammar = self._make_object_with_mock_grammar()
|
||||||
obj.maybe_init_reasoning(True)
|
obj.maybe_init_reasoning(True)
|
||||||
|
|
||||||
# Accept 3 thinking tokens then think_end_id
|
|
||||||
obj.accept_token(10)
|
obj.accept_token(10)
|
||||||
obj.accept_token(11)
|
obj.accept_token(11)
|
||||||
obj.accept_token(12)
|
obj.accept_token(12)
|
||||||
obj.accept_token(7) # think_end_id → tokens_after_end = 0
|
obj.accept_token(7)
|
||||||
|
|
||||||
self.assertTrue(obj._is_generation())
|
self.assertTrue(obj._is_generation())
|
||||||
self.assertEqual(obj.tokens_after_end, 0)
|
self.assertEqual(obj.tokens_after_end, 0)
|
||||||
@@ -296,7 +352,7 @@ class TestReasonerGrammarObjectRollback(unittest.TestCase):
|
|||||||
# 2 thinking tokens + think_end + 3 generation tokens
|
# 2 thinking tokens + think_end + 3 generation tokens
|
||||||
obj.accept_token(10) # think
|
obj.accept_token(10) # think
|
||||||
obj.accept_token(11) # think
|
obj.accept_token(11) # think
|
||||||
obj.accept_token(7) # think_end_id
|
obj.accept_token(7)
|
||||||
obj.accept_token(20) # gen 1
|
obj.accept_token(20) # gen 1
|
||||||
obj.accept_token(21) # gen 2
|
obj.accept_token(21) # gen 2
|
||||||
obj.accept_token(22) # gen 3
|
obj.accept_token(22) # gen 3
|
||||||
@@ -315,7 +371,7 @@ class TestReasonerGrammarObjectRollback(unittest.TestCase):
|
|||||||
obj.maybe_init_reasoning(True)
|
obj.maybe_init_reasoning(True)
|
||||||
|
|
||||||
obj.accept_token(10) # think
|
obj.accept_token(10) # think
|
||||||
obj.accept_token(7) # think_end_id
|
obj.accept_token(7)
|
||||||
obj.accept_token(20) # gen 1
|
obj.accept_token(20) # gen 1
|
||||||
obj.accept_token(21) # gen 2
|
obj.accept_token(21) # gen 2
|
||||||
|
|
||||||
@@ -344,7 +400,7 @@ class TestReasonerGrammarObjectRollback(unittest.TestCase):
|
|||||||
obj.maybe_init_reasoning(True)
|
obj.maybe_init_reasoning(True)
|
||||||
|
|
||||||
obj.accept_token(10)
|
obj.accept_token(10)
|
||||||
obj.accept_token(7) # think_end_id → GENERATION
|
obj.accept_token(7)
|
||||||
obj.accept_token(20)
|
obj.accept_token(20)
|
||||||
|
|
||||||
self.assertEqual(obj.tokens_in_think, 1)
|
self.assertEqual(obj.tokens_in_think, 1)
|
||||||
@@ -370,6 +426,26 @@ class TestReasonerGrammarObjectRollback(unittest.TestCase):
|
|||||||
self.assertEqual(copy.tokens_after_end, -1)
|
self.assertEqual(copy.tokens_after_end, -1)
|
||||||
self.assertTrue(copy._is_thinking())
|
self.assertTrue(copy._is_thinking())
|
||||||
|
|
||||||
|
def test_multi_token_marker_survives_rollback(self):
|
||||||
|
obj = ReasonerGrammarObject(grammar=None, think_end_ids=[2, 3])
|
||||||
|
obj.maybe_init_reasoning(True)
|
||||||
|
obj.accept_token(2)
|
||||||
|
obj.accept_token(9)
|
||||||
|
obj.rollback(1)
|
||||||
|
obj.accept_token(3)
|
||||||
|
self.assertTrue(obj._is_generation())
|
||||||
|
|
||||||
|
obj.rollback(1)
|
||||||
|
self.assertTrue(obj._is_thinking())
|
||||||
|
self.assertEqual(obj._matched_think_end_tokens, 1)
|
||||||
|
|
||||||
|
def test_self_overlapping_marker_is_matched(self):
|
||||||
|
obj = ReasonerGrammarObject(grammar=None, think_end_ids=[2, 2, 3])
|
||||||
|
obj.maybe_init_reasoning(True)
|
||||||
|
for token in (2, 2, 2, 3):
|
||||||
|
obj.accept_token(token)
|
||||||
|
self.assertTrue(obj._is_generation())
|
||||||
|
|
||||||
|
|
||||||
class TestReasonerGrammarObjectFillVocabMask(unittest.TestCase):
|
class TestReasonerGrammarObjectFillVocabMask(unittest.TestCase):
|
||||||
"""Tests for fill_vocab_mask behavior in different states."""
|
"""Tests for fill_vocab_mask behavior in different states."""
|
||||||
@@ -383,7 +459,7 @@ class TestReasonerGrammarObjectFillVocabMask(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
obj = ReasonerGrammarObject(
|
obj = ReasonerGrammarObject(
|
||||||
grammar=inner_grammar,
|
grammar=inner_grammar,
|
||||||
think_end_id=7,
|
think_end_ids=[7],
|
||||||
think_excluded_token_ids=[3, 5],
|
think_excluded_token_ids=[3, 5],
|
||||||
max_think_tokens=-1,
|
max_think_tokens=-1,
|
||||||
enable_token_filter=True,
|
enable_token_filter=True,
|
||||||
@@ -411,7 +487,7 @@ class TestReasonerGrammarObjectFillVocabMask(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
obj = ReasonerGrammarObject(
|
obj = ReasonerGrammarObject(
|
||||||
grammar=inner_grammar,
|
grammar=inner_grammar,
|
||||||
think_end_id=7,
|
think_end_ids=[7],
|
||||||
think_excluded_token_ids=[3, 5],
|
think_excluded_token_ids=[3, 5],
|
||||||
max_think_tokens=-1,
|
max_think_tokens=-1,
|
||||||
enable_token_filter=True,
|
enable_token_filter=True,
|
||||||
@@ -424,7 +500,7 @@ class TestReasonerGrammarObjectFillVocabMask(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
obj.maybe_init_reasoning(True)
|
obj.maybe_init_reasoning(True)
|
||||||
obj.accept_token(10)
|
obj.accept_token(10)
|
||||||
obj.accept_token(7) # think_end_id → GENERATION
|
obj.accept_token(7)
|
||||||
|
|
||||||
mask = obj.allocate_vocab_mask(64, 1, "cpu")
|
mask = obj.allocate_vocab_mask(64, 1, "cpu")
|
||||||
obj.fill_vocab_mask(mask, 0)
|
obj.fill_vocab_mask(mask, 0)
|
||||||
@@ -435,7 +511,7 @@ class TestReasonerGrammarObjectFillVocabMask(unittest.TestCase):
|
|||||||
inner_grammar = MagicMock()
|
inner_grammar = MagicMock()
|
||||||
obj = ReasonerGrammarObject(
|
obj = ReasonerGrammarObject(
|
||||||
grammar=inner_grammar,
|
grammar=inner_grammar,
|
||||||
think_end_id=7,
|
think_end_ids=[7],
|
||||||
think_excluded_token_ids=None,
|
think_excluded_token_ids=None,
|
||||||
max_think_tokens=-1,
|
max_think_tokens=-1,
|
||||||
enable_token_filter=False,
|
enable_token_filter=False,
|
||||||
@@ -463,7 +539,7 @@ class TestReasonerGrammarObjectCurrentToken(unittest.TestCase):
|
|||||||
inner_grammar.is_terminated.return_value = False
|
inner_grammar.is_terminated.return_value = False
|
||||||
obj = ReasonerGrammarObject(
|
obj = ReasonerGrammarObject(
|
||||||
grammar=inner_grammar,
|
grammar=inner_grammar,
|
||||||
think_end_id=7,
|
think_end_ids=[7],
|
||||||
think_excluded_token_ids=None,
|
think_excluded_token_ids=None,
|
||||||
max_think_tokens=-1,
|
max_think_tokens=-1,
|
||||||
enable_token_filter=False,
|
enable_token_filter=False,
|
||||||
@@ -480,7 +556,7 @@ class TestReasonerGrammarObjectCurrentToken(unittest.TestCase):
|
|||||||
obj, inner_grammar = self._make_object_with_mock_grammar()
|
obj, inner_grammar = self._make_object_with_mock_grammar()
|
||||||
obj.maybe_init_reasoning(True)
|
obj.maybe_init_reasoning(True)
|
||||||
obj.accept_token(10) # thinking token
|
obj.accept_token(10) # thinking token
|
||||||
obj.accept_token(7) # think_end_id -> GENERATION
|
obj.accept_token(7)
|
||||||
obj.accept_token(58) # generation token "["
|
obj.accept_token(58) # generation token "["
|
||||||
self.assertEqual(obj.current_token, 58)
|
self.assertEqual(obj.current_token, 58)
|
||||||
|
|
||||||
@@ -497,7 +573,7 @@ class TestReasonerGrammarObjectCurrentToken(unittest.TestCase):
|
|||||||
must not be re-accepted; with current_token tracked, the guard skips."""
|
must not be re-accepted; with current_token tracked, the guard skips."""
|
||||||
obj, inner_grammar = self._make_object_with_mock_grammar()
|
obj, inner_grammar = self._make_object_with_mock_grammar()
|
||||||
obj.maybe_init_reasoning(True)
|
obj.maybe_init_reasoning(True)
|
||||||
obj.accept_token(7) # think_end_id -> GENERATION
|
obj.accept_token(7)
|
||||||
obj.accept_token(58) # "[" accepted into inner grammar
|
obj.accept_token(58) # "[" accepted into inner grammar
|
||||||
obj.accept_token(4913) # '{"' accepted into inner grammar
|
obj.accept_token(4913) # '{"' accepted into inner grammar
|
||||||
inner_grammar.accept_token.reset_mock()
|
inner_grammar.accept_token.reset_mock()
|
||||||
|
|||||||
@@ -438,6 +438,68 @@ class TestChatCompletionRequest(unittest.TestCase):
|
|||||||
self.assertEqual(name, "VoiceNote")
|
self.assertEqual(name, "VoiceNote")
|
||||||
self.assertEqual(strict, True)
|
self.assertEqual(strict, True)
|
||||||
|
|
||||||
|
def test_schema_derived_strict_false_constraint_gated_on_renderer(self):
|
||||||
|
"""A `strict` field on the user's model doubles as the protocol switch.
|
||||||
|
|
||||||
|
set_json_schema pops `strict` out of the schema's properties and feeds
|
||||||
|
its default into response_format. strict=False drops the sampling
|
||||||
|
constraint only when the renderer forwards response_format to the
|
||||||
|
model; otherwise the schema would be silently ignored, so the
|
||||||
|
constraint stays installed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
class Note(BaseModel):
|
||||||
|
title: str
|
||||||
|
strict: bool = False
|
||||||
|
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="test-model",
|
||||||
|
messages=[{"role": "user", "content": "Return JSON"}],
|
||||||
|
response_format={
|
||||||
|
"type": "json_schema",
|
||||||
|
"schema": Note.model_json_schema(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(request.response_format.json_schema.strict, False)
|
||||||
|
self.assertNotIn(
|
||||||
|
"strict", request.response_format.json_schema.schema_["properties"]
|
||||||
|
)
|
||||||
|
sampling_params = request.to_sampling_params(
|
||||||
|
stop=[], model_generation_config={}
|
||||||
|
)
|
||||||
|
self.assertIn("json_schema", sampling_params)
|
||||||
|
sampling_params = request.to_sampling_params(
|
||||||
|
stop=[],
|
||||||
|
model_generation_config={},
|
||||||
|
renderer_handles_response_format=True,
|
||||||
|
)
|
||||||
|
self.assertNotIn("json_schema", sampling_params)
|
||||||
|
|
||||||
|
def test_non_strict_response_format_constraint_gated_on_renderer(self):
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="test-model",
|
||||||
|
messages=[{"role": "user", "content": "Return JSON"}],
|
||||||
|
response_format={
|
||||||
|
"type": "json_schema",
|
||||||
|
"json_schema": {
|
||||||
|
"name": "answer",
|
||||||
|
"schema": {"type": "object"},
|
||||||
|
"strict": False,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
sampling_params = request.to_sampling_params(
|
||||||
|
stop=[], model_generation_config={}
|
||||||
|
)
|
||||||
|
self.assertIn("json_schema", sampling_params)
|
||||||
|
sampling_params = request.to_sampling_params(
|
||||||
|
stop=[],
|
||||||
|
model_generation_config={},
|
||||||
|
renderer_handles_response_format=True,
|
||||||
|
)
|
||||||
|
self.assertNotIn("json_schema", sampling_params)
|
||||||
|
|
||||||
|
|
||||||
class TestModelSerialization(unittest.TestCase):
|
class TestModelSerialization(unittest.TestCase):
|
||||||
"""Test model serialization with hidden states"""
|
"""Test model serialization with hidden states"""
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from sglang.srt.entrypoints.openai.serving_chat import (
|
|||||||
OpenAIServingChat,
|
OpenAIServingChat,
|
||||||
normalize_tool_content,
|
normalize_tool_content,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.function_call.kimik3_format import TOOLS_CLOSE
|
||||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
from sglang.srt.parser.template_detection import ReasoningToggleConfig
|
from sglang.srt.parser.template_detection import ReasoningToggleConfig
|
||||||
from sglang.srt.utils import get_or_create_event_loop
|
from sglang.srt.utils import get_or_create_event_loop
|
||||||
@@ -301,12 +302,40 @@ class ServingChatTestCase(unittest.TestCase):
|
|||||||
[],
|
[],
|
||||||
[],
|
[],
|
||||||
None,
|
None,
|
||||||
|
require_reasoning=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
adapted, _ = self.chat._convert_to_internal_request(req)
|
adapted, _ = self.chat._convert_to_internal_request(req)
|
||||||
|
|
||||||
self.assertTrue(adapted.require_reasoning)
|
self.assertTrue(adapted.require_reasoning)
|
||||||
|
|
||||||
|
def test_process_messages_records_template_reasoning_state(self):
|
||||||
|
self.chat.default_chat_template_kwargs = {"thinking": True}
|
||||||
|
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||||||
|
toggle_param="thinking", default_enabled=False
|
||||||
|
)
|
||||||
|
self.chat.reasoning_parser = "deepseek-v3"
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="x",
|
||||||
|
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||||||
|
)
|
||||||
|
rendered = MessageProcessingResult(
|
||||||
|
prompt="prompt",
|
||||||
|
prompt_ids=[1, 2, 3],
|
||||||
|
image_data=None,
|
||||||
|
audio_data=None,
|
||||||
|
video_data=None,
|
||||||
|
modalities=[],
|
||||||
|
stop=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
self.chat, "_apply_conversation_template", return_value=rendered
|
||||||
|
):
|
||||||
|
processed = self.chat._process_messages(request, is_multimodal=False)
|
||||||
|
|
||||||
|
self.assertTrue(processed.require_reasoning)
|
||||||
|
|
||||||
def test_kimi_tool_call_respects_explicit_reasoning_disable(self):
|
def test_kimi_tool_call_respects_explicit_reasoning_disable(self):
|
||||||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||||||
toggle_param="thinking", default_enabled=True
|
toggle_param="thinking", default_enabled=True
|
||||||
@@ -632,6 +661,260 @@ class ServingChatTestCase(unittest.TestCase):
|
|||||||
parser.get_structure_constraint.call_args.kwargs["thinking_mode"]
|
parser.get_structure_constraint.call_args.kwargs["thinking_mode"]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_kimi_k3_constraint_failure_keeps_native_stop_format(self):
|
||||||
|
self.template_manager.chat_template_name = None
|
||||||
|
self.template_manager.jinja_template_content_format = "string"
|
||||||
|
self.chat.chat_encoding_spec = "kimi_k3"
|
||||||
|
self.chat.tool_call_parser = "kimi_k3"
|
||||||
|
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||||||
|
tool = {
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "get_weather",
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"location": {"type": "string"}},
|
||||||
|
"required": ["location"],
|
||||||
|
},
|
||||||
|
"strict": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
cases = (
|
||||||
|
(None, [TOOLS_CLOSE]),
|
||||||
|
("USER_STOP", ["USER_STOP", TOOLS_CLOSE]),
|
||||||
|
(["USER_STOP"], ["USER_STOP", TOOLS_CLOSE]),
|
||||||
|
([TOOLS_CLOSE], [TOOLS_CLOSE]),
|
||||||
|
)
|
||||||
|
for request_stop, expected in cases:
|
||||||
|
with (
|
||||||
|
self.subTest(request_stop=request_stop),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.entrypoints.openai.serving_chat.FunctionCallParser"
|
||||||
|
) as parser_cls,
|
||||||
|
):
|
||||||
|
parser = parser_cls.return_value
|
||||||
|
parser.detector.eot_token = TOOLS_CLOSE
|
||||||
|
parser.detector.parses_required_natively.return_value = False
|
||||||
|
parser.get_structure_constraint.return_value = None
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="x",
|
||||||
|
messages=[{"role": "user", "content": "Weather in Paris?"}],
|
||||||
|
tools=[tool],
|
||||||
|
tool_choice="required",
|
||||||
|
stop=request_stop,
|
||||||
|
)
|
||||||
|
original_stop = (
|
||||||
|
list(request.stop)
|
||||||
|
if isinstance(request.stop, list)
|
||||||
|
else request.stop
|
||||||
|
)
|
||||||
|
|
||||||
|
result = self.chat._process_messages(request, is_multimodal=False)
|
||||||
|
|
||||||
|
self.assertEqual(result.stop, expected)
|
||||||
|
self.assertEqual(request.stop, original_stop)
|
||||||
|
self.assertIsNone(result.tool_call_constraint)
|
||||||
|
|
||||||
|
def test_kimi_k3_tool_call_stop_is_scoped_to_active_tools(self):
|
||||||
|
self.template_manager.chat_template_name = None
|
||||||
|
self.template_manager.jinja_template_content_format = "string"
|
||||||
|
self.chat.chat_encoding_spec = "kimi_k3"
|
||||||
|
self.chat.tool_call_parser = "kimi_k3"
|
||||||
|
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="x",
|
||||||
|
messages=[{"role": "user", "content": "Weather in Paris?"}],
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "get_weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
tool_choice="none",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = self.chat._process_messages(request, is_multimodal=False)
|
||||||
|
|
||||||
|
self.assertIsNone(result.stop)
|
||||||
|
|
||||||
|
def test_kimi_k3_encoder_receives_wire_request_fields(self):
|
||||||
|
self.template_manager.chat_template_name = None
|
||||||
|
self.chat.chat_encoding_spec = "kimi_k3"
|
||||||
|
self.tm.model_config.is_multimodal = True
|
||||||
|
self.tm.tokenizer.apply_chat_template.return_value = [7, 8, 9]
|
||||||
|
tool = {
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="x",
|
||||||
|
messages=[
|
||||||
|
{
|
||||||
|
"role": "developer",
|
||||||
|
"content": "<|kimi_image_placeholder|>",
|
||||||
|
"tools": [tool],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{
|
||||||
|
"type": "text",
|
||||||
|
"text": "Explain <|kimi_image_placeholder|>",
|
||||||
|
},
|
||||||
|
{"type": "image_url", "image_url": {"url": "image-1"}},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"reasoning_content": "Inspect <|kimi_image_placeholder|>",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": "call-1",
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "inspect",
|
||||||
|
"arguments": {
|
||||||
|
"source": "<|kimi_image_placeholder|>",
|
||||||
|
"nested": ["<|kimi_image_placeholder|>"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
tools=[tool],
|
||||||
|
tool_choice="required",
|
||||||
|
response_format={
|
||||||
|
"type": "json_schema",
|
||||||
|
"json_schema": {
|
||||||
|
"name": "answer",
|
||||||
|
"schema": {"type": "object"},
|
||||||
|
"strict": False,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
result = self.chat._process_messages(request, is_multimodal=True)
|
||||||
|
|
||||||
|
call = self.tm.tokenizer.apply_chat_template.call_args
|
||||||
|
rendered_messages = call.args[0]
|
||||||
|
self.assertEqual(rendered_messages[0]["role"], "system")
|
||||||
|
self.assertEqual(
|
||||||
|
rendered_messages[0]["content"], "<| kimi_image_placeholder |>"
|
||||||
|
)
|
||||||
|
self.assertNotIn("strict", rendered_messages[0]["tools"][0]["function"])
|
||||||
|
self.assertEqual(
|
||||||
|
rendered_messages[1]["content"][0]["text"],
|
||||||
|
"Explain <| kimi_image_placeholder |>",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
rendered_messages[2]["reasoning_content"],
|
||||||
|
"Inspect <| kimi_image_placeholder |>",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
rendered_messages[2]["tool_calls"][0]["function"]["arguments"],
|
||||||
|
{
|
||||||
|
"source": "<| kimi_image_placeholder |>",
|
||||||
|
"nested": ["<| kimi_image_placeholder |>"],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
self.assertEqual(call.kwargs["image_prompts"], ["<|media_pad|>"])
|
||||||
|
self.assertEqual(call.kwargs["tool_choice"], "required")
|
||||||
|
self.assertNotIn("strict", call.kwargs["tools"][0]["function"])
|
||||||
|
self.assertEqual(
|
||||||
|
call.kwargs["response_format"]["json_schema"]["schema"],
|
||||||
|
{"type": "object"},
|
||||||
|
)
|
||||||
|
self.assertNotIn("schema_", call.kwargs["response_format"]["json_schema"])
|
||||||
|
self.assertEqual(result.prompt_ids, [7, 8, 9])
|
||||||
|
self.assertEqual(result.image_data[0].url, "image-1")
|
||||||
|
|
||||||
|
def test_kimi_k3_neutralizes_text_only_assistant_history(self):
|
||||||
|
self.template_manager.chat_template_name = None
|
||||||
|
self.chat.chat_encoding_spec = "kimi_k3"
|
||||||
|
self.tm.tokenizer.apply_chat_template.return_value = [1, 2, 3]
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="x",
|
||||||
|
messages=[
|
||||||
|
{"role": "user", "content": "Run it"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"reasoning_content": "Read <|kimi_image_placeholder|>",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": "call-1",
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "shell",
|
||||||
|
"arguments": "not-json <|kimi_image_placeholder|>",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
self.chat._process_messages(request, is_multimodal=False)
|
||||||
|
|
||||||
|
messages = self.tm.tokenizer.apply_chat_template.call_args.args[0]
|
||||||
|
kwargs = self.tm.tokenizer.apply_chat_template.call_args.kwargs
|
||||||
|
self.assertEqual(messages[-1]["role"], "assistant")
|
||||||
|
self.assertEqual(
|
||||||
|
messages[-1]["reasoning_content"],
|
||||||
|
"Read <| kimi_image_placeholder |>",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
messages[-1]["tool_calls"][0]["function"]["arguments"],
|
||||||
|
"not-json <| kimi_image_placeholder |>",
|
||||||
|
)
|
||||||
|
self.assertNotIn("image_prompts", kwargs)
|
||||||
|
|
||||||
|
def test_message_tools_participate_in_validation_across_encodings(self):
|
||||||
|
tool = {
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "weather",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": "", "tools": [tool]},
|
||||||
|
{"role": "user", "content": "Weather?"},
|
||||||
|
]
|
||||||
|
for chat_encoding_spec in (None, "dsv4", "dsv32", "kimi_k3"):
|
||||||
|
with self.subTest(chat_encoding_spec=chat_encoding_spec):
|
||||||
|
self.chat.chat_encoding_spec = chat_encoding_spec
|
||||||
|
automatic = ChatCompletionRequest(
|
||||||
|
model="x", messages=messages, tool_choice=None
|
||||||
|
)
|
||||||
|
self.assertEqual(automatic.tool_choice, "auto")
|
||||||
|
self.assertIsNone(self.chat._validate_request(automatic))
|
||||||
|
|
||||||
|
required = ChatCompletionRequest(
|
||||||
|
model="x", messages=messages, tool_choice="required"
|
||||||
|
)
|
||||||
|
self.assertIsNone(self.chat._validate_request(required))
|
||||||
|
|
||||||
|
duplicate = ChatCompletionRequest(
|
||||||
|
model="x",
|
||||||
|
messages=messages,
|
||||||
|
tools=[tool],
|
||||||
|
tool_choice="required",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
self.chat._validate_request(duplicate),
|
||||||
|
"Tool names must be unique across request and message tools.",
|
||||||
|
)
|
||||||
|
|
||||||
def test_jinja_rejects_non_object_tool_call_arguments(self):
|
def test_jinja_rejects_non_object_tool_call_arguments(self):
|
||||||
"""History tool call arguments must parse to a JSON object."""
|
"""History tool call arguments must parse to a JSON object."""
|
||||||
self.template_manager.chat_template_name = None
|
self.template_manager.chat_template_name = None
|
||||||
@@ -1256,6 +1539,19 @@ class ServingChatTestCase(unittest.TestCase):
|
|||||||
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||||||
self.assertEqual(serving_chat.chat_encoding_spec, "dsv4")
|
self.assertEqual(serving_chat.chat_encoding_spec, "dsv4")
|
||||||
|
|
||||||
|
def test_kimi_k3_encoding_detection(self):
|
||||||
|
from sglang.srt.parser.template_manager import TemplateManager
|
||||||
|
|
||||||
|
tm = _MockTokenizerManager()
|
||||||
|
tm.model_config.hf_config.architectures = ["KimiK3ForConditionalGeneration"]
|
||||||
|
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||||||
|
self.assertEqual(serving_chat.chat_encoding_spec, "kimi_k3")
|
||||||
|
|
||||||
|
tm.model_config.hf_config.architectures = ["LlamaForCausalLM"]
|
||||||
|
tm.server_args.tool_call_parser = "kimi_k3"
|
||||||
|
serving_chat = OpenAIServingChat(tm, TemplateManager())
|
||||||
|
self.assertEqual(serving_chat.chat_encoding_spec, "kimi_k3")
|
||||||
|
|
||||||
# ------------- dsv4 task + latest_reminder -------------
|
# ------------- dsv4 task + latest_reminder -------------
|
||||||
def test_dsv4_task_field_schema(self):
|
def test_dsv4_task_field_schema(self):
|
||||||
"""Top-level `task` accepts the 6 DS task tokens and rejects others."""
|
"""Top-level `task` accepts the 6 DS task tokens and rejects others."""
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from sglang.srt.entrypoints.openai.protocol import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.entrypoints.openai.serving_responses import OpenAIServingResponses
|
from sglang.srt.entrypoints.openai.serving_responses import OpenAIServingResponses
|
||||||
from sglang.srt.function_call.core_types import ToolCallItem
|
from sglang.srt.function_call.core_types import ToolCallItem
|
||||||
|
from sglang.srt.parser.template_detection import ReasoningToggleConfig
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
|
||||||
@@ -205,6 +206,105 @@ class ChatToolForwardingTestCase(unittest.TestCase):
|
|||||||
result = asyncio.run(serving.create_responses(request, raw_request=None))
|
result = asyncio.run(serving.create_responses(request, raw_request=None))
|
||||||
self.assertEqual(getattr(result, "status_code", None), 400)
|
self.assertEqual(getattr(result, "status_code", None), 400)
|
||||||
|
|
||||||
|
def test_kimi_k3_request_uses_chat_encoder_fields(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.chat_encoding_spec = "kimi_k3"
|
||||||
|
serving.default_chat_template_kwargs = {}
|
||||||
|
serving.template_manager.chat_template_name = None
|
||||||
|
serving.tokenizer_manager.tokenizer.apply_chat_template.return_value = [4, 5, 6]
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="Explain <|kimi_image_placeholder|>",
|
||||||
|
tools=[
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "lookup",
|
||||||
|
"parameters": {"type": "object"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
tool_choice="required",
|
||||||
|
reasoning={"effort": "high"},
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
_, request_prompts, engine_prompts, _ = asyncio.run(
|
||||||
|
serving._make_request(request, None, serving.tokenizer_manager.tokenizer)
|
||||||
|
)
|
||||||
|
|
||||||
|
call = serving.tokenizer_manager.tokenizer.apply_chat_template.call_args
|
||||||
|
self.assertEqual(
|
||||||
|
call.args[0][0]["content"], "Explain <| kimi_image_placeholder |>"
|
||||||
|
)
|
||||||
|
self.assertEqual(call.kwargs["thinking_effort"], "high")
|
||||||
|
self.assertEqual(call.kwargs["tool_choice"], "required")
|
||||||
|
self.assertEqual(call.kwargs["tools"][0]["function"]["name"], "lookup")
|
||||||
|
self.assertEqual(request_prompts, [[4, 5, 6]])
|
||||||
|
self.assertEqual(engine_prompts, [[4, 5, 6]])
|
||||||
|
|
||||||
|
|
||||||
|
class ReasoningRequestForwardingTestCase(unittest.TestCase):
|
||||||
|
def test_create_responses_uses_processed_reasoning_state(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.reasoning_parser = "deepseek-r1"
|
||||||
|
serving.default_chat_template_kwargs = {"thinking": False}
|
||||||
|
serving.template_manager.reasoning_config = ReasoningToggleConfig(
|
||||||
|
toggle_param="thinking", default_enabled=True
|
||||||
|
)
|
||||||
|
rendered = MessageProcessingResult(
|
||||||
|
prompt="prompt",
|
||||||
|
prompt_ids=[1, 2, 3],
|
||||||
|
image_data=None,
|
||||||
|
audio_data=None,
|
||||||
|
video_data=None,
|
||||||
|
modalities=[],
|
||||||
|
stop=[],
|
||||||
|
)
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
async def fake_generate(
|
||||||
|
request_id,
|
||||||
|
request_prompt,
|
||||||
|
adapted_request,
|
||||||
|
sampling_params,
|
||||||
|
context,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
captured["adapted_request"] = adapted_request
|
||||||
|
context.append_output(
|
||||||
|
{
|
||||||
|
"text": "done",
|
||||||
|
"meta_info": {
|
||||||
|
"prompt_tokens": 3,
|
||||||
|
"completion_tokens": 1,
|
||||||
|
"cached_tokens": 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
yield context
|
||||||
|
|
||||||
|
serving._generate_with_builtin_tools = fake_generate
|
||||||
|
request = ResponsesRequest(
|
||||||
|
model="x",
|
||||||
|
input="answer",
|
||||||
|
request_id="resp_reasoning",
|
||||||
|
store=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
serving, "_apply_conversation_template", return_value=rendered
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.entrypoints.openai.serving_responses.ReasoningParser"
|
||||||
|
) as parser_cls,
|
||||||
|
):
|
||||||
|
parser_cls.return_value.parse_non_stream.return_value = (None, "done")
|
||||||
|
response = asyncio.run(serving.create_responses(request))
|
||||||
|
|
||||||
|
self.assertEqual(response.status, "completed")
|
||||||
|
self.assertFalse(captured["adapted_request"].require_reasoning)
|
||||||
|
self.assertFalse(parser_cls.call_args.kwargs["force_reasoning"])
|
||||||
|
|
||||||
|
|
||||||
class InputItemNormalizationTestCase(unittest.TestCase):
|
class InputItemNormalizationTestCase(unittest.TestCase):
|
||||||
def test_function_call_becomes_assistant_tool_call(self):
|
def test_function_call_becomes_assistant_tool_call(self):
|
||||||
@@ -293,6 +393,7 @@ class FullResponseUsageTestCase(unittest.TestCase):
|
|||||||
tokenizer=serving.tokenizer_manager.tokenizer,
|
tokenizer=serving.tokenizer_manager.tokenizer,
|
||||||
request_metadata=metadata,
|
request_metadata=metadata,
|
||||||
created_time=123,
|
created_time=123,
|
||||||
|
require_reasoning=False,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -432,6 +533,7 @@ class OutputItemsTestCase(unittest.TestCase):
|
|||||||
self._function_tool_request(),
|
self._function_tool_request(),
|
||||||
"raw model output with <tool_call>",
|
"raw model output with <tool_call>",
|
||||||
tokenizer=Mock(),
|
tokenizer=Mock(),
|
||||||
|
require_reasoning=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
tool_calls = [
|
tool_calls = [
|
||||||
@@ -464,7 +566,10 @@ class OutputItemsTestCase(unittest.TestCase):
|
|||||||
[fake_call],
|
[fake_call],
|
||||||
)
|
)
|
||||||
output_items = serving._make_response_output_items(
|
output_items = serving._make_response_output_items(
|
||||||
self._function_tool_request(), "raw model output", tokenizer=Mock()
|
self._function_tool_request(),
|
||||||
|
"raw model output",
|
||||||
|
tokenizer=Mock(),
|
||||||
|
require_reasoning=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
types = [type(item).__name__ for item in output_items]
|
types = [type(item).__name__ for item in output_items]
|
||||||
@@ -489,7 +594,7 @@ class OutputItemsTestCase(unittest.TestCase):
|
|||||||
raw = '[{"name": "get_weather", "parameters": {"city": "Beijing"}}]'
|
raw = '[{"name": "get_weather", "parameters": {"city": "Beijing"}}]'
|
||||||
|
|
||||||
output_items = serving._make_response_output_items(
|
output_items = serving._make_response_output_items(
|
||||||
request, raw, tokenizer=Mock()
|
request, raw, tokenizer=Mock(), require_reasoning=False
|
||||||
)
|
)
|
||||||
|
|
||||||
tool_calls = [
|
tool_calls = [
|
||||||
@@ -524,7 +629,10 @@ class OutputItemsTestCase(unittest.TestCase):
|
|||||||
"sglang.srt.entrypoints.openai.serving_responses.FunctionCallParser"
|
"sglang.srt.entrypoints.openai.serving_responses.FunctionCallParser"
|
||||||
) as parser_cls:
|
) as parser_cls:
|
||||||
output_items = serving._make_response_output_items(
|
output_items = serving._make_response_output_items(
|
||||||
request, "just a plain answer", tokenizer=Mock()
|
request,
|
||||||
|
"just a plain answer",
|
||||||
|
tokenizer=Mock(),
|
||||||
|
require_reasoning=False,
|
||||||
)
|
)
|
||||||
parser_cls.assert_not_called()
|
parser_cls.assert_not_called()
|
||||||
|
|
||||||
|
|||||||
@@ -20,9 +20,10 @@ register_cpu_ci(est_time=4, suite="base-a-test-cpu")
|
|||||||
|
|
||||||
|
|
||||||
class _StreamFixture:
|
class _StreamFixture:
|
||||||
def __init__(self, serving, request):
|
def __init__(self, serving, request, *, require_reasoning=False):
|
||||||
self.serving = serving
|
self.serving = serving
|
||||||
self.request = request
|
self.request = request
|
||||||
|
self.require_reasoning = require_reasoning
|
||||||
self.request_metadata = RequestResponseMetadata(request_id=request.request_id)
|
self.request_metadata = RequestResponseMetadata(request_id=request.request_id)
|
||||||
|
|
||||||
def run(self, chunks):
|
def run(self, chunks):
|
||||||
@@ -39,6 +40,7 @@ class _StreamFixture:
|
|||||||
model_name="x",
|
model_name="x",
|
||||||
tokenizer=Mock(),
|
tokenizer=Mock(),
|
||||||
request_metadata=self.request_metadata,
|
request_metadata=self.request_metadata,
|
||||||
|
require_reasoning=self.require_reasoning,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -60,6 +62,20 @@ def _engine_chunk(text, completion_tokens, *, finish=False):
|
|||||||
|
|
||||||
|
|
||||||
class NonHarmonyStreamTestCase(unittest.TestCase):
|
class NonHarmonyStreamTestCase(unittest.TestCase):
|
||||||
|
def test_reasoning_parser_uses_processed_reasoning_state(self):
|
||||||
|
serving = make_serving()
|
||||||
|
serving.reasoning_parser = "deepseek-r1"
|
||||||
|
request = ResponsesRequest(model="x", input="hi", stream=True, store=False)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.entrypoints.openai.serving_responses.ReasoningParser"
|
||||||
|
) as parser_cls:
|
||||||
|
parser_cls.return_value.parse_stream_chunk.return_value = (None, "done")
|
||||||
|
fixture = _StreamFixture(serving, request, require_reasoning=True)
|
||||||
|
fixture.run([_engine_chunk("done", 1, finish=True)])
|
||||||
|
|
||||||
|
self.assertTrue(parser_cls.call_args.kwargs["force_reasoning"])
|
||||||
|
|
||||||
def test_emits_typed_sse_events_in_order(self):
|
def test_emits_typed_sse_events_in_order(self):
|
||||||
serving = make_serving()
|
serving = make_serving()
|
||||||
serving.reasoning_parser = None
|
serving.reasoning_parser = None
|
||||||
|
|||||||
@@ -0,0 +1,810 @@
|
|||||||
|
import json
|
||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import xgrammar as xgr
|
||||||
|
from xgrammar.testing import _is_grammar_accept_string
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
|
ChatCompletionRequest,
|
||||||
|
Function,
|
||||||
|
Tool,
|
||||||
|
ToolChoice,
|
||||||
|
ToolChoiceFuncName,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import ToolStrictLevel, envs
|
||||||
|
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
||||||
|
from sglang.srt.function_call.kimik3_detector import KimiK3Detector
|
||||||
|
from sglang.srt.function_call.kimik3_format import (
|
||||||
|
ARGUMENT_CLOSE,
|
||||||
|
CALL_CLOSE,
|
||||||
|
THINK_CLOSE,
|
||||||
|
TOOLS_CLOSE,
|
||||||
|
TOOLS_OPEN,
|
||||||
|
)
|
||||||
|
from sglang.srt.function_call.kimik3_structural_tag import (
|
||||||
|
get_kimik3_auto_tool_call_structural_tag,
|
||||||
|
get_kimik3_structural_tag,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_CLOSE_TOKEN = "<|close|>"
|
||||||
|
_CLOSE_TOKEN_ID = 256
|
||||||
|
_TOKENIZER_INFO = xgr.TokenizerInfo(
|
||||||
|
[bytes([token_id]) for token_id in range(256)] + [_CLOSE_TOKEN.encode()]
|
||||||
|
)
|
||||||
|
_TOKEN_COMPILER = xgr.GrammarCompiler(_TOKENIZER_INFO, cache_enabled=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _tool(name="weather", strict=True):
|
||||||
|
return Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name=name,
|
||||||
|
strict=strict,
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"city": {
|
||||||
|
"type": "string",
|
||||||
|
"pattern": "[A-Z][A-Za-z ]+",
|
||||||
|
},
|
||||||
|
"days": {
|
||||||
|
"type": "integer",
|
||||||
|
"minimum": 1,
|
||||||
|
"maximum": 10,
|
||||||
|
},
|
||||||
|
"unit": {
|
||||||
|
"type": ["string", "null"],
|
||||||
|
"enum": ["celsius", "fahrenheit", None],
|
||||||
|
},
|
||||||
|
"metadata": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"source": {"type": "string"}},
|
||||||
|
"required": ["source"],
|
||||||
|
"additionalProperties": False,
|
||||||
|
},
|
||||||
|
"tags": {
|
||||||
|
"type": "array",
|
||||||
|
"items": {"type": "string"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["city", "days"],
|
||||||
|
"additionalProperties": False,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _argument(key, argument_type, value):
|
||||||
|
return (
|
||||||
|
f'<|open|>argument key="{key}" type="{argument_type}"<|sep|>'
|
||||||
|
f"{value}{ARGUMENT_CLOSE}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _call(name, index, *arguments):
|
||||||
|
return (
|
||||||
|
f'<|open|>call tool="{name}" index="{index}"<|sep|>'
|
||||||
|
+ "".join(arguments)
|
||||||
|
+ CALL_CLOSE
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _tools_section(*calls):
|
||||||
|
return TOOLS_OPEN + "".join(calls) + TOOLS_CLOSE
|
||||||
|
|
||||||
|
|
||||||
|
def _grammar(tools, tool_choice="auto", thinking_mode=False, parallel_tool_calls=True):
|
||||||
|
structural_tag = get_kimik3_structural_tag(
|
||||||
|
tools,
|
||||||
|
tool_choice=tool_choice,
|
||||||
|
thinking_mode=thinking_mode,
|
||||||
|
parallel_tool_calls=parallel_tool_calls,
|
||||||
|
)
|
||||||
|
return xgr.Grammar.from_structural_tag(structural_tag)
|
||||||
|
|
||||||
|
|
||||||
|
def _accepts(grammar, value):
|
||||||
|
return _is_grammar_accept_string(grammar, value)
|
||||||
|
|
||||||
|
|
||||||
|
def _encode_with_close_token(value):
|
||||||
|
token_ids = []
|
||||||
|
start = 0
|
||||||
|
while (index := value.find(_CLOSE_TOKEN, start)) != -1:
|
||||||
|
token_ids.extend(value[start:index].encode())
|
||||||
|
token_ids.append(_CLOSE_TOKEN_ID)
|
||||||
|
start = index + len(_CLOSE_TOKEN)
|
||||||
|
token_ids.extend(value[start:].encode())
|
||||||
|
return token_ids
|
||||||
|
|
||||||
|
|
||||||
|
def _token_accepts(structural_tag, value):
|
||||||
|
compiled = _TOKEN_COMPILER.compile_structural_tag(structural_tag)
|
||||||
|
matcher = xgr.GrammarMatcher(compiled)
|
||||||
|
for token_id in _encode_with_close_token(value):
|
||||||
|
if not matcher.accept_token(token_id):
|
||||||
|
return False
|
||||||
|
return matcher.is_completed()
|
||||||
|
|
||||||
|
|
||||||
|
def _valid_weather_call(index=1):
|
||||||
|
return _call(
|
||||||
|
"weather",
|
||||||
|
index,
|
||||||
|
_argument("city", "string", "San Francisco"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
_argument("unit", "string", "celsius"),
|
||||||
|
_argument("metadata", "object", '{"source":"forecast"}'),
|
||||||
|
_argument("tags", "array", '["coastal","windy"]'),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_strict_schema_accepts_native_xtml_values():
|
||||||
|
grammar = _grammar([_tool()], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(_valid_weather_call()))
|
||||||
|
assert _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(
|
||||||
|
_call(
|
||||||
|
"weather",
|
||||||
|
1,
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "1"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"arguments",
|
||||||
|
[
|
||||||
|
(_argument("city", "string", "paris"), _argument("days", "number", "3")),
|
||||||
|
(
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "11"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
_argument("city", "number", "3"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
_argument("unit", "string", "kelvin"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
_argument("tags", "array", "[coastal]"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
_argument("unknown", "string", "value"),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_strict_schema_rejects_invalid_parameters(arguments):
|
||||||
|
grammar = _grammar([_tool()], tool_choice="required")
|
||||||
|
|
||||||
|
assert not _accepts(grammar, _tools_section(_call("weather", 1, *arguments)))
|
||||||
|
|
||||||
|
|
||||||
|
def test_required_allows_response_prefix_but_requires_tools():
|
||||||
|
grammar = _grammar([_tool()], tool_choice="required")
|
||||||
|
response = "<|open|>response<|sep|>Checking." "<|close|>response<|sep|>"
|
||||||
|
|
||||||
|
assert _accepts(grammar, response + _tools_section(_valid_weather_call()))
|
||||||
|
assert not _accepts(grammar, response)
|
||||||
|
|
||||||
|
|
||||||
|
def test_auto_allows_plain_response_or_multiple_tool_calls():
|
||||||
|
grammar = _grammar([_tool(), _tool("forecast")])
|
||||||
|
plain = (
|
||||||
|
"<|open|>response<|sep|>No tool needed."
|
||||||
|
"<|close|>response<|sep|><|close|>message<|sep|>"
|
||||||
|
)
|
||||||
|
calls = _tools_section(
|
||||||
|
_valid_weather_call(),
|
||||||
|
_call(
|
||||||
|
"forecast",
|
||||||
|
2,
|
||||||
|
_argument("city", "string", "London"),
|
||||||
|
_argument("days", "number", "2"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _accepts(grammar, plain)
|
||||||
|
assert _accepts(grammar, calls)
|
||||||
|
|
||||||
|
|
||||||
|
def test_named_tool_choice_forces_only_the_selected_tool():
|
||||||
|
grammar = _grammar(
|
||||||
|
[_tool(), _tool("forecast")],
|
||||||
|
tool_choice=ToolChoice(function=ToolChoiceFuncName(name="forecast")),
|
||||||
|
)
|
||||||
|
forecast_call = _call(
|
||||||
|
"forecast",
|
||||||
|
1,
|
||||||
|
_argument("city", "string", "London"),
|
||||||
|
_argument("days", "number", "2"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(forecast_call))
|
||||||
|
assert not _accepts(grammar, _tools_section(_valid_weather_call()))
|
||||||
|
|
||||||
|
|
||||||
|
def test_function_call_parser_uses_native_tag_for_named_tool_choice():
|
||||||
|
tool_choice = ToolChoice(function=ToolChoiceFuncName(name="forecast"))
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(), _tool("forecast")], "kimi_k3"
|
||||||
|
).get_structure_constraint(tool_choice)
|
||||||
|
assert constraint is not None
|
||||||
|
grammar = xgr.Grammar.from_structural_tag(constraint[1])
|
||||||
|
forecast_call = _call(
|
||||||
|
"forecast",
|
||||||
|
1,
|
||||||
|
_argument("city", "string", "London"),
|
||||||
|
_argument("days", "number", "2"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(forecast_call))
|
||||||
|
assert not _accepts(grammar, _tools_section(_valid_weather_call()))
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_strict_tool_keeps_xtml_structure_and_loose_parameters():
|
||||||
|
grammar = _grammar([_tool(strict=False)], tool_choice="required")
|
||||||
|
call = _call(
|
||||||
|
"weather",
|
||||||
|
1,
|
||||||
|
_argument("custom", "array", '["x",1]'),
|
||||||
|
_argument("other", "string", "raw text"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(call))
|
||||||
|
assert not _accepts(grammar, _tools_section("unstructured"))
|
||||||
|
|
||||||
|
|
||||||
|
def test_strict_schema_supports_refs_and_mixed_unions():
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="convert",
|
||||||
|
strict=True,
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"$defs": {
|
||||||
|
"mode": {
|
||||||
|
"type": "string",
|
||||||
|
"enum": ["fast", "safe"],
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"properties": {
|
||||||
|
"mode": {"$ref": "#/$defs/mode"},
|
||||||
|
"value": {
|
||||||
|
"anyOf": [
|
||||||
|
{"type": "string", "enum": ["auto"]},
|
||||||
|
{
|
||||||
|
"type": "integer",
|
||||||
|
"minimum": 2,
|
||||||
|
"maximum": 3,
|
||||||
|
},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["mode", "value"],
|
||||||
|
"additionalProperties": False,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
grammar = _grammar([tool], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(
|
||||||
|
_call(
|
||||||
|
"convert",
|
||||||
|
1,
|
||||||
|
_argument("mode", "string", "fast"),
|
||||||
|
_argument("value", "number", "2"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
assert not _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(
|
||||||
|
_call(
|
||||||
|
"convert",
|
||||||
|
1,
|
||||||
|
_argument("mode", "string", "unsafe"),
|
||||||
|
_argument("value", "number", "1"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_strict_schema_handles_number_enums_and_all_of_integer_constraints():
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="score",
|
||||||
|
strict=True,
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"value": {
|
||||||
|
"type": "number",
|
||||||
|
"enum": [1, 1.5],
|
||||||
|
},
|
||||||
|
"count": {
|
||||||
|
"allOf": [
|
||||||
|
{"type": "number"},
|
||||||
|
{"type": "integer", "minimum": 1},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["value", "count"],
|
||||||
|
"additionalProperties": False,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
grammar = _grammar([tool], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(
|
||||||
|
_call(
|
||||||
|
"score",
|
||||||
|
1,
|
||||||
|
_argument("value", "number", "1"),
|
||||||
|
_argument("count", "number", "2"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
assert not _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(
|
||||||
|
_call(
|
||||||
|
"score",
|
||||||
|
1,
|
||||||
|
_argument("value", "number", "2"),
|
||||||
|
_argument("count", "number", "1.5"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_strict_schema_preserves_additional_properties_default():
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="annotate",
|
||||||
|
strict=True,
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"label": {"type": "string"},
|
||||||
|
},
|
||||||
|
"required": ["label"],
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
grammar = _grammar([tool], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(
|
||||||
|
_call(
|
||||||
|
"annotate",
|
||||||
|
1,
|
||||||
|
_argument("label", "string", "sample"),
|
||||||
|
_argument("confidence", "number", "0.9"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_dynamic_argument_key_compiles_without_xgrammar_unicode_warning(capfd):
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="annotate",
|
||||||
|
strict=True,
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
grammar = _grammar([tool], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(_call("annotate", 1, _argument("置信度", "number", "0.9"))),
|
||||||
|
)
|
||||||
|
assert "Negative Character class" not in capfd.readouterr().err
|
||||||
|
|
||||||
|
|
||||||
|
def test_strict_empty_object_accepts_no_arguments_only():
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="ping",
|
||||||
|
strict=True,
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {},
|
||||||
|
"additionalProperties": False,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
grammar = _grammar([tool], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(_call("ping", 1)))
|
||||||
|
assert not _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(_call("ping", 1, _argument("unexpected", "string", "value"))),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_tool_strict_level_controls_native_tag_parameter_schema():
|
||||||
|
invalid_call = _tools_section(
|
||||||
|
_call(
|
||||||
|
"weather",
|
||||||
|
264,
|
||||||
|
_argument("city", "string", "paris"),
|
||||||
|
_argument("days", "number", "99"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
empty_call = _tools_section(_call("weather", 264))
|
||||||
|
|
||||||
|
with envs.SGLANG_TOOL_STRICT_LEVEL.override(ToolStrictLevel.OFF.value):
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(strict=False)], "kimi_k3"
|
||||||
|
).get_structure_constraint("auto")
|
||||||
|
assert constraint is not None
|
||||||
|
assert _token_accepts(constraint[1], invalid_call)
|
||||||
|
assert not _token_accepts(constraint[1], empty_call)
|
||||||
|
|
||||||
|
with envs.SGLANG_TOOL_STRICT_LEVEL.override(ToolStrictLevel.FUNCTION.value):
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(strict=False)], "kimi_k3"
|
||||||
|
).get_structure_constraint("auto")
|
||||||
|
assert constraint is not None
|
||||||
|
grammar = xgr.Grammar.from_structural_tag(constraint[1])
|
||||||
|
assert _accepts(grammar, invalid_call)
|
||||||
|
assert _accepts(grammar, empty_call)
|
||||||
|
|
||||||
|
with envs.SGLANG_TOOL_STRICT_LEVEL.override(ToolStrictLevel.PARAMETER.value):
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(strict=False)], "kimi_k3"
|
||||||
|
).get_structure_constraint("auto")
|
||||||
|
assert constraint is not None
|
||||||
|
assert not _accepts(
|
||||||
|
xgr.Grammar.from_structural_tag(constraint[1]), invalid_call
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_auto_hook_constrains_all_calls_and_requires_nonempty_values():
|
||||||
|
structural_tag = get_kimik3_auto_tool_call_structural_tag([_tool(strict=False)])
|
||||||
|
assert structural_tag is not None
|
||||||
|
first = _call(
|
||||||
|
"weather",
|
||||||
|
3,
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
)
|
||||||
|
second = _call(
|
||||||
|
"weather",
|
||||||
|
264,
|
||||||
|
_argument("city", "string", "London"),
|
||||||
|
_argument("days", "number", "2"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _token_accepts(structural_tag, _tools_section(first, second))
|
||||||
|
assert not _token_accepts(
|
||||||
|
structural_tag,
|
||||||
|
_tools_section(first, _call("weather", 264)),
|
||||||
|
)
|
||||||
|
assert not _token_accepts(
|
||||||
|
structural_tag,
|
||||||
|
_tools_section(
|
||||||
|
_call(
|
||||||
|
"weather",
|
||||||
|
3,
|
||||||
|
_argument("city", "string", ""),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
assert not _token_accepts(structural_tag, _tools_section(_call("weather", 3)))
|
||||||
|
|
||||||
|
|
||||||
|
def test_auto_hook_rejects_unknown_or_unclosed_calls():
|
||||||
|
structural_tag = get_kimik3_auto_tool_call_structural_tag([_tool(strict=False)])
|
||||||
|
assert structural_tag is not None
|
||||||
|
call = _call(
|
||||||
|
"weather",
|
||||||
|
49,
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
)
|
||||||
|
unknown = _call(
|
||||||
|
"forecast",
|
||||||
|
49,
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert not _token_accepts(structural_tag, _tools_section(unknown))
|
||||||
|
assert not _token_accepts(
|
||||||
|
structural_tag, TOOLS_OPEN + call.removesuffix(CALL_CLOSE)
|
||||||
|
)
|
||||||
|
assert not _token_accepts(structural_tag, TOOLS_OPEN + call)
|
||||||
|
|
||||||
|
|
||||||
|
def test_auto_hook_does_not_swallow_parser_visible_closes():
|
||||||
|
structural_tag = get_kimik3_auto_tool_call_structural_tag([_tool(strict=False)])
|
||||||
|
assert structural_tag is not None
|
||||||
|
output = (
|
||||||
|
TOOLS_OPEN
|
||||||
|
+ '<|open|>call tool="weather" index="23"<|sep|>'
|
||||||
|
+ _argument("city", "string", "Paris")
|
||||||
|
+ '<|open|>argument key="days" type="number"<|sep|>'
|
||||||
|
+ ARGUMENT_CLOSE
|
||||||
|
+ CALL_CLOSE
|
||||||
|
+ "3"
|
||||||
|
+ ARGUMENT_CLOSE
|
||||||
|
+ CALL_CLOSE
|
||||||
|
+ TOOLS_CLOSE
|
||||||
|
)
|
||||||
|
|
||||||
|
parsed = KimiK3Detector().detect_and_parse(output, [_tool(strict=False)])
|
||||||
|
|
||||||
|
assert json.loads(parsed.calls[0].parameters) == {"city": "Paris", "days": ""}
|
||||||
|
assert not _token_accepts(structural_tag, output)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"tool_strict_level",
|
||||||
|
[
|
||||||
|
ToolStrictLevel.OFF,
|
||||||
|
ToolStrictLevel.FUNCTION,
|
||||||
|
ToolStrictLevel.PARAMETER,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_parallel_tool_calls_false_rejects_second_call(tool_strict_level):
|
||||||
|
with envs.SGLANG_TOOL_STRICT_LEVEL.override(tool_strict_level.value):
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(strict=False)], "kimi_k3"
|
||||||
|
).get_structure_constraint("auto", parallel_tool_calls=False)
|
||||||
|
assert constraint is not None
|
||||||
|
first = _call(
|
||||||
|
"weather",
|
||||||
|
165,
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
_argument("days", "number", "3"),
|
||||||
|
)
|
||||||
|
second = _call(
|
||||||
|
"weather",
|
||||||
|
166,
|
||||||
|
_argument("city", "string", "London"),
|
||||||
|
_argument("days", "number", "2"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _token_accepts(constraint[1], _tools_section(first))
|
||||||
|
assert not _token_accepts(constraint[1], _tools_section(first, second))
|
||||||
|
assert not _token_accepts(
|
||||||
|
constraint[1], _tools_section(first) + _tools_section(second)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parallel_tool_calls_true_constrains_every_call():
|
||||||
|
grammar = _grammar(
|
||||||
|
[_tool(strict=False)],
|
||||||
|
parallel_tool_calls=True,
|
||||||
|
)
|
||||||
|
first = _call(
|
||||||
|
"weather",
|
||||||
|
7,
|
||||||
|
_argument("city", "string", "Paris"),
|
||||||
|
)
|
||||||
|
second = _call(
|
||||||
|
"weather",
|
||||||
|
19,
|
||||||
|
_argument("other", "array", '["loose"]'),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(first, second))
|
||||||
|
|
||||||
|
|
||||||
|
def test_parameter_level_constrains_every_parallel_call():
|
||||||
|
with envs.SGLANG_TOOL_STRICT_LEVEL.override(ToolStrictLevel.PARAMETER.value):
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(strict=False)], "kimi_k3"
|
||||||
|
).get_structure_constraint("auto")
|
||||||
|
assert constraint is not None
|
||||||
|
grammar = xgr.Grammar.from_structural_tag(constraint[1])
|
||||||
|
first = _valid_weather_call(index=3)
|
||||||
|
valid_second = _call(
|
||||||
|
"weather",
|
||||||
|
49,
|
||||||
|
_argument("city", "string", "London"),
|
||||||
|
_argument("days", "number", "2"),
|
||||||
|
)
|
||||||
|
invalid_second = _call(
|
||||||
|
"weather",
|
||||||
|
49,
|
||||||
|
_argument("city", "string", "london"),
|
||||||
|
_argument("days", "number", "99"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(first, valid_second))
|
||||||
|
assert not _accepts(grammar, _tools_section(first, invalid_second))
|
||||||
|
|
||||||
|
|
||||||
|
def test_strict_tool_without_parameters_compiles_to_empty_arguments():
|
||||||
|
"""SGLANG_TOOL_STRICT_LEVEL=2 marks every tool strict, including tools
|
||||||
|
that declare no parameters; the grammar build must not fail for them."""
|
||||||
|
tool = Tool(type="function", function=Function(name="ping", strict=True))
|
||||||
|
grammar = _grammar([tool], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(grammar, _tools_section(_call("ping", 1)))
|
||||||
|
assert not _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(_call("ping", 1, _argument("x", "string", "y"))),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parameter_level_keeps_constraint_with_no_parameter_tool():
|
||||||
|
"""A single no-parameter tool must not poison the whole request: a build
|
||||||
|
error here was swallowed and required fell back to a JSON-only grammar
|
||||||
|
the K3 parser cannot read."""
|
||||||
|
tools = [
|
||||||
|
_tool(strict=False),
|
||||||
|
Tool(type="function", function=Function(name="ping")),
|
||||||
|
]
|
||||||
|
with envs.SGLANG_TOOL_STRICT_LEVEL.override(ToolStrictLevel.PARAMETER.value):
|
||||||
|
constraint = FunctionCallParser(tools, "kimi_k3").get_structure_constraint(
|
||||||
|
"required"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert constraint is not None
|
||||||
|
assert constraint[0] == "structural_tag"
|
||||||
|
|
||||||
|
|
||||||
|
def test_all_of_number_branches_do_not_narrow_to_integer():
|
||||||
|
"""allOf with only number branches was intersected down to integer,
|
||||||
|
silently dropping non-integer enum values."""
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="scale",
|
||||||
|
strict=True,
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"factor": {"allOf": [{"type": "number"}], "enum": [1.5, 2]},
|
||||||
|
},
|
||||||
|
"required": ["factor"],
|
||||||
|
"additionalProperties": False,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
grammar = _grammar([tool], tool_choice="required")
|
||||||
|
|
||||||
|
assert _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(_call("scale", 1, _argument("factor", "number", "1.5"))),
|
||||||
|
)
|
||||||
|
assert not _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(_call("scale", 1, _argument("factor", "number", "3"))),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_auto_hook_serializes_into_sampling_parameters():
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(strict=False)], "kimi_k3"
|
||||||
|
).get_structure_constraint("auto")
|
||||||
|
assert constraint is not None
|
||||||
|
request = ChatCompletionRequest(
|
||||||
|
model="test",
|
||||||
|
messages=[{"role": "user", "content": "Weather?"}],
|
||||||
|
max_completion_tokens=16,
|
||||||
|
)
|
||||||
|
|
||||||
|
sampling_params = request.to_sampling_params(
|
||||||
|
stop=[],
|
||||||
|
model_generation_config={},
|
||||||
|
tool_call_constraint=constraint,
|
||||||
|
)
|
||||||
|
|
||||||
|
serialized = json.loads(sampling_params["structural_tag"])
|
||||||
|
assert serialized["type"] == "structural_tag"
|
||||||
|
assert serialized["format"]["type"] == "triggered_tags"
|
||||||
|
|
||||||
|
|
||||||
|
def test_auto_hook_forces_one_typed_property_when_none_are_required():
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="search",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"query": {"type": "string"},
|
||||||
|
"limit": {"type": "integer"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
structural_tag = get_kimik3_auto_tool_call_structural_tag([tool])
|
||||||
|
assert structural_tag is not None
|
||||||
|
|
||||||
|
assert _token_accepts(
|
||||||
|
structural_tag,
|
||||||
|
_tools_section(_call("search", 1, _argument("limit", "number", "3"))),
|
||||||
|
)
|
||||||
|
assert not _token_accepts(structural_tag, _tools_section(_call("search", 1)))
|
||||||
|
|
||||||
|
|
||||||
|
def test_auto_hook_keeps_structure_for_ambiguous_required_argument_type():
|
||||||
|
tool = Tool(
|
||||||
|
type="function",
|
||||||
|
function=Function(
|
||||||
|
name="lookup",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"key": {"type": ["string", "integer"]},
|
||||||
|
},
|
||||||
|
"required": ["key"],
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
structural_tag = get_kimik3_auto_tool_call_structural_tag([tool])
|
||||||
|
assert structural_tag is not None
|
||||||
|
grammar = xgr.Grammar.from_structural_tag(structural_tag)
|
||||||
|
|
||||||
|
assert _accepts(
|
||||||
|
grammar,
|
||||||
|
_tools_section(_call("lookup", 9254, _argument("key", "number", "3"))),
|
||||||
|
)
|
||||||
|
assert not _accepts(
|
||||||
|
grammar,
|
||||||
|
TOOLS_OPEN + _call("lookup", 9254, _argument("key", "number", "3")),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parameter_level_applies_to_other_model_native_tags():
|
||||||
|
with envs.SGLANG_TOOL_STRICT_LEVEL.override(ToolStrictLevel.PARAMETER.value):
|
||||||
|
constraint = FunctionCallParser(
|
||||||
|
[_tool(strict=False)], "kimi_k2"
|
||||||
|
).get_structure_constraint("auto")
|
||||||
|
|
||||||
|
assert constraint is not None
|
||||||
|
serialized = constraint[1].model_dump_json()
|
||||||
|
assert '"properties"' in serialized
|
||||||
|
assert '"city"' in serialized
|
||||||
|
|
||||||
|
|
||||||
|
def test_reasoning_prefix_is_owned_by_exactly_one_layer():
|
||||||
|
tool = _tool()
|
||||||
|
wrapped_by_xgrammar = _grammar([tool], tool_choice="required", thinking_mode=True)
|
||||||
|
post_reasoning_only = _grammar([tool], tool_choice="required", thinking_mode=False)
|
||||||
|
output = "reasoning" + THINK_CLOSE + _tools_section(_valid_weather_call())
|
||||||
|
|
||||||
|
assert _accepts(wrapped_by_xgrammar, output)
|
||||||
|
assert not _accepts(post_reasoning_only, output)
|
||||||
|
assert _accepts(post_reasoning_only, _tools_section(_valid_weather_call()))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__]))
|
||||||
@@ -66,7 +66,7 @@ def _make_processor() -> SchedulerBatchResultProcessor:
|
|||||||
enable_overlap=False,
|
enable_overlap=False,
|
||||||
enable_overlap_mlx=False,
|
enable_overlap_mlx=False,
|
||||||
server_args=SimpleNamespace(enable_metrics=False),
|
server_args=SimpleNamespace(enable_metrics=False),
|
||||||
model_config=SimpleNamespace(think_end_id=None),
|
model_config=SimpleNamespace(think_end_ids=None),
|
||||||
token_to_kv_pool_allocator=None,
|
token_to_kv_pool_allocator=None,
|
||||||
tree_cache=None,
|
tree_cache=None,
|
||||||
hisparse_coordinator=None,
|
hisparse_coordinator=None,
|
||||||
@@ -134,5 +134,19 @@ class TestSpecV2GrammarTruncation(CustomTestCase):
|
|||||||
self.assertEqual(req.kv_committed_len, 3)
|
self.assertEqual(req.kv_committed_len, 3)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReasoningTokenAccounting(CustomTestCase):
|
||||||
|
def test_multi_token_end_can_span_decode_steps(self):
|
||||||
|
req = _make_req(terminate_after=99)
|
||||||
|
req.require_reasoning = True
|
||||||
|
processor = _make_processor()
|
||||||
|
processor.model_config.think_end_ids = [7, 8]
|
||||||
|
|
||||||
|
processor._maybe_update_reasoning_tokens(req, [10, 7])
|
||||||
|
processor._maybe_update_reasoning_tokens(req, [8, 11])
|
||||||
|
|
||||||
|
self.assertEqual(req.reasoning_tokens, 3)
|
||||||
|
self.assertTrue(req._is_reasoning_over)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,162 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.srt.function_call.kimik3_format import (
|
||||||
|
MESSAGE_CLOSE,
|
||||||
|
RESPONSE_CLOSE,
|
||||||
|
RESPONSE_OPEN,
|
||||||
|
THINK_CLOSE,
|
||||||
|
THINK_OPEN,
|
||||||
|
TOOLS_CLOSE,
|
||||||
|
TOOLS_OPEN,
|
||||||
|
)
|
||||||
|
from sglang.srt.parser.reasoning_parser import KimiK3Detector, ReasoningParser
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _stream(detector: KimiK3Detector, chunks: list[str]) -> tuple[str, str]:
|
||||||
|
reasoning = ""
|
||||||
|
content = ""
|
||||||
|
for chunk in chunks:
|
||||||
|
result = detector.parse_streaming_increment(chunk)
|
||||||
|
reasoning += result.reasoning_text
|
||||||
|
content += result.normal_text
|
||||||
|
return reasoning, content
|
||||||
|
|
||||||
|
|
||||||
|
def _chunks(text: str, size: int) -> list[str]:
|
||||||
|
return [text[index : index + size] for index in range(0, len(text), size)]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("text", "reasoning", "content"),
|
||||||
|
[
|
||||||
|
(
|
||||||
|
f"{THINK_OPEN}deep thought{THINK_CLOSE}"
|
||||||
|
f"{RESPONSE_OPEN}the answer{RESPONSE_CLOSE}{MESSAGE_CLOSE}",
|
||||||
|
"deep thought",
|
||||||
|
"the answer",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
f"thinking...{THINK_CLOSE}{RESPONSE_OPEN}done{RESPONSE_CLOSE}",
|
||||||
|
"thinking...",
|
||||||
|
"done",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
f"{RESPONSE_OPEN}plain reply{RESPONSE_CLOSE}{MESSAGE_CLOSE}",
|
||||||
|
"",
|
||||||
|
"plain reply",
|
||||||
|
),
|
||||||
|
("still going", "still going", ""),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_non_stream_reasoning_channels(text: str, reasoning: str, content: str) -> None:
|
||||||
|
detector = KimiK3Detector(force_reasoning=True)
|
||||||
|
result = detector.detect_and_parse(text)
|
||||||
|
assert result.reasoning_text == reasoning
|
||||||
|
assert result.normal_text == content
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_stream_tools_channel_passthrough() -> None:
|
||||||
|
tools_channel = (
|
||||||
|
f'{TOOLS_OPEN}<|open|>call tool="python" index="1"<|sep|>'
|
||||||
|
"<|close|>call<|sep|>"
|
||||||
|
f"{TOOLS_CLOSE}"
|
||||||
|
)
|
||||||
|
detector = KimiK3Detector(force_reasoning=True)
|
||||||
|
result = detector.detect_and_parse(
|
||||||
|
f"thought{THINK_CLOSE}{RESPONSE_OPEN}reply{RESPONSE_CLOSE}{tools_channel}"
|
||||||
|
)
|
||||||
|
assert result.reasoning_text == "thought"
|
||||||
|
assert result.normal_text == f"reply{tools_channel}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_stream_recovers_missing_think_separator() -> None:
|
||||||
|
detector = KimiK3Detector(force_reasoning=True)
|
||||||
|
result = detector.detect_and_parse(
|
||||||
|
f"thought{THINK_CLOSE.removesuffix('<|sep|>')}{RESPONSE_OPEN}"
|
||||||
|
f"reply{RESPONSE_CLOSE}"
|
||||||
|
)
|
||||||
|
assert result.reasoning_text == "thought"
|
||||||
|
assert result.normal_text == "reply"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("text", "reasoning", "content"),
|
||||||
|
[
|
||||||
|
("deep thought<|close|>", "deep thought", ""),
|
||||||
|
("deep thought<|close|>think", "deep thought", ""),
|
||||||
|
(f"{THINK_CLOSE}<|open|>", "", ""),
|
||||||
|
(f"{THINK_CLOSE}<|open|>response", "", ""),
|
||||||
|
(
|
||||||
|
f"{THINK_CLOSE}{RESPONSE_OPEN}the answer<|close|>response",
|
||||||
|
"",
|
||||||
|
"the answer",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
f"{THINK_CLOSE}{RESPONSE_OPEN}the answer{RESPONSE_CLOSE}<|close|>message",
|
||||||
|
"",
|
||||||
|
"the answer",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_non_stream_strips_partial_marker_suffixes(
|
||||||
|
text: str, reasoning: str, content: str
|
||||||
|
) -> None:
|
||||||
|
result = KimiK3Detector(force_reasoning=True).detect_and_parse(text)
|
||||||
|
assert result.reasoning_text == reasoning
|
||||||
|
assert result.normal_text == content
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_stream_preserves_non_marker_angle_bracket_suffix() -> None:
|
||||||
|
result = KimiK3Detector(force_reasoning=True).detect_and_parse(
|
||||||
|
f"{THINK_CLOSE}{RESPONSE_OPEN}answer <3"
|
||||||
|
)
|
||||||
|
assert result.normal_text == "answer <3"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("chunk_size", [1, 4, 13])
|
||||||
|
def test_streaming_split_markers(chunk_size: int) -> None:
|
||||||
|
detector = KimiK3Detector(force_reasoning=True)
|
||||||
|
text = (
|
||||||
|
f"{THINK_OPEN}deep thought{THINK_CLOSE}"
|
||||||
|
f"{RESPONSE_OPEN}the answer{RESPONSE_CLOSE}{MESSAGE_CLOSE}"
|
||||||
|
)
|
||||||
|
reasoning, content = _stream(detector, _chunks(text, chunk_size))
|
||||||
|
assert reasoning == "deep thought"
|
||||||
|
assert content == "the answer"
|
||||||
|
|
||||||
|
|
||||||
|
def test_streaming_tools_channel_passthrough() -> None:
|
||||||
|
tools_channel = (
|
||||||
|
f'{TOOLS_OPEN}<|open|>call tool="python" index="1"<|sep|>'
|
||||||
|
"<|close|>call<|sep|>"
|
||||||
|
f"{TOOLS_CLOSE}"
|
||||||
|
)
|
||||||
|
detector = KimiK3Detector(force_reasoning=True)
|
||||||
|
text = f"thought{THINK_CLOSE}{RESPONSE_OPEN}reply{RESPONSE_CLOSE}{tools_channel}"
|
||||||
|
reasoning, content = _stream(detector, _chunks(text, 5))
|
||||||
|
assert reasoning == "thought"
|
||||||
|
assert content == f"reply{tools_channel}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_streaming_recovers_missing_think_separator() -> None:
|
||||||
|
detector = KimiK3Detector(force_reasoning=True)
|
||||||
|
text = (
|
||||||
|
f"thought{THINK_CLOSE.removesuffix('<|sep|>')}{RESPONSE_OPEN}"
|
||||||
|
f"reply{RESPONSE_CLOSE}"
|
||||||
|
)
|
||||||
|
reasoning, content = _stream(detector, _chunks(text, 3))
|
||||||
|
assert reasoning == "thought"
|
||||||
|
assert content == "reply"
|
||||||
|
|
||||||
|
|
||||||
|
def test_reasoning_parser_registration() -> None:
|
||||||
|
assert isinstance(ReasoningParser("kimi_k3").detector, KimiK3Detector)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__]))
|
||||||
@@ -15,7 +15,7 @@ from sglang.srt.parser.template_detection import (
|
|||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(2.0, "base-a-test-cpu")
|
register_cpu_ci(est_time=2.0, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
class _DummyTokenizer:
|
class _DummyTokenizer:
|
||||||
@@ -872,6 +872,34 @@ class TestResolveAutoParsers(unittest.TestCase):
|
|||||||
self.assertEqual(args.reasoning_parser, "deepseek-v4")
|
self.assertEqual(args.reasoning_parser, "deepseek-v4")
|
||||||
self.assertEqual(args.tool_call_parser, "deepseekv4")
|
self.assertEqual(args.tool_call_parser, "deepseekv4")
|
||||||
|
|
||||||
|
def test_kimi_k3_arch_without_chat_template_uses_custom_encoder(self):
|
||||||
|
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
|
||||||
|
tokenizer = _DummyTokenizer([])
|
||||||
|
config = SimpleNamespace(
|
||||||
|
architectures=["KimiK3ForConditionalGeneration"], model_type="kimi_k3"
|
||||||
|
)
|
||||||
|
|
||||||
|
with _patch_hf_transformers_utils(
|
||||||
|
Mock(return_value=tokenizer), Mock(return_value=config)
|
||||||
|
):
|
||||||
|
resolve_auto_parsers(args)
|
||||||
|
|
||||||
|
self.assertEqual(args.reasoning_parser, "kimi_k3")
|
||||||
|
self.assertEqual(args.tool_call_parser, "kimi_k3")
|
||||||
|
|
||||||
|
def test_kimi_k3_model_type_without_architecture_uses_custom_encoder(self):
|
||||||
|
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
|
||||||
|
tokenizer = _DummyTokenizer([])
|
||||||
|
config = SimpleNamespace(architectures=None, model_type="kimi_k3")
|
||||||
|
|
||||||
|
with _patch_hf_transformers_utils(
|
||||||
|
Mock(return_value=tokenizer), Mock(return_value=config)
|
||||||
|
):
|
||||||
|
resolve_auto_parsers(args)
|
||||||
|
|
||||||
|
self.assertEqual(args.reasoning_parser, "kimi_k3")
|
||||||
|
self.assertEqual(args.tool_call_parser, "kimi_k3")
|
||||||
|
|
||||||
def test_deepseek_arch_fallback_runs_when_tokenizer_load_fails(self):
|
def test_deepseek_arch_fallback_runs_when_tokenizer_load_fails(self):
|
||||||
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
|
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
|
||||||
config = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"])
|
config = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"])
|
||||||
|
|||||||
Reference in New Issue
Block a user